mirror of
https://github.com/wassname/TTS.git
synced 2026-09-09 11:16:00 +08:00
Make style
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user