from __future__ import annotations import argparse from pathlib import Path import numpy as np import onnxruntime as ort import soundfile as sf from huggingface_hub import hf_hub_download SOURCE_REPO = "FunAudioLLM/Fun-CosyVoice3-0.5B-2512" SOURCE_ONNX = "speech_tokenizer_v3.onnx" def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser(description="Audio to token sequence via CosyVoice3 ONNX") p.add_argument("--audio", required=True, help="Input wav path") p.add_argument("--onnx", default=None, help="Local onnx path; if omitted download from source repo") p.add_argument("--sample-rate", type=int, default=16000) p.add_argument("--n-mels", type=int, default=128) p.add_argument("--win-length", type=int, default=400) p.add_argument("--hop-length", type=int, default=160) p.add_argument("--n-fft", type=int, default=512) return p.parse_args() def load_audio(path: Path, target_sr: int) -> np.ndarray: wav, sr = sf.read(str(path), dtype="float32") if wav.ndim == 2: wav = wav.mean(axis=1) wav = np.asarray(wav, dtype=np.float32) if sr != target_sr: wav = resample_linear(wav, sr, target_sr) return wav def resample_linear(x: np.ndarray, src_sr: int, tgt_sr: int) -> np.ndarray: if src_sr == tgt_sr or x.size < 2: return x.astype(np.float32, copy=False) new_len = max(1, int(round(x.shape[0] * float(tgt_sr) / float(src_sr)))) old = np.arange(x.shape[0], dtype=np.float64) new = np.linspace(0, x.shape[0] - 1, num=new_len, dtype=np.float64) return np.interp(new, old, x.astype(np.float64)).astype(np.float32) def hz_to_mel(hz: np.ndarray) -> np.ndarray: return 2595.0 * np.log10(1.0 + hz / 700.0) def mel_to_hz(mel: np.ndarray) -> np.ndarray: return 700.0 * (10.0 ** (mel / 2595.0) - 1.0) def mel_filterbank(sr: int, n_fft: int, n_mels: int, f_min: float = 0.0, f_max: float | None = None) -> np.ndarray: if f_max is None: f_max = sr / 2.0 mel_min = hz_to_mel(np.array([f_min]))[0] mel_max = hz_to_mel(np.array([f_max]))[0] mels = np.linspace(mel_min, mel_max, n_mels + 2) hz = mel_to_hz(mels) bins = np.floor((n_fft + 1) * hz / sr).astype(int) fb = np.zeros((n_mels, n_fft // 2 + 1), dtype=np.float32) for i in range(1, n_mels + 1): left, center, right = bins[i - 1], bins[i], bins[i + 1] if center > left: fb[i - 1, left:center] = (np.arange(left, center) - left) / float(center - left) if right > center: fb[i - 1, center:right] = (right - np.arange(center, right)) / float(right - center) return fb def extract_logmel(wav: np.ndarray, sr: int, n_fft: int, win_length: int, hop_length: int, n_mels: int) -> np.ndarray: if wav.shape[0] < win_length: wav = np.pad(wav, (0, win_length - wav.shape[0])) frames = 1 + (wav.shape[0] - win_length) // hop_length if frames <= 0: frames = 1 pad_len = win_length + hop_length * (frames - 1) if pad_len > wav.shape[0]: wav = np.pad(wav, (0, pad_len - wav.shape[0])) idx = np.arange(win_length)[None, :] + hop_length * np.arange(frames)[:, None] framed = wav[idx] window = np.hanning(win_length).astype(np.float32) spectrum = np.fft.rfft(framed * window[None, :], n=n_fft, axis=1) power = (np.abs(spectrum) ** 2).astype(np.float32) fb = mel_filterbank(sr=sr, n_fft=n_fft, n_mels=n_mels) mel = power @ fb.T logmel = np.log(np.clip(mel, 1e-10, None)).astype(np.float32) return logmel.T[None, :, :] def resolve_onnx(path: str | None) -> str: if path: return path return hf_hub_download(repo_id=SOURCE_REPO, filename=SOURCE_ONNX) def main() -> None: args = parse_args() audio = load_audio(Path(args.audio), target_sr=args.sample_rate) feats = extract_logmel( wav=audio, sr=args.sample_rate, n_fft=args.n_fft, win_length=args.win_length, hop_length=args.hop_length, n_mels=args.n_mels, ) feats_len = np.asarray([feats.shape[-1]], dtype=np.int32) onnx_path = resolve_onnx(args.onnx) sess = ort.InferenceSession(onnx_path, providers=["CPUExecutionProvider"]) input_names = [i.name for i in sess.get_inputs()] feats_name = next((n for n in input_names if "feat" in n.lower()), input_names[0]) len_name = next((n for n in input_names if "length" in n.lower()), input_names[-1]) outputs = sess.run(None, {feats_name: feats.astype(np.float32), len_name: feats_len}) indices = np.asarray(outputs[0]) tokens = indices[0].tolist() if indices.ndim >= 2 else indices.tolist() print(tokens) if __name__ == "__main__": main()