| """AX board inference: acoustic (CPU/ONNX) + BigVGAN (NPU/axmodel). |
| |
| Run on the AX board with: |
| python3 infer_board.py --text "你好" --acoustic acoustic_female.onnx \ |
| --vocoder bigvgan_base.axmodel --output out.wav |
| """ |
|
|
| import argparse |
| import os |
| import sys |
|
|
| import numpy as np |
|
|
| sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) |
| from cn_frontend import text_to_sequence |
|
|
| MAX_TEXT = 256 |
| MAX_MEL = 1024 |
| VOC_CHUNK = 512 |
| VOC_OVERLAP = 48 |
| MAX_SENT_TOKENS = 250 |
| MEL_GATE_LO = -3.5 |
| MEL_GATE_HI = -2.2 |
| LENGTH_SCALE = 0.7 |
| TAIL_STRETCH_OLD = 20 |
| TAIL_STRETCH_NEW = 45 |
|
|
|
|
| def load_acoustic(path): |
| """Load acoustic ONNX (CPU, fp32).""" |
| import onnxruntime as ort |
| so = ort.SessionOptions() |
| so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL |
| return ort.InferenceSession(path, so, providers=["CPUExecutionProvider"]) |
|
|
|
|
| def load_vocoder(path): |
| """Load BigVGAN axmodel (NPU via axengine).""" |
| if path.endswith(".onnx"): |
| import onnxruntime as ort |
| so = ort.SessionOptions() |
| so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL |
| return ort.InferenceSession(path, so, providers=["CPUExecutionProvider"]) |
| import axengine as axe |
| try: |
| return axe.InferenceSession(path, providers=["AxEngineExecutionProvider"]) |
| except Exception: |
| return axe.InferenceSession( |
| path, providers=["RemoteAXExecutionProvider"], |
| provider_options={"host": os.environ.get("AX_BOARD", "127.0.0.1"), |
| "port": "18500"}) |
|
|
|
|
| def split_sentences(text, max_tokens=MAX_SENT_TOKENS): |
| """Split long text into sentences that fit the acoustic model.""" |
| ids = text_to_sequence(text) |
| tokens = [0] |
| for pid in ids: |
| tokens.append(pid) |
| tokens.append(0) |
| if len(tokens) <= max_tokens: |
| return [text] |
| |
| import re |
| parts = re.split(r"([,。!?;:、,.!?;:])", text) |
| sents, cur = [], "" |
| for p in parts: |
| cand = cur + p |
| cand_tokens = len([0]) + 2 * len(text_to_sequence(cand)) |
| if len(text_to_sequence(cand)) * 2 + 1 > max_tokens and cur: |
| sents.append(cur) |
| cur = p |
| else: |
| cur = cand |
| if cur: |
| sents.append(cur) |
| |
| out = [] |
| for s in sents: |
| while len(text_to_sequence(s)) * 2 + 1 > max_tokens: |
| cut = max(1, (max_tokens - 1) // 2) |
| out.append(s[:cut]) |
| s = s[cut:] |
| if s: |
| out.append(s) |
| return [s for s in out if s.strip()] |
|
|
|
|
| def text_to_inputs(text, noise_scale, seed): |
| np.random.seed(seed) |
| ids = text_to_sequence(text) |
| tokens = [0] |
| for pid in ids: |
| tokens.append(pid) |
| tokens.append(0) |
| n = min(len(tokens), MAX_TEXT) |
| x = np.zeros((1, MAX_TEXT), dtype=np.int64) |
| x[0, :n] = tokens[:n] |
| x_lengths = np.array([n], dtype=np.int64) |
| noise_z = np.random.randn(1, 192, MAX_MEL).astype(np.float32) |
| if noise_scale != 0.3: |
| noise_z = noise_z * (noise_scale / 0.3) |
| return x, x_lengths, noise_z |
|
|
|
|
| def acoustic_to_mel(acoustic, x, x_lengths, noise_z): |
| mel, y_lengths = acoustic.run(None, { |
| "x": x, "x_lengths": x_lengths, "noise_z": noise_z, |
| }) |
| return mel[:, :, :int(y_lengths[0])] |
|
|
|
|
| def mel_soft_gate(mel, thr_lo=MEL_GATE_LO, thr_hi=MEL_GATE_HI, floor=-11.5): |
| """压平预测 mel 的弱帧(静音/停顿)以消除渲染底噪。 |
| 只调整帧能量 < thr_lo 的帧(向 floor 收敛)和过渡区, |
| 强语音帧(>= thr_hi)完全保持。""" |
| out = mel.copy() |
| fe = out.max(axis=1) |
| for i in range(mel.shape[2]): |
| f = fe[0, i] |
| if f < thr_lo: |
| target = floor |
| elif f < thr_hi: |
| t = (f - thr_lo) / (thr_hi - thr_lo) |
| target = floor * (1 - t) + f * t |
| else: |
| continue |
| out[:, :, i] = out[:, :, i] - f + target |
| return out |
|
|
|
|
| def mel_sharpen(mel, alpha=0.5, k=5): |
| """Spectral contrast boost along the mel-band axis: m + alpha*(m - smooth(m)).""" |
| if alpha <= 0 or k <= 1: |
| return mel |
| m = mel[0] if mel.ndim == 3 else mel |
| pad = k // 2 |
| mp = np.pad(m, ((pad, pad), (0, 0)), mode="reflect") |
| smooth = np.zeros_like(m) |
| for i in range(k): |
| smooth += mp[i:i + m.shape[0]] |
| smooth /= k |
| sharp = m + alpha * (m - smooth) |
| return (sharp if mel.ndim == 2 else sharp[None]).astype(np.float32) |
|
|
|
|
| def tail_stretch(mel, n_old=TAIL_STRETCH_OLD, n_new=TAIL_STRETCH_NEW): |
| """句尾 mel 拉伸:最后 n_old 帧线性插值到 n_new 帧。 |
| 模型对句尾音节 duration 预测偏短(如 11 帧 vs 参考 43 帧), |
| 拉伸后尾字清晰完整。""" |
| T = mel.shape[2] |
| if T <= n_old: |
| return mel |
| tail = mel[:, :, -n_old:] |
| xo = np.linspace(0, 1, n_old) |
| xn = np.linspace(0, 1, n_new) |
| nt = np.stack([np.interp(xn, xo, tail[0, c]) for c in range(mel.shape[1])]) |
| return np.concatenate([mel[:, :, :-n_old], nt[None].astype(np.float32)], axis=2) |
|
|
|
|
| def vocoder_chunked(vocoder, mel): |
| """Run BigVGAN on arbitrary-length mel via overlapped chunks + crossfade.""" |
| mel = mel.astype(np.float32) |
| T = mel.shape[2] |
| if T <= VOC_CHUNK: |
| m = np.pad(mel, ((0, 0), (0, 0), (0, VOC_CHUNK - T)), |
| constant_values=-11.5) |
| w = vocoder.run(None, {"mel": m})[0][0, 0] |
| return w[:T * 256] |
| hop = VOC_CHUNK - VOC_OVERLAP |
| starts = list(range(0, T, hop)) |
| frame = 256 |
| wav = np.zeros(T * frame) |
| for i, s in enumerate(starts): |
| v = min(VOC_CHUNK, T - s) |
| block = mel[:, :, s:s + VOC_CHUNK] |
| if block.shape[2] < VOC_CHUNK: |
| block = np.pad(block, ((0, 0), (0, 0), (0, VOC_CHUNK - block.shape[2])), |
| constant_values=-11.5) |
| c = vocoder.run(None, {"mel": block})[0][0, 0] |
| c = c[:v * frame] |
| base = s * frame |
| ov = 0 |
| if i > 0: |
| ov = min(VOC_OVERLAP * frame, base, len(c)) |
| fade = np.linspace(0, 1, ov) |
| wav[base - ov:base] = wav[base - ov:base] * (1 - fade) + c[:ov] * fade |
| wav[base + ov:base + len(c)] = c[ov:] |
| return wav |
|
|
|
|
| def compress_pauses(wav, sr=24000, min_pause_ms=90, target_ms=60): |
| """Compress long silences to reduce choppy rhythm.""" |
| win, hop = int(sr * 0.01), int(sr * 0.005) |
| n = (len(wav) - win) // hop |
| e = np.array([np.sqrt(np.mean(wav[i*hop:i*hop+win]**2)) for i in range(n)]) |
| thr = max(e.max() * 0.12, 0.008) |
| sil = e < thr |
| |
| runs = [] |
| i = 0 |
| while i < len(sil): |
| if sil[i]: |
| j = i |
| while j < len(sil) and sil[j]: |
| j += 1 |
| runs.append((i, j)) |
| i = j |
| else: |
| i += 1 |
| keep = np.ones(len(wav), dtype=bool) |
| for s, e_ in runs: |
| dur_ms = (e_ - s) * hop / sr * 1000 |
| if dur_ms > min_pause_ms: |
| s0, s1 = s * hop, min(e_ * hop + win, len(wav)) |
| target = int(target_ms / 1000 * sr) |
| |
| if s1 - s0 > target: |
| keep[s0 + target:s1] = False |
| out = wav[keep] |
| return out |
|
|
|
|
| def synthesize(acoustic, vocoder, text, noise_scale=0.0, seed=0): |
| sents = split_sentences(text) |
| pieces = [] |
| for s in sents: |
| x, x_lengths, noise_z = text_to_inputs(s, noise_scale, seed) |
| mel = acoustic_to_mel(acoustic, x, x_lengths, noise_z) |
| mel = mel_sharpen(mel) |
| mel = mel_soft_gate(mel) |
| mel = tail_stretch(mel) |
| pieces.append(vocoder_chunked(vocoder, mel)) |
| if len(pieces) == 1: |
| return pieces[0], pieces[0].shape[0] // 256 |
| |
| gap = np.zeros(int(24000 * 0.2)) |
| wav = pieces[0] |
| for p in pieces[1:]: |
| wav = np.concatenate([wav, gap, p]) |
| return wav, wav.shape[0] // 256 |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--text", required=True) |
| parser.add_argument("--acoustic", default="export/acoustic_female.onnx") |
| parser.add_argument("--vocoder", default="export/axmodel/bigvgan_base.axmodel") |
| parser.add_argument("--output", default="board_out.wav") |
| parser.add_argument("--noise_scale", type=float, default=0.0) |
| args = parser.parse_args() |
|
|
| import soundfile as sf |
| acoustic = load_acoustic(args.acoustic) |
| vocoder = load_vocoder(args.vocoder) |
| wav, T = synthesize(acoustic, vocoder, args.text, |
| noise_scale=args.noise_scale) |
| wav = wav / (np.abs(wav).max() + 1e-8) * 0.95 |
| sf.write(args.output, wav, 24000) |
| print(f"saved: {args.output} ({len(wav)/24000:.2f}s, mel_frames={T})") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|