| """ |
| Inflect-Nano-v2 TTS Engine for AX650 NPU3 |
| 纯 NPU 推理 + 轻量 CPU(numpy + onnxruntime),无 PyTorch 依赖。 |
| """ |
| from __future__ import annotations |
| import sys, re, math |
| from pathlib import Path |
| import numpy as np |
| import onnxruntime as ort |
| import axengine as axe |
|
|
| PKG = Path(__file__).resolve().parent |
| sys.path.insert(0, str(PKG)) |
|
|
| from inflect_vits_frontend import run_vits_frontend |
| from text import cleaned_text_to_sequence |
| from text.symbols import symbols |
|
|
|
|
| def sequence_mask(length, max_length=None): |
| if max_length is None: |
| max_length = length.max() |
| x = np.arange(max_length, dtype=length.dtype) |
| return x[np.newaxis, :] < length[:, np.newaxis] |
|
|
|
|
| def generate_path(duration, mask): |
| """duration: [b, 1, t_x], mask: [b, 1, t_y, t_x] |
| Fully vectorized — matches PyTorch commons.generate_path exactly.""" |
| b, _, t_y, t_x = mask.shape |
| w = np.cumsum(duration, axis=-1).astype(np.int64) |
| starts = np.zeros_like(w) |
| starts[:, :, 1:] = w[:, :, :-1] |
| idx = np.arange(t_y, dtype=np.int64).reshape(1, 1, t_y, 1) |
| path = ((idx >= starts[:, :, np.newaxis, :]) & |
| (idx < w[:, :, np.newaxis, :])).astype(np.float32) |
| return path * mask |
|
|
|
|
| class InflectTTSEngine: |
| """Inflect-Nano-v2 TTS — AX650 NPU3 only. 纯 numpy+onnx,零 torch。""" |
|
|
| def __init__(self, model_dir: str | Path | None = None): |
| if model_dir: |
| self.root = Path(model_dir) |
| else: |
| self.root = PKG.parent.parent / "models" |
|
|
| config_path = str(self.root / "config.json") |
| import json |
| with open(config_path) as f: |
| cfg = json.load(f) |
|
|
| self.sample_rate = cfg["data"]["sampling_rate"] |
| self.hop_length = cfg["data"]["hop_length"] |
| self.add_blank = cfg["data"]["add_blank"] |
| self.inter_channels = cfg["model"]["inter_channels"] |
| self.hidden_channels = cfg["model"]["hidden_channels"] |
| self.max_tokens = 200 |
| self.max_mel = 500 |
|
|
| |
| self.emb_weight = np.load(str(self.root / "emb_weight.npy")) |
| |
| ort_providers = ["CPUExecutionProvider"] |
| self.dp_sess = ort.InferenceSession(str(self.root / "dp.onnx"), providers=ort_providers) |
| |
|
|
| |
| enc_ax = str(self.root / "inflect_encoder.axmodel") |
| dec_ax = str(self.root / "inflect_decoder.axmodel") |
| self.enc_session = axe.InferenceSession(enc_ax) |
| self.dec_session = axe.InferenceSession(dec_ax) |
|
|
| def _embed(self, tokens: np.ndarray) -> np.ndarray: |
| """token ids [1, T] → x_emb [1, H, T]""" |
| emb = self.emb_weight[tokens[0]] * math.sqrt(self.hidden_channels) |
| return emb.T[np.newaxis, :, :].astype(np.float32) |
|
|
| def _encode(self, x_emb: np.ndarray, tlen: int) -> tuple: |
| """NPU encoder: x_emb → m_p, logs_p, x, x_mask (trimmed to tlen)""" |
| MAX = self.max_tokens |
| x_pad = np.zeros((1, self.hidden_channels, MAX), dtype=np.float32) |
| x_pad[:, :, :tlen] = x_emb[:, :, :tlen] |
|
|
| out = self.enc_session.run(None, { |
| "x_emb": x_pad, |
| "lengths": np.array([tlen], dtype=np.int32), |
| }) |
| return (out[0][:, :, :tlen], out[1][:, :, :tlen], |
| out[2][:, :, :tlen], out[3][:, :, :tlen]) |
|
|
| def _duration_align(self, x: np.ndarray, x_mask: np.ndarray, |
| speed: float) -> tuple: |
| """DP + alignment on CPU (onnxruntime)""" |
| MAX = self.max_tokens |
| tlen = x.shape[2] |
| x_pad = np.pad(x, ((0, 0), (0, 0), (0, MAX - tlen))).astype(np.float32) |
| m_pad = np.pad(x_mask, ((0, 0), (0, 0), (0, MAX - tlen))).astype(np.float32) |
|
|
| logw = self.dp_sess.run(None, {"x": x_pad, "x_mask": m_pad})[0] |
| logw = logw[:, :, :tlen] |
|
|
| w = np.exp(logw) * x_mask * (1.0 / speed) |
| w_ceil = np.ceil(w) |
| y_len = max(int(np.sum(w_ceil)), 1) |
| y_mask = sequence_mask(np.array([y_len]), None).astype(np.float32)[:, np.newaxis, :] |
| attn_mask = x_mask[:, :, np.newaxis, :] * y_mask[:, :, :, np.newaxis] |
| attn = generate_path(w_ceil, attn_mask) |
| return attn, y_mask, y_len, logw |
|
|
| def _expand(self, m_p, logs_p, attn): |
| """Expand via argmax gather — uses np.take to avoid mixed-indexing transpose.""" |
| a = attn[0, 0] |
| token_idx = np.argmax(a, axis=1) |
| m_exp = np.take(m_p[0], token_idx, axis=-1)[np.newaxis, :, :] |
| l_exp = np.take(logs_p[0], token_idx, axis=-1)[np.newaxis, :, :] |
| return m_exp, l_exp |
|
|
| def _decode(self, z_p: np.ndarray, y_mask: np.ndarray, mel_len: int) -> np.ndarray: |
| """NPU decoder: z_p + y_mask → waveform""" |
| MAX = self.max_mel |
| zp = np.zeros((1, self.inter_channels, MAX), dtype=np.float32) |
| ym = np.zeros((1, 1, MAX), dtype=np.float32) |
| zp[:, :, :mel_len] = z_p[:, :, :mel_len] |
| ym[:, :, :mel_len] = y_mask[:, :, :mel_len] |
|
|
| out = self.dec_session.run(None, { |
| "z_p": zp, |
| "y_mask": ym, |
| }) |
| return out[0][0, 0, :mel_len * self.hop_length] |
|
|
| def synthesize(self, text: str, speed: float = 1.0, variation: float = 0.667, |
| seed: int = 0): |
| normalized = " ".join(text.split()) |
| if not normalized: |
| raise ValueError("Text must not be empty.") |
|
|
| sentences = [p.strip() for p in re.split(r"(?<=[.!?;:])\s+", normalized) if p.strip()] |
| if not sentences: |
| sentences = [normalized] |
|
|
| pieces = [] |
| for idx, chunk in enumerate(sentences): |
| if idx > 0: |
| pause = round(self.sample_rate * 0.08) |
| pieces.append(np.zeros(pause, dtype=np.float32)) |
|
|
| phonemes = run_vits_frontend(chunk).phoneme_text |
| seq = cleaned_text_to_sequence(phonemes) |
| if self.add_blank: |
| seq = intersperse(seq, 0) |
| if not seq: |
| continue |
|
|
| tokens = np.array([seq], dtype=np.int64) |
| tlen = tokens.shape[1] |
| np.random.seed(seed + idx) |
|
|
| |
| x_emb = self._embed(tokens) |
|
|
| |
| m_p, logs_p, x, x_mask = self._encode(x_emb, tlen) |
|
|
| |
| attn, y_mask, mel_len, _ = self._duration_align(x, x_mask, speed) |
|
|
| |
| m_p_exp, logs_p_exp = self._expand(m_p, logs_p, attn) |
|
|
| |
| z_p = m_p_exp + np.random.default_rng(seed + idx).standard_normal(m_p_exp.shape, dtype=np.float32) * np.exp(logs_p_exp) * variation |
|
|
| |
| waveform = self._decode(z_p, y_mask, mel_len) |
| pieces.append(waveform) |
|
|
| waveform = np.clip(np.concatenate(pieces), -1.0, 1.0) |
| return self.sample_rate, waveform |
|
|
| def save(self, text: str, output: str | Path, **kwargs): |
| import soundfile as sf |
| dest = Path(output) |
| dest.parent.mkdir(parents=True, exist_ok=True) |
| sr, wav = self.synthesize(text, **kwargs) |
| sf.write(dest, wav, sr) |
| return dest |
|
|
|
|
| def intersperse(lst, item): |
| result = [item] * (len(lst) * 2 + 1) |
| result[1::2] = lst |
| return result |
|
|