mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
allow inference to input and output style vectors
This commit is contained in:
@@ -1,41 +1,48 @@
|
||||
# FakeYou StyleTTS 2 Inference
|
||||
|
||||
import argparse
|
||||
import numpy as np
|
||||
import soundfile as sf
|
||||
from styletts2importable import compute_style, inference
|
||||
from tortoise.utils.text import split_and_recombine_text
|
||||
|
||||
def clsynthesize(text_file_path, voice_path, vcsteps):
|
||||
# Load the text from the file
|
||||
with open(text_file_path, 'r', encoding="utf-8") as file:
|
||||
text = file.read()
|
||||
|
||||
# Load the custom voice
|
||||
custom_voice = compute_style(voice_path)
|
||||
|
||||
# Split and recombine long text into manageable chunks
|
||||
texts = split_and_recombine_text(text)
|
||||
|
||||
audios = []
|
||||
for t in texts:
|
||||
# Generate audio for each text chunk
|
||||
audios.append(inference(t, custom_voice, alpha=0.3, beta=0.7, diffusion_steps=vcsteps, embedding_scale=1))
|
||||
|
||||
return 24000, np.concatenate(audios)
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description='CLI for StyleTTS Voice Synthesis')
|
||||
parser.add_argument('--input', type=str, required=True, help='Path to the text file')
|
||||
parser.add_argument('--voice', type=str, required=True, help='Path to the custom voice audio file')
|
||||
parser.add_argument('--vcsteps', type=int, required=True, help='Number of diffusion steps')
|
||||
parser.add_argument('--output', type=str, required=True, help='Output path for the synthesized audio file')
|
||||
args = parser.parse_args()
|
||||
|
||||
sample_rate, audio = clsynthesize(args.input, args.voice, args.vcsteps)
|
||||
|
||||
# Save the synthesized audio to the specified output path
|
||||
sf.write(args.output, audio, sample_rate)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
# FakeYou StyleTTS 2 Inference
|
||||
|
||||
import os
|
||||
import torch
|
||||
import argparse
|
||||
import numpy as np
|
||||
import soundfile as sf
|
||||
from styletts2importable import compute_style, inference
|
||||
from tortoise.utils.text import split_and_recombine_text
|
||||
|
||||
def clsynthesize(text_file_path, voice_path, vcsteps, style_npz_path=None, npz_output_path=None):
|
||||
# Load the text from the file
|
||||
with open(text_file_path, 'r', encoding="utf-8") as file:
|
||||
text = file.read()
|
||||
|
||||
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
||||
|
||||
if style_npz_path and os.path.exists(style_npz_path):
|
||||
# Load style vector from NPZ file
|
||||
style_vector_np = np.load(style_npz_path)['style_vector']
|
||||
custom_voice = torch.from_numpy(style_vector_np).float().to(device)
|
||||
elif voice_path:
|
||||
# Compute and save style vector
|
||||
custom_voice = compute_style(voice_path, save_npz=True, npz_path=npz_output_path or "default_output_path.npz").to(device)
|
||||
else:
|
||||
raise ValueError("Either a voice path or an input style NPZ file must be provided")
|
||||
|
||||
texts = split_and_recombine_text(text)
|
||||
audios = [inference(t, custom_voice, alpha=0.3, beta=0.7, diffusion_steps=vcsteps, embedding_scale=1) for t in texts]
|
||||
|
||||
return 24000, np.concatenate(audios)
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description='CLI for StyleTTS Voice Synthesis')
|
||||
parser.add_argument('--input', type=str, required=True, help='Path to the text file')
|
||||
parser.add_argument('--voice', type=str, help='Path to the custom voice audio file (optional if using --input-style-npz)')
|
||||
parser.add_argument('--vcsteps', type=int, required=True, help='Number of diffusion steps')
|
||||
parser.add_argument('--output', type=str, help='Output path for the synthesized audio file (optional)')
|
||||
parser.add_argument('--input-style-npz', type=str, help='Path to the input style vector NPZ file')
|
||||
parser.add_argument('--output-style-npz', type=str, help='Output path for the style vector NPZ file')
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.output:
|
||||
sample_rate, audio = clsynthesize(args.input, args.voice, args.vcsteps, args.input_style_npz, args.output_style_npz)
|
||||
sf.write(args.output, audio, sample_rate)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -1,385 +1,388 @@
|
||||
import torch
|
||||
from time import strftime
|
||||
import os, sys, time
|
||||
import platform
|
||||
|
||||
print("Env vars:")
|
||||
print(os.environ)
|
||||
|
||||
|
||||
def print_gpu_info():
|
||||
print('========================================')
|
||||
print('Python interpreter', sys.executable)
|
||||
print('PyTorch version', torch.__version__)
|
||||
print('CUDA Available?', torch.cuda.is_available())
|
||||
print('CUDA Device count', torch.cuda.device_count())
|
||||
print('CUDA threads for parallelizing CPU operations', torch.get_num_threads())
|
||||
print('CUDA architectures library was compiled for', torch.cuda.get_arch_list())
|
||||
|
||||
#try:
|
||||
# from tensorflow.python.client import device_lib
|
||||
# print('local devices', str(device_lib.list_local_devices()).replace("\n", "\n "))
|
||||
#except ImportError:
|
||||
# print('no tensorflow - cannot list devices')
|
||||
# pass
|
||||
print('========================================', flush=True)
|
||||
|
||||
|
||||
print_gpu_info(
|
||||
|
||||
)
|
||||
|
||||
print("NLTK")
|
||||
import nltk
|
||||
nltk.download('punkt')
|
||||
print("SCIPY")
|
||||
from scipy.io.wavfile import write
|
||||
print("TORCH STUFF")
|
||||
import torch
|
||||
print("START")
|
||||
torch.manual_seed(0)
|
||||
torch.backends.cudnn.benchmark = False
|
||||
torch.backends.cudnn.deterministic = True
|
||||
|
||||
import random
|
||||
random.seed(0)
|
||||
|
||||
import numpy as np
|
||||
np.random.seed(0)
|
||||
|
||||
# load packages
|
||||
import time
|
||||
import random
|
||||
import yaml
|
||||
from munch import Munch
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
import torchaudio
|
||||
import librosa
|
||||
from nltk.tokenize import word_tokenize
|
||||
|
||||
from models import *
|
||||
from utils import *
|
||||
from text_utils import TextCleaner
|
||||
textclenaer = TextCleaner()
|
||||
|
||||
|
||||
to_mel = torchaudio.transforms.MelSpectrogram(
|
||||
n_mels=80, n_fft=2048, win_length=1200, hop_length=300)
|
||||
mean, std = -4, 4
|
||||
|
||||
def length_to_mask(lengths):
|
||||
mask = torch.arange(lengths.max()).unsqueeze(0).expand(lengths.shape[0], -1).type_as(lengths)
|
||||
mask = torch.gt(mask+1, lengths.unsqueeze(1))
|
||||
return mask
|
||||
|
||||
def preprocess(wave):
|
||||
wave_tensor = torch.from_numpy(wave).float()
|
||||
mel_tensor = to_mel(wave_tensor)
|
||||
mel_tensor = (torch.log(1e-5 + mel_tensor.unsqueeze(0)) - mean) / std
|
||||
return mel_tensor
|
||||
|
||||
def compute_style(path):
|
||||
wave, sr = librosa.load(path, sr=24000)
|
||||
audio, index = librosa.effects.trim(wave, top_db=30)
|
||||
if sr != 24000:
|
||||
audio = librosa.resample(audio, sr, 24000)
|
||||
mel_tensor = preprocess(audio).to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_s = model.style_encoder(mel_tensor.unsqueeze(1))
|
||||
ref_p = model.predictor_encoder(mel_tensor.unsqueeze(1))
|
||||
|
||||
return torch.cat([ref_s, ref_p], dim=1)
|
||||
|
||||
device = 'cpu'
|
||||
if torch.cuda.is_available():
|
||||
device = 'cuda'
|
||||
elif torch.backends.mps.is_available():
|
||||
print("MPS would be available but cannot be used rn")
|
||||
# device = 'mps'
|
||||
|
||||
import phonemizer
|
||||
global_phonemizer = phonemizer.backend.EspeakBackend(language='en-us', preserve_punctuation=True, with_stress=True)
|
||||
|
||||
config = yaml.safe_load(open("Models/LibriTTS/config.yml"))
|
||||
|
||||
# load pretrained ASR model
|
||||
ASR_config = config.get('ASR_config', False)
|
||||
ASR_path = config.get('ASR_path', False)
|
||||
text_aligner = load_ASR_models(ASR_path, ASR_config)
|
||||
|
||||
# load pretrained F0 model
|
||||
F0_path = config.get('F0_path', False)
|
||||
pitch_extractor = load_F0_models(F0_path)
|
||||
|
||||
# load BERT model
|
||||
from Utils.PLBERT.util import load_plbert
|
||||
BERT_path = config.get('PLBERT_dir', False)
|
||||
plbert = load_plbert(BERT_path)
|
||||
|
||||
model_params = recursive_munch(config['model_params'])
|
||||
model = build_model(model_params, text_aligner, pitch_extractor, plbert)
|
||||
_ = [model[key].eval() for key in model]
|
||||
_ = [model[key].to(device) for key in model]
|
||||
|
||||
params_whole = torch.load("Models/LibriTTS/epochs_2nd_00020.pth", map_location='cpu')
|
||||
params = params_whole['net']
|
||||
|
||||
for key in model:
|
||||
if key in params:
|
||||
print('%s loaded' % key)
|
||||
try:
|
||||
model[key].load_state_dict(params[key])
|
||||
except:
|
||||
from collections import OrderedDict
|
||||
state_dict = params[key]
|
||||
new_state_dict = OrderedDict()
|
||||
for k, v in state_dict.items():
|
||||
name = k[7:] # remove `module.`
|
||||
new_state_dict[name] = v
|
||||
# load params
|
||||
model[key].load_state_dict(new_state_dict, strict=False)
|
||||
# except:
|
||||
# _load(params[key], model[key])
|
||||
_ = [model[key].eval() for key in model]
|
||||
|
||||
from Modules.diffusion.sampler import DiffusionSampler, ADPM2Sampler, KarrasSchedule
|
||||
|
||||
sampler = DiffusionSampler(
|
||||
model.diffusion.diffusion,
|
||||
sampler=ADPM2Sampler(),
|
||||
sigma_schedule=KarrasSchedule(sigma_min=0.0001, sigma_max=3.0, rho=9.0), # empirical parameters
|
||||
clamp=False
|
||||
)
|
||||
|
||||
def inference(text, ref_s, alpha = 0.3, beta = 0.7, diffusion_steps=5, embedding_scale=1, use_gruut=False):
|
||||
text = text.strip()
|
||||
ps = global_phonemizer.phonemize([text])
|
||||
ps = word_tokenize(ps[0])
|
||||
ps = ' '.join(ps)
|
||||
tokens = textclenaer(ps)
|
||||
tokens.insert(0, 0)
|
||||
tokens = torch.LongTensor(tokens).to(device).unsqueeze(0)
|
||||
|
||||
with torch.no_grad():
|
||||
input_lengths = torch.LongTensor([tokens.shape[-1]]).to(device)
|
||||
text_mask = length_to_mask(input_lengths).to(device)
|
||||
|
||||
t_en = model.text_encoder(tokens, input_lengths, text_mask)
|
||||
bert_dur = model.bert(tokens, attention_mask=(~text_mask).int())
|
||||
d_en = model.bert_encoder(bert_dur).transpose(-1, -2)
|
||||
|
||||
s_pred = sampler(noise = torch.randn((1, 256)).unsqueeze(1).to(device),
|
||||
embedding=bert_dur,
|
||||
embedding_scale=embedding_scale,
|
||||
features=ref_s, # reference from the same speaker as the embedding
|
||||
num_steps=diffusion_steps).squeeze(1)
|
||||
|
||||
|
||||
s = s_pred[:, 128:]
|
||||
ref = s_pred[:, :128]
|
||||
|
||||
ref = alpha * ref + (1 - alpha) * ref_s[:, :128]
|
||||
s = beta * s + (1 - beta) * ref_s[:, 128:]
|
||||
|
||||
d = model.predictor.text_encoder(d_en,
|
||||
s, input_lengths, text_mask)
|
||||
|
||||
x, _ = model.predictor.lstm(d)
|
||||
duration = model.predictor.duration_proj(x)
|
||||
|
||||
duration = torch.sigmoid(duration).sum(axis=-1)
|
||||
pred_dur = torch.round(duration.squeeze()).clamp(min=1)
|
||||
|
||||
|
||||
pred_aln_trg = torch.zeros(input_lengths, int(pred_dur.sum().data))
|
||||
c_frame = 0
|
||||
for i in range(pred_aln_trg.size(0)):
|
||||
pred_aln_trg[i, c_frame:c_frame + int(pred_dur[i].data)] = 1
|
||||
c_frame += int(pred_dur[i].data)
|
||||
|
||||
# encode prosody
|
||||
en = (d.transpose(-1, -2) @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(en)
|
||||
asr_new[:, :, 0] = en[:, :, 0]
|
||||
asr_new[:, :, 1:] = en[:, :, 0:-1]
|
||||
en = asr_new
|
||||
|
||||
F0_pred, N_pred = model.predictor.F0Ntrain(en, s)
|
||||
|
||||
asr = (t_en @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(asr)
|
||||
asr_new[:, :, 0] = asr[:, :, 0]
|
||||
asr_new[:, :, 1:] = asr[:, :, 0:-1]
|
||||
asr = asr_new
|
||||
|
||||
out = model.decoder(asr,
|
||||
F0_pred, N_pred, ref.squeeze().unsqueeze(0))
|
||||
|
||||
|
||||
return out.squeeze().cpu().numpy()[..., :-50] # weird pulse at the end of the model, need to be fixed later
|
||||
|
||||
def LFinference(text, s_prev, ref_s, alpha = 0.3, beta = 0.7, t = 0.7, diffusion_steps=5, embedding_scale=1, use_gruut=False):
|
||||
text = text.strip()
|
||||
ps = global_phonemizer.phonemize([text])
|
||||
ps = word_tokenize(ps[0])
|
||||
ps = ' '.join(ps)
|
||||
ps = ps.replace('``', '"')
|
||||
ps = ps.replace("''", '"')
|
||||
|
||||
tokens = textclenaer(ps)
|
||||
tokens.insert(0, 0)
|
||||
tokens = torch.LongTensor(tokens).to(device).unsqueeze(0)
|
||||
|
||||
with torch.no_grad():
|
||||
input_lengths = torch.LongTensor([tokens.shape[-1]]).to(device)
|
||||
text_mask = length_to_mask(input_lengths).to(device)
|
||||
|
||||
t_en = model.text_encoder(tokens, input_lengths, text_mask)
|
||||
bert_dur = model.bert(tokens, attention_mask=(~text_mask).int())
|
||||
d_en = model.bert_encoder(bert_dur).transpose(-1, -2)
|
||||
|
||||
s_pred = sampler(noise = torch.randn((1, 256)).unsqueeze(1).to(device),
|
||||
embedding=bert_dur,
|
||||
embedding_scale=embedding_scale,
|
||||
features=ref_s, # reference from the same speaker as the embedding
|
||||
num_steps=diffusion_steps).squeeze(1)
|
||||
|
||||
if s_prev is not None:
|
||||
# convex combination of previous and current style
|
||||
s_pred = t * s_prev + (1 - t) * s_pred
|
||||
|
||||
s = s_pred[:, 128:]
|
||||
ref = s_pred[:, :128]
|
||||
|
||||
ref = alpha * ref + (1 - alpha) * ref_s[:, :128]
|
||||
s = beta * s + (1 - beta) * ref_s[:, 128:]
|
||||
|
||||
s_pred = torch.cat([ref, s], dim=-1)
|
||||
|
||||
d = model.predictor.text_encoder(d_en,
|
||||
s, input_lengths, text_mask)
|
||||
|
||||
x, _ = model.predictor.lstm(d)
|
||||
duration = model.predictor.duration_proj(x)
|
||||
|
||||
duration = torch.sigmoid(duration).sum(axis=-1)
|
||||
pred_dur = torch.round(duration.squeeze()).clamp(min=1)
|
||||
|
||||
|
||||
pred_aln_trg = torch.zeros(input_lengths, int(pred_dur.sum().data))
|
||||
c_frame = 0
|
||||
for i in range(pred_aln_trg.size(0)):
|
||||
pred_aln_trg[i, c_frame:c_frame + int(pred_dur[i].data)] = 1
|
||||
c_frame += int(pred_dur[i].data)
|
||||
|
||||
# encode prosody
|
||||
en = (d.transpose(-1, -2) @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(en)
|
||||
asr_new[:, :, 0] = en[:, :, 0]
|
||||
asr_new[:, :, 1:] = en[:, :, 0:-1]
|
||||
en = asr_new
|
||||
|
||||
F0_pred, N_pred = model.predictor.F0Ntrain(en, s)
|
||||
|
||||
asr = (t_en @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(asr)
|
||||
asr_new[:, :, 0] = asr[:, :, 0]
|
||||
asr_new[:, :, 1:] = asr[:, :, 0:-1]
|
||||
asr = asr_new
|
||||
|
||||
out = model.decoder(asr,
|
||||
F0_pred, N_pred, ref.squeeze().unsqueeze(0))
|
||||
|
||||
|
||||
return out.squeeze().cpu().numpy()[..., :-100], s_pred # weird pulse at the end of the model, need to be fixed later
|
||||
|
||||
def STinference(text, ref_s, ref_text, alpha = 0.3, beta = 0.7, diffusion_steps=5, embedding_scale=1, use_gruut=False):
|
||||
text = text.strip()
|
||||
ps = global_phonemizer.phonemize([text])
|
||||
ps = word_tokenize(ps[0])
|
||||
ps = ' '.join(ps)
|
||||
|
||||
tokens = textclenaer(ps)
|
||||
tokens.insert(0, 0)
|
||||
tokens = torch.LongTensor(tokens).to(device).unsqueeze(0)
|
||||
|
||||
ref_text = ref_text.strip()
|
||||
ps = global_phonemizer.phonemize([ref_text])
|
||||
ps = word_tokenize(ps[0])
|
||||
ps = ' '.join(ps)
|
||||
|
||||
ref_tokens = textclenaer(ps)
|
||||
ref_tokens.insert(0, 0)
|
||||
ref_tokens = torch.LongTensor(ref_tokens).to(device).unsqueeze(0)
|
||||
|
||||
|
||||
with torch.no_grad():
|
||||
input_lengths = torch.LongTensor([tokens.shape[-1]]).to(device)
|
||||
text_mask = length_to_mask(input_lengths).to(device)
|
||||
|
||||
t_en = model.text_encoder(tokens, input_lengths, text_mask)
|
||||
bert_dur = model.bert(tokens, attention_mask=(~text_mask).int())
|
||||
d_en = model.bert_encoder(bert_dur).transpose(-1, -2)
|
||||
|
||||
ref_input_lengths = torch.LongTensor([ref_tokens.shape[-1]]).to(device)
|
||||
ref_text_mask = length_to_mask(ref_input_lengths).to(device)
|
||||
ref_bert_dur = model.bert(ref_tokens, attention_mask=(~ref_text_mask).int())
|
||||
s_pred = sampler(noise = torch.randn((1, 256)).unsqueeze(1).to(device),
|
||||
embedding=bert_dur,
|
||||
embedding_scale=embedding_scale,
|
||||
features=ref_s, # reference from the same speaker as the embedding
|
||||
num_steps=diffusion_steps).squeeze(1)
|
||||
|
||||
|
||||
s = s_pred[:, 128:]
|
||||
ref = s_pred[:, :128]
|
||||
|
||||
ref = alpha * ref + (1 - alpha) * ref_s[:, :128]
|
||||
s = beta * s + (1 - beta) * ref_s[:, 128:]
|
||||
|
||||
d = model.predictor.text_encoder(d_en,
|
||||
s, input_lengths, text_mask)
|
||||
|
||||
x, _ = model.predictor.lstm(d)
|
||||
duration = model.predictor.duration_proj(x)
|
||||
|
||||
duration = torch.sigmoid(duration).sum(axis=-1)
|
||||
pred_dur = torch.round(duration.squeeze()).clamp(min=1)
|
||||
|
||||
|
||||
pred_aln_trg = torch.zeros(input_lengths, int(pred_dur.sum().data))
|
||||
c_frame = 0
|
||||
for i in range(pred_aln_trg.size(0)):
|
||||
pred_aln_trg[i, c_frame:c_frame + int(pred_dur[i].data)] = 1
|
||||
c_frame += int(pred_dur[i].data)
|
||||
|
||||
# encode prosody
|
||||
en = (d.transpose(-1, -2) @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(en)
|
||||
asr_new[:, :, 0] = en[:, :, 0]
|
||||
asr_new[:, :, 1:] = en[:, :, 0:-1]
|
||||
en = asr_new
|
||||
|
||||
F0_pred, N_pred = model.predictor.F0Ntrain(en, s)
|
||||
|
||||
asr = (t_en @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(asr)
|
||||
asr_new[:, :, 0] = asr[:, :, 0]
|
||||
asr_new[:, :, 1:] = asr[:, :, 0:-1]
|
||||
asr = asr_new
|
||||
|
||||
out = model.decoder(asr,
|
||||
F0_pred, N_pred, ref.squeeze().unsqueeze(0))
|
||||
|
||||
|
||||
return out.squeeze().cpu().numpy()[..., :-50] # weird pulse at the end of the model, need to be fixed later
|
||||
import torch
|
||||
from time import strftime
|
||||
import os, sys, time
|
||||
import platform
|
||||
|
||||
print("Env vars:")
|
||||
print(os.environ)
|
||||
|
||||
|
||||
def print_gpu_info():
|
||||
print('========================================')
|
||||
print('Python interpreter', sys.executable)
|
||||
print('PyTorch version', torch.__version__)
|
||||
print('CUDA Available?', torch.cuda.is_available())
|
||||
print('CUDA Device count', torch.cuda.device_count())
|
||||
print('CUDA threads for parallelizing CPU operations', torch.get_num_threads())
|
||||
print('CUDA architectures library was compiled for', torch.cuda.get_arch_list())
|
||||
|
||||
#try:
|
||||
# from tensorflow.python.client import device_lib
|
||||
# print('local devices', str(device_lib.list_local_devices()).replace("\n", "\n "))
|
||||
#except ImportError:
|
||||
# print('no tensorflow - cannot list devices')
|
||||
# pass
|
||||
print('========================================', flush=True)
|
||||
|
||||
|
||||
print_gpu_info(
|
||||
|
||||
)
|
||||
|
||||
print("NLTK")
|
||||
import nltk
|
||||
nltk.download('punkt')
|
||||
print("SCIPY")
|
||||
from scipy.io.wavfile import write
|
||||
print("START")
|
||||
torch.manual_seed(0)
|
||||
torch.backends.cudnn.benchmark = False
|
||||
torch.backends.cudnn.deterministic = True
|
||||
|
||||
import random
|
||||
random.seed(0)
|
||||
|
||||
import numpy as np
|
||||
np.random.seed(0)
|
||||
|
||||
# load packages
|
||||
import time
|
||||
import random
|
||||
import yaml
|
||||
from munch import Munch
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
import torchaudio
|
||||
import librosa
|
||||
from nltk.tokenize import word_tokenize
|
||||
|
||||
from models import *
|
||||
from utils import *
|
||||
from text_utils import TextCleaner
|
||||
textclenaer = TextCleaner()
|
||||
|
||||
|
||||
to_mel = torchaudio.transforms.MelSpectrogram(
|
||||
n_mels=80, n_fft=2048, win_length=1200, hop_length=300)
|
||||
mean, std = -4, 4
|
||||
|
||||
def length_to_mask(lengths):
|
||||
mask = torch.arange(lengths.max()).unsqueeze(0).expand(lengths.shape[0], -1).type_as(lengths)
|
||||
mask = torch.gt(mask+1, lengths.unsqueeze(1))
|
||||
return mask
|
||||
|
||||
def preprocess(wave):
|
||||
wave_tensor = torch.from_numpy(wave).float()
|
||||
mel_tensor = to_mel(wave_tensor)
|
||||
mel_tensor = (torch.log(1e-5 + mel_tensor.unsqueeze(0)) - mean) / std
|
||||
return mel_tensor
|
||||
|
||||
def compute_style(path, save_npz=False, npz_path=None):
|
||||
wave, sr = librosa.load(path, sr=24000)
|
||||
audio, index = librosa.effects.trim(wave, top_db=30)
|
||||
if sr != 24000:
|
||||
audio = librosa.resample(audio, sr, 24000)
|
||||
mel_tensor = preprocess(audio).to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_s = model.style_encoder(mel_tensor.unsqueeze(1))
|
||||
ref_p = model.predictor_encoder(mel_tensor.unsqueeze(1))
|
||||
style_vector = torch.cat([ref_s, ref_p], dim=1)
|
||||
|
||||
# Save style vector to NPZ file if requested
|
||||
if save_npz and npz_path:
|
||||
np.savez(npz_path, style_vector=style_vector.cpu().numpy())
|
||||
|
||||
return style_vector
|
||||
|
||||
device = 'cpu'
|
||||
if torch.cuda.is_available():
|
||||
device = 'cuda'
|
||||
elif torch.backends.mps.is_available():
|
||||
print("MPS would be available but cannot be used rn")
|
||||
# device = 'mps'
|
||||
|
||||
import phonemizer
|
||||
global_phonemizer = phonemizer.backend.EspeakBackend(language='en-us', preserve_punctuation=True, with_stress=True)
|
||||
|
||||
config = yaml.safe_load(open("Models/LibriTTS/config.yml"))
|
||||
|
||||
# load pretrained ASR model
|
||||
ASR_config = config.get('ASR_config', False)
|
||||
ASR_path = config.get('ASR_path', False)
|
||||
text_aligner = load_ASR_models(ASR_path, ASR_config)
|
||||
|
||||
# load pretrained F0 model
|
||||
F0_path = config.get('F0_path', False)
|
||||
pitch_extractor = load_F0_models(F0_path)
|
||||
|
||||
# load BERT model
|
||||
from Utils.PLBERT.util import load_plbert
|
||||
BERT_path = config.get('PLBERT_dir', False)
|
||||
plbert = load_plbert(BERT_path)
|
||||
|
||||
model_params = recursive_munch(config['model_params'])
|
||||
model = build_model(model_params, text_aligner, pitch_extractor, plbert)
|
||||
_ = [model[key].eval() for key in model]
|
||||
_ = [model[key].to(device) for key in model]
|
||||
|
||||
params_whole = torch.load("Models/LibriTTS/epochs_2nd_00020.pth", map_location='cpu')
|
||||
params = params_whole['net']
|
||||
|
||||
for key in model:
|
||||
if key in params:
|
||||
print('%s loaded' % key)
|
||||
try:
|
||||
model[key].load_state_dict(params[key])
|
||||
except:
|
||||
from collections import OrderedDict
|
||||
state_dict = params[key]
|
||||
new_state_dict = OrderedDict()
|
||||
for k, v in state_dict.items():
|
||||
name = k[7:] # remove `module.`
|
||||
new_state_dict[name] = v
|
||||
# load params
|
||||
model[key].load_state_dict(new_state_dict, strict=False)
|
||||
# except:
|
||||
# _load(params[key], model[key])
|
||||
_ = [model[key].eval() for key in model]
|
||||
|
||||
from Modules.diffusion.sampler import DiffusionSampler, ADPM2Sampler, KarrasSchedule
|
||||
|
||||
sampler = DiffusionSampler(
|
||||
model.diffusion.diffusion,
|
||||
sampler=ADPM2Sampler(),
|
||||
sigma_schedule=KarrasSchedule(sigma_min=0.0001, sigma_max=3.0, rho=9.0), # empirical parameters
|
||||
clamp=False
|
||||
)
|
||||
|
||||
def inference(text, ref_s, alpha = 0.3, beta = 0.7, diffusion_steps=5, embedding_scale=1, use_gruut=False):
|
||||
text = text.strip()
|
||||
ps = global_phonemizer.phonemize([text])
|
||||
ps = word_tokenize(ps[0])
|
||||
ps = ' '.join(ps)
|
||||
tokens = textclenaer(ps)
|
||||
tokens.insert(0, 0)
|
||||
tokens = torch.LongTensor(tokens).to(device).unsqueeze(0)
|
||||
|
||||
with torch.no_grad():
|
||||
input_lengths = torch.LongTensor([tokens.shape[-1]]).to(device)
|
||||
text_mask = length_to_mask(input_lengths).to(device)
|
||||
|
||||
t_en = model.text_encoder(tokens, input_lengths, text_mask)
|
||||
bert_dur = model.bert(tokens, attention_mask=(~text_mask).int())
|
||||
d_en = model.bert_encoder(bert_dur).transpose(-1, -2)
|
||||
|
||||
s_pred = sampler(noise = torch.randn((1, 256)).unsqueeze(1).to(device),
|
||||
embedding=bert_dur,
|
||||
embedding_scale=embedding_scale,
|
||||
features=ref_s, # reference from the same speaker as the embedding
|
||||
num_steps=diffusion_steps).squeeze(1)
|
||||
|
||||
|
||||
s = s_pred[:, 128:]
|
||||
ref = s_pred[:, :128]
|
||||
|
||||
ref = alpha * ref + (1 - alpha) * ref_s[:, :128]
|
||||
s = beta * s + (1 - beta) * ref_s[:, 128:]
|
||||
|
||||
d = model.predictor.text_encoder(d_en,
|
||||
s, input_lengths, text_mask)
|
||||
|
||||
x, _ = model.predictor.lstm(d)
|
||||
duration = model.predictor.duration_proj(x)
|
||||
|
||||
duration = torch.sigmoid(duration).sum(axis=-1)
|
||||
pred_dur = torch.round(duration.squeeze()).clamp(min=1)
|
||||
|
||||
|
||||
pred_aln_trg = torch.zeros(input_lengths, int(pred_dur.sum().data))
|
||||
c_frame = 0
|
||||
for i in range(pred_aln_trg.size(0)):
|
||||
pred_aln_trg[i, c_frame:c_frame + int(pred_dur[i].data)] = 1
|
||||
c_frame += int(pred_dur[i].data)
|
||||
|
||||
# encode prosody
|
||||
en = (d.transpose(-1, -2) @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(en)
|
||||
asr_new[:, :, 0] = en[:, :, 0]
|
||||
asr_new[:, :, 1:] = en[:, :, 0:-1]
|
||||
en = asr_new
|
||||
|
||||
F0_pred, N_pred = model.predictor.F0Ntrain(en, s)
|
||||
|
||||
asr = (t_en @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(asr)
|
||||
asr_new[:, :, 0] = asr[:, :, 0]
|
||||
asr_new[:, :, 1:] = asr[:, :, 0:-1]
|
||||
asr = asr_new
|
||||
|
||||
out = model.decoder(asr,
|
||||
F0_pred, N_pred, ref.squeeze().unsqueeze(0))
|
||||
|
||||
|
||||
return out.squeeze().cpu().numpy()[..., :-50] # weird pulse at the end of the model, need to be fixed later
|
||||
|
||||
def LFinference(text, s_prev, ref_s, alpha = 0.3, beta = 0.7, t = 0.7, diffusion_steps=5, embedding_scale=1, use_gruut=False):
|
||||
text = text.strip()
|
||||
ps = global_phonemizer.phonemize([text])
|
||||
ps = word_tokenize(ps[0])
|
||||
ps = ' '.join(ps)
|
||||
ps = ps.replace('``', '"')
|
||||
ps = ps.replace("''", '"')
|
||||
|
||||
tokens = textclenaer(ps)
|
||||
tokens.insert(0, 0)
|
||||
tokens = torch.LongTensor(tokens).to(device).unsqueeze(0)
|
||||
|
||||
with torch.no_grad():
|
||||
input_lengths = torch.LongTensor([tokens.shape[-1]]).to(device)
|
||||
text_mask = length_to_mask(input_lengths).to(device)
|
||||
|
||||
t_en = model.text_encoder(tokens, input_lengths, text_mask)
|
||||
bert_dur = model.bert(tokens, attention_mask=(~text_mask).int())
|
||||
d_en = model.bert_encoder(bert_dur).transpose(-1, -2)
|
||||
|
||||
s_pred = sampler(noise = torch.randn((1, 256)).unsqueeze(1).to(device),
|
||||
embedding=bert_dur,
|
||||
embedding_scale=embedding_scale,
|
||||
features=ref_s, # reference from the same speaker as the embedding
|
||||
num_steps=diffusion_steps).squeeze(1)
|
||||
|
||||
if s_prev is not None:
|
||||
# convex combination of previous and current style
|
||||
s_pred = t * s_prev + (1 - t) * s_pred
|
||||
|
||||
s = s_pred[:, 128:]
|
||||
ref = s_pred[:, :128]
|
||||
|
||||
ref = alpha * ref + (1 - alpha) * ref_s[:, :128]
|
||||
s = beta * s + (1 - beta) * ref_s[:, 128:]
|
||||
|
||||
s_pred = torch.cat([ref, s], dim=-1)
|
||||
|
||||
d = model.predictor.text_encoder(d_en,
|
||||
s, input_lengths, text_mask)
|
||||
|
||||
x, _ = model.predictor.lstm(d)
|
||||
duration = model.predictor.duration_proj(x)
|
||||
|
||||
duration = torch.sigmoid(duration).sum(axis=-1)
|
||||
pred_dur = torch.round(duration.squeeze()).clamp(min=1)
|
||||
|
||||
|
||||
pred_aln_trg = torch.zeros(input_lengths, int(pred_dur.sum().data))
|
||||
c_frame = 0
|
||||
for i in range(pred_aln_trg.size(0)):
|
||||
pred_aln_trg[i, c_frame:c_frame + int(pred_dur[i].data)] = 1
|
||||
c_frame += int(pred_dur[i].data)
|
||||
|
||||
# encode prosody
|
||||
en = (d.transpose(-1, -2) @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(en)
|
||||
asr_new[:, :, 0] = en[:, :, 0]
|
||||
asr_new[:, :, 1:] = en[:, :, 0:-1]
|
||||
en = asr_new
|
||||
|
||||
F0_pred, N_pred = model.predictor.F0Ntrain(en, s)
|
||||
|
||||
asr = (t_en @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(asr)
|
||||
asr_new[:, :, 0] = asr[:, :, 0]
|
||||
asr_new[:, :, 1:] = asr[:, :, 0:-1]
|
||||
asr = asr_new
|
||||
|
||||
out = model.decoder(asr,
|
||||
F0_pred, N_pred, ref.squeeze().unsqueeze(0))
|
||||
|
||||
|
||||
return out.squeeze().cpu().numpy()[..., :-100], s_pred # weird pulse at the end of the model, need to be fixed later
|
||||
|
||||
def STinference(text, ref_s, ref_text, alpha = 0.3, beta = 0.7, diffusion_steps=5, embedding_scale=1, use_gruut=False):
|
||||
text = text.strip()
|
||||
ps = global_phonemizer.phonemize([text])
|
||||
ps = word_tokenize(ps[0])
|
||||
ps = ' '.join(ps)
|
||||
|
||||
tokens = textclenaer(ps)
|
||||
tokens.insert(0, 0)
|
||||
tokens = torch.LongTensor(tokens).to(device).unsqueeze(0)
|
||||
|
||||
ref_text = ref_text.strip()
|
||||
ps = global_phonemizer.phonemize([ref_text])
|
||||
ps = word_tokenize(ps[0])
|
||||
ps = ' '.join(ps)
|
||||
|
||||
ref_tokens = textclenaer(ps)
|
||||
ref_tokens.insert(0, 0)
|
||||
ref_tokens = torch.LongTensor(ref_tokens).to(device).unsqueeze(0)
|
||||
|
||||
|
||||
with torch.no_grad():
|
||||
input_lengths = torch.LongTensor([tokens.shape[-1]]).to(device)
|
||||
text_mask = length_to_mask(input_lengths).to(device)
|
||||
|
||||
t_en = model.text_encoder(tokens, input_lengths, text_mask)
|
||||
bert_dur = model.bert(tokens, attention_mask=(~text_mask).int())
|
||||
d_en = model.bert_encoder(bert_dur).transpose(-1, -2)
|
||||
|
||||
ref_input_lengths = torch.LongTensor([ref_tokens.shape[-1]]).to(device)
|
||||
ref_text_mask = length_to_mask(ref_input_lengths).to(device)
|
||||
ref_bert_dur = model.bert(ref_tokens, attention_mask=(~ref_text_mask).int())
|
||||
s_pred = sampler(noise = torch.randn((1, 256)).unsqueeze(1).to(device),
|
||||
embedding=bert_dur,
|
||||
embedding_scale=embedding_scale,
|
||||
features=ref_s, # reference from the same speaker as the embedding
|
||||
num_steps=diffusion_steps).squeeze(1)
|
||||
|
||||
|
||||
s = s_pred[:, 128:]
|
||||
ref = s_pred[:, :128]
|
||||
|
||||
ref = alpha * ref + (1 - alpha) * ref_s[:, :128]
|
||||
s = beta * s + (1 - beta) * ref_s[:, 128:]
|
||||
|
||||
d = model.predictor.text_encoder(d_en,
|
||||
s, input_lengths, text_mask)
|
||||
|
||||
x, _ = model.predictor.lstm(d)
|
||||
duration = model.predictor.duration_proj(x)
|
||||
|
||||
duration = torch.sigmoid(duration).sum(axis=-1)
|
||||
pred_dur = torch.round(duration.squeeze()).clamp(min=1)
|
||||
|
||||
|
||||
pred_aln_trg = torch.zeros(input_lengths, int(pred_dur.sum().data))
|
||||
c_frame = 0
|
||||
for i in range(pred_aln_trg.size(0)):
|
||||
pred_aln_trg[i, c_frame:c_frame + int(pred_dur[i].data)] = 1
|
||||
c_frame += int(pred_dur[i].data)
|
||||
|
||||
# encode prosody
|
||||
en = (d.transpose(-1, -2) @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(en)
|
||||
asr_new[:, :, 0] = en[:, :, 0]
|
||||
asr_new[:, :, 1:] = en[:, :, 0:-1]
|
||||
en = asr_new
|
||||
|
||||
F0_pred, N_pred = model.predictor.F0Ntrain(en, s)
|
||||
|
||||
asr = (t_en @ pred_aln_trg.unsqueeze(0).to(device))
|
||||
if model_params.decoder.type == "hifigan":
|
||||
asr_new = torch.zeros_like(asr)
|
||||
asr_new[:, :, 0] = asr[:, :, 0]
|
||||
asr_new[:, :, 1:] = asr[:, :, 0:-1]
|
||||
asr = asr_new
|
||||
|
||||
out = model.decoder(asr,
|
||||
F0_pred, N_pred, ref.squeeze().unsqueeze(0))
|
||||
|
||||
|
||||
return out.squeeze().cpu().numpy()[..., :-50] # weird pulse at the end of the model, need to be fixed later
|
||||
Reference in New Issue
Block a user