MCPcopy Create free account
hub / github.com/MoonInTheRiver/DiffSinger / forward

Method forward

usr/diff/shallow_diffusion_tts.py:313–341  ·  view source on GitHub ↗
(self, txt_tokens, mel2ph=None, spk_embed=None,
                ref_mels=None, f0=None, uv=None, energy=None, infer=False)

Source from the content-addressed store, hash-verified

311 return loss
312
313 def forward(self, txt_tokens, mel2ph=None, spk_embed=None,
314 ref_mels=None, f0=None, uv=None, energy=None, infer=False):
315 b, *_, device = *txt_tokens.shape, txt_tokens.device
316 ret = self.fs2(txt_tokens, mel2ph, spk_embed, ref_mels, f0, uv, energy,
317 skip_decoder=(not infer), infer=infer)
318 cond = ret['decoder_inp'].transpose(1, 2)
319
320 if not infer:
321 t = torch.randint(0, self.K_step, (b,), device=device).long()
322 x = ref_mels
323 x = self.norm_spec(x)
324 x = x.transpose(1, 2)[:, None, :, :] # [B, 1, M, T]
325 ret['diff_loss'] = self.p_losses(x, t, cond)
326 # nonpadding = (mel2ph != 0).float()
327 # ret['diff_loss'] = self.p_losses(x, t, cond, nonpadding=nonpadding)
328 else:
329 ret['fs2_mel'] = ret['mel_out']
330 fs2_mels = ret['mel_out']
331 t = self.K_step
332 fs2_mels = self.norm_spec(fs2_mels)
333 fs2_mels = fs2_mels.transpose(1, 2)[:, None, :, :]
334
335 x = self.q_sample(x_start=fs2_mels, t=torch.tensor([t - 1], device=device).long())
336 for i in tqdm(reversed(range(0, t)), desc='sample time step', total=t):
337 x = self.p_sample(x, torch.full((b,), i, device=device, dtype=torch.long), cond)
338 x = x[:, 0].transpose(1, 2)
339 ret['mel_out'] = self.denorm_spec(x)
340
341 return ret
342
343 def norm_spec(self, x):
344 return (x - self.spec_min) / (self.spec_max - self.spec_min) * 2 - 1

Callers

nothing calls this directly

Calls 5

norm_specMethod · 0.95
p_lossesMethod · 0.95
q_sampleMethod · 0.95
p_sampleMethod · 0.95
denorm_specMethod · 0.95

Tested by

no test coverage detected