(
self,
prompt: str,
duration: int,
infer_steps: int,
guidance_scale: float,
omega_scale: float,
seed: Optional[int],
)
| 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 |
no test coverage detected