Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions src/diffsinger_task.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -93,7 +93,7 @@ def validation_step(self, sample, batch_idx):
else:
gt_f0 = denorm_f0(sample['f0'], sample['uv'], hparams)
pred_f0 = model_out.get('f0_denorm')
self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], is_mel=True, gt_f0=gt_f0, f0=pred_f0)
self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], gt_f0=gt_f0, pred_f0=pred_f0)
self.plot_mel(batch_idx, sample['mels'], model_out['mel_out'], name=f'diffmel_{batch_idx}')
self.plot_mel(batch_idx, sample['mels'], model_out['fs2_mel'], name=f'fs2mel_{batch_idx}')
return outputs
Expand DownExpand Up@@ -200,7 +200,7 @@ def validation_step(self, sample, batch_idx):
else:
gt_f0 = denorm_f0(sample['f0'], sample['uv'], hparams)
pred_f0 = model_out.get('f0_denorm')
self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], is_mel=True, gt_f0=gt_f0, f0=pred_f0)
self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], gt_f0=gt_f0, pred_f0=pred_f0)
self.plot_mel(batch_idx, sample['mels'], model_out['mel_out'], name=f'diffmel_{batch_idx}')
self.plot_mel(batch_idx, sample['mels'], fs2_mel, name=f'fs2mel_{batch_idx}')
return outputs
Expand DownExpand Up@@ -341,7 +341,7 @@ def validation_step(self, sample, batch_idx):
else:
gt_f0 = denorm_f0(sample['f0'], sample['uv'], hparams)
pred_f0 = model_out.get('f0_denorm')
self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], is_mel=True, gt_f0=gt_f0, f0=pred_f0)
self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], gt_f0=gt_f0, pred_f0=pred_f0)
self.plot_mel(batch_idx, sample['mels'], model_out['mel_out'], name=f'diffmel_{batch_idx}')
#self.plot_mel(batch_idx, sample['mels'], model_out['fs2_mel'], name=f'fs2mel_{batch_idx}')
if hparams['use_pitch_embed']:
Expand Down
24 changes: 13 additions & 11 deletions src/diffspeech_task.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -19,6 +19,7 @@ def __init__(self):
super(DiffSpeechTask, self).__init__()
self.dataset_cls = FastSpeechDataset
self.vocoder: BaseVocoder = get_vocoder_cls(hparams)()
self.logged_gt_wav = set()

def build_tts_model(self):
mel_bins = hparams['audio_num_mel_bins']
Expand DownExpand Up@@ -102,21 +103,22 @@ def validation_step(self, sample, batch_idx):
model_out = self.model(
txt_tokens, spk_embed=spk_embed, mel2ph=mel2ph, f0=f0, uv=uv, energy=energy, ref_mels=None, infer=True)
gt_f0 = denorm_f0(sample['f0'], sample['uv'], hparams)
self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], is_mel=True, gt_f0=gt_f0, f0=model_out.get('f0_denorm'))
self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], gt_f0=gt_f0,
pred_f0=model_out.get('f0_denorm'))
self.plot_mel(batch_idx, sample['mels'], model_out['mel_out'])
return outputs

############
# validation plots
############
def plot_wav(self, batch_idx, gt_wav, wav_out, is_mel=False, gt_f0=None, f0=None, name=None):
gt_wav = gt_wav[0].cpu().numpy()
wav_out = wav_out[0].cpu().numpy()
def plot_wav(self, batch_idx, gt_mel, pred_mel, gt_f0=None, pred_f0=None):
gt_mel = gt_mel[0].cpu().numpy()
pred_mel = pred_mel[0].cpu().numpy()
gt_f0 = gt_f0[0].cpu().numpy()
f0 = f0[0].cpu().numpy()
if is_mel:
gt_wav = self.vocoder.spec2wav(gt_wav, f0=gt_f0)
wav_out = self.vocoder.spec2wav(wav_out, f0=f0)
self.logger.experiment.add_audio(f'gt_{batch_idx}', gt_wav, sample_rate=hparams['audio_sample_rate'], global_step=self.global_step)
self.logger.experiment.add_audio(f'wav_{batch_idx}', wav_out, sample_rate=hparams['audio_sample_rate'], global_step=self.global_step)

pred_f0 = pred_f0[0].cpu().numpy()
if batch_idx not in self.logged_gt_wav:
gt_wav = self.vocoder.spec2wav(gt_mel, f0=gt_f0)
self.logger.experiment.add_audio(f'gt_{batch_idx}', gt_wav, sample_rate=hparams['audio_sample_rate'], global_step=self.global_step)
self.logged_gt_wav.add(batch_idx)
pred_wav = self.vocoder.spec2wav(pred_mel, f0=pred_f0)
self.logger.experiment.add_audio(f'pred_{batch_idx}', pred_wav, sample_rate=hparams['audio_sample_rate'], global_step=self.global_step)
2 changes: 1 addition & 1 deletion src/naive_task.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -82,7 +82,7 @@ def validation_step(self, sample, batch_idx):
else:
gt_f0 = denorm_f0(sample['f0'], sample['uv'], hparams)
pred_f0 = gt_f0
self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], is_mel=True, gt_f0=gt_f0, f0=pred_f0)
self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], gt_f0=gt_f0, pred_f0=pred_f0)
self.plot_mel(batch_idx, sample['mels'], model_out['mel_out'], name=f'diffmel_{batch_idx}')

return outputs
2 changes: 1 addition & 1 deletion src/vocoders/base_vocoder.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,7 @@ def get_vocoder_cls(hparams):


class BaseVocoder:
def spec2wav(self, mel):
def spec2wav(self, mel, **kwargs):
"""

:param mel: [T, 80]
Expand Down
2 changes: 1 addition & 1 deletion tts/tasks/fs2.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -285,10 +285,10 @@ def add_energy_loss(self, energy_pred, energy, losses):
# validation plots
############
def plot_mel(self, batch_idx, spec, spec_out, name=None):
spec_cat = torch.cat([spec, spec_out], -1)
name = f'mel_{batch_idx}' if name is None else name
vmin = hparams['mel_vmin']
vmax = hparams['mel_vmax']
spec_cat = torch.cat([(spec_out - spec).abs() + vmin, spec, spec_out], -1)
self.logger.experiment.add_figure(name, spec_to_figure(spec_cat[0], vmin, vmax), self.global_step)

def plot_dur(self, batch_idx, sample, model_out):
Expand Down
3 changes: 2 additions & 1 deletion utils/plot.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -8,8 +8,9 @@
def spec_to_figure(spec, vmin=None, vmax=None):
if isinstance(spec, torch.Tensor):
spec = spec.cpu().numpy()
fig = plt.figure(figsize=(12, 6))
fig = plt.figure(figsize=(12, 9))
plt.pcolor(spec.T, vmin=vmin, vmax=vmax)
plt.tight_layout()
return fig


Expand Down