From 4982d3c995621d1d68a11c58ddc81b23bedf22ee Mon Sep 17 00:00:00 2001 From: yqzhishen Date: Fri, 3 Mar 2023 13:11:46 +0800 Subject: [PATCH 1/6] Save gt wav only once to reduce log size --- src/diffsinger_task.py | 6 +++--- src/diffspeech_task.py | 23 +++++++++++++---------- src/naive_task.py | 2 +- src/vocoders/base_vocoder.py | 2 +- 4 files changed, 18 insertions(+), 15 deletions(-) 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..114159135 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,23 @@ 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, pred, gt_f0=None, pred_f0=None, name=None): + gt = gt[0].cpu().numpy() + pred = pred[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 = self.vocoder.spec2wav(gt, f0=gt_f0) + self.logger.experiment.add_audio(f'gt_{batch_idx}', gt, sample_rate=hparams['audio_sample_rate'], global_step=self.global_step) + self.logged_gt_wav.add(batch_idx) + pred = self.vocoder.spec2wav(pred, f0=pred_f0) + self.logger.experiment.add_audio(f'wav_{batch_idx}', pred, 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] From b96c66208c1eb5ce043624f0abf7f03764421db3 Mon Sep 17 00:00:00 2001 From: yqzhishen Date: Fri, 3 Mar 2023 13:14:16 +0800 Subject: [PATCH 2/6] Change name --- src/diffspeech_task.py | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/src/diffspeech_task.py b/src/diffspeech_task.py index 114159135..b94f7a0cb 100644 --- a/src/diffspeech_task.py +++ b/src/diffspeech_task.py @@ -111,15 +111,14 @@ def validation_step(self, sample, batch_idx): ############ # validation plots ############ - def plot_wav(self, batch_idx, gt, pred, gt_f0=None, pred_f0=None, name=None): - gt = gt[0].cpu().numpy() - pred = pred[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() pred_f0 = pred_f0[0].cpu().numpy() if batch_idx not in self.logged_gt_wav: - gt = self.vocoder.spec2wav(gt, f0=gt_f0) - self.logger.experiment.add_audio(f'gt_{batch_idx}', gt, sample_rate=hparams['audio_sample_rate'], global_step=self.global_step) + 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 = self.vocoder.spec2wav(pred, f0=pred_f0) - self.logger.experiment.add_audio(f'wav_{batch_idx}', pred, sample_rate=hparams['audio_sample_rate'], global_step=self.global_step) - + 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) From 3f272d92555d9efad6a91d86584da496f9158270 Mon Sep 17 00:00:00 2001 From: yqzhishen Date: Fri, 3 Mar 2023 18:45:39 +0800 Subject: [PATCH 3/6] Add mel difference to plot --- tts/tasks/fs2.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tts/tasks/fs2.py b/tts/tasks/fs2.py index 50c4e21ab..812404118 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, 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): From 899d5d7e5b4c2c9ba5098cc08d5584c319b1988e Mon Sep 17 00:00:00 2001 From: yqzhishen Date: Fri, 3 Mar 2023 19:06:40 +0800 Subject: [PATCH 4/6] Fix color --- tts/tasks/fs2.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tts/tasks/fs2.py b/tts/tasks/fs2.py index 812404118..1b6690b74 100644 --- a/tts/tasks/fs2.py +++ b/tts/tasks/fs2.py @@ -288,7 +288,7 @@ def plot_mel(self, batch_idx, spec, spec_out, name=None): 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, spec, spec_out], -1) + 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): From 02587e07d56e70aaf4044e3791cf98fc47949313 Mon Sep 17 00:00:00 2001 From: yqzhishen Date: Fri, 3 Mar 2023 19:23:04 +0800 Subject: [PATCH 5/6] Change plot size --- utils/plot.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/utils/plot.py b/utils/plot.py index bdca62a8c..6aec44a27 100644 --- a/utils/plot.py +++ b/utils/plot.py @@ -8,7 +8,7 @@ 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) return fig From fa55526f893eb3534199319559eeef24aa560447 Mon Sep 17 00:00:00 2001 From: yqzhishen Date: Fri, 3 Mar 2023 20:11:19 +0800 Subject: [PATCH 6/6] Tight layout --- utils/plot.py | 1 + 1 file changed, 1 insertion(+) diff --git a/utils/plot.py b/utils/plot.py index 6aec44a27..0e2a9ec82 100644 --- a/utils/plot.py +++ b/utils/plot.py @@ -10,6 +10,7 @@ def spec_to_figure(spec, vmin=None, vmax=None): spec = spec.cpu().numpy() fig = plt.figure(figsize=(12, 9)) plt.pcolor(spec.T, vmin=vmin, vmax=vmax) + plt.tight_layout() return fig