mirror of
https://github.com/wassname/TTS.git
synced 2026-09-09 11:16:00 +08:00
Batch update after data-loss
This commit is contained in:
@@ -0,0 +1,142 @@
|
||||
import os
|
||||
import unittest
|
||||
import numpy as np
|
||||
import torch as T
|
||||
from TTS.utils.audio import AudioProcessor
|
||||
from TTS.utils.generic_utils import load_config
|
||||
|
||||
file_path = os.path.dirname(os.path.realpath(__file__))
|
||||
INPUTPATH = os.path.join(file_path, 'inputs')
|
||||
OUTPATH = os.path.join(file_path, "outputs/audio_tests")
|
||||
os.makedirs(OUTPATH, exist_ok=True)
|
||||
|
||||
c = load_config(os.path.join(file_path, 'test_config.json'))
|
||||
|
||||
|
||||
class TestAudio(unittest.TestCase):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(TestAudio, self).__init__(*args, **kwargs)
|
||||
self.ap = AudioProcessor(**c.audio)
|
||||
|
||||
def test_audio_synthesis(self):
|
||||
""" 1. load wav
|
||||
2. set normalization parameters
|
||||
3. extract mel-spec
|
||||
4. invert to wav and save the output
|
||||
"""
|
||||
print(" > Sanity check for the process wav -> mel -> wav")
|
||||
|
||||
def _test(max_norm, signal_norm, symmetric_norm, clip_norm):
|
||||
self.ap.max_norm = max_norm
|
||||
self.ap.signal_norm = signal_norm
|
||||
self.ap.symmetric_norm = symmetric_norm
|
||||
self.ap.clip_norm = clip_norm
|
||||
wav = self.ap.load_wav(INPUTPATH + "/example_1.wav")
|
||||
mel = self.ap.melspectrogram(wav)
|
||||
wav_ = self.ap.inv_mel_spectrogram(mel)
|
||||
file_name = "/audio_test-melspec_max_norm_{}-signal_norm_{}-symmetric_{}-clip_norm_{}.wav"\
|
||||
.format(max_norm, signal_norm, symmetric_norm, clip_norm)
|
||||
print(" | > Creating wav file at : ", file_name)
|
||||
self.ap.save_wav(wav_, OUTPATH + file_name)
|
||||
|
||||
# maxnorm = 1.0
|
||||
_test(1., False, False, False)
|
||||
_test(1., True, False, False)
|
||||
_test(1., True, True, False)
|
||||
_test(1., True, False, True)
|
||||
_test(1., True, True, True)
|
||||
# maxnorm = 4.0
|
||||
_test(4., False, False, False)
|
||||
_test(4., True, False, False)
|
||||
_test(4., True, True, False)
|
||||
_test(4., True, False, True)
|
||||
_test(4., True, True, True)
|
||||
|
||||
def test_normalize(self):
|
||||
"""Check normalization and denormalization for range values and consistency """
|
||||
print(" > Testing normalization and denormalization.")
|
||||
wav = self.ap.load_wav(INPUTPATH + "/example_1.wav")
|
||||
self.ap.signal_norm = False
|
||||
x = self.ap.melspectrogram(wav)
|
||||
x_old = x
|
||||
|
||||
self.ap.signal_norm = True
|
||||
self.ap.symmetric_norm = False
|
||||
self.ap.clip_norm = False
|
||||
self.ap.max_norm = 4.0
|
||||
x_norm = self.ap._normalize(x)
|
||||
print(x_norm.max(), " -- ", x_norm.min())
|
||||
assert (x_old - x).sum() == 0
|
||||
# check value range
|
||||
assert x_norm.max() <= self.ap.max_norm + 1, x_norm.max()
|
||||
assert x_norm.min() >= 0 - 1, x_norm.min()
|
||||
# check denorm.
|
||||
x_ = self.ap._denormalize(x_norm)
|
||||
assert (x - x_).sum() < 1e-3, (x - x_).mean()
|
||||
|
||||
self.ap.signal_norm = True
|
||||
self.ap.symmetric_norm = False
|
||||
self.ap.clip_norm = True
|
||||
self.ap.max_norm = 4.0
|
||||
x_norm = self.ap._normalize(x)
|
||||
print(x_norm.max(), " -- ", x_norm.min())
|
||||
assert (x_old - x).sum() == 0
|
||||
# check value range
|
||||
assert x_norm.max() <= self.ap.max_norm, x_norm.max()
|
||||
assert x_norm.min() >= 0, x_norm.min()
|
||||
# check denorm.
|
||||
x_ = self.ap._denormalize(x_norm)
|
||||
assert (x - x_).sum() < 1e-3, (x - x_).mean()
|
||||
|
||||
self.ap.signal_norm = True
|
||||
self.ap.symmetric_norm = True
|
||||
self.ap.clip_norm = False
|
||||
self.ap.max_norm = 4.0
|
||||
x_norm = self.ap._normalize(x)
|
||||
print(x_norm.max(), " -- ", x_norm.min())
|
||||
assert (x_old - x).sum() == 0
|
||||
# check value range
|
||||
assert x_norm.max() <= self.ap.max_norm + 1, x_norm.max()
|
||||
assert x_norm.min() >= -self.ap.max_norm - 2, x_norm.min()
|
||||
assert x_norm.min() <= 0, x_norm.min()
|
||||
# check denorm.
|
||||
x_ = self.ap._denormalize(x_norm)
|
||||
assert (x - x_).sum() < 1e-3, (x - x_).mean()
|
||||
|
||||
self.ap.signal_norm = True
|
||||
self.ap.symmetric_norm = True
|
||||
self.ap.clip_norm = True
|
||||
self.ap.max_norm = 4.0
|
||||
x_norm = self.ap._normalize(x)
|
||||
print(x_norm.max(), " -- ", x_norm.min())
|
||||
assert (x_old - x).sum() == 0
|
||||
# check value range
|
||||
assert x_norm.max() <= self.ap.max_norm, x_norm.max()
|
||||
assert x_norm.min() >= -self.ap.max_norm, x_norm.min()
|
||||
assert x_norm.min() <= 0, x_norm.min()
|
||||
# check denorm.
|
||||
x_ = self.ap._denormalize(x_norm)
|
||||
assert (x - x_).sum() < 1e-3, (x - x_).mean()
|
||||
|
||||
self.ap.signal_norm = True
|
||||
self.ap.symmetric_norm = False
|
||||
self.ap.max_norm = 1.0
|
||||
x_norm = self.ap._normalize(x)
|
||||
print(x_norm.max(), " -- ", x_norm.min())
|
||||
assert (x_old - x).sum() == 0
|
||||
assert x_norm.max() <= self.ap.max_norm, x_norm.max()
|
||||
assert x_norm.min() >= 0, x_norm.min()
|
||||
x_ = self.ap._denormalize(x_norm)
|
||||
assert (x - x_).sum() < 1e-3
|
||||
|
||||
self.ap.signal_norm = True
|
||||
self.ap.symmetric_norm = True
|
||||
self.ap.max_norm = 1.0
|
||||
x_norm = self.ap._normalize(x)
|
||||
print(x_norm.max(), " -- ", x_norm.min())
|
||||
assert (x_old - x).sum() == 0
|
||||
assert x_norm.max() <= self.ap.max_norm, x_norm.max()
|
||||
assert x_norm.min() >= -self.ap.max_norm, x_norm.min()
|
||||
assert x_norm.min() < 0, x_norm.min()
|
||||
x_ = self.ap._denormalize(x_norm)
|
||||
assert (x - x_).sum() < 1e-3
|
||||
+407
-208
@@ -1,50 +1,49 @@
|
||||
import os
|
||||
import unittest
|
||||
import shutil
|
||||
import numpy as np
|
||||
|
||||
from torch.utils.data import DataLoader
|
||||
from TTS.utils.generic_utils import load_config
|
||||
from TTS.utils.audio import AudioProcessor
|
||||
from TTS.datasets import LJSpeech, Kusal
|
||||
from TTS.datasets import TTSDataset, TTSDatasetCached, TTSDatasetMemory
|
||||
from TTS.datasets.preprocess import ljspeech, tts_cache
|
||||
|
||||
file_path = os.path.dirname(os.path.realpath(__file__))
|
||||
OUTPATH = os.path.join(file_path, "outputs/loader_tests/")
|
||||
os.makedirs(OUTPATH, exist_ok=True)
|
||||
c = load_config(os.path.join(file_path, 'test_config.json'))
|
||||
ok_kusal = os.path.exists(c.data_path_Kusal)
|
||||
ok_ljspeech = os.path.exists(c.data_path_LJSpeech)
|
||||
ok_ljspeech = os.path.exists(c.data_path)
|
||||
|
||||
|
||||
class TestLJSpeechDataset(unittest.TestCase):
|
||||
class TestTTSDataset(unittest.TestCase):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(TestLJSpeechDataset, self).__init__(*args, **kwargs)
|
||||
super(TestTTSDataset, self).__init__(*args, **kwargs)
|
||||
self.max_loader_iter = 4
|
||||
self.ap = AudioProcessor(
|
||||
sample_rate=c.sample_rate,
|
||||
num_mels=c.num_mels,
|
||||
min_level_db=c.min_level_db,
|
||||
frame_shift_ms=c.frame_shift_ms,
|
||||
frame_length_ms=c.frame_length_ms,
|
||||
ref_level_db=c.ref_level_db,
|
||||
num_freq=c.num_freq,
|
||||
power=c.power,
|
||||
preemphasis=c.preemphasis)
|
||||
self.ap = AudioProcessor(**c.audio)
|
||||
|
||||
def _create_dataloader(self, batch_size, r, bgs):
|
||||
dataset = TTSDataset.MyDataset(
|
||||
c.data_path,
|
||||
'metadata.csv',
|
||||
r,
|
||||
c.text_cleaner,
|
||||
preprocessor=ljspeech,
|
||||
ap=self.ap,
|
||||
batch_group_size=bgs,
|
||||
min_seq_len=c.min_seq_len)
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=False,
|
||||
collate_fn=dataset.collate_fn,
|
||||
drop_last=True,
|
||||
num_workers=c.num_loader_workers)
|
||||
return dataloader, dataset
|
||||
|
||||
def test_loader(self):
|
||||
if ok_ljspeech:
|
||||
dataset = LJSpeech.MyDataset(
|
||||
os.path.join(c.data_path_LJSpeech),
|
||||
os.path.join(c.data_path_LJSpeech, 'metadata.csv'),
|
||||
c.r,
|
||||
c.text_cleaner,
|
||||
ap=self.ap,
|
||||
min_seq_len=c.min_seq_len)
|
||||
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
shuffle=True,
|
||||
collate_fn=dataset.collate_fn,
|
||||
drop_last=True,
|
||||
num_workers=c.num_loader_workers)
|
||||
dataloader, dataset = self._create_dataloader(2, c.r, 0)
|
||||
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
@@ -63,29 +62,158 @@ class TestLJSpeechDataset(unittest.TestCase):
|
||||
" !! Negative values in text_input: {}".format(check_count)
|
||||
# TODO: more assertion here
|
||||
assert linear_input.shape[0] == c.batch_size
|
||||
assert linear_input.shape[2] == self.ap.num_freq
|
||||
assert mel_input.shape[0] == c.batch_size
|
||||
assert mel_input.shape[2] == c.num_mels
|
||||
assert mel_input.shape[2] == c.audio['num_mels']
|
||||
# check normalization ranges
|
||||
if self.ap.symmetric_norm:
|
||||
assert mel_input.max() <= self.ap.max_norm
|
||||
assert mel_input.min() >= -self.ap.max_norm
|
||||
assert mel_input.min() < 0
|
||||
else:
|
||||
assert mel_input.max() <= self.ap.max_norm
|
||||
assert mel_input.min() >= 0
|
||||
|
||||
def test_batch_group_shuffle(self):
|
||||
if ok_ljspeech:
|
||||
dataset = LJSpeech.MyDataset(
|
||||
os.path.join(c.data_path_LJSpeech),
|
||||
os.path.join(c.data_path_LJSpeech, 'metadata.csv'),
|
||||
c.r,
|
||||
c.text_cleaner,
|
||||
ap=self.ap,
|
||||
batch_group_size=16,
|
||||
min_seq_len=c.min_seq_len)
|
||||
dataloader, dataset = self._create_dataloader(2, c.r, 16)
|
||||
last_length = 0
|
||||
frames = dataset.items
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
break
|
||||
text_input = data[0]
|
||||
text_lengths = data[1]
|
||||
linear_input = data[2]
|
||||
mel_input = data[3]
|
||||
mel_lengths = data[4]
|
||||
stop_target = data[5]
|
||||
item_idx = data[6]
|
||||
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
shuffle=True,
|
||||
collate_fn=dataset.collate_fn,
|
||||
drop_last=True,
|
||||
num_workers=c.num_loader_workers)
|
||||
avg_length = mel_lengths.numpy().mean()
|
||||
assert avg_length >= last_length
|
||||
dataloader.dataset.sort_items()
|
||||
assert frames[0] != dataloader.dataset.items[0]
|
||||
|
||||
frames = dataset.frames
|
||||
def test_padding_and_spec(self):
|
||||
if ok_ljspeech:
|
||||
dataloader, dataset = self._create_dataloader(1, 1, 0)
|
||||
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
break
|
||||
text_input = data[0]
|
||||
text_lengths = data[1]
|
||||
linear_input = data[2]
|
||||
mel_input = data[3]
|
||||
mel_lengths = data[4]
|
||||
stop_target = data[5]
|
||||
item_idx = data[6]
|
||||
|
||||
# check mel_spec consistency
|
||||
wav = self.ap.load_wav(item_idx[0])
|
||||
mel = self.ap.melspectrogram(wav)
|
||||
mel_dl = mel_input[0].cpu().numpy()
|
||||
assert (
|
||||
abs(mel.T).astype("float32") - abs(mel_dl[:-1])).sum() == 0
|
||||
|
||||
# check mel-spec correctness
|
||||
mel_spec = mel_input[0].cpu().numpy()
|
||||
wav = self.ap.inv_mel_spectrogram(mel_spec.T)
|
||||
self.ap.save_wav(wav, OUTPATH + '/mel_inv_dataloader.wav')
|
||||
shutil.copy(item_idx[0], OUTPATH + '/mel_target_dataloader.wav')
|
||||
|
||||
# check linear-spec
|
||||
linear_spec = linear_input[0].cpu().numpy()
|
||||
wav = self.ap.inv_spectrogram(linear_spec.T)
|
||||
self.ap.save_wav(wav, OUTPATH + '/linear_inv_dataloader.wav')
|
||||
shutil.copy(item_idx[0], OUTPATH + '/linear_target_dataloader.wav')
|
||||
|
||||
# check the last time step to be zero padded
|
||||
assert linear_input[0, -1].sum() == 0
|
||||
assert linear_input[0, -2].sum() != 0
|
||||
assert mel_input[0, -1].sum() == 0
|
||||
assert mel_input[0, -2].sum() != 0
|
||||
assert stop_target[0, -1] == 1
|
||||
assert stop_target[0, -2] == 0
|
||||
assert stop_target.sum() == 1
|
||||
assert len(mel_lengths.shape) == 1
|
||||
assert mel_lengths[0] == linear_input[0].shape[0]
|
||||
assert mel_lengths[0] == mel_input[0].shape[0]
|
||||
|
||||
# Test for batch size 2
|
||||
dataloader, dataset = self._create_dataloader(2, 1, 0)
|
||||
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
break
|
||||
text_input = data[0]
|
||||
text_lengths = data[1]
|
||||
linear_input = data[2]
|
||||
mel_input = data[3]
|
||||
mel_lengths = data[4]
|
||||
stop_target = data[5]
|
||||
item_idx = data[6]
|
||||
|
||||
if mel_lengths[0] > mel_lengths[1]:
|
||||
idx = 0
|
||||
else:
|
||||
idx = 1
|
||||
|
||||
# check the first item in the batch
|
||||
assert linear_input[idx, -1].sum() == 0
|
||||
assert linear_input[idx, -2].sum() != 0, linear_input
|
||||
assert mel_input[idx, -1].sum() == 0
|
||||
assert mel_input[idx, -2].sum() != 0, mel_input
|
||||
assert stop_target[idx, -1] == 1
|
||||
assert stop_target[idx, -2] == 0
|
||||
assert stop_target[idx].sum() == 1
|
||||
assert len(mel_lengths.shape) == 1
|
||||
assert mel_lengths[idx] == mel_input[idx].shape[0]
|
||||
assert mel_lengths[idx] == linear_input[idx].shape[0]
|
||||
|
||||
# check the second itme in the batch
|
||||
assert linear_input[1 - idx, -1].sum() == 0
|
||||
assert mel_input[1 - idx, -1].sum() == 0
|
||||
assert stop_target[1 - idx, -1] == 1
|
||||
assert len(mel_lengths.shape) == 1
|
||||
|
||||
# check batch conditions
|
||||
assert (linear_input * stop_target.unsqueeze(2)).sum() == 0
|
||||
assert (mel_input * stop_target.unsqueeze(2)).sum() == 0
|
||||
|
||||
|
||||
class TestTTSDatasetCached(unittest.TestCase):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(TestTTSDatasetCached, self).__init__(*args, **kwargs)
|
||||
self.max_loader_iter = 4
|
||||
self.c = load_config(os.path.join(c.data_path_cache, 'config.json'))
|
||||
self.ap = AudioProcessor(**self.c.audio)
|
||||
|
||||
def _create_dataloader(self, batch_size, r, bgs):
|
||||
|
||||
dataset = TTSDatasetCached.MyDataset(
|
||||
c.data_path_cache,
|
||||
'tts_metadata.csv',
|
||||
r,
|
||||
c.text_cleaner,
|
||||
preprocessor=tts_cache,
|
||||
ap=self.ap,
|
||||
batch_group_size=bgs,
|
||||
min_seq_len=c.min_seq_len)
|
||||
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=False,
|
||||
collate_fn=dataset.collate_fn,
|
||||
drop_last=True,
|
||||
num_workers=c.num_loader_workers)
|
||||
return dataloader, dataset
|
||||
|
||||
def test_loader(self):
|
||||
if ok_ljspeech:
|
||||
dataloader, dataset = self._create_dataloader(2, c.r, 0)
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
break
|
||||
@@ -102,32 +230,21 @@ class TestLJSpeechDataset(unittest.TestCase):
|
||||
assert check_count == 0, \
|
||||
" !! Negative values in text_input: {}".format(check_count)
|
||||
# TODO: more assertion here
|
||||
assert linear_input.shape[0] == c.batch_size
|
||||
assert mel_input.shape[0] == c.batch_size
|
||||
assert mel_input.shape[2] == c.num_mels
|
||||
dataloader.dataset.sort_frames()
|
||||
assert frames[0] != dataloader.dataset.frames[0]
|
||||
assert mel_input.shape[2] == c.audio['num_mels']
|
||||
|
||||
if self.ap.symmetric_norm:
|
||||
assert mel_input.max() <= self.ap.max_norm
|
||||
assert mel_input.min() >= -self.ap.max_norm
|
||||
assert mel_input.min() < 0
|
||||
else:
|
||||
assert mel_input.max() <= self.ap.max_norm
|
||||
assert mel_input.min() >= 0
|
||||
|
||||
def test_padding(self):
|
||||
def test_batch_group_shuffle(self):
|
||||
if ok_ljspeech:
|
||||
dataset = LJSpeech.MyDataset(
|
||||
os.path.join(c.data_path_LJSpeech),
|
||||
os.path.join(c.data_path_LJSpeech, 'metadata.csv'),
|
||||
1,
|
||||
c.text_cleaner,
|
||||
ap=self.ap,
|
||||
min_seq_len=c.min_seq_len)
|
||||
|
||||
# Test for batch size 1
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=1,
|
||||
shuffle=False,
|
||||
collate_fn=dataset.collate_fn,
|
||||
drop_last=True,
|
||||
num_workers=c.num_loader_workers)
|
||||
|
||||
dataloader, dataset = self._create_dataloader(2, c.r, 16)
|
||||
frames = dataset.items
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
break
|
||||
@@ -139,11 +256,51 @@ class TestLJSpeechDataset(unittest.TestCase):
|
||||
stop_target = data[5]
|
||||
item_idx = data[6]
|
||||
|
||||
neg_values = text_input[text_input < 0]
|
||||
check_count = len(neg_values)
|
||||
assert check_count == 0, \
|
||||
" !! Negative values in text_input: {}".format(check_count)
|
||||
# TODO: more assertion here
|
||||
assert mel_input.shape[0] == c.batch_size
|
||||
assert mel_input.shape[2] == c.audio['num_mels']
|
||||
dataloader.dataset.sort_items()
|
||||
assert frames[0] != dataloader.dataset.items[0]
|
||||
|
||||
def test_padding_and_spec(self):
|
||||
if ok_ljspeech:
|
||||
dataloader, dataset = self._create_dataloader(1, 1, 0)
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
break
|
||||
text_input = data[0]
|
||||
text_lengths = data[1]
|
||||
linear_input = data[2]
|
||||
mel_input = data[3]
|
||||
mel_lengths = data[4]
|
||||
stop_target = data[5]
|
||||
item_idx = data[6]
|
||||
|
||||
# check mel_spec consistency
|
||||
if item_idx[0].split('.')[-1] == 'npy':
|
||||
wav = np.load(item_idx[0])
|
||||
else:
|
||||
wav = self.ap.load_wav(item_idx[0])
|
||||
mel = self.ap.melspectrogram(wav)
|
||||
mel_dl = mel_input[0].cpu().numpy()
|
||||
assert (abs(mel.T).astype("float32") - abs(
|
||||
mel_dl[:-1])).sum() == 0, (
|
||||
abs(mel.T).astype("float32") - abs(mel_dl[:-1])).sum()
|
||||
|
||||
# check mel-spec correctness
|
||||
mel_spec = mel_input[0].cpu().numpy()
|
||||
wav = self.ap.inv_mel_spectrogram(mel_spec.T)
|
||||
self.ap.save_wav(wav,
|
||||
OUTPATH + '/mel_inv_dataloader_cache.wav')
|
||||
shutil.copy(item_idx[0], OUTPATH + '/mel_target_dataloader_cache.wav')
|
||||
|
||||
# check the last time step to be zero padded
|
||||
assert mel_input[0, -1].sum() == 0
|
||||
assert mel_input[0, -2].sum() != 0
|
||||
assert linear_input[0, -1].sum() == 0
|
||||
assert linear_input[0, -2].sum() != 0
|
||||
assert stop_target[0, -1] == 1
|
||||
assert stop_target[0, -2] == 0
|
||||
assert stop_target.sum() == 1
|
||||
@@ -151,14 +308,7 @@ class TestLJSpeechDataset(unittest.TestCase):
|
||||
assert mel_lengths[0] == mel_input[0].shape[0]
|
||||
|
||||
# Test for batch size 2
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
shuffle=False,
|
||||
collate_fn=dataset.collate_fn,
|
||||
drop_last=False,
|
||||
num_workers=c.num_loader_workers)
|
||||
|
||||
dataloader, dataset = self._create_dataloader(2, 1, 0)
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
break
|
||||
@@ -178,8 +328,6 @@ class TestLJSpeechDataset(unittest.TestCase):
|
||||
# check the first item in the batch
|
||||
assert mel_input[idx, -1].sum() == 0
|
||||
assert mel_input[idx, -2].sum() != 0, mel_input
|
||||
assert linear_input[idx, -1].sum() == 0
|
||||
assert linear_input[idx, -2].sum() != 0
|
||||
assert stop_target[idx, -1] == 1
|
||||
assert stop_target[idx, -2] == 0
|
||||
assert stop_target[idx].sum() == 1
|
||||
@@ -188,151 +336,202 @@ class TestLJSpeechDataset(unittest.TestCase):
|
||||
|
||||
# check the second itme in the batch
|
||||
assert mel_input[1 - idx, -1].sum() == 0
|
||||
assert linear_input[1 - idx, -1].sum() == 0
|
||||
assert stop_target[1 - idx, -1] == 1
|
||||
assert len(mel_lengths.shape) == 1
|
||||
|
||||
# check batch conditions
|
||||
assert (mel_input * stop_target.unsqueeze(2)).sum() == 0
|
||||
assert (linear_input * stop_target.unsqueeze(2)).sum() == 0
|
||||
|
||||
|
||||
class TestKusalDataset(unittest.TestCase):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(TestKusalDataset, self).__init__(*args, **kwargs)
|
||||
self.max_loader_iter = 4
|
||||
self.ap = AudioProcessor(
|
||||
sample_rate=c.sample_rate,
|
||||
num_mels=c.num_mels,
|
||||
min_level_db=c.min_level_db,
|
||||
frame_shift_ms=c.frame_shift_ms,
|
||||
frame_length_ms=c.frame_length_ms,
|
||||
ref_level_db=c.ref_level_db,
|
||||
num_freq=c.num_freq,
|
||||
power=c.power,
|
||||
preemphasis=c.preemphasis)
|
||||
# class TestTTSDatasetMemory(unittest.TestCase):
|
||||
# def __init__(self, *args, **kwargs):
|
||||
# super(TestTTSDatasetMemory, self).__init__(*args, **kwargs)
|
||||
# self.max_loader_iter = 4
|
||||
# self.c = load_config(os.path.join(c.data_path_cache, 'config.json'))
|
||||
# self.ap = AudioProcessor(**c.audio)
|
||||
|
||||
def test_loader(self):
|
||||
if ok_kusal:
|
||||
dataset = Kusal.MyDataset(
|
||||
os.path.join(c.data_path_Kusal),
|
||||
os.path.join(c.data_path_Kusal, 'prompts.txt'),
|
||||
c.r,
|
||||
c.text_cleaner,
|
||||
ap=self.ap,
|
||||
min_seq_len=c.min_seq_len)
|
||||
# def test_loader(self):
|
||||
# if ok_ljspeech:
|
||||
# dataset = TTSDatasetMemory.MyDataset(
|
||||
# c.data_path_cache,
|
||||
# 'tts_metadata.csv',
|
||||
# c.r,
|
||||
# c.text_cleaner,
|
||||
# preprocessor=tts_cache,
|
||||
# ap=self.ap,
|
||||
# min_seq_len=c.min_seq_len)
|
||||
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
shuffle=True,
|
||||
collate_fn=dataset.collate_fn,
|
||||
drop_last=True,
|
||||
num_workers=c.num_loader_workers)
|
||||
# dataloader = DataLoader(
|
||||
# dataset,
|
||||
# batch_size=2,
|
||||
# shuffle=True,
|
||||
# collate_fn=dataset.collate_fn,
|
||||
# drop_last=True,
|
||||
# num_workers=c.num_loader_workers)
|
||||
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
break
|
||||
text_input = data[0]
|
||||
text_lengths = data[1]
|
||||
linear_input = data[2]
|
||||
mel_input = data[3]
|
||||
mel_lengths = data[4]
|
||||
stop_target = data[5]
|
||||
item_idx = data[6]
|
||||
# for i, data in enumerate(dataloader):
|
||||
# if i == self.max_loader_iter:
|
||||
# break
|
||||
# text_input = data[0]
|
||||
# text_lengths = data[1]
|
||||
# linear_input = data[2]
|
||||
# mel_input = data[3]
|
||||
# mel_lengths = data[4]
|
||||
# stop_target = data[5]
|
||||
# item_idx = data[6]
|
||||
|
||||
neg_values = text_input[text_input < 0]
|
||||
check_count = len(neg_values)
|
||||
assert check_count == 0, \
|
||||
" !! Negative values in text_input: {}".format(check_count)
|
||||
# TODO: more assertion here
|
||||
assert linear_input.shape[0] == c.batch_size
|
||||
assert mel_input.shape[0] == c.batch_size
|
||||
assert mel_input.shape[2] == c.num_mels
|
||||
# neg_values = text_input[text_input < 0]
|
||||
# check_count = len(neg_values)
|
||||
# assert check_count == 0, \
|
||||
# " !! Negative values in text_input: {}".format(check_count)
|
||||
# # check mel-spec shape
|
||||
# assert mel_input.shape[0] == c.batch_size
|
||||
# assert mel_input.shape[2] == c.audio['num_mels']
|
||||
# assert mel_input.max() <= self.ap.max_norm
|
||||
# # check data range
|
||||
# if self.ap.symmetric_norm:
|
||||
# assert mel_input.max() <= self.ap.max_norm
|
||||
# assert mel_input.min() >= -self.ap.max_norm
|
||||
# assert mel_input.min() < 0
|
||||
# else:
|
||||
# assert mel_input.max() <= self.ap.max_norm
|
||||
# assert mel_input.min() >= 0
|
||||
|
||||
def test_padding(self):
|
||||
if ok_kusal:
|
||||
dataset = Kusal.MyDataset(
|
||||
os.path.join(c.data_path_Kusal),
|
||||
os.path.join(c.data_path_Kusal, 'prompts.txt'),
|
||||
1,
|
||||
c.text_cleaner,
|
||||
ap=self.ap,
|
||||
min_seq_len=c.min_seq_len)
|
||||
# def test_batch_group_shuffle(self):
|
||||
# if ok_ljspeech:
|
||||
# dataset = TTSDatasetMemory.MyDataset(
|
||||
# c.data_path_cache,
|
||||
# 'tts_metadata.csv',
|
||||
# c.r,
|
||||
# c.text_cleaner,
|
||||
# preprocessor=ljspeech,
|
||||
# ap=self.ap,
|
||||
# batch_group_size=16,
|
||||
# min_seq_len=c.min_seq_len)
|
||||
|
||||
# Test for batch size 1
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=1,
|
||||
shuffle=False,
|
||||
collate_fn=dataset.collate_fn,
|
||||
drop_last=True,
|
||||
num_workers=c.num_loader_workers)
|
||||
# dataloader = DataLoader(
|
||||
# dataset,
|
||||
# batch_size=2,
|
||||
# shuffle=True,
|
||||
# collate_fn=dataset.collate_fn,
|
||||
# drop_last=True,
|
||||
# num_workers=c.num_loader_workers)
|
||||
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
break
|
||||
text_input = data[0]
|
||||
text_lengths = data[1]
|
||||
linear_input = data[2]
|
||||
mel_input = data[3]
|
||||
mel_lengths = data[4]
|
||||
stop_target = data[5]
|
||||
item_idx = data[6]
|
||||
# frames = dataset.items
|
||||
# for i, data in enumerate(dataloader):
|
||||
# if i == self.max_loader_iter:
|
||||
# break
|
||||
# text_input = data[0]
|
||||
# text_lengths = data[1]
|
||||
# linear_input = data[2]
|
||||
# mel_input = data[3]
|
||||
# mel_lengths = data[4]
|
||||
# stop_target = data[5]
|
||||
# item_idx = data[6]
|
||||
|
||||
# check the last time step to be zero padded
|
||||
assert mel_input[0, -1].sum() == 0
|
||||
# assert mel_input[0, -2].sum() != 0
|
||||
assert linear_input[0, -1].sum() == 0
|
||||
# assert linear_input[0, -2].sum() != 0
|
||||
assert stop_target[0, -1] == 1
|
||||
assert stop_target[0, -2] == 0
|
||||
assert stop_target.sum() == 1
|
||||
assert len(mel_lengths.shape) == 1
|
||||
assert mel_lengths[0] == mel_input[0].shape[0]
|
||||
# neg_values = text_input[text_input < 0]
|
||||
# check_count = len(neg_values)
|
||||
# assert check_count == 0, \
|
||||
# " !! Negative values in text_input: {}".format(check_count)
|
||||
# assert mel_input.shape[0] == c.batch_size
|
||||
# assert mel_input.shape[2] == c.audio['num_mels']
|
||||
# dataloader.dataset.sort_items()
|
||||
# assert frames[0] != dataloader.dataset.items[0]
|
||||
|
||||
# Test for batch size 2
|
||||
dataloader = DataLoader(
|
||||
dataset,
|
||||
batch_size=2,
|
||||
shuffle=False,
|
||||
collate_fn=dataset.collate_fn,
|
||||
drop_last=False,
|
||||
num_workers=c.num_loader_workers)
|
||||
# def test_padding_and_spec(self):
|
||||
# if ok_ljspeech:
|
||||
# dataset = TTSDatasetMemory.MyDataset(
|
||||
# c.data_path_cache,
|
||||
# 'tts_meta_data.csv',
|
||||
# 1,
|
||||
# c.text_cleaner,
|
||||
# preprocessor=ljspeech,
|
||||
# ap=self.ap,
|
||||
# min_seq_len=c.min_seq_len)
|
||||
|
||||
for i, data in enumerate(dataloader):
|
||||
if i == self.max_loader_iter:
|
||||
break
|
||||
text_input = data[0]
|
||||
text_lengths = data[1]
|
||||
linear_input = data[2]
|
||||
mel_input = data[3]
|
||||
mel_lengths = data[4]
|
||||
stop_target = data[5]
|
||||
item_idx = data[6]
|
||||
# # Test for batch size 1
|
||||
# dataloader = DataLoader(
|
||||
# dataset,
|
||||
# batch_size=1,
|
||||
# shuffle=False,
|
||||
# collate_fn=dataset.collate_fn,
|
||||
# drop_last=True,
|
||||
# num_workers=c.num_loader_workers)
|
||||
|
||||
if mel_lengths[0] > mel_lengths[1]:
|
||||
idx = 0
|
||||
else:
|
||||
idx = 1
|
||||
# for i, data in enumerate(dataloader):
|
||||
# if i == self.max_loader_iter:
|
||||
# break
|
||||
# text_input = data[0]
|
||||
# text_lengths = data[1]
|
||||
# linear_input = data[2]
|
||||
# mel_input = data[3]
|
||||
# mel_lengths = data[4]
|
||||
# stop_target = data[5]
|
||||
# item_idx = data[6]
|
||||
|
||||
# check the first item in the batch
|
||||
assert mel_input[idx, -1].sum() == 0
|
||||
assert mel_input[idx, -2].sum() != 0, mel_input
|
||||
assert linear_input[idx, -1].sum() == 0
|
||||
assert linear_input[idx, -2].sum() != 0
|
||||
assert stop_target[idx, -1] == 1
|
||||
assert stop_target[idx, -2] == 0
|
||||
assert stop_target[idx].sum() == 1
|
||||
assert len(mel_lengths.shape) == 1
|
||||
assert mel_lengths[idx] == mel_input[idx].shape[0]
|
||||
# # check mel_spec consistency
|
||||
# if item_idx[0].split('.')[-1] == 'npy':
|
||||
# wav = np.load(item_idx[0])
|
||||
# else:
|
||||
# wav = self.ap.load_wav(item_idx[0])
|
||||
# mel = self.ap.melspectrogram(wav)
|
||||
# mel_dl = mel_input[0].cpu().numpy()
|
||||
# assert (
|
||||
# abs(mel.T).astype("float32") - abs(mel_dl[:-1])).sum() == 0
|
||||
|
||||
# check the second itme in the batch
|
||||
assert mel_input[1 - idx, -1].sum() == 0
|
||||
assert linear_input[1 - idx, -1].sum() == 0
|
||||
assert stop_target[1 - idx, -1] == 1
|
||||
assert len(mel_lengths.shape) == 1
|
||||
# # check mel-spec correctness
|
||||
# mel_spec = mel_input[0].cpu().numpy()
|
||||
# wav = self.ap.inv_mel_spectrogram(mel_spec.T)
|
||||
# self.ap.save_wav(wav, OUTPATH + '/mel_inv_dataloader_memo.wav')
|
||||
# shutil.copy(item_idx[0], OUTPATH + '/mel_target_dataloader_memo.wav')
|
||||
|
||||
# check batch conditions
|
||||
assert (mel_input * stop_target.unsqueeze(2)).sum() == 0
|
||||
assert (linear_input * stop_target.unsqueeze(2)).sum() == 0
|
||||
# # check the last time step to be zero padded
|
||||
# assert mel_input[0, -1].sum() == 0
|
||||
# assert mel_input[0, -2].sum() != 0
|
||||
# assert stop_target[0, -1] == 1
|
||||
# assert stop_target[0, -2] == 0
|
||||
# assert stop_target.sum() == 1
|
||||
# assert len(mel_lengths.shape) == 1
|
||||
# assert mel_lengths[0] == mel_input[0].shape[0]
|
||||
|
||||
# # Test for batch size 2
|
||||
# dataloader = DataLoader(
|
||||
# dataset,
|
||||
# batch_size=2,
|
||||
# shuffle=False,
|
||||
# collate_fn=dataset.collate_fn,
|
||||
# drop_last=False,
|
||||
# num_workers=c.num_loader_workers)
|
||||
|
||||
# for i, data in enumerate(dataloader):
|
||||
# if i == self.max_loader_iter:
|
||||
# break
|
||||
# text_input = data[0]
|
||||
# text_lengths = data[1]
|
||||
# linear_input = data[2]
|
||||
# mel_input = data[3]
|
||||
# mel_lengths = data[4]
|
||||
# stop_target = data[5]
|
||||
# item_idx = data[6]
|
||||
|
||||
# if mel_lengths[0] > mel_lengths[1]:
|
||||
# idx = 0
|
||||
# else:
|
||||
# idx = 1
|
||||
|
||||
# # check the first item in the batch
|
||||
# assert mel_input[idx, -1].sum() == 0
|
||||
# assert mel_input[idx, -2].sum() != 0, mel_input
|
||||
# assert stop_target[idx, -1] == 1
|
||||
# assert stop_target[idx, -2] == 0
|
||||
# assert stop_target[idx].sum() == 1
|
||||
# assert len(mel_lengths.shape) == 1
|
||||
# assert mel_lengths[idx] == mel_input[idx].shape[0]
|
||||
|
||||
# # check the second itme in the batch
|
||||
# assert mel_input[1 - idx, -1].sum() == 0
|
||||
# assert stop_target[1 - idx, -1] == 1
|
||||
# assert len(mel_lengths.shape) == 1
|
||||
|
||||
# # check batch conditions
|
||||
# assert (mel_input * stop_target.unsqueeze(2)).sum() == 0
|
||||
|
||||
@@ -21,8 +21,8 @@ c = load_config(os.path.join(file_path, 'test_config.json'))
|
||||
class TacotronTrainTest(unittest.TestCase):
|
||||
def test_train_step(self):
|
||||
input = torch.randint(0, 24, (8, 128)).long().to(device)
|
||||
mel_spec = torch.rand(8, 30, c.num_mels).to(device)
|
||||
linear_spec = torch.rand(8, 30, c.num_freq).to(device)
|
||||
mel_spec = torch.rand(8, 30, c.audio['num_mels']).to(device)
|
||||
linear_spec = torch.rand(8, 30, c.audio['num_freq']).to(device)
|
||||
mel_lengths = torch.randint(20, 30, (8, )).long().to(device)
|
||||
stop_targets = torch.zeros(8, 30, 1).float().to(device)
|
||||
|
||||
@@ -35,7 +35,7 @@ class TacotronTrainTest(unittest.TestCase):
|
||||
|
||||
criterion = L1LossMasked().to(device)
|
||||
criterion_st = nn.BCELoss().to(device)
|
||||
model = Tacotron(c.embedding_size, c.num_freq, c.num_mels,
|
||||
model = Tacotron(c.embedding_size, c.audio['num_freq'], c.audio['num_mels'],
|
||||
c.r).to(device)
|
||||
model.train()
|
||||
model_ref = copy.deepcopy(model)
|
||||
|
||||
+21
-17
@@ -1,16 +1,25 @@
|
||||
{
|
||||
"num_mels": 80,
|
||||
"num_freq": 1025,
|
||||
"sample_rate": 22050,
|
||||
"frame_length_ms": 50,
|
||||
"frame_shift_ms": 12.5,
|
||||
"preemphasis": 0.97,
|
||||
"min_level_db": -100,
|
||||
"ref_level_db": 20,
|
||||
"audio":{
|
||||
"audio_processor": "audio", // to use dictate different audio processors, if available.
|
||||
"num_mels": 80, // size of the mel spec frame.
|
||||
"num_freq": 1025, // number of stft frequency levels. Size of the linear spectogram frame.
|
||||
"sample_rate": 22050, // wav sample-rate. If different than the original data, it is resampled.
|
||||
"frame_length_ms": 50, // stft window length in ms.
|
||||
"frame_shift_ms": 12.5, // stft window hop-lengh in ms.
|
||||
"preemphasis": 0.97, // pre-emphasis to reduce spec noise and make it more structured. If 0.0, no -pre-emphasis.
|
||||
"min_level_db": -100, // normalization range
|
||||
"ref_level_db": 20, // reference level db, theoretically 20db is the sound of air.
|
||||
"power": 1.5, // value to sharpen wav signals after GL algorithm.
|
||||
"griffin_lim_iters": 30,// #griffin-lim iterations. 30-60 is a good range. Larger the value, slower the generation.
|
||||
"signal_norm": true, // normalize the spec values in range [0, 1]
|
||||
"symmetric_norm": true, // move normalization to range [-1, 1]
|
||||
"clip_norm": true, // clip normalized values into the range.
|
||||
"max_norm": 4, // scale normalization to range [-max_norm, max_norm] or [0, max_norm]
|
||||
"mel_fmin": 95, // minimum freq level for mel-spec. ~50 for male and ~95 for female voices. Tune for dataset!!
|
||||
"mel_fmax": 7600 // maximum freq level for mel-spec. Tune for dataset!!
|
||||
},
|
||||
"hidden_size": 128,
|
||||
"embedding_size": 256,
|
||||
"min_mel_freq": null,
|
||||
"max_mel_freq": null,
|
||||
"text_cleaner": "english_cleaners",
|
||||
|
||||
"epochs": 2000,
|
||||
@@ -21,16 +30,11 @@
|
||||
"r": 5,
|
||||
"mk": 1.0,
|
||||
"priority_freq": false,
|
||||
|
||||
|
||||
"griffin_lim_iters": 60,
|
||||
"power": 1.5,
|
||||
|
||||
"num_loader_workers": 4,
|
||||
|
||||
"save_step": 200,
|
||||
"data_path_LJSpeech": "/home/erogol/Data/LJSpeech-1.1",
|
||||
"data_path_Kusal": "/home/erogol/Data/Kusal",
|
||||
"data_path": "/home/erogol/Data/LJSpeech-1.1/",
|
||||
"data_path_cache": "/home/erogol/Data/LJSpeech-1.1/tts_cache/",
|
||||
"output_path": "result",
|
||||
"min_seq_len": 0,
|
||||
"log_dir": "/home/erogol/projects/TTS/logs/"
|
||||
|
||||
Reference in New Issue
Block a user