commit d49d20dd810ed2dbcef6c926921b072713366654
parent 78decc2ddd05519780153a8d09f68884e8d11fc6
Author: vin <git@vineetk.net>
Date: Sat, 23 Aug 2025 02:26:35 -0400
adapt asr.py for multiple file inputs
Diffstat:
| M | asr.py | | | 62 | ++++++++++++++++++++++++++++++-------------------------------- |
1 file changed, 30 insertions(+), 32 deletions(-)
diff --git a/asr.py b/asr.py
@@ -6,6 +6,23 @@ import json
import numpy as np
from tqdm import tqdm
+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)
+
+chunk_duration_seconds = 20
+sample_rate = 16000
+
# converts TimestampedResult to JSON-serializable dictionary
def result_to_dict(result, offset_seconds):
token_map = [
@@ -24,37 +41,18 @@ def result_to_dict(result, offset_seconds):
}
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)
-
- 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}")
+ tqdm.write(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
- for i in tqdm(range(num_chunks), desc="Transcribing"):
+ for i in tqdm(range(num_chunks), desc="Transcribing", position=1, leave=False):
start_sample = i * chunk_size_samples
end_sample = start_sample + chunk_size_samples
chunk = audio[start_sample:end_sample]
@@ -73,21 +71,21 @@ def main(input_file, output_file):
for result in results_list:
full_transcript_data.append(result_to_dict(result, chunk_offset_seconds))
except Exception as e:
- print(f" Error processing chunk {i + 1}: {e}")
+ tqdm.write(f" Error processing chunk {i + 1}: {e}")
- print(f"\nTranscription complete. Saving results to {output_file}...")
with open(output_file, "w") as f:
json.dump(full_transcript_data, f, indent=2)
- 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>")
+ if len(sys.argv) < 2:
+ print("Usage: python asr.py input_audio_file1 [input_audio_file2 [input_audio_file3 [...]]]")
sys.exit(1)
- input_path = sys.argv[1]
- output_path = sys.argv[2]
- main(input_path, output_path)
+ all_inputs = sys.argv[1:]
+ inputs = [f for f in all_inputs if not os.path.exists(os.path.splitext(f)[0] + ".json")]
+
+ for i in tqdm(inputs, desc="Processing files", position=0):
+ main(i, os.path.splitext(i)[0] + ".json")
+
+# os._exit(0) workaround for ONNX Runtime bug
+os._exit(0)