cosyvoice3-speech-tokenizer-pt / audio_to_tokens_example.py
wookee3's picture
Update README with audio-to-token example and add helper script
9517782 verified
Raw
History Blame Contribute Delete
4.68 kB
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()