mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
Test Harness Tested
This commit is contained in:
@@ -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")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user