(self, text, audio_token, audio_token_len, text_token, text_token_len, embeddings=None,
prompt_text=torch.zeros(1, 0, dtype=torch.int32),
llm_prompt_audio_token=torch.zeros(1, 0, dtype=torch.int32),
flow_prompt_audio_token=torch.zeros(1, 0, dtype=torch.int32),
prompt_audio_feat=torch.zeros(1, 0, 80), sample_rate=48000, duration_to_gen = 30, task="continuation", trim = True, stream=False, **kwargs)
| 306 | yield {'music_audio': this_music_audio.cpu()} |
| 307 | torch.cuda.synchronize() |
| 308 | def batch_inference(self, text, audio_token, audio_token_len, text_token, text_token_len, embeddings=None, |
| 309 | prompt_text=torch.zeros(1, 0, dtype=torch.int32), |
| 310 | llm_prompt_audio_token=torch.zeros(1, 0, dtype=torch.int32), |
| 311 | flow_prompt_audio_token=torch.zeros(1, 0, dtype=torch.int32), |
| 312 | prompt_audio_feat=torch.zeros(1, 0, 80), sample_rate=48000, duration_to_gen = 30, task="continuation", trim = True, stream=False, **kwargs): |
| 313 | batch_size = text_token.shape[0] |
| 314 | inference_kwargs = { |
| 315 | 'text': text_token, |
| 316 | 'text_len': torch.tensor(text_token_len, dtype=torch.int32).to(self.device), |
| 317 | 'audio_token': audio_token, |
| 318 | 'audio_token_len': torch.tensor(audio_token_len, dtype=torch.int32).to(self.device), |
| 319 | 'prompt_text': torch.zeros(batch_size, 0, dtype=torch.int32).to(self.device), |
| 320 | 'prompt_text_len': torch.tensor(prompt_text.shape[1], dtype=torch.int32).to(self.device), |
| 321 | 'prompt_audio_token': torch.zeros(batch_size, 0, dtype=torch.int32).to(self.device), |
| 322 | 'prompt_audio_token_len': torch.tensor([llm_prompt_audio_token.shape[1]], dtype=torch.int32).to(self.device), |
| 323 | 'embeddings': embeddings, |
| 324 | 'duration_to_gen': duration_to_gen, |
| 325 | 'task': task |
| 326 | } |
| 327 | music_audios = [] |
| 328 | with autocast(device_type='cuda', enabled=self.fp16, dtype=self.dtype, cache_enabled=True): |
| 329 | data = self.llm.batch_inference(**inference_kwargs) |
| 330 | for i in range(data.shape[0]): |
| 331 | this_uuid = str(uuid.uuid1()) |
| 332 | this_music_token = data[i][data[i]!=0].unsqueeze(0) |
| 333 | if self.fast: |
| 334 | this_music_audio = self.semantictoken2wav(token=this_music_token) |
| 335 | else: |
| 336 | music_audio = self.token2wav(token=this_music_token, |
| 337 | token_len=torch.tensor([data[i][data[i]!=0].shape[0]]), |
| 338 | uuid=this_uuid, |
| 339 | sample_rate=sample_rate, |
| 340 | finalize=False) |
| 341 | music_audios.append({"music_audio":music_audio, "text":text[i]}) |
| 342 | |
| 343 | return music_audios |
no test coverage detected