Make style

This commit is contained in:
Eren Gölge
2022-02-25 11:26:59 +01:00
parent b3ed6ff6b7
commit 1f0c8179da
13 changed files with 42 additions and 29 deletions
+1 -1
View File
@@ -63,7 +63,7 @@ class TestTTSDataset(unittest.TestCase):
max_text_len=c.max_text_len,
min_audio_len=c.min_audio_len,
max_audio_len=c.max_audio_len,
start_by_longest=start_by_longest
start_by_longest=start_by_longest,
)
dataloader = DataLoader(
dataset,
+3 -3
View File
@@ -1,6 +1,6 @@
import torch as T
from TTS.tts.utils.helpers import average_over_durations, generate_path, segment, sequence_mask, rand_segments
from TTS.tts.utils.helpers import average_over_durations, generate_path, rand_segments, segment, sequence_mask
def average_over_durations_test(): # pylint: disable=no-self-use
@@ -57,12 +57,12 @@ def rand_segments_test():
assert segments.shape == (2, 3, 3)
assert all(seg_idxs >= 0), seg_idxs
try:
segments, _ = rand_segments(x, x_lens, segment_size=5)
segments, _ = rand_segments(x, x_lens, segment_size=5)
raise Exception("Should have failed")
except:
pass
x_lens_back = x_lens.clone()
segments, seg_idxs= rand_segments(x, x_lens.clone(), segment_size=5, pad_short=True, let_short_samples=True)
segments, seg_idxs = rand_segments(x, x_lens.clone(), segment_size=5, pad_short=True, let_short_samples=True)
assert segments.shape == (2, 3, 5)
assert all(seg_idxs >= 0), seg_idxs
assert all(x_lens_back == x_lens)