commit cc2d05f027ef2672a96d1b780e5af44b03375af8
parent 968414d2d38f3cb30ac4e1af3151a0842c113366
Author: vin <git@vineetk.net>
Date: Fri, 22 Aug 2025 16:53:10 -0400
improve asr.py
Diffstat:
| M | asr.py | | | 119 | +++++++++++++++++++++++++++++++++++++++++++------------------------------------ |
1 file changed, 65 insertions(+), 54 deletions(-)
diff --git a/asr.py b/asr.py
@@ -1,68 +1,79 @@
import librosa
-import soundfile as sf
import onnx_asr
import os
import sys
-#import numpy as np
+import json
+import numpy as np
from tqdm import tqdm
-print("Loading ASR model...")
-providers = [
- "ROCMExecutionProvider",
- #"CUDAExecutionProvider",
- "CPUExecutionProvider",
-]
-model = onnx_asr.load_model(
- "nemo-parakeet-tdt-0.6b-v3",
- providers=providers,
- #quantization="int8"
-).with_timestamps()
+# converts TimestampedResult to JSON-serializable dictionary
+def result_to_dict(result, offset_seconds):
+ return {
+ "text": result.text,
+ "tokens": result.tokens,
+ "timestamps": [ts + offset_seconds for ts in result.timestamps]
+ }
-# parakeet needs 16khz mono wav as input
-input_file = "input.wav"
-chunk_duration_seconds = 20
-sample_rate = 16000
-
-print(f"Loading and resampling {input_file}...")
-try:
- audio, sr = librosa.load(input_file, sr=sample_rate, mono=True)
-except Exception as e:
- print(f"Error loading audio file: {e}")
- exit()
-
-# process audio in chunks
-chunk_size_samples = int(chunk_duration_seconds * sample_rate)
-full_transcript = []
-temp_dir = "temp_chunks"
-os.makedirs(temp_dir, exist_ok=True)
+def main(input_file, output_file):
+ print("Loading ASR model...")
+ providers = [
+ "ROCMExecutionProvider",
+ "CPUExecutionProvider",
+ ]
+ try:
+ model = onnx_asr.load_model(
+ "nemo-parakeet-tdt-0.6b-v3",
+ providers=providers,
+ ).with_timestamps()
+ except Exception as e:
+ print(f"Failed to load the ASR model. Ensure ONNX Runtime for ROCm is installed correctly. Error: {e}")
+ sys.exit(1)
-print("Starting transcription process in chunks...")
-num_chunks = (len(audio) + chunk_size_samples - 1) // chunk_size_samples
+ chunk_duration_seconds = 20
+ sample_rate = 16000
-for i in tqdm(range(num_chunks), desc="Transcribing"):
- start_sample = i * chunk_size_samples
- end_sample = start_sample + chunk_size_samples
- chunk = audio[start_sample:end_sample]
-
- temp_chunk_file = os.path.join(temp_dir, f"chunk_{i}.wav")
- sf.write(temp_chunk_file, chunk, sample_rate)
-
+ print(f"Loading and resampling {input_file}...")
try:
- transcript = model.recognize(temp_chunk_file)
- if transcript:
- full_transcript.append(transcript)
+ audio, sr = librosa.load(input_file, sr=sample_rate, mono=True)
except Exception as e:
- print(f" Error processing chunk {i + 1}: {e}")
- finally:
- os.remove(temp_chunk_file)
+ print(f"Error loading audio file: {e}")
+ sys.exit(1)
+
+ chunk_size_samples = int(chunk_duration_seconds * sample_rate)
+ full_transcript_data = []
+
+ print("Starting transcription process in chunks...")
+ num_chunks = (len(audio) + chunk_size_samples - 1) // chunk_size_samples
-os.rmdir(temp_dir)
+ for i in tqdm(range(num_chunks), desc="Transcribing"):
+ start_sample = i * chunk_size_samples
+ end_sample = start_sample + chunk_size_samples
+ chunk = audio[start_sample:end_sample]
+
+ chunk_offset_seconds = start_sample / sample_rate
+
+ try:
+ transcript_results = model.recognize(np.copy(chunk))
+
+ if transcript_results:
+ for result in transcript_results:
+ full_transcript_data.append(result_to_dict(result, chunk_offset_seconds))
+ except Exception as e:
+ print(f" Error processing chunk {i + 1}: {e}")
-print("\n" + "="*30)
-print(" FINAL TRANSCRIPT")
-print("="*30)
-print(full_transcript)
+ print(f"\nTranscription complete. Saving results to {output_file}...")
+ with open(output_file, "w") as f:
+ json.dump(full_transcript_data, f, indent=2)
-# onnxruntime has some bug where it doesn't exit properly, aborting instead
-# so hard exit instead (sys.exit does graceful exit)
-os._exit(0)
+ print("Done.")
+ # os._exit(0) workaround for ONNX Runtime bug
+ os._exit(0)
+
+if __name__ == "__main__":
+ if len(sys.argv) < 3:
+ print("Usage: python asr.py <input_audio_file> <output_json_file>")
+ sys.exit(1)
+
+ input_path = sys.argv[1]
+ output_path = sys.argv[2]
+ main(input_path, output_path)