Download lgtm/inference.py from polyskill/LGTM: direct link, hf CLI and curl.
- Browser
- Download file 6.57 kB
-
https://huggingface.co/polyskill/LGTM/resolve/main/lgtm/inference.py
- Command line
-
hf download hf://polyskill/LGTM/lgtm/inference.py
-
curl -L -o inference.py https://huggingface.co/polyskill/LGTM/resolve/main/lgtm/inference.py
6.57 kB
| """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") | |
| 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) | |
| 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) | |
| 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) | |
| 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) | |
| def save_wav(wav, path): | |
| sf.write(path, wav, SAMPLE_RATE) | |