ZipVoice.AXERA / scripts /common_infer.py
HY-2012's picture
Upload the AX630C inference workflow.
e405f21 verified
Raw
History Blame Contribute Delete
4.66 kB
from __future__ import annotations
from pathlib import Path
import numpy as np
def load_tokenizer(repo_dir: str | Path):
from scripts.local_tokenizer import LocalEmiliaTokenizer
token_file = Path(repo_dir) / "resources" / "zipvoice_hf" / "zipvoice" / "tokens.txt"
if not token_file.is_file():
raise FileNotFoundError(f"tokens.txt not found: {token_file}")
return LocalEmiliaTokenizer(token_file=str(token_file))
def extract_prompt_features(
prompt_wav: str | Path,
repo_dir: str | Path,
sampling_rate: int,
feat_scale: float,
target_rms: float,
):
from scripts.local_audio import LocalVocosFbank, load_prompt_wav, rms_norm
import torch
wav = load_prompt_wav(prompt_wav, sampling_rate=sampling_rate)
wav, prompt_rms = rms_norm(wav, target_rms)
extractor = LocalVocosFbank()
features = extractor.extract(wav, sampling_rate=sampling_rate)
if not isinstance(features, torch.Tensor):
features = torch.from_numpy(features)
features = features.unsqueeze(0) * feat_scale
return features.cpu().numpy().astype(np.float32), float(prompt_rms)
def load_vocoder(repo_dir: str | Path):
from scripts.local_audio import load_local_vocos
import torch
vocoder_dir = Path(repo_dir) / "resources" / "vocos-mel-24khz"
if not (vocoder_dir / "config.yaml").is_file() or not (
vocoder_dir / "pytorch_model.bin"
).is_file():
raise FileNotFoundError(f"Local Vocos files not found in {vocoder_dir}")
vocoder = load_local_vocos(vocoder_dir)
vocoder = vocoder.to(torch.device("cpu"))
vocoder.eval()
return vocoder
def vocoder_decode_loaded(
vocoder,
features: np.ndarray,
feat_scale: float,
target_rms: float,
prompt_rms: float,
) -> np.ndarray:
from scripts.local_audio import rms_norm
import torch
feat_tensor = torch.from_numpy(features).float().permute(0, 2, 1) / feat_scale
with torch.no_grad():
wav = vocoder.decode(feat_tensor).squeeze(1).clamp(-1, 1)
wav = rms_norm(wav, target_rms)[0]
if prompt_rms < target_rms:
wav = wav * prompt_rms / target_rms
return wav.squeeze().cpu().numpy()
def load_axmodel_vocoder(model_path: str):
"""Load vocoder axmodel for inference."""
import axengine as axe
session = axe.InferenceSession(model_path)
return session
def axmodel_vocoder_decode(session, features: np.ndarray, feat_scale: float,
target_rms: float, prompt_rms: float) -> np.ndarray:
"""Decode mel features to audio using axmodel vocoder + IRFFT."""
import math
# features shape: (1, T, 100) or (T, 100), squeeze batch dim
features = np.squeeze(features)
if features.ndim == 2:
features = features.T # (T, 100) → (100, T)
T = features.shape[1]
n_mels = features.shape[0]
n_fft, hop = 1024, 256
T_model = 620
# Undo feat_scale and pad to [1, n_mels, T_model]
inv_scale = 1.0 / feat_scale if feat_scale != 0 else 1.0
mel_input = np.zeros((1, n_mels, T_model), dtype=np.float32)
t_actual = min(T, T_model)
for t in range(t_actual):
mel_input[0, :, t] = features[:, t] * inv_scale
# Run axmodel: mel → (real, imag)
real, imag = session.run(None, {"mel": mel_input})
real, imag = real.squeeze(0), imag.squeeze(0) # [T_model, n_freqs]
# IRFFT + overlap-add (matching C++ vocoder)
n_freqs = n_fft // 2 + 1
window = 0.5 * (1.0 - np.cos(2.0 * math.pi * np.arange(n_fft) / (n_fft - 1)))
window_sq = window ** 2
out_len = (T - 1) * hop + n_fft
audio = np.zeros(out_len, dtype=np.float32)
envelope = np.zeros(out_len, dtype=np.float32)
for t_idx in range(T):
spec = np.zeros(n_fft, dtype=np.complex64)
spec[0] = real[t_idx, 0]
for k in range(1, n_freqs - 1):
spec[k] = real[t_idx, k] + 1j * imag[t_idx, k]
spec[n_fft - k] = real[t_idx, k] - 1j * imag[t_idx, k]
spec[n_freqs - 1] = real[t_idx, n_freqs - 1]
ifft_out = np.fft.irfft(spec, n=n_fft).real
pos = t_idx * hop
for n in range(n_fft):
p = pos + n
if p < out_len:
audio[p] += ifft_out[n] * window[n]
envelope[p] += window_sq[n]
audio /= np.maximum(envelope, 1e-10)
pad = n_fft // 2
audio = audio[pad:out_len - pad].astype(np.float32)
# RMS normalize (numpy version)
rms = np.sqrt(np.mean(audio ** 2))
if rms < target_rms and rms > 1e-10:
audio = audio * (target_rms / rms)
if prompt_rms < target_rms:
audio = audio * (prompt_rms / target_rms)
return audio