diff --git a/src/diffsinger_task.py b/src/diffsinger_task.py index 19a0196e2..5c4802001 100644 --- a/src/diffsinger_task.py +++ b/src/diffsinger_task.py @@ -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 @@ -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 @@ -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']: diff --git a/src/diffspeech_task.py b/src/diffspeech_task.py index 12b0f41c2..b94f7a0cb 100644 --- a/src/diffspeech_task.py +++ b/src/diffspeech_task.py @@ -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'] @@ -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) diff --git a/src/naive_task.py b/src/naive_task.py index e1cb256aa..7e6983cc9 100644 --- a/src/naive_task.py +++ b/src/naive_task.py @@ -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 diff --git a/src/vocoders/base_vocoder.py b/src/vocoders/base_vocoder.py index fe49a9e4f..1f6c2a277 100644 --- a/src/vocoders/base_vocoder.py +++ b/src/vocoders/base_vocoder.py @@ -20,7 +20,7 @@ def get_vocoder_cls(hparams): class BaseVocoder: - def spec2wav(self, mel): + def spec2wav(self, mel, **kwargs): """ :param mel: [T, 80] diff --git a/tts/tasks/fs2.py b/tts/tasks/fs2.py index 50c4e21ab..1b6690b74 100644 --- a/tts/tasks/fs2.py +++ b/tts/tasks/fs2.py @@ -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): diff --git a/utils/plot.py b/utils/plot.py index bdca62a8c..0e2a9ec82 100644 --- a/utils/plot.py +++ b/utils/plot.py @@ -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