diff options
Diffstat (limited to 'tools/whistle-eval/parakeet_baseline.py')
| -rw-r--r-- | tools/whistle-eval/parakeet_baseline.py | 67 |
1 files changed, 67 insertions, 0 deletions
diff --git a/tools/whistle-eval/parakeet_baseline.py b/tools/whistle-eval/parakeet_baseline.py new file mode 100644 index 0000000..16987de --- /dev/null +++ b/tools/whistle-eval/parakeet_baseline.py @@ -0,0 +1,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()) |
