File size: 4,644 Bytes
b0c7e09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
"""
Minimal streaming inference example for Silero VAD v6 (MLX, 16 kHz).

Reads a 16 kHz mono PCM-16 WAV file, runs streaming VAD, and prints per-chunk
speech probabilities for the first 25 chunks plus the speech ratio at threshold
0.5 across the whole clip.

Usage:
    uv pip install mlx safetensors numpy huggingface_hub
    uv run python example.py path/to/audio_16k_mono.wav
"""
from __future__ import annotations

import sys
import wave
from pathlib import Path

import mlx.core as mx
import numpy as np
from huggingface_hub import hf_hub_download
from safetensors.numpy import load_file


CHUNK_SAMPLES = 512
CONTEXT_SAMPLES = 64
STFT_HOP = 128
STFT_PAD_RIGHT = 64
LSTM_HIDDEN = 128
SAMPLE_RATE = 16_000

KEYS = [
    "vad_16k.stft_conv.weight",
    "vad_16k.conv1.weight", "vad_16k.conv1.bias",
    "vad_16k.conv2.weight", "vad_16k.conv2.bias",
    "vad_16k.conv3.weight", "vad_16k.conv3.bias",
    "vad_16k.conv4.weight", "vad_16k.conv4.bias",
    "vad_16k.lstm.Wx", "vad_16k.lstm.Wh", "vad_16k.lstm.bias",
    "vad_16k.final_conv.weight", "vad_16k.final_conv.bias",
]


def load_weights(path: str) -> dict[str, mx.array]:
    raw = load_file(path)
    return {k: mx.array(raw[k]) for k in KEYS}


def reflect_pad_right_576(samples_576: np.ndarray) -> np.ndarray:
    """PyTorch ReflectionPad1d(right=64): out[L+i] = in[L-2-i] for i in 0..63."""
    L = samples_576.shape[-1]
    pad = np.empty(STFT_PAD_RIGHT, dtype=samples_576.dtype)
    for i in range(STFT_PAD_RIGHT):
        pad[i] = samples_576[L - 2 - i]
    return np.concatenate([samples_576, pad], axis=-1)


def predict_chunk(audio_640: mx.array, h: mx.array, c: mx.array, w: dict
                   ) -> tuple[mx.array, mx.array, mx.array]:
    z = mx.conv1d(audio_640, w["vad_16k.stft_conv.weight"], stride=STFT_HOP, padding=0)
    real, imag = z[:, :, :129], z[:, :, 129:]
    x = mx.sqrt(real * real + imag * imag + 1e-12)
    cfgs = [(1, 1), (2, 1), (2, 1), (1, 1)]
    for i, (stride, padding) in enumerate(cfgs, start=1):
        x = mx.conv1d(x, w[f"vad_16k.conv{i}.weight"], stride=stride, padding=padding)
        x = x + w[f"vad_16k.conv{i}.bias"]
        x = mx.maximum(x, 0)
    feat = x[:, 0, :]
    gates = feat @ w["vad_16k.lstm.Wx"].T + h @ w["vad_16k.lstm.Wh"].T + w["vad_16k.lstm.bias"]
    i_g, f_g, g_g, o_g = mx.split(gates, 4, axis=-1)
    c_new = mx.sigmoid(f_g) * c + mx.sigmoid(i_g) * mx.tanh(g_g)
    h_new = mx.sigmoid(o_g) * mx.tanh(c_new)
    dec = mx.conv1d(
        mx.maximum(h_new, 0)[:, None, :],
        w["vad_16k.final_conv.weight"], stride=1, padding=0,
    )
    dec = dec + w["vad_16k.final_conv.bias"]
    return mx.sigmoid(dec), h_new, c_new


def load_wav_16k_mono(path: str) -> np.ndarray:
    with wave.open(path, "rb") as w:
        if w.getframerate() != SAMPLE_RATE:
            raise SystemExit(f"expected 16 kHz, got {w.getframerate()}")
        if w.getnchannels() != 1:
            raise SystemExit(f"expected mono, got {w.getnchannels()} channels")
        if w.getsampwidth() != 2:
            raise SystemExit(f"expected 16-bit PCM, got {w.getsampwidth()*8}-bit")
        n = w.getnframes()
        raw = w.readframes(n)
    return (np.frombuffer(raw, dtype=np.int16).astype(np.float32) / 32768.0).copy()


def main() -> None:
    if len(sys.argv) != 2:
        print(__doc__, file=sys.stderr)
        sys.exit(1)

    weights_path = hf_hub_download(repo_id="mlx-community/silero-vad-v6",
                                    filename="model.safetensors")
    w = load_weights(weights_path)

    samples = load_wav_16k_mono(sys.argv[1])
    n_chunks = samples.shape[0] // CHUNK_SAMPLES
    h = mx.zeros((1, LSTM_HIDDEN))
    c = mx.zeros((1, LSTM_HIDDEN))
    context = np.zeros(CONTEXT_SAMPLES, dtype=np.float32)

    n_speech = 0
    threshold = 0.5
    print(f"chunks={n_chunks}, duration={samples.shape[0]/SAMPLE_RATE:.2f}s")
    print(f"{'idx':>5} {'time(s)':>8} {'p(speech)':>10}")
    for ci in range(n_chunks):
        chunk = samples[ci * CHUNK_SAMPLES : (ci + 1) * CHUNK_SAMPLES]
        merged = np.concatenate([context, chunk])
        padded = reflect_pad_right_576(merged)
        audio = mx.array(padded[None, :, None])
        prob, h, c = predict_chunk(audio, h, c, w)
        mx.eval(prob, h, c)
        p = float(prob[0, 0, 0].item())
        if p >= threshold:
            n_speech += 1
        if ci < 25:
            print(f"{ci:>5d} {ci*CHUNK_SAMPLES/SAMPLE_RATE:>8.3f} {p:>10.4f}")
        context = chunk[-CONTEXT_SAMPLES:].copy()

    print(f"\nspeech chunks (threshold {threshold}): {n_speech}/{n_chunks} "
          f"({n_speech/n_chunks*100:.1f}%)")


if __name__ == "__main__":
    main()