summaryrefslogtreecommitdiff
path: root/tools/whistle-eval/parakeet_baseline.py
blob: 16987de73a9c3954a0c6e394e88e6f31ee1db1f3 (plain)
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
#!/usr/bin/env python3
"""Run NVIDIA Parakeet TDT 0.6B v3 (multilingual) through sherpa-onnx.

Usage:
    python parakeet_baseline.py <audio.wav>

Requires a virtualenv with `sherpa-onnx` installed and the int8 model under
./models/parakeet-v3-int8/ (see README.md).
"""
import argparse
import os
import sys
import wave

import numpy as np
import sherpa_onnx


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("audio")
    args = parser.parse_args()

    model_dir = os.path.join(os.path.dirname(__file__), "models", "parakeet-v3-int8")
    encoder = os.path.join(model_dir, "encoder.int8.onnx")
    decoder = os.path.join(model_dir, "decoder.int8.onnx")
    joiner = os.path.join(model_dir, "joiner.int8.onnx")
    tokens = os.path.join(model_dir, "tokens.txt")
    for path in (encoder, decoder, joiner, tokens):
        if not os.path.isfile(path):
            print(f"missing {path}", file=sys.stderr)
            return 2

    recognizer = sherpa_onnx.OfflineRecognizer.from_transducer(
        encoder=encoder,
        decoder=decoder,
        joiner=joiner,
        tokens=tokens,
        num_threads=2,
        sample_rate=16000,
        feature_dim=80,
        decoding_method="greedy_search",
        model_type="nemo_transducer",
        debug=False,
    )

    with wave.open(args.audio, "rb") as wav:
        sample_rate = wav.getframerate()
        channels = wav.getnchannels()
        width = wav.getsampwidth()
        frames = wav.readframes(wav.getnframes())
    if width != 2:
        print("only 16-bit PCM WAV is supported", file=sys.stderr)
        return 2
    samples = np.frombuffer(frames, dtype=np.int16).astype(np.float32) / 32768.0
    if channels > 1:
        samples = samples.reshape(-1, channels).mean(axis=1)

    stream = recognizer.create_stream()
    stream.accept_waveform(sample_rate, samples)
    recognizer.decode_stream(stream)
    print(stream.result.text)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())