diff --git a/tts/StyleTTS2Train/test_harness.py b/tts/StyleTTS2Train/test_harness.py index 0c502b6..e202ef6 100644 --- a/tts/StyleTTS2Train/test_harness.py +++ b/tts/StyleTTS2Train/test_harness.py @@ -2,7 +2,8 @@ import wandb from pathlib import Path from collections import defaultdict from style_tts2 import StyleTTS2 - +import torch +import torchaudio class CharacterData: def __init__(self,full_path,char_name,file_name,ref_text) -> None: self.char_name = char_name @@ -23,7 +24,7 @@ class TestHarness: ) self.test_text = ["This is!","Hello world testing!""This is a slightly... longer test.""This is a longer test, of style TTS Two!", "This should be a lengthy sentence."] - self.test_characters = [] + self.test_characters = self.get_characters() def log_values(self,loss_gen_all, d_loss, loss_ce, loss_dur, loss_lm, loss_norm_rec, @@ -55,17 +56,29 @@ class TestHarness: saved_path = callback(file_input,text,checkpoint_name) return saved_path - def test(self,check_point:Path): - model = StyleTTS2() + def test(self,check_point:Path,top=1): + model = StyleTTS2(check_point) - model.STinference(text="",ref_text="",) - for char_name in self.test_characters.keys(): - ls = self.test_characters[char_name] + char_name_ls = TestHarness.get_characters() + + for char_name in char_name_ls.keys(): + char_sample = char_name_ls[char_name] - for sample in ls: - s_ref = model.compute_style(sample.full_path) - - + for sample in char_sample: + s_ref = Path(sample.full_path) + s_ref = model.compute_style(s_ref) + + for i,text in enumerate(self.test_text): + # we shouldn't set the embedding scale high on some text / input audio lengths it will break the audio + mono_wav = model.STinference(text=text, + ref_text=sample.ref_text, + ref_s=s_ref, + diffusion_steps=10,alpha=0.45, beta=0.6, embedding_scale=1.0) + stereo_wav = torch.tensor(mono_wav).repeat(2,1) + torchaudio.save(f"{char_name}|{sample.file_name}",stereo_wav,sample_rate=24000) + count += 1 + if count == top: + break @staticmethod def get_lines_from_file(file_path): with open(file_path,'r') as f: @@ -83,16 +96,17 @@ class TestHarness: full_path = comp[0] pieces = full_path.split("/") char_name = pieces[1] - file_wav = pieces[2] + file_name = pieces[2] ref_text = comp[1] - print(f"{full_path}|{char_name}| {file_wav} | {ref_text}") - + base_path = "TestHarnessV2" + print(f"{char_name}|{file_name}|{ref_text}") if char_name in characters: - characters[char_name].append(CharacterData(char_name=char_name,file_name=file_wav,full_path=full_path,ref_text=ref_text)) + characters[char_name].append(CharacterData(char_name=char_name,file_name=file_name,full_path= base_path + full_path,ref_text=ref_text)) else: characters[char_name] = [] - characters[char_name].append(CharacterData(char_name=char_name,file_name=file_wav,full_path=full_path,ref_text=ref_text)) + characters[char_name].append(CharacterData(char_name=char_name,file_name=file_name,full_path= base_path + full_path,ref_text=ref_text)) + return characters def done(self): wandb.finish() @@ -101,7 +115,6 @@ if __name__ == "__main__": harness = TestHarness() #harness.sample_ood_voice(file="testing.wav",sample_rate=24000,character="Raiden",caption="What I just Said") #harness.sample_list_of_characters(file_path="./output/") - - print(TestHarness.get_lines_from_file("TestHarness/test.txt")) - - TestHarness.get_characters() \ No newline at end of file + harness.test("Models/LibriTTS/epochs_2nd_00020.pth") + + \ No newline at end of file