From 32974dd6a9f30fad1eb3652a6637bcc607bbc763 Mon Sep 17 00:00:00 2001 From: WeberJulian Date: Tue, 13 Jul 2021 16:04:42 +0200 Subject: [PATCH 1/6] Fix test sentences synthesis --- TTS/trainer.py | 6 +++--- TTS/tts/models/base_tts.py | 12 ++++++------ 2 files changed, 9 insertions(+), 9 deletions(-) diff --git a/TTS/trainer.py b/TTS/trainer.py index c56be140..bbd9665a 100644 --- a/TTS/trainer.py +++ b/TTS/trainer.py @@ -764,11 +764,11 @@ class Trainer: """Run test and log the results. Test run must be defined by the model. Model must return figures and audios to be logged by the Tensorboard.""" if hasattr(self.model, "test_run"): - if hasattr(self.eval_loader.load_test_samples): + if hasattr(self.eval_loader, "load_test_samples"): samples = self.eval_loader.load_test_samples(1) figures, audios = self.model.test_run(samples) else: - figures, audios = self.model.test_run() + figures, audios = self.model.test_run(use_cuda=self.use_cuda, ap=self.ap) self.tb_logger.tb_test_audios(self.total_steps_done, audios, self.config.audio["sample_rate"]) self.tb_logger.tb_test_figures(self.total_steps_done, figures) @@ -790,7 +790,7 @@ class Trainer: self.train_epoch() if self.config.run_eval: self.eval_epoch() - if epoch >= self.config.test_delay_epochs and self.args.rank < 0: + if epoch >= self.config.test_delay_epochs and self.args.rank <= 0: self.test_run() self.c_logger.print_epoch_end( epoch, self.keep_avg_eval.avg_values if self.config.run_eval else self.keep_avg_train.avg_values diff --git a/TTS/tts/models/base_tts.py b/TTS/tts/models/base_tts.py index 2ec268d6..64c0ba6f 100644 --- a/TTS/tts/models/base_tts.py +++ b/TTS/tts/models/base_tts.py @@ -200,7 +200,7 @@ class BaseTTS(BaseModel): ) return loader - def test_run(self) -> Tuple[Dict, Dict]: + def test_run(self, use_cuda=True, ap=None) -> Tuple[Dict, Dict]: """Generic test run for `tts` models used by `Trainer`. You can override this for a different behaviour. @@ -212,14 +212,14 @@ class BaseTTS(BaseModel): test_audios = {} test_figures = {} test_sentences = self.config.test_sentences - aux_inputs = self._get_aux_inputs() + aux_inputs = self.get_aux_input() for idx, sen in enumerate(test_sentences): wav, alignment, model_outputs, _ = synthesis( - self.model, + self, sen, self.config, - self.use_cuda, - self.ap, + use_cuda, + ap, speaker_id=aux_inputs["speaker_id"], d_vector=aux_inputs["d_vector"], style_wav=aux_inputs["style_wav"], @@ -229,6 +229,6 @@ class BaseTTS(BaseModel): ).values() test_audios["{}-audio".format(idx)] = wav - test_figures["{}-prediction".format(idx)] = plot_spectrogram(model_outputs, self.ap, output_fig=False) + test_figures["{}-prediction".format(idx)] = plot_spectrogram(model_outputs, ap, output_fig=False) test_figures["{}-alignment".format(idx)] = plot_alignment(alignment, output_fig=False) return test_figures, test_audios From 7d92b309465b58005450b1cc6459ef6e119c8453 Mon Sep 17 00:00:00 2001 From: WeberJulian Date: Tue, 13 Jul 2021 23:00:34 +0200 Subject: [PATCH 2/6] Fix tests --- TTS/trainer.py | 11 +++++++---- TTS/tts/models/base_tts.py | 4 ++-- TTS/vocoder/models/wavegrad.py | 5 ++++- TTS/vocoder/models/wavernn.py | 10 ++++++---- 4 files changed, 19 insertions(+), 11 deletions(-) diff --git a/TTS/trainer.py b/TTS/trainer.py index bbd9665a..b2494bad 100644 --- a/TTS/trainer.py +++ b/TTS/trainer.py @@ -22,6 +22,7 @@ from torch.utils.data import DataLoader from TTS.config import load_config, register_config from TTS.tts.datasets import load_meta_data from TTS.tts.models import setup_model as setup_tts_model +from TTS.vocoder.models.wavegrad import Wavegrad from TTS.tts.utils.text.symbols import parse_symbols from TTS.utils.audio import AudioProcessor from TTS.utils.callbacks import TrainerCallback @@ -764,11 +765,13 @@ class Trainer: """Run test and log the results. Test run must be defined by the model. Model must return figures and audios to be logged by the Tensorboard.""" if hasattr(self.model, "test_run"): - if hasattr(self.eval_loader, "load_test_samples"): - samples = self.eval_loader.load_test_samples(1) - figures, audios = self.model.test_run(samples) + if isinstance(self.model, Wavegrad): + return None # TODO: Fix inference on WaveGrad + elif hasattr(self.eval_loader.dataset, "load_test_samples"): + samples = self.eval_loader.dataset.load_test_samples(1) + figures, audios = self.model.test_run(self.ap, samples, None, self.use_cuda) else: - figures, audios = self.model.test_run(use_cuda=self.use_cuda, ap=self.ap) + figures, audios = self.model.test_run(self.ap, self.use_cuda) self.tb_logger.tb_test_audios(self.total_steps_done, audios, self.config.audio["sample_rate"]) self.tb_logger.tb_test_figures(self.total_steps_done, figures) diff --git a/TTS/tts/models/base_tts.py b/TTS/tts/models/base_tts.py index 64c0ba6f..a30c5f02 100644 --- a/TTS/tts/models/base_tts.py +++ b/TTS/tts/models/base_tts.py @@ -70,7 +70,7 @@ class BaseTTS(BaseModel): def get_aux_input(self, **kwargs) -> Dict: """Prepare and return `aux_input` used by `forward()`""" - pass + return {"speaker_id": None, "style_wav": None, "d_vector": None} def format_batch(self, batch: Dict) -> Dict: """Generic batch formatting for `TTSDataset`. @@ -200,7 +200,7 @@ class BaseTTS(BaseModel): ) return loader - def test_run(self, use_cuda=True, ap=None) -> Tuple[Dict, Dict]: + def test_run(self, ap, use_cuda) -> Tuple[Dict, Dict]: """Generic test run for `tts` models used by `Trainer`. You can override this for a different behaviour. diff --git a/TTS/vocoder/models/wavegrad.py b/TTS/vocoder/models/wavegrad.py index 03d5160e..7781b5f5 100644 --- a/TTS/vocoder/models/wavegrad.py +++ b/TTS/vocoder/models/wavegrad.py @@ -261,13 +261,16 @@ class Wavegrad(BaseModel): def eval_log(self, ap: AudioProcessor, batch: Dict, outputs: Dict) -> Tuple[Dict, np.ndarray]: return None, None - def test_run(self, ap: AudioProcessor, samples: List[Dict], ouputs: Dict): # pylint: disable=unused-argument + def test_run(self, ap: AudioProcessor, samples: List[Dict], ouputs: Dict, use_cuda): # pylint: disable=unused-argument # setup noise schedule and inference noise_schedule = self.config["test_noise_schedule"] betas = np.linspace(noise_schedule["min_val"], noise_schedule["max_val"], noise_schedule["num_steps"]) self.compute_noise_level(betas) for sample in samples: + sample = self.format_batch(sample) x = sample["input"] + if use_cuda: + x = x.cuda() y = sample["waveform"] # compute voice y_pred = self.inference(x) diff --git a/TTS/vocoder/models/wavernn.py b/TTS/vocoder/models/wavernn.py index a5d89d5a..12a29a72 100644 --- a/TTS/vocoder/models/wavernn.py +++ b/TTS/vocoder/models/wavernn.py @@ -322,7 +322,7 @@ class Wavernn(BaseVocoder): with torch.no_grad(): if isinstance(mels, np.ndarray): - mels = torch.FloatTensor(mels).type_as(mels) + mels = torch.FloatTensor(mels) if mels.ndim == 2: mels = mels.unsqueeze(0) @@ -571,12 +571,14 @@ class Wavernn(BaseVocoder): @torch.no_grad() def test_run( - self, ap: AudioProcessor, samples: List[Dict], output: Dict # pylint: disable=unused-argument + self, ap: AudioProcessor, samples: List[Dict], output: Dict, use_cuda # pylint: disable=unused-argument ) -> Tuple[Dict, Dict]: figures = {} audios = {} for idx, sample in enumerate(samples): - x = sample["input"] + x = torch.FloatTensor(sample[0]) + if use_cuda: + x = x.cuda() y_hat = self.inference(x, self.config.batched, self.config.target_samples, self.config.overlap_samples) x_hat = ap.melspectrogram(y_hat) figures.update( @@ -585,7 +587,7 @@ class Wavernn(BaseVocoder): f"test_{idx}/prediction": plot_spectrogram(x_hat.T), } ) - audios.update({f"test_{idx}/audio", y_hat}) + audios.update({f"test_{idx}/audio": y_hat}) return figures, audios @staticmethod From c79a82ed07b79f0b802cc3f08dd2594da4a81547 Mon Sep 17 00:00:00 2001 From: WeberJulian Date: Tue, 13 Jul 2021 23:12:18 +0200 Subject: [PATCH 3/6] refix linter --- TTS/trainer.py | 7 ++++--- TTS/tts/models/glow_tts.py | 2 +- TTS/vocoder/models/__init__.py | 4 ++-- TTS/vocoder/models/wavegrad.py | 4 +++- TTS/vocoder/models/wavernn.py | 2 +- tests/test_speaker_encoder_train.py | 2 ++ 6 files changed, 13 insertions(+), 8 deletions(-) diff --git a/TTS/trainer.py b/TTS/trainer.py index b2494bad..fd316e78 100644 --- a/TTS/trainer.py +++ b/TTS/trainer.py @@ -22,7 +22,6 @@ from torch.utils.data import DataLoader from TTS.config import load_config, register_config from TTS.tts.datasets import load_meta_data from TTS.tts.models import setup_model as setup_tts_model -from TTS.vocoder.models.wavegrad import Wavegrad from TTS.tts.utils.text.symbols import parse_symbols from TTS.utils.audio import AudioProcessor from TTS.utils.callbacks import TrainerCallback @@ -41,6 +40,7 @@ from TTS.utils.logging import ConsoleLogger, TensorboardLogger from TTS.utils.trainer_utils import get_optimizer, get_scheduler, is_apex_available, setup_torch_training_env from TTS.vocoder.datasets.preprocess import load_wav_data, load_wav_feat_data from TTS.vocoder.models import setup_model as setup_vocoder_model +from TTS.vocoder.models.wavegrad import Wavegrad if platform.system() != "Windows": # https://github.com/pytorch/pytorch/issues/973 @@ -766,14 +766,15 @@ class Trainer: Model must return figures and audios to be logged by the Tensorboard.""" if hasattr(self.model, "test_run"): if isinstance(self.model, Wavegrad): - return None # TODO: Fix inference on WaveGrad - elif hasattr(self.eval_loader.dataset, "load_test_samples"): + return None # TODO: Fix inference on WaveGrad + if hasattr(self.eval_loader.dataset, "load_test_samples"): samples = self.eval_loader.dataset.load_test_samples(1) figures, audios = self.model.test_run(self.ap, samples, None, self.use_cuda) else: figures, audios = self.model.test_run(self.ap, self.use_cuda) self.tb_logger.tb_test_audios(self.total_steps_done, audios, self.config.audio["sample_rate"]) self.tb_logger.tb_test_figures(self.total_steps_done, figures) + return None def _fit(self) -> None: """🏃 train -> evaluate -> test for the number of epochs.""" diff --git a/TTS/tts/models/glow_tts.py b/TTS/tts/models/glow_tts.py index 9f235fad..b3bceb09 100755 --- a/TTS/tts/models/glow_tts.py +++ b/TTS/tts/models/glow_tts.py @@ -113,7 +113,7 @@ class GlowTTS(BaseTTS): @staticmethod def compute_outputs(attn, o_mean, o_log_scale, x_mask): - """ Compute and format the mode outputs with the given alignment map""" + """Compute and format the mode outputs with the given alignment map""" y_mean = torch.matmul(attn.squeeze(1).transpose(1, 2), o_mean.transpose(1, 2)).transpose( 1, 2 ) # [b, t', t], [b, t, d] -> [b, d, t'] diff --git a/TTS/vocoder/models/__init__.py b/TTS/vocoder/models/__init__.py index 9479095e..7c209af4 100644 --- a/TTS/vocoder/models/__init__.py +++ b/TTS/vocoder/models/__init__.py @@ -31,7 +31,7 @@ def setup_model(config: Coqpit): def setup_generator(c): - """ TODO: use config object as arguments""" + """TODO: use config object as arguments""" print(" > Generator Model: {}".format(c.generator_model)) MyModel = importlib.import_module("TTS.vocoder.models." + c.generator_model.lower()) MyModel = getattr(MyModel, to_camel(c.generator_model)) @@ -94,7 +94,7 @@ def setup_generator(c): def setup_discriminator(c): - """ TODO: use config objekt as arguments""" + """TODO: use config objekt as arguments""" print(" > Discriminator Model: {}".format(c.discriminator_model)) if "parallel_wavegan" in c.discriminator_model: MyModel = importlib.import_module("TTS.vocoder.models.parallel_wavegan_discriminator") diff --git a/TTS/vocoder/models/wavegrad.py b/TTS/vocoder/models/wavegrad.py index 7781b5f5..01b47a20 100644 --- a/TTS/vocoder/models/wavegrad.py +++ b/TTS/vocoder/models/wavegrad.py @@ -261,7 +261,9 @@ class Wavegrad(BaseModel): def eval_log(self, ap: AudioProcessor, batch: Dict, outputs: Dict) -> Tuple[Dict, np.ndarray]: return None, None - def test_run(self, ap: AudioProcessor, samples: List[Dict], ouputs: Dict, use_cuda): # pylint: disable=unused-argument + def test_run( + self, ap: AudioProcessor, samples: List[Dict], ouputs: Dict, use_cuda + ): # pylint: disable=unused-argument # setup noise schedule and inference noise_schedule = self.config["test_noise_schedule"] betas = np.linspace(noise_schedule["min_val"], noise_schedule["max_val"], noise_schedule["num_steps"]) diff --git a/TTS/vocoder/models/wavernn.py b/TTS/vocoder/models/wavernn.py index 12a29a72..90eee58e 100644 --- a/TTS/vocoder/models/wavernn.py +++ b/TTS/vocoder/models/wavernn.py @@ -571,7 +571,7 @@ class Wavernn(BaseVocoder): @torch.no_grad() def test_run( - self, ap: AudioProcessor, samples: List[Dict], output: Dict, use_cuda # pylint: disable=unused-argument + self, ap: AudioProcessor, samples: List[Dict], output: Dict, use_cuda # pylint: disable=unused-argument ) -> Tuple[Dict, Dict]: figures = {} audios = {} diff --git a/tests/test_speaker_encoder_train.py b/tests/test_speaker_encoder_train.py index 4419a00f..7901fe5a 100644 --- a/tests/test_speaker_encoder_train.py +++ b/tests/test_speaker_encoder_train.py @@ -6,6 +6,7 @@ from tests import get_device_id, get_tests_output_path, run_cli from TTS.config.shared_configs import BaseAudioConfig from TTS.speaker_encoder.speaker_encoder_config import SpeakerEncoderConfig + def run_test_train(): command = ( f"CUDA_VISIBLE_DEVICES='{get_device_id()}' python TTS/bin/train_encoder.py --config_path {config_path} " @@ -17,6 +18,7 @@ def run_test_train(): ) run_cli(command) + config_path = os.path.join(get_tests_output_path(), "test_speaker_encoder_config.json") output_path = os.path.join(get_tests_output_path(), "train_outputs") From 25832eb97bad61e56e15299eec2462c08b2798c5 Mon Sep 17 00:00:00 2001 From: WeberJulian Date: Thu, 15 Jul 2021 11:38:45 +0200 Subject: [PATCH 4/6] Changes for review --- TTS/trainer.py | 4 ++-- TTS/tts/models/base_tts.py | 4 ++-- TTS/vocoder/models/wavegrad.py | 7 ++----- TTS/vocoder/models/wavernn.py | 7 +++---- 4 files changed, 9 insertions(+), 13 deletions(-) diff --git a/TTS/trainer.py b/TTS/trainer.py index fd316e78..12f43563 100644 --- a/TTS/trainer.py +++ b/TTS/trainer.py @@ -769,9 +769,9 @@ class Trainer: return None # TODO: Fix inference on WaveGrad if hasattr(self.eval_loader.dataset, "load_test_samples"): samples = self.eval_loader.dataset.load_test_samples(1) - figures, audios = self.model.test_run(self.ap, samples, None, self.use_cuda) + figures, audios = self.model.test_run(self.ap, samples, None) else: - figures, audios = self.model.test_run(self.ap, self.use_cuda) + figures, audios = self.model.test_run(self.ap) self.tb_logger.tb_test_audios(self.total_steps_done, audios, self.config.audio["sample_rate"]) self.tb_logger.tb_test_figures(self.total_steps_done, figures) return None diff --git a/TTS/tts/models/base_tts.py b/TTS/tts/models/base_tts.py index a30c5f02..561b76fb 100644 --- a/TTS/tts/models/base_tts.py +++ b/TTS/tts/models/base_tts.py @@ -200,7 +200,7 @@ class BaseTTS(BaseModel): ) return loader - def test_run(self, ap, use_cuda) -> Tuple[Dict, Dict]: + def test_run(self, ap) -> Tuple[Dict, Dict]: """Generic test run for `tts` models used by `Trainer`. You can override this for a different behaviour. @@ -218,7 +218,7 @@ class BaseTTS(BaseModel): self, sen, self.config, - use_cuda, + "cuda" in str(next(self.parameters()).device), ap, speaker_id=aux_inputs["speaker_id"], d_vector=aux_inputs["d_vector"], diff --git a/TTS/vocoder/models/wavegrad.py b/TTS/vocoder/models/wavegrad.py index 01b47a20..9249f81c 100644 --- a/TTS/vocoder/models/wavegrad.py +++ b/TTS/vocoder/models/wavegrad.py @@ -261,9 +261,7 @@ class Wavegrad(BaseModel): def eval_log(self, ap: AudioProcessor, batch: Dict, outputs: Dict) -> Tuple[Dict, np.ndarray]: return None, None - def test_run( - self, ap: AudioProcessor, samples: List[Dict], ouputs: Dict, use_cuda - ): # pylint: disable=unused-argument + def test_run(self, ap: AudioProcessor, samples: List[Dict], ouputs: Dict): # pylint: disable=unused-argument # setup noise schedule and inference noise_schedule = self.config["test_noise_schedule"] betas = np.linspace(noise_schedule["min_val"], noise_schedule["max_val"], noise_schedule["num_steps"]) @@ -271,8 +269,7 @@ class Wavegrad(BaseModel): for sample in samples: sample = self.format_batch(sample) x = sample["input"] - if use_cuda: - x = x.cuda() + x = x.to(next(self.parameters()).device) y = sample["waveform"] # compute voice y_pred = self.inference(x) diff --git a/TTS/vocoder/models/wavernn.py b/TTS/vocoder/models/wavernn.py index 90eee58e..c2e47120 100644 --- a/TTS/vocoder/models/wavernn.py +++ b/TTS/vocoder/models/wavernn.py @@ -322,7 +322,7 @@ class Wavernn(BaseVocoder): with torch.no_grad(): if isinstance(mels, np.ndarray): - mels = torch.FloatTensor(mels) + mels = torch.FloatTensor(mels).to(str(next(self.parameters()).device)) if mels.ndim == 2: mels = mels.unsqueeze(0) @@ -571,14 +571,13 @@ class Wavernn(BaseVocoder): @torch.no_grad() def test_run( - self, ap: AudioProcessor, samples: List[Dict], output: Dict, use_cuda # pylint: disable=unused-argument + self, ap: AudioProcessor, samples: List[Dict], output: Dict # pylint: disable=unused-argument ) -> Tuple[Dict, Dict]: figures = {} audios = {} for idx, sample in enumerate(samples): x = torch.FloatTensor(sample[0]) - if use_cuda: - x = x.cuda() + x = x.to(next(self.parameters()).device) y_hat = self.inference(x, self.config.batched, self.config.target_samples, self.config.overlap_samples) x_hat = ap.melspectrogram(y_hat) figures.update( From 58cc414477df5dfdd897c56a860bcfb308002b7b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Eren=20G=C3=B6lge?= Date: Fri, 16 Jul 2021 13:02:25 +0200 Subject: [PATCH 5/6] Fix WaveGrad `test_run` --- TTS/trainer.py | 2 -- TTS/vocoder/datasets/wavegrad_dataset.py | 15 ++++++++++++++- TTS/vocoder/models/wavegrad.py | 19 ++++++++++++------- 3 files changed, 26 insertions(+), 10 deletions(-) diff --git a/TTS/trainer.py b/TTS/trainer.py index 12f43563..f3f45ebd 100644 --- a/TTS/trainer.py +++ b/TTS/trainer.py @@ -765,8 +765,6 @@ class Trainer: """Run test and log the results. Test run must be defined by the model. Model must return figures and audios to be logged by the Tensorboard.""" if hasattr(self.model, "test_run"): - if isinstance(self.model, Wavegrad): - return None # TODO: Fix inference on WaveGrad if hasattr(self.eval_loader.dataset, "load_test_samples"): samples = self.eval_loader.dataset.load_test_samples(1) figures, audios = self.model.test_run(self.ap, samples, None) diff --git a/TTS/vocoder/datasets/wavegrad_dataset.py b/TTS/vocoder/datasets/wavegrad_dataset.py index d99fc417..05e0fae8 100644 --- a/TTS/vocoder/datasets/wavegrad_dataset.py +++ b/TTS/vocoder/datasets/wavegrad_dataset.py @@ -2,6 +2,7 @@ import glob import os import random from multiprocessing import Manager +from typing import List, Tuple import numpy as np import torch @@ -67,7 +68,19 @@ class WaveGradDataset(Dataset): item = self.load_item(idx) return item - def load_test_samples(self, num_samples): + def load_test_samples(self, num_samples: int) -> List[Tuple]: + """Return test samples. + + Args: + num_samples (int): Number of samples to return. + + Returns: + List[Tuple]: melspectorgram and audio. + + Shapes: + - melspectrogram (Tensor): :math:`[C, T]` + - audio (Tensor): :math:`[T_audio]` + """ samples = [] return_segments = self.return_segments self.return_segments = False diff --git a/TTS/vocoder/models/wavegrad.py b/TTS/vocoder/models/wavegrad.py index 9249f81c..22d2a015 100644 --- a/TTS/vocoder/models/wavegrad.py +++ b/TTS/vocoder/models/wavegrad.py @@ -124,11 +124,16 @@ class Wavegrad(BaseModel): @torch.no_grad() def inference(self, x, y_n=None): - """x: B x D X T""" + """ + Shapes: + x: :math:`[B, C , T]` + y_n: :math:`[B, 1, T]` + """ if y_n is None: - y_n = torch.randn(x.shape[0], 1, self.hop_len * x.shape[-1], dtype=torch.float32).to(x) + y_n = torch.randn(x.shape[0], 1, self.hop_len * x.shape[-1]) else: - y_n = torch.FloatTensor(y_n).unsqueeze(0).unsqueeze(0).to(x) + y_n = torch.FloatTensor(y_n).unsqueeze(0).unsqueeze(0) + y_n = y_n.type_as(x) sqrt_alpha_hat = self.noise_level.to(x) for n in range(len(self.alpha) - 1, -1, -1): y_n = self.c1[n] * (y_n - self.c2[n] * self.forward(y_n, x, sqrt_alpha_hat[n].repeat(x.shape[0]))) @@ -267,10 +272,10 @@ class Wavegrad(BaseModel): betas = np.linspace(noise_schedule["min_val"], noise_schedule["max_val"], noise_schedule["num_steps"]) self.compute_noise_level(betas) for sample in samples: - sample = self.format_batch(sample) - x = sample["input"] - x = x.to(next(self.parameters()).device) - y = sample["waveform"] + x = sample[0] + x = x[None, : , :].to(next(self.parameters()).device) + y = sample[1] + y = y[None, :] # compute voice y_pred = self.inference(x) # compute spectrograms From 05c75aa9d5358065f86fc321a2c843b0f41add38 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Eren=20G=C3=B6lge?= Date: Fri, 16 Jul 2021 13:37:38 +0200 Subject: [PATCH 6/6] Fix linter issues --- TTS/trainer.py | 2 -- TTS/vocoder/models/wavegrad.py | 2 +- 2 files changed, 1 insertion(+), 3 deletions(-) diff --git a/TTS/trainer.py b/TTS/trainer.py index f3f45ebd..903aee5f 100644 --- a/TTS/trainer.py +++ b/TTS/trainer.py @@ -40,7 +40,6 @@ from TTS.utils.logging import ConsoleLogger, TensorboardLogger from TTS.utils.trainer_utils import get_optimizer, get_scheduler, is_apex_available, setup_torch_training_env from TTS.vocoder.datasets.preprocess import load_wav_data, load_wav_feat_data from TTS.vocoder.models import setup_model as setup_vocoder_model -from TTS.vocoder.models.wavegrad import Wavegrad if platform.system() != "Windows": # https://github.com/pytorch/pytorch/issues/973 @@ -772,7 +771,6 @@ class Trainer: figures, audios = self.model.test_run(self.ap) self.tb_logger.tb_test_audios(self.total_steps_done, audios, self.config.audio["sample_rate"]) self.tb_logger.tb_test_figures(self.total_steps_done, figures) - return None def _fit(self) -> None: """🏃 train -> evaluate -> test for the number of epochs.""" diff --git a/TTS/vocoder/models/wavegrad.py b/TTS/vocoder/models/wavegrad.py index 22d2a015..d2983be2 100644 --- a/TTS/vocoder/models/wavegrad.py +++ b/TTS/vocoder/models/wavegrad.py @@ -273,7 +273,7 @@ class Wavegrad(BaseModel): self.compute_noise_level(betas) for sample in samples: x = sample[0] - x = x[None, : , :].to(next(self.parameters()).device) + x = x[None, :, :].to(next(self.parameters()).device) y = sample[1] y = y[None, :] # compute voice