(
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,
)
| 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, |
no test coverage detected