Pashto-stt / app.py
codewithjarair's picture
Update app.py
8704a27 verified
Raw
History Blame Contribute Delete
3.98 kB
import os
import time
import spaces # MUST be before torch (ZeroGPU)
import gradio as gr
import torch
import numpy as np
from transformers import WhisperForConditionalGeneration, WhisperProcessor
# ================= CONFIG =================
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
DTYPE = torch.float16 if DEVICE == "cuda" else torch.float32
SR = 16000
# ================= MODELS =================
MODELS = {
"Sabtain (Stable)": "Sabtain-Dev/STT-Whisper-Pashto",
"ihanif (Experimental)": "ihanif/exp_009_small_cv24_vanilla"
}
DEFAULT_MODEL = "Sabtain (Stable)"
# model cache
_cache = {}
# ================= LOAD MODEL =================
def load_model(model_id):
if model_id not in _cache:
print(f"[Loading] {model_id}")
model = WhisperForConditionalGeneration.from_pretrained(
model_id,
torch_dtype=DTYPE,
low_cpu_mem_usage=True
).to(DEVICE)
model.eval()
processor = WhisperProcessor.from_pretrained(model_id)
_cache[model_id] = (model, processor)
return _cache[model_id]
# ================= AUDIO =================
def load_audio(path):
import librosa
audio, _ = librosa.load(path, sr=SR, mono=True)
return audio
# ================= DECODER FIX =================
def get_decoder_ids(model, processor):
# use model config if available (important for ihanif model)
if hasattr(model.config, "forced_decoder_ids") and model.config.forced_decoder_ids:
return model.config.forced_decoder_ids
# fallback for Sabtain model
return processor.get_decoder_prompt_ids(
language="pashto",
task="transcribe"
)
# ================= INFERENCE =================
@spaces.GPU
def transcribe(audio_path, model_name):
if not audio_path:
return "", "⚠️ Upload audio first"
model_id = MODELS[model_name]
model, processor = load_model(model_id)
audio = load_audio(audio_path)
inputs = processor(audio, sampling_rate=SR, return_tensors="pt")
input_features = inputs.input_features.to(DEVICE, dtype=DTYPE)
forced_ids = get_decoder_ids(model, processor)
start = time.time()
with torch.no_grad():
pred_ids = model.generate(
input_features,
forced_decoder_ids=forced_ids,
max_length=256
)
text = processor.batch_decode(
pred_ids,
skip_special_tokens=True
)[0]
elapsed = round(time.time() - start, 2)
return text.strip(), f"⏱️ {elapsed}s | {model_name}"
# ================= DOWNLOAD =================
def download_txt(text):
if not text:
return None
path = "/tmp/transcription.txt"
with open(path, "w", encoding="utf-8") as f:
f.write(text)
return path
# ================= UI =================
def build_ui():
with gr.Blocks(theme=gr.themes.Soft()) as demo:
gr.Markdown("## 🎙️ Pashto Speech-to-Text")
model_select = gr.Dropdown(
choices=list(MODELS.keys()),
value=DEFAULT_MODEL,
label="Model"
)
audio = gr.Audio(
sources=["upload", "microphone"],
type="filepath",
label="Audio"
)
btn = gr.Button("Transcribe", variant="primary")
output = gr.Textbox(
label="Transcription",
lines=5
)
status = gr.Markdown()
download = gr.DownloadButton("Download .txt")
btn.click(
transcribe,
inputs=[audio, model_select],
outputs=[output, status]
)
download.click(
download_txt,
inputs=output,
outputs=download
)
return demo
# ================= RUN =================
if __name__ == "__main__":
demo = build_ui()
demo.queue().launch(
server_name="0.0.0.0",
server_port=int(os.environ.get("PORT", 7860)),
show_error=True
)