summaryrefslogtreecommitdiff
path: root/tools/whistle-eval/whisper_baseline.py
diff options
context:
space:
mode:
Diffstat (limited to 'tools/whistle-eval/whisper_baseline.py')
-rw-r--r--tools/whistle-eval/whisper_baseline.py65
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())