mirror of
https://github.com/wassname/TTS.git
synced 2026-09-12 12:11:47 +08:00
Make lint
This commit is contained in:
@@ -1,8 +1,6 @@
|
||||
import copy
|
||||
import os
|
||||
import unittest
|
||||
from TTS.tts.utils.speakers import SpeakerManager
|
||||
from TTS.utils.logging.tensorboard_logger import TensorboardLogger
|
||||
|
||||
import torch
|
||||
from torch import optim
|
||||
@@ -11,7 +9,9 @@ from tests import get_tests_data_path, get_tests_input_path, get_tests_output_pa
|
||||
from TTS.tts.configs.glow_tts_config import GlowTTSConfig
|
||||
from TTS.tts.layers.losses import GlowTTSLoss
|
||||
from TTS.tts.models.glow_tts import GlowTTS
|
||||
from TTS.tts.utils.speakers import SpeakerManager
|
||||
from TTS.utils.audio import AudioProcessor
|
||||
from TTS.utils.logging.tensorboard_logger import TensorboardLogger
|
||||
|
||||
# pylint: disable=unused-variable
|
||||
|
||||
@@ -31,7 +31,8 @@ def count_parameters(model):
|
||||
|
||||
|
||||
class TestGlowTTS(unittest.TestCase):
|
||||
def _create_inputs(self):
|
||||
@staticmethod
|
||||
def _create_inputs():
|
||||
input_dummy = torch.randint(0, 24, (8, 128)).long().to(device)
|
||||
input_lengths = torch.randint(100, 129, (8,)).long().to(device)
|
||||
input_lengths[-1] = 128
|
||||
@@ -40,7 +41,8 @@ class TestGlowTTS(unittest.TestCase):
|
||||
speaker_ids = torch.randint(0, 5, (8,)).long().to(device)
|
||||
return input_dummy, input_lengths, mel_spec, mel_lengths, speaker_ids
|
||||
|
||||
def _check_parameter_changes(self, model, model_ref):
|
||||
@staticmethod
|
||||
def _check_parameter_changes(model, model_ref):
|
||||
count = 0
|
||||
for param, param_ref in zip(model.parameters(), model_ref.parameters()):
|
||||
assert (param != param_ref).any(), "param {} with shape {} not updated!! \n{}\n{}".format(
|
||||
@@ -166,7 +168,7 @@ class TestGlowTTS(unittest.TestCase):
|
||||
|
||||
def _assert_inference_outputs(self, outputs, input_dummy, mel_spec):
|
||||
output_shape = outputs["model_outputs"].shape
|
||||
self.assertEqual(outputs["model_outputs"].shape[::2] , mel_spec.shape[::2])
|
||||
self.assertEqual(outputs["model_outputs"].shape[::2], mel_spec.shape[::2])
|
||||
self.assertEqual(outputs["logdet"], None)
|
||||
self.assertEqual(outputs["y_mean"].shape, output_shape)
|
||||
self.assertEqual(outputs["y_log_scale"].shape, output_shape)
|
||||
@@ -185,7 +187,12 @@ class TestGlowTTS(unittest.TestCase):
|
||||
def test_inference_with_d_vector(self):
|
||||
input_dummy, input_lengths, mel_spec, mel_lengths, speaker_ids = self._create_inputs()
|
||||
d_vector = torch.rand(8, 256).to(device)
|
||||
config = GlowTTSConfig(num_chars=32, use_d_vector_file=True, d_vector_dim=256, d_vector_file=os.path.join(get_tests_data_path(), "dummy_speakers.json"))
|
||||
config = GlowTTSConfig(
|
||||
num_chars=32,
|
||||
use_d_vector_file=True,
|
||||
d_vector_dim=256,
|
||||
d_vector_file=os.path.join(get_tests_data_path(), "dummy_speakers.json"),
|
||||
)
|
||||
model = GlowTTS.init_from_config(config, verbose=False).to(device)
|
||||
model.eval()
|
||||
outputs = model.inference(input_dummy, {"x_lengths": input_lengths, "d_vectors": d_vector})
|
||||
@@ -268,7 +275,9 @@ class TestGlowTTS(unittest.TestCase):
|
||||
model = GlowTTS.init_from_config(config, verbose=False).to(device)
|
||||
model.run_data_dep_init = False
|
||||
model.train()
|
||||
logger = TensorboardLogger(log_dir=os.path.join(get_tests_output_path(), "dummy_glow_tts_logs"), model_name = "glow_tts_test_train_log")
|
||||
logger = TensorboardLogger(
|
||||
log_dir=os.path.join(get_tests_output_path(), "dummy_glow_tts_logs"), model_name="glow_tts_test_train_log"
|
||||
)
|
||||
criterion = model.get_criterion()
|
||||
outputs, _ = model.train_step(batch, criterion)
|
||||
model.train_log(batch, outputs, logger, None, 1)
|
||||
@@ -316,14 +325,23 @@ class TestGlowTTS(unittest.TestCase):
|
||||
self.assertTrue(model.num_speakers == 2)
|
||||
self.assertTrue(hasattr(model, "emb_g"))
|
||||
|
||||
config = GlowTTSConfig(num_chars=32, num_speakers=2, use_speaker_embedding=True, speakers_file=os.path.join(get_tests_data_path(), "ljspeech", "speakers.json"))
|
||||
config = GlowTTSConfig(
|
||||
num_chars=32,
|
||||
num_speakers=2,
|
||||
use_speaker_embedding=True,
|
||||
speakers_file=os.path.join(get_tests_data_path(), "ljspeech", "speakers.json"),
|
||||
)
|
||||
model = GlowTTS.init_from_config(config, verbose=False).to(device)
|
||||
self.assertTrue(model.num_speakers == 10)
|
||||
self.assertTrue(hasattr(model, "emb_g"))
|
||||
|
||||
config = GlowTTSConfig(num_chars=32, use_d_vector_file=True, d_vector_dim=256, d_vector_file=os.path.join(get_tests_data_path(), "dummy_speakers.json"))
|
||||
config = GlowTTSConfig(
|
||||
num_chars=32,
|
||||
use_d_vector_file=True,
|
||||
d_vector_dim=256,
|
||||
d_vector_file=os.path.join(get_tests_data_path(), "dummy_speakers.json"),
|
||||
)
|
||||
model = GlowTTS.init_from_config(config, verbose=False).to(device)
|
||||
self.assertTrue(model.num_speakers == 1)
|
||||
self.assertTrue(not hasattr(model, "emb_g"))
|
||||
self.assertTrue(model.c_in_channels == config.d_vector_dim)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user