Chaka-ASR / examples /transcribe.py
caesar-abrham's picture
Add Chaka-ASR model card, inference example, and upstream credits
047d4ce verified
Raw History Blame Contribute Delete
1.71 kB
"""Transcribe a short WAV or FLAC recording with Chaka-ASR."""
import argparse
from math import gcd
import numpy as np
import soundfile as sf
import torch
from scipy.signal import resample_poly
from transformers import AutoProcessor, AutoModelForCTC
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("audio", help="Path to a WAV or FLAC recording")
parser.add_argument("--revision", default="main", help="Model commit or branch")
args = parser.parse_args()
audio, sample_rate = sf.read(args.audio, dtype="float32", always_2d=True)
audio = audio.mean(axis=1)
if audio.size == 0:
raise ValueError("Audio recording is empty")
if not np.isfinite(audio).all():
raise ValueError("Audio contains non-finite samples")
processor = AutoProcessor.from_pretrained("caesar-abrham/Chaka-ASR", revision=args.revision)
model = AutoModelForCTC.from_pretrained("caesar-abrham/Chaka-ASR", revision=args.revision)
target_rate = processor.feature_extractor.sampling_rate
if sample_rate != target_rate:
divisor = gcd(sample_rate, target_rate)
audio = resample_poly(audio, target_rate // divisor, sample_rate // divisor).astype("float32")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device).eval()
inputs = processor(audio, sampling_rate=target_rate, return_tensors="pt")
inputs = {key: value.to(device) for key, value in inputs.items()}
with torch.inference_mode():
predicted = model(**inputs).logits.argmax(dim=-1)
print(processor.batch_decode(predicted.cpu())[0])
if __name__ == "__main__":
main()