Download infer_v2.py from SupraLabs/SupraTTS-0.1-Beta: direct link, hf CLI and curl.
- Browser
- Download file 4.38 kB
-
https://huggingface.co/SupraLabs/SupraTTS-0.1-Beta/resolve/main/infer_v2.py
- Command line
-
hf download hf://SupraLabs/SupraTTS-0.1-Beta/infer_v2.py
-
curl -L -o infer_v2.py https://huggingface.co/SupraLabs/SupraTTS-0.1-Beta/resolve/main/infer_v2.py
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 | |
| 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() |