(
self,
duration,
encoder_text_hidden_states,
text_attention_mask,
speaker_embds,
lyric_token_ids,
lyric_mask,
random_generators=None,
infer_steps=60,
guidance_scale=15.0,
omega_scale=10.0,
)
| 615 | |
| 616 | @torch.no_grad() |
| 617 | def diffusion_process( |
| 618 | self, |
| 619 | duration, |
| 620 | encoder_text_hidden_states, |
| 621 | text_attention_mask, |
| 622 | speaker_embds, |
| 623 | lyric_token_ids, |
| 624 | lyric_mask, |
| 625 | random_generators=None, |
| 626 | infer_steps=60, |
| 627 | guidance_scale=15.0, |
| 628 | omega_scale=10.0, |
| 629 | ): |
| 630 | |
| 631 | do_classifier_free_guidance = True |
| 632 | if guidance_scale == 0.0 or guidance_scale == 1.0: |
| 633 | do_classifier_free_guidance = False |
| 634 | |
| 635 | device = encoder_text_hidden_states.device |
| 636 | dtype = encoder_text_hidden_states.dtype |
| 637 | bsz = encoder_text_hidden_states.shape[0] |
| 638 | |
| 639 | scheduler = FlowMatchEulerDiscreteScheduler( |
| 640 | num_train_timesteps=1000, |
| 641 | shift=3.0, |
| 642 | ) |
| 643 | |
| 644 | frame_length = int(duration * 44100 / 512 / 8) |
| 645 | timesteps, num_inference_steps = retrieve_timesteps( |
| 646 | scheduler, num_inference_steps=infer_steps, device=device, timesteps=None |
| 647 | ) |
| 648 | |
| 649 | target_latents = randn_tensor( |
| 650 | shape=(bsz, 8, 16, frame_length), |
| 651 | generator=random_generators, |
| 652 | device=device, |
| 653 | dtype=dtype, |
| 654 | ) |
| 655 | attention_mask = torch.ones(bsz, frame_length, device=device, dtype=dtype) |
| 656 | if do_classifier_free_guidance: |
| 657 | attention_mask = torch.cat([attention_mask] * 2, dim=0) |
| 658 | encoder_text_hidden_states = torch.cat( |
| 659 | [ |
| 660 | encoder_text_hidden_states, |
| 661 | torch.zeros_like(encoder_text_hidden_states), |
| 662 | ], |
| 663 | 0, |
| 664 | ) |
| 665 | text_attention_mask = torch.cat([text_attention_mask] * 2, dim=0) |
| 666 | |
| 667 | speaker_embds = torch.cat( |
| 668 | [speaker_embds, torch.zeros_like(speaker_embds)], 0 |
| 669 | ) |
| 670 | |
| 671 | lyric_token_ids = torch.cat( |
| 672 | [lyric_token_ids, torch.zeros_like(lyric_token_ids)], 0 |
| 673 | ) |
| 674 | lyric_mask = torch.cat([lyric_mask, torch.zeros_like(lyric_mask)], 0) |
no test coverage detected