MCPcopy Create free account
hub / github.com/ace-step/ACE-Step / decode

Method decode

acestep/models/ace_step_transformer.py:414–525  ·  view source on GitHub ↗
(
        self,
        hidden_states: torch.Tensor,
        attention_mask: torch.Tensor,
        encoder_hidden_states: torch.Tensor,
        encoder_hidden_mask: torch.Tensor,
        timestep: Optional[torch.Tensor],
        ssl_hidden_states: Optional[List[torch.Tensor]] = None,
        output_length: int = 0,
        block_controlnet_hidden_states: Optional[
            Union[List[torch.Tensor], torch.Tensor]
        ] = None,
        controlnet_scale: Union[float, torch.Tensor] = 1.0,
        return_dict: bool = True,
    )

Source from the content-addressed store, hash-verified

412 return encoder_hidden_states, encoder_hidden_mask
413
414 def decode(
415 self,
416 hidden_states: torch.Tensor,
417 attention_mask: torch.Tensor,
418 encoder_hidden_states: torch.Tensor,
419 encoder_hidden_mask: torch.Tensor,
420 timestep: Optional[torch.Tensor],
421 ssl_hidden_states: Optional[List[torch.Tensor]] = None,
422 output_length: int = 0,
423 block_controlnet_hidden_states: Optional[
424 Union[List[torch.Tensor], torch.Tensor]
425 ] = None,
426 controlnet_scale: Union[float, torch.Tensor] = 1.0,
427 return_dict: bool = True,
428 ):
429
430 embedded_timestep = self.timestep_embedder(
431 self.time_proj(timestep).to(dtype=hidden_states.dtype)
432 )
433 temb = self.t_block(embedded_timestep)
434
435 hidden_states = self.proj_in(hidden_states)
436
437 # controlnet logic
438 if block_controlnet_hidden_states is not None:
439 control_condi = cross_norm(hidden_states, block_controlnet_hidden_states)
440 hidden_states = hidden_states + control_condi * controlnet_scale
441
442 inner_hidden_states = []
443
444 rotary_freqs_cis = self.rotary_emb(
445 hidden_states, seq_len=hidden_states.shape[1]
446 )
447 encoder_rotary_freqs_cis = self.rotary_emb(
448 encoder_hidden_states, seq_len=encoder_hidden_states.shape[1]
449 )
450
451 for index_block, block in enumerate(self.transformer_blocks):
452
453 if self.training and self.gradient_checkpointing:
454
455 hidden_states = torch.utils.checkpoint.checkpoint(
456 block,
457 hidden_states=hidden_states,
458 attention_mask=attention_mask,
459 encoder_hidden_states=encoder_hidden_states,
460 encoder_attention_mask=encoder_hidden_mask,
461 rotary_freqs_cis=rotary_freqs_cis,
462 rotary_freqs_cis_cross=encoder_rotary_freqs_cis,
463 temb=temb,
464 use_reentrant=False,
465 )
466
467 else:
468 hidden_states = block(
469 hidden_states=hidden_states,
470 attention_mask=attention_mask,
471 encoder_hidden_states=encoder_hidden_states,

Callers 1

forwardMethod · 0.95

Calls 2

cross_normFunction · 0.85

Tested by

no test coverage detected