LGTM / lgtm /inference.py
qnguyen3's picture
LGTM: PyTorch + ONNX weights and inference code
409d4fb verified
Raw History Blame Contribute Delete
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")
@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)