yangrongzhao
Inflect-Nano-v2 for AX650 NPU3
a44ca9d
Raw
History Blame Contribute Delete
7.41 kB
"""
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"] # 128
self.hidden_channels = cfg["model"]["hidden_channels"] # 72
self.max_tokens = 200
self.max_mel = 500
# ---- CPU 组件 ----
self.emb_weight = np.load(str(self.root / "emb_weight.npy"))
# ONNX Runtime sessions
ort_providers = ["CPUExecutionProvider"]
self.dp_sess = ort.InferenceSession(str(self.root / "dp.onnx"), providers=ort_providers)
# ONNX encoder (for CPU-side DP input prep; actual encoder runs on NPU)
# ---- NPU sessions ----
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) # [T, H]
return emb.T[np.newaxis, :, :].astype(np.float32) # [1, H, T]
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] # [t_y, t_x]
token_idx = np.argmax(a, axis=1) # [t_y]
m_exp = np.take(m_p[0], token_idx, axis=-1)[np.newaxis, :, :] # [1, C, t_y]
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)
# CPU: Embedding
x_emb = self._embed(tokens)
# NPU: Encoder
m_p, logs_p, x, x_mask = self._encode(x_emb, tlen)
# CPU: Duration + Alignment
attn, y_mask, mel_len, _ = self._duration_align(x, x_mask, speed)
# CPU: Expand
m_p_exp, logs_p_exp = self._expand(m_p, logs_p, attn)
# Sampling
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
# NPU: Decoder
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