diff options
| author | Tom Cooks <tommasogagliardi+github@gmail.com> | 2026-10-10 15:42:04 -0400 |
|---|---|---|
| committer | Tom Cooks <tommasogagliardi+github@gmail.com> | 2026-10-10 15:42:04 -0400 |
| commit | 901bf58c5cf9151942f8c4cd4e1427a02dec4b52 (patch) | |
| tree | e9e312fb9893881f142a6f79e30417cd0c3f2be8 /tools/whistle-eval/whisper_baseline.py | |
| parent | 6db1714d0354d1e1105e330ce84c46a54487c3bb (diff) | |
| download | reccoon-901bf58c5cf9151942f8c4cd4e1427a02dec4b52.tar.gz | |
RECCoon 1.0.0: rebrand to Tom Cooks, Parakeet-only, GPLv3, F-Droid-ready
Diffstat (limited to 'tools/whistle-eval/whisper_baseline.py')
| -rw-r--r-- | tools/whistle-eval/whisper_baseline.py | 65 |
1 files changed, 65 insertions, 0 deletions
diff --git a/tools/whistle-eval/whisper_baseline.py b/tools/whistle-eval/whisper_baseline.py new file mode 100644 index 0000000..9be7014 --- /dev/null +++ b/tools/whistle-eval/whisper_baseline.py @@ -0,0 +1,65 @@ +#!/usr/bin/env python3 +"""Run the same sherpa-onnx multilingual Whisper model that the app uses. + +Usage: + python whisper_baseline.py <audio.wav> [tiny|base] [language] + +Requires a virtualenv with `sherpa-onnx` installed (see README.md). The model +files are expected under ./models/<size>/. +""" +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") + parser.add_argument("size", nargs="?", default="base", choices=["tiny", "base"]) + parser.add_argument("language", nargs="?", default="en") + args = parser.parse_args() + + model_dir = os.path.join(os.path.dirname(__file__), "models", args.size) + encoder = os.path.join(model_dir, "encoder.int8.onnx") + decoder = os.path.join(model_dir, "decoder.int8.onnx") + tokens = os.path.join(model_dir, "tokens.txt") + for path in (encoder, decoder, tokens): + if not os.path.isfile(path): + print(f"missing {path}", file=sys.stderr) + return 2 + + recognizer = sherpa_onnx.OfflineRecognizer.from_whisper( + encoder=encoder, + decoder=decoder, + tokens=tokens, + language=args.language, + task="transcribe", + num_threads=2, + decoding_method="greedy_search", + 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()) |
