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

Method generate_audio

trainer-api.py:155–211  ·  view source on GitHub ↗
(
        self,
        prompt: str,
        duration: int,
        infer_steps: int,
        guidance_scale: float,
        omega_scale: float,
        seed: Optional[int],
    )

Source from the content-addressed store, hash-verified

153 return target_latents
154
155 def generate_audio(
156 self,
157 prompt: str,
158 duration: int,
159 infer_steps: int,
160 guidance_scale: float,
161 omega_scale: float,
162 seed: Optional[int],
163 ):
164 # Set random seed
165 if seed is not None:
166 random.seed(seed)
167 torch.manual_seed(seed)
168 else:
169 seed = random.randint(0, 2**32 - 1)
170 random.seed(seed)
171 torch.manual_seed(seed)
172
173 generator = torch.Generator(device=self.device).manual_seed(seed)
174
175 # Get text embeddings
176 encoder_text_hidden_states, text_attention_mask = self.get_text_embeddings(
177 [prompt], self.device
178 )
179
180 # Dummy speaker embeddings and lyrics (since not provided in API request)
181 bsz = 1
182 speaker_embds = torch.zeros(bsz, 512, device=self.device, dtype=encoder_text_hidden_states.dtype)
183 lyric_token_ids = torch.zeros(bsz, 256, device=self.device, dtype=torch.long)
184 lyric_mask = torch.zeros(bsz, 256, device=self.device, dtype=torch.long)
185
186 # Run diffusion process
187 pred_latents = self.diffusion_process(
188 duration=duration,
189 encoder_text_hidden_states=encoder_text_hidden_states,
190 text_attention_mask=text_attention_mask,
191 speaker_embds=speaker_embds,
192 lyric_token_ids=lyric_token_ids,
193 lyric_mask=lyric_mask,
194 random_generator=generator,
195 infer_steps=infer_steps,
196 guidance_scale=guidance_scale,
197 omega_scale=omega_scale,
198 )
199
200 # Decode latents to audio
201 audio_lengths = torch.tensor([int(duration * 44100)], device=self.device)
202 sr, pred_wavs = self.dcae.decode(pred_latents, audio_lengths=audio_lengths, sr=48000)
203
204 # Save audio
205 timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
206 output_dir = "generated_audio"
207 os.makedirs(output_dir, exist_ok=True)
208 audio_path = f"{output_dir}/generated_{timestamp}_{seed}.wav"
209 torchaudio.save(audio_path, pred_wavs.float().cpu(), sr)
210
211 return audio_path, sr, seed
212

Callers 1

generate_musicFunction · 0.80

Calls 3

get_text_embeddingsMethod · 0.95
diffusion_processMethod · 0.95
decodeMethod · 0.45

Tested by

no test coverage detected