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

Method diffusion_process

trainer-api.py:76–153  ·  view source on GitHub ↗
(
        self,
        duration,
        encoder_text_hidden_states,
        text_attention_mask,
        speaker_embds,
        lyric_token_ids,
        lyric_mask,
        random_generator=None,
        infer_steps=60,
        guidance_scale=15.0,
        omega_scale=10.0,
    )

Source from the content-addressed store, hash-verified

74 return last_hidden_states, attention_mask
75
76 def diffusion_process(
77 self,
78 duration,
79 encoder_text_hidden_states,
80 text_attention_mask,
81 speaker_embds,
82 lyric_token_ids,
83 lyric_mask,
84 random_generator=None,
85 infer_steps=60,
86 guidance_scale=15.0,
87 omega_scale=10.0,
88 ):
89 do_classifier_free_guidance = guidance_scale > 1.0
90 device = encoder_text_hidden_states.device
91 dtype = encoder_text_hidden_states.dtype
92 bsz = encoder_text_hidden_states.shape[0]
93
94 timesteps, num_inference_steps = retrieve_timesteps(
95 self.scheduler, num_inference_steps=infer_steps, device=device
96 )
97
98 frame_length = int(duration * 44100 / 512 / 8)
99 target_latents = randn_tensor(
100 shape=(bsz, 8, 16, frame_length),
101 generator=random_generator,
102 device=device,
103 dtype=dtype,
104 )
105 attention_mask = torch.ones(bsz, frame_length, device=device, dtype=dtype)
106
107 if do_classifier_free_guidance:
108 attention_mask = torch.cat([attention_mask] * 2, dim=0)
109 encoder_text_hidden_states = torch.cat(
110 [encoder_text_hidden_states, torch.zeros_like(encoder_text_hidden_states)],
111 0,
112 )
113 text_attention_mask = torch.cat([text_attention_mask] * 2, dim=0)
114 speaker_embds = torch.cat([speaker_embds, torch.zeros_like(speaker_embds)], 0)
115 lyric_token_ids = torch.cat([lyric_token_ids, torch.zeros_like(lyric_token_ids)], 0)
116 lyric_mask = torch.cat([lyric_mask, torch.zeros_like(lyric_mask)], 0)
117
118 momentum_buffer = MomentumBuffer()
119
120 for t in timesteps:
121 latent_model_input = (
122 torch.cat([target_latents] * 2) if do_classifier_free_guidance else target_latents
123 )
124 timestep = t.expand(latent_model_input.shape[0])
125 with torch.no_grad():
126 noise_pred = self.transformers(
127 hidden_states=latent_model_input,
128 attention_mask=attention_mask,
129 encoder_text_hidden_states=encoder_text_hidden_states,
130 text_attention_mask=text_attention_mask,
131 speaker_embeds=speaker_embds,
132 lyric_token_idx=lyric_token_ids,
133 lyric_mask=lyric_mask,

Callers 1

generate_audioMethod · 0.95

Calls 3

MomentumBufferClass · 0.90
apg_forwardFunction · 0.90
stepMethod · 0.45

Tested by

no test coverage detected