File size: 4,681 Bytes
9517782
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
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()