Test Harness Tested

This commit is contained in:
arensc
2024-03-08 11:33:47 -05:00
parent b20d0e0d8c
commit d191a9eb3b
+33 -20
View File
@@ -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()
harness.test("Models/LibriTTS/epochs_2nd_00020.pth")