ahmed
Deploy ecapa quran md3 checkpoint
575b70b
Raw
History Blame Contribute Delete
3.8 kB
from __future__ import annotations
import os
from pathlib import Path
import gradio as gr
try:
import spaces
except ImportError:
# The spaces package is injected by Hugging Face ZeroGPU. Keep the same app
# runnable locally, where requesting a managed GPU is unnecessary.
class _LocalSpaces:
@staticmethod
def GPU(*decorator_args, **decorator_kwargs):
def decorator(function):
return function
return decorator
spaces = _LocalSpaces()
from quran_reciter_id.predict import predict
def _default_checkpoint_dir() -> str:
configured = os.getenv("CHECKPOINT_DIR")
if configured:
return configured
if Path("runs/ecapa_quran_md3/model.pt").is_file():
return "runs/ecapa_quran_md3"
if Path("checkpoint/model.pt").is_file():
return "checkpoint"
local_runs = sorted(Path("runs").glob("*/model.pt"))
return str(local_runs[-1].parent) if local_runs else "checkpoint"
@spaces.GPU(duration=120)
def identify_reciter(
audio_file: str | None,
video_file: str | None,
run_dir_text: str,
device: str,
top_k: int,
) -> tuple[str, list[list[object]]]:
media_file = audio_file or video_file
if not media_file:
raise gr.Error("Record audio, upload audio, or upload an MP4 video first.")
run_dir = Path(run_dir_text).expanduser()
checkpoint_path = run_dir / "model.pt"
if not checkpoint_path.is_file():
raise gr.Error(f"Checkpoint not found: {checkpoint_path}")
try:
result = predict(
run_dir=run_dir,
audio_file=Path(media_file),
top_k=int(top_k),
device=device.strip() or "auto",
)
except Exception as exc:
raise gr.Error(f"Prediction failed: {exc}") from exc
if result["is_unknown"]:
verdict = "### Unknown speaker\nNo known reciter passed the acceptance thresholds."
else:
verdict = f"### Prediction: `{result['prediction']}`"
candidates = [
[
candidate["reciter_id"],
round(candidate["probability"] * 100, 2),
round(candidate["centroid_similarity"], 4),
]
for candidate in result["top_candidates"]
]
return verdict, candidates
with gr.Blocks(title="Quran Reciter ID") as demo:
gr.Markdown(
"# Quran Reciter Identification\n"
"Record a recitation or upload audio to test a trained checkpoint."
)
with gr.Row():
with gr.Column(scale=2):
audio = gr.Audio(
label="Audio",
sources=["microphone", "upload"],
type="filepath",
)
video = gr.Video(
label="MP4 video (the audio track will be analyzed)",
sources=["upload"],
format="mp4",
)
with gr.Column(scale=1):
run_dir = gr.Textbox(
label="Checkpoint directory",
value=_default_checkpoint_dir(),
)
device = gr.Textbox(label="Device", value="auto")
top_k = gr.Slider(1, 10, value=5, step=1, label="Top candidates")
identify = gr.Button("Identify reciter", variant="primary")
verdict = gr.Markdown()
candidates = gr.Dataframe(
headers=["Reciter", "Probability (%)", "Centroid similarity"],
datatype=["str", "number", "number"],
interactive=False,
label="Top candidates",
)
identify.click(
fn=identify_reciter,
inputs=[audio, video, run_dir, device, top_k],
outputs=[verdict, candidates],
api_name="identify_reciter",
)
demo.queue(default_concurrency_limit=1)
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0", share=True)