whisper-quran-v1 / handler.py
Cursor Agent
harden endpoint handler decode and queue behavior
dcbfa13 unverified
Raw
History Blame Contribute Delete
10.3 kB
import base64
import binascii
import os
import queue
import threading
import time
import warnings
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
import torch
from transformers import pipeline
from transformers.pipelines.audio_utils import ffmpeg_read
SAMPLE_RATE = 16000
MAX_QUEUE_DEPTH = 32
MIN_AUDIO_SECONDS = 0.2
MAX_AUDIO_SECONDS = 35.0
@dataclass
class _QueuedRequest:
audio: np.ndarray
params: Dict[str, Any]
event: threading.Event = field(default_factory=threading.Event)
result: Optional[Dict[str, Any]] = None
error: Optional[Exception] = None
class EndpointHandler:
"""
Custom HF Inference Endpoint handler with micro-batching.
Why this exists:
- default ASR handling is effectively one-request-at-a-time on many setups
- this handler coalesces near-simultaneous requests into one GPU forward pass
- response shape is compatible with callers expecting {text, chunks}
"""
def __init__(self, path: str = ""):
model_path = path or "wasimlhr/whisper-quran-v1"
use_cuda = torch.cuda.is_available()
torch_dtype = torch.float16 if use_cuda else torch.float32
device = 0 if use_cuda else -1
self._pipe = pipeline(
task="automatic-speech-recognition",
model=model_path,
device=device,
torch_dtype=torch_dtype,
)
self._batch_window_ms = self._read_int("ASR_BATCH_WINDOW_MS", 35, 1, 200)
self._max_batch_size = self._read_int("ASR_MAX_BATCH_SIZE", 4, 1, 16)
self._request_timeout_s = float(os.getenv("ASR_REQUEST_TIMEOUT_S", "45"))
self._max_queue_depth = self._read_int("ASR_MAX_QUEUE_DEPTH", MAX_QUEUE_DEPTH, 4, 256)
self._max_audio_seconds = float(os.getenv("ASR_MAX_AUDIO_SECONDS", str(MAX_AUDIO_SECONDS)))
self._min_audio_seconds = float(os.getenv("ASR_MIN_AUDIO_SECONDS", str(MIN_AUDIO_SECONDS)))
self._queue: "queue.Queue[_QueuedRequest]" = queue.Queue()
self._worker = threading.Thread(target=self._drain_loop, daemon=True)
self._worker.start()
print(
f"[handler] initialized model={model_path} device={device} "
f"batch_window_ms={self._batch_window_ms} max_batch_size={self._max_batch_size} "
f"max_queue_depth={self._max_queue_depth}"
)
@staticmethod
def _read_int(name: str, default: int, min_v: int, max_v: int) -> int:
try:
value = int(os.getenv(name, str(default)))
except ValueError:
value = default
return max(min_v, min(max_v, value))
def __call__(self, data: Any) -> Dict[str, Any]:
payload, raw_params = self._extract_payload_and_params(data)
audio = self._decode_audio(payload)
params = self._normalize_params(raw_params)
if self._queue.qsize() >= self._max_queue_depth:
raise RuntimeError("ASR endpoint busy (queue saturated), please retry shortly")
req = _QueuedRequest(audio=audio, params=params)
self._queue.put(req)
if not req.event.wait(timeout=self._request_timeout_s):
raise TimeoutError("ASR request timed out while waiting in handler queue")
if req.error is not None:
raise RuntimeError(f"ASR request failed: {req.error}")
return req.result or {"text": "", "chunks": []}
def _drain_loop(self) -> None:
while True:
first = self._queue.get()
batch = [first]
deadline = time.perf_counter() + (self._batch_window_ms / 1000.0)
while len(batch) < self._max_batch_size:
timeout = deadline - time.perf_counter()
if timeout <= 0:
break
try:
batch.append(self._queue.get(timeout=timeout))
except queue.Empty:
break
self._process_batch(batch)
def _process_batch(self, batch: List[_QueuedRequest]) -> None:
groups: Dict[Tuple[Any, ...], List[_QueuedRequest]] = {}
for req in batch:
groups.setdefault(self._group_key(req.params), []).append(req)
for group in groups.values():
params = group[0].params
pipeline_inputs = [{"array": r.audio, "sampling_rate": SAMPLE_RATE} for r in group]
try:
# Some transformers versions still emit internal deprecation warnings
# while this pipeline API migrates. We explicitly suppress that known
# warning noise in endpoint logs and rely on supported kwargs.
with warnings.catch_warnings():
warnings.filterwarnings(
"ignore",
message="The input name `inputs` is deprecated.*",
category=FutureWarning,
)
outputs = self._pipe(
pipeline_inputs,
return_timestamps=params["return_timestamps"],
batch_size=len(group),
chunk_length_s=params["chunk_length_s"],
generate_kwargs={
"language": params["language"],
"task": params["task"],
"temperature": params["temperature"],
},
)
if isinstance(outputs, dict):
outputs = [outputs]
for req, out in zip(group, outputs):
req.result = self._format_output(out)
req.event.set()
except Exception as exc: # noqa: BLE001
for req in group:
req.error = exc
req.event.set()
@staticmethod
def _group_key(params: Dict[str, Any]) -> Tuple[Any, ...]:
return (
params["language"],
params["task"],
params["return_timestamps"],
params["chunk_length_s"],
params["temperature"],
)
@staticmethod
def _format_output(output: Dict[str, Any]) -> Dict[str, Any]:
text = str(output.get("text", "")).strip()
chunks_out: List[Dict[str, Any]] = []
for ch in output.get("chunks", []) or []:
if not isinstance(ch, dict):
continue
ctext = str(ch.get("text", "")).strip()
ts = ch.get("timestamp")
if isinstance(ts, (list, tuple)) and len(ts) >= 2:
start, end = ts[0], ts[1]
else:
start, end = None, None
chunks_out.append({"text": ctext, "start": start, "end": end, "timestamp": [start, end]})
return {"text": text, "chunks": chunks_out}
@staticmethod
def _normalize_params(params: Dict[str, Any]) -> Dict[str, Any]:
params = params or {}
return {
"language": str(params.get("language", "ar")),
"task": str(params.get("task", "transcribe")),
"return_timestamps": bool(params.get("return_timestamps", True)),
"chunk_length_s": float(params.get("chunk_length_s", 20)),
"temperature": float(params.get("temperature", 0.0)),
}
@staticmethod
def _extract_payload_and_params(data: Any) -> Tuple[Any, Dict[str, Any]]:
if isinstance(data, (bytes, bytearray)):
return bytes(data), {}
if isinstance(data, dict):
params = data.get("parameters", {}) or {}
if "inputs" in data:
return data["inputs"], params
if "audio" in data:
return data["audio"], params
raise ValueError("Expected 'inputs' or 'audio' in request body")
raise TypeError(f"Unsupported request type: {type(data)}")
def _decode_audio(self, payload: Any) -> np.ndarray:
if isinstance(payload, dict) and "array" in payload:
arr = np.asarray(payload["array"], dtype=np.float32)
if arr.ndim > 1:
arr = np.squeeze(arr)
if arr.ndim != 1 or arr.size == 0:
raise ValueError("Audio array payload must be a non-empty 1D array")
duration_s = arr.size / float(SAMPLE_RATE)
if duration_s < self._min_audio_seconds:
raise ValueError(
f"Audio too short ({duration_s:.3f}s), need >= {self._min_audio_seconds:.1f}s"
)
if duration_s > self._max_audio_seconds:
max_samples = int(self._max_audio_seconds * SAMPLE_RATE)
arr = arr[:max_samples]
return arr
if isinstance(payload, str):
if payload.startswith("data:") and "," in payload:
payload = payload.split(",", 1)[1]
try:
audio_bytes = base64.b64decode(payload, validate=True)
except (binascii.Error, ValueError) as exc:
raise ValueError("String payload must be base64-encoded audio bytes") from exc
elif isinstance(payload, (bytes, bytearray)):
audio_bytes = bytes(payload)
else:
raise TypeError(f"Unsupported audio payload type: {type(payload)}")
if not audio_bytes:
raise ValueError("Audio payload is empty")
try:
audio = ffmpeg_read(audio_bytes, SAMPLE_RATE)
except Exception as exc: # noqa: BLE001
raise ValueError(
"Audio decode failed: malformed or unsupported audio payload"
) from exc
arr = np.asarray(audio, dtype=np.float32)
if arr.ndim > 1:
arr = np.squeeze(arr)
if arr.ndim != 1 or arr.size == 0:
raise ValueError("Decoded audio is empty or invalid")
duration_s = arr.size / float(SAMPLE_RATE)
if duration_s < self._min_audio_seconds:
raise ValueError(
f"Audio too short ({duration_s:.3f}s), need >= {self._min_audio_seconds:.1f}s"
)
if duration_s > self._max_audio_seconds:
max_samples = int(self._max_audio_seconds * SAMPLE_RATE)
arr = arr[:max_samples]
return arr