Text-to-Speech
PyTorch
English
tts
glow-tts
hifi-gan
coqui-tts
supratts
SupraTTS-0.1-Beta / infer_v2.py
LH-Tech-AI's picture
Update infer_v2.py
edf6a4c verified
Raw History Blame Contribute Delete
4.38 kB
import math
from contextlib import nullcontext
import torch
from torch import nn
from TTS.config import load_config
from TTS.tts.configs.glow_tts_config import GlowTTSConfig
from TTS.tts.layers.vits.stochastic_duration_predictor import StochasticDurationPredictor
from TTS.tts.models.glow_tts import GlowTTS
from TTS.tts.utils.helpers import generate_path, sequence_mask
from TTS.tts.utils.text.tokenizer import TTSTokenizer
from TTS.utils.audio import AudioProcessor
from TTS.vocoder.models import setup_model as setup_vocoder
try:
from monotonic_alignment_search import maximum_path
except ImportError:
try:
from TTS.tts.utils.monotonic_align import maximum_path
except ImportError:
from TTS.tts.utils.helpers import maximum_path
def _fp32():
return torch.autocast("cuda", enabled=False)
class _DPTap(nn.Module):
def __init__(self):
super().__init__()
self.cache = None
def forward(self, x, x_mask, *args, **kwargs):
self.cache = x
return torch.zeros_like(x_mask)
class FlareGlowTTS(GlowTTS):
def __init__(self, config, ap=None, tokenizer=None, speaker_manager=None,
noise_scale_dp=0.8, sdp_dropout=0.5, dur_loss_weight=1.0, decoder_fp32=False):
super().__init__(config, ap, tokenizer, speaker_manager)
self.encoder.duration_predictor = _DPTap()
h = config.hidden_channels_enc
self.sdp = StochasticDurationPredictor(h, h, 3, sdp_dropout, 4)
self.noise_scale_dp = noise_scale_dp
self.dur_loss_weight = dur_loss_weight
self.decoder_fp32 = decoder_fp32
self.mas_fn = maximum_path
@torch.no_grad()
def inference(self, x, aux_input=None):
aux_input = aux_input or {}
x_lengths = aux_input.get("x_lengths")
if x_lengths is None:
x_lengths = torch.full((x.shape[0],), x.shape[1], dtype=torch.long, device=x.device)
o_mean, o_log_scale, _, x_mask = self.encoder(x, x_lengths, g=None)
x_dp = self.encoder.duration_predictor.cache
logw = self.sdp(x_dp.float(), x_mask.float(), reverse=True, noise_scale=self.noise_scale_dp)
w = torch.exp(logw) * x_mask * self.length_scale
w_ceil = torch.clamp_min(torch.ceil(w), 1) * x_mask
y_lengths = torch.clamp_min(torch.sum(w_ceil, [1, 2]), self.num_squeeze).long()
y_lengths = (y_lengths // self.num_squeeze) * self.num_squeeze
y_mask = sequence_mask(y_lengths, int(y_lengths.max())).unsqueeze(1).to(x_mask.dtype)
attn_mask = x_mask.unsqueeze(-1) * y_mask.unsqueeze(2)
attn = generate_path(w_ceil.squeeze(1), attn_mask.squeeze(1))
attn_t = attn.transpose(1, 2)
y_mean = torch.matmul(attn_t, o_mean.transpose(1, 2)).transpose(1, 2)
y_log_scale = torch.matmul(attn_t, o_log_scale.transpose(1, 2)).transpose(1, 2)
z = (y_mean + torch.exp(y_log_scale) * torch.randn_like(y_mean) * self.inference_noise_scale) * y_mask
y, _ = self.decoder(z, y_mask, g=None, reverse=True)
return {"model_outputs": y.transpose(1, 2), "alignments": attn.transpose(1, 2), "durations": w_ceil}
def main():
TTS_CKPT, VOC_CKPT, OUT = "model.pth", "vocoder.pth", "output.wav"
TEXT = "This is the first sample generated by SupraTTS zero point one Beta... Have fun trying it out."
NOISE, NOISE_DP, LENGTH_SCALE, SEED = 0.333, 0.8, 1.0, 0
torch.manual_seed(SEED)
config = GlowTTSConfig()
config.load_json("config.json")
ap = AudioProcessor(config.audio)
tokenizer, config = TTSTokenizer.init_from_config(config)
model = FlareGlowTTS(config, ap, tokenizer=tokenizer, noise_scale_dp=NOISE_DP)
model.load_checkpoint(config, TTS_CKPT, eval=True)
model.inference_noise_scale, model.length_scale = NOISE, LENGTH_SCALE
model.cuda().eval()
vconfig = load_config("vocoder_config.json")
vocoder = setup_vocoder(vconfig)
vocoder.load_checkpoint(vconfig, VOC_CKPT, eval=True)
vocoder.cuda().eval()
x = torch.LongTensor(tokenizer.text_to_ids(TEXT))[None].cuda()
with torch.inference_mode():
out = model.inference(x, aux_input={"x_lengths": torch.LongTensor([x.shape[1]]).cuda()})
mel = out["model_outputs"].transpose(1, 2)
wav = vocoder.inference(mel).squeeze().cpu().numpy()
ap.save_wav(wav, OUT)
if __name__ == "__main__":
main()