mirror of
https://github.com/wassname/TTS.git
synced 2026-09-09 11:16:00 +08:00
REBASED: Transform Speaker Encoder in a Generic Encoder and Implement Emotion Encoder training support (#1349)
* Rename Speaker encoder module to encoder * Add a generic emotion dataset formatter * Transform the Speaker Encoder dataset to a generic dataset and create emotion encoder config * Add class map in emotion config * Add Base encoder config * Add evaluation encoder script * Fix the bug in plot_embeddings * Enable Weight decay for encoder training * Add argumnet to disable storage * Add Perfect Sampler and remove storage * Add evaluation during encoder training * Fix lint checks * Remove useless config parameter * Active evaluation in speaker encoder test and use multispeaker dataset for this test * Unit tests fixs * Remove useless tests for speedup the aux_tests * Use get_optimizer in Encoder * Add BaseEncoder Class * Fix the unitests * Add Perfect Batch Sampler unit test * Add compute encoder accuracy in a function
This commit is contained in:
@@ -8,6 +8,7 @@ from TTS.config.shared_configs import BaseDatasetConfig
|
||||
from TTS.tts.datasets import load_tts_samples
|
||||
from TTS.tts.utils.languages import get_language_balancer_weights
|
||||
from TTS.tts.utils.speakers import get_speaker_balancer_weights
|
||||
from TTS.encoder.utils.samplers import PerfectBatchSampler
|
||||
|
||||
# Fixing random state to avoid random fails
|
||||
torch.manual_seed(0)
|
||||
@@ -82,3 +83,51 @@ class TestSamplers(unittest.TestCase):
|
||||
spk2 += 1
|
||||
|
||||
assert is_balanced(spk1, spk2), "Speaker Weighted sampler is supposed to be balanced"
|
||||
|
||||
def test_perfect_sampler(self): # pylint: disable=no-self-use
|
||||
classes = set()
|
||||
for item in train_samples:
|
||||
classes.add(item["speaker_name"])
|
||||
|
||||
sampler = PerfectBatchSampler(
|
||||
train_samples,
|
||||
classes,
|
||||
batch_size=2 * 3, # total batch size
|
||||
num_classes_in_batch=2,
|
||||
label_key="speaker_name",
|
||||
shuffle=False,
|
||||
drop_last=True)
|
||||
batchs = functools.reduce(lambda a, b: a + b, [list(sampler) for i in range(100)])
|
||||
for batch in batchs:
|
||||
spk1, spk2 = 0, 0
|
||||
# for in each batch
|
||||
for index in batch:
|
||||
if train_samples[index]["speaker_name"] == "ljspeech-0":
|
||||
spk1 += 1
|
||||
else:
|
||||
spk2 += 1
|
||||
assert spk1 == spk2, "PerfectBatchSampler is supposed to be perfectly balanced"
|
||||
|
||||
def test_perfect_sampler_shuffle(self): # pylint: disable=no-self-use
|
||||
classes = set()
|
||||
for item in train_samples:
|
||||
classes.add(item["speaker_name"])
|
||||
|
||||
sampler = PerfectBatchSampler(
|
||||
train_samples,
|
||||
classes,
|
||||
batch_size=2 * 3, # total batch size
|
||||
num_classes_in_batch=2,
|
||||
label_key="speaker_name",
|
||||
shuffle=True,
|
||||
drop_last=False)
|
||||
batchs = functools.reduce(lambda a, b: a + b, [list(sampler) for i in range(100)])
|
||||
for batch in batchs:
|
||||
spk1, spk2 = 0, 0
|
||||
# for in each batch
|
||||
for index in batch:
|
||||
if train_samples[index]["speaker_name"] == "ljspeech-0":
|
||||
spk1 += 1
|
||||
else:
|
||||
spk2 += 1
|
||||
assert spk1 == spk2, "PerfectBatchSampler is supposed to be perfectly balanced"
|
||||
|
||||
Reference in New Issue
Block a user