"""PyTorch inference for LGTM. from lgtm import LGTMTTS tts = LGTMTTS.from_pretrained("path/or/hf-repo") # downloads from the Hub if needed wav = tts.synthesize("Xin chào!", lang="vi", voice="F1") tts.save_wav(wav, "out.wav") voice = tts.clone_voice("reference.wav") # zero-shot voice from 5-15 s of audio wav = tts.synthesize("Hello there.", lang="en", voice=voice) """ import json import os import re import numpy as np import soundfile as sf import torch import torch.nn.functional as F import torchaudio from safetensors.torch import load_file from .model import LGTM from .text import AVAILABLE_LANGS, TextProcessor SAMPLE_RATE = 44100 def load_voice_style(path, device="cpu"): d = json.load(open(path)) ttl = torch.tensor(np.array(d["style_ttl"]["data"], np.float32).reshape(d["style_ttl"]["dims"])) dp = torch.tensor(np.array(d["style_dp"]["data"], np.float32).reshape(d["style_dp"]["dims"])) return ttl.to(device), dp.to(device) def save_voice_style(path, voice): ttl, dp = voice d = {"style_ttl": {"data": ttl.detach().float().cpu().numpy().tolist(), "dims": list(ttl.shape), "type": "float32"}, "style_dp": {"data": dp.detach().float().cpu().numpy().tolist(), "dims": list(dp.shape), "type": "float32"}} json.dump(d, open(path, "w")) def _prepare_reference(w, top_db=40.0, pad_s=0.1, target_rms_db=-23.0, peak=0.95): """Trim leading/trailing silence and normalise loudness (same as the training data).""" hop = SAMPLE_RATE // 100 n = len(w) // hop if n >= 3: db = 10 * torch.log10(w[: n * hop].view(n, hop).pow(2).mean(1) + 1e-10) active = torch.nonzero(db > db.max() - top_db).flatten() if len(active): pad = int(pad_s * SAMPLE_RATE) w = w[max(int(active[0]) * hop - pad, 0): min((int(active[-1]) + 1) * hop + pad, len(w))] w = w * (10 ** (target_rms_db / 20) / (w.pow(2).mean().sqrt() + 1e-6)) m = w.abs().max() return w * (peak / m) if m > peak else w def split_text(text, max_len=300): """Split long text into sentence chunks of at most max_len characters.""" chunks = [] for para in [p.strip() for p in re.split(r"\n\s*\n+", text.strip()) if p.strip()]: cur = "" for sent in re.split(r"(?<=[.!?。!?])\s+", para): if cur and len(cur) + len(sent) + 1 > max_len: chunks.append(cur) cur = sent else: cur = f"{cur} {sent}".strip() if cur: chunks.append(cur) return chunks class LGTMTTS: def __init__(self, model_dir, device=None): self.device = device or ("cuda" if torch.cuda.is_available() else "cpu") cfg = json.load(open(os.path.join(model_dir, "config.json"))) self.model = LGTM(cfg) self.model.load_state_dict(load_file(os.path.join(model_dir, "pytorch", "model.safetensors"))) self.model.to(self.device).eval() self.tp = TextProcessor(os.path.join(model_dir, "unicode_indexer.json")) self.voice_dir = os.path.join(model_dir, "voice_styles") @classmethod def from_pretrained(cls, repo_or_dir, device=None): if not os.path.isdir(repo_or_dir): from huggingface_hub import snapshot_download repo_or_dir = snapshot_download(repo_or_dir, allow_patterns=["config.json", "unicode_indexer.json", "voice_styles/*", "pytorch/*"]) return cls(repo_or_dir, device) @property def voices(self): return sorted(f[:-5] for f in os.listdir(self.voice_dir) if f.endswith(".json")) def _voice(self, voice): if isinstance(voice, str): path = voice if voice.endswith(".json") else os.path.join(self.voice_dir, f"{voice}.json") return load_voice_style(path, self.device) return voice[0].to(self.device), voice[1].to(self.device) @torch.no_grad() def clone_voice(self, wav_path, max_seconds=15.0): """Voice style from a reference recording (5-15 s of clean speech works best).""" wav, sr = sf.read(wav_path, dtype="float32") if wav.ndim == 2: wav = wav.mean(1) w = torch.from_numpy(wav) if sr != SAMPLE_RATE: w = torchaudio.functional.resample(w, sr, SAMPLE_RATE) w = _prepare_reference(w)[: int(max_seconds * SAMPLE_RATE)] lat = self.model.ae.encode_ttl(w[None].to(self.device)) mask = torch.ones(1, 1, lat.shape[-1], device=self.device) return self.model.ttl.style_encoder(lat, mask), self.model.dp.style_encoder(lat, mask) @torch.no_grad() def _synth_batch(self, texts, langs, voice, steps, speed, cfg_scale): s_ttl, s_dp = voice b = len(texts) ids, mask = self.tp.batch(texts, langs, device=self.device) s_ttl, s_dp = s_ttl.expand(b, -1, -1), s_dp.expand(b, -1, -1) m = self.model dur = m.dp(ids, mask, s_dp) / speed text_emb = m.ttl.encode_text(ids, mask, s_ttl) chunk = m.ae.hop * m.ae.ccf wav_len = (dur * SAMPLE_RATE).long() lat_len = (wav_len + chunk - 1) // chunk L = int(lat_len.max()) lmask = (torch.arange(L, device=self.device)[None] < lat_len[:, None]).float().unsqueeze(1) x = torch.randn(b, m.ae.ldim * m.ae.ccf, L, device=self.device) * lmask total = torch.full((b,), float(steps), device=self.device) for i in range(steps): x = m.ttl.euler_step(x, torch.full((b,), float(i), device=self.device), total, text_emb, mask, s_ttl, lmask, cfg_scale) wav = m.ae.decode_ttl(x) return [wav[k, : wav_len[k]].float().cpu().numpy() for k in range(b)] def synthesize(self, text, lang="en", voice="F1", steps=8, speed=1.05, cfg_scale=4.0, silence=0.3): """Text -> 44.1 kHz float32 waveform. Long text is split into sentence chunks. voice: preset name (F1..F5, M1..M5), a .json path, or the output of clone_voice().""" if lang not in AVAILABLE_LANGS: raise ValueError(f"unsupported language {lang!r}") v = self._voice(voice) chunks = split_text(text, 120 if lang in ("ja", "ko") else 300) wavs = self._synth_batch(chunks, [lang] * len(chunks), v, steps, speed, cfg_scale) gap = np.zeros(int(silence * SAMPLE_RATE), np.float32) out = [] for i, w in enumerate(wavs): out += [w] + ([gap] if i < len(wavs) - 1 else []) return np.concatenate(out) @staticmethod def save_wav(wav, path): sf.write(path, wav, SAMPLE_RATE)