host_model_gra / app.py
masumtechnonext's picture
Upload app.py with huggingface_hub
b50ecc9 verified
Raw
History Blame Contribute Delete
2.69 kB
import json
import os
import gradio as gr
import spaces
import torch
import torchaudio
from huggingface_hub import hf_hub_download
from transformers import Wav2Vec2FeatureExtractor, Wav2Vec2ForSequenceClassification
MODEL_ID = os.environ.get("HF_ASR_MODEL", "masumtechnonext/wav2vec2-arabic-letter-verifier")
HF_TOKEN = os.environ.get("HF_TOKEN")
SAMPLE_RATE = 16000
feature_extractor = Wav2Vec2FeatureExtractor.from_pretrained(MODEL_ID, token=HF_TOKEN)
model = Wav2Vec2ForSequenceClassification.from_pretrained(MODEL_ID, token=HF_TOKEN)
model.eval()
id2label = {int(k): v for k, v in model.config.id2label.items()}
calibration_path = hf_hub_download(MODEL_ID, "calibration.json", token=HF_TOKEN)
with open(calibration_path) as f:
THRESHOLD = json.load(f)["confidence_threshold"]
UNKNOWN_ID = next(i for i, label in id2label.items() if label == "Unknown")
LETTERS = sorted(label for label in id2label.values() if label != "Unknown")
@spaces.GPU
def predict(audio, target_letter):
if audio is None:
return "Record or upload audio first.", ""
device = "cuda" if torch.cuda.is_available() else "cpu"
model.to(device)
sample_rate, waveform = audio
waveform = torch.tensor(waveform, dtype=torch.float32)
if waveform.ndim > 1:
waveform = waveform.mean(dim=-1)
if sample_rate != SAMPLE_RATE:
waveform = torchaudio.functional.resample(waveform, sample_rate, SAMPLE_RATE)
inputs = feature_extractor(waveform.numpy(), sampling_rate=SAMPLE_RATE, return_tensors="pt")
inputs = {k: v.to(device) for k, v in inputs.items()}
with torch.no_grad():
logits = model(**inputs).logits[0]
probs = torch.softmax(logits, dim=-1)
pred_id = int(torch.argmax(probs))
confidence = float(probs[pred_id])
accepted = confidence >= THRESHOLD and pred_id != UNKNOWN_ID
predicted_label = id2label[pred_id] if accepted else "Unrecognized"
prediction = f"{predicted_label} ({confidence:.1%} confidence)"
if not target_letter:
return prediction, "Pick an expected letter to verify."
is_correct = accepted and predicted_label == target_letter
verdict = "✅ Correct" if is_correct else "❌ Incorrect"
return prediction, verdict
demo = gr.Interface(
fn=predict,
inputs=[
gr.Audio(sources=["microphone", "upload"], type="numpy", label="Speak the letter"),
gr.Dropdown(choices=LETTERS, label="Expected letter", value=LETTERS[0]),
],
outputs=[
gr.Textbox(label="Prediction"),
gr.Textbox(label="Verification"),
],
title="Arabic Letter Verifier",
description=f"Model: {MODEL_ID}",
)
if __name__ == "__main__":
demo.launch()