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

Method encode

acestep/models/ace_step_transformer.py:376–412  ·  view source on GitHub ↗
(
        self,
        encoder_text_hidden_states: Optional[torch.Tensor] = None,
        text_attention_mask: Optional[torch.LongTensor] = None,
        speaker_embeds: Optional[torch.FloatTensor] = None,
        lyric_token_idx: Optional[torch.LongTensor] = None,
        lyric_mask: Optional[torch.LongTensor] = None,
    )

Source from the content-addressed store, hash-verified

374 return prompt_prenet_out
375
376 def encode(
377 self,
378 encoder_text_hidden_states: Optional[torch.Tensor] = None,
379 text_attention_mask: Optional[torch.LongTensor] = None,
380 speaker_embeds: Optional[torch.FloatTensor] = None,
381 lyric_token_idx: Optional[torch.LongTensor] = None,
382 lyric_mask: Optional[torch.LongTensor] = None,
383 ):
384
385 bs = encoder_text_hidden_states.shape[0]
386 device = encoder_text_hidden_states.device
387
388 # speaker embedding
389 encoder_spk_hidden_states = self.speaker_embedder(speaker_embeds).unsqueeze(1)
390 speaker_mask = torch.ones(bs, 1, device=device)
391
392 # genre embedding
393 encoder_text_hidden_states = self.genre_embedder(encoder_text_hidden_states)
394
395 # lyric
396 encoder_lyric_hidden_states = self.forward_lyric_encoder(
397 lyric_token_idx=lyric_token_idx,
398 lyric_mask=lyric_mask,
399 )
400
401 encoder_hidden_states = torch.cat(
402 [
403 encoder_spk_hidden_states,
404 encoder_text_hidden_states,
405 encoder_lyric_hidden_states,
406 ],
407 dim=1,
408 )
409 encoder_hidden_mask = torch.cat(
410 [speaker_mask, text_attention_mask, lyric_mask], dim=1
411 )
412 return encoder_hidden_states, encoder_hidden_mask
413
414 def decode(
415 self,

Callers 1

forwardMethod · 0.95

Calls 1

forward_lyric_encoderMethod · 0.95

Tested by

no test coverage detected