sinhala-flow / core.py
sadew12's picture
Upload 4 files
b1b29c3 verified
Raw
History Blame Contribute Delete
9.16 kB
from __future__ import annotations
import base64
import binascii
import io
import json
import os
import re
import threading
import time
import unicodedata
from typing import Any
import numpy as np
import soundfile as sf
from scipy.signal import resample_poly
APP_VERSION = "0.2.0"
MAX_DURATION_SECONDS = 30.0
MAX_AUDIO_BYTES = 10 * 1024 * 1024
TARGET_SAMPLE_RATE = 16_000
DEFAULT_MODEL_ID = "seniruk/whisper-small-si"
DEFAULT_MODEL_REVISION = "9de25732250c1a3f4edf1d1eec4a17f3014e2655"
MODEL_ID = os.getenv("MODEL_ID", DEFAULT_MODEL_ID)
MODEL_REVISION = os.getenv("MODEL_REVISION", DEFAULT_MODEL_REVISION)
class ApiError(Exception):
def __init__(self, code: str, message: str) -> None:
super().__init__(message)
self.code = code
self.message = message
class TranscriptionEngine:
"""Lazy CPU pipeline shared by all calls in the free Space."""
def __init__(self) -> None:
self._pipeline: Any | None = None
self._load_lock = threading.Lock()
self._inference_lock = threading.Lock()
@property
def ready(self) -> bool:
return self._pipeline is not None
def _ensure_loaded(self) -> None:
if self._pipeline is not None:
return
with self._load_lock:
if self._pipeline is not None:
return
import torch
from transformers import AutoModelForSpeechSeq2Seq, AutoProcessor, pipeline
model = AutoModelForSpeechSeq2Seq.from_pretrained(
MODEL_ID,
revision=MODEL_REVISION,
dtype=torch.float32,
low_cpu_mem_usage=True,
use_safetensors=True,
)
processor = AutoProcessor.from_pretrained(MODEL_ID, revision=MODEL_REVISION)
self._pipeline = pipeline(
"automatic-speech-recognition",
model=model,
tokenizer=processor.tokenizer,
feature_extractor=processor.feature_extractor,
device=-1,
)
def transcribe(self, samples: np.ndarray) -> str:
self._ensure_loaded()
assert self._pipeline is not None
with self._inference_lock:
result = self._pipeline(
{"array": samples, "sampling_rate": TARGET_SAMPLE_RATE},
generate_kwargs={"language": "si", "task": "transcribe"},
)
return str(result.get("text", ""))
engine = TranscriptionEngine()
def health() -> dict[str, Any]:
return {
"status": "ready" if engine.ready else "warming",
"model_id": MODEL_ID,
"model_revision": MODEL_REVISION,
"max_duration_seconds": MAX_DURATION_SECONDS,
"queue_waiting": 0,
"version": APP_VERSION,
"transport": "gradio",
}
def transcribe_base64(audio_base64: str, language: str = "si", dictionary: str = "{}") -> dict[str, Any]:
started = time.perf_counter()
try:
if language.lower().strip() != "si":
raise ApiError("UNSUPPORTED_LANGUAGE", "Only Sinhala (si) is supported.")
replacements = parse_dictionary(dictionary)
samples = decode_base64_audio(audio_base64)
raw_text = engine.transcribe(samples)
transcript = normalize_transcript(raw_text, replacements)
if not transcript:
raise ApiError("NO_SPEECH", "No Sinhala speech was recognized.")
return {
"transcript": transcript,
"processing_seconds": round(time.perf_counter() - started, 3),
"model_id": MODEL_ID,
"model_revision": MODEL_REVISION,
"error_code": None,
"message": None,
}
except ApiError as exc:
return {
"transcript": "",
"processing_seconds": round(time.perf_counter() - started, 3),
"model_id": MODEL_ID,
"model_revision": MODEL_REVISION,
"error_code": exc.code,
"message": exc.message,
}
except Exception:
return {
"transcript": "",
"processing_seconds": round(time.perf_counter() - started, 3),
"model_id": MODEL_ID,
"model_revision": MODEL_REVISION,
"error_code": "MODEL_UNAVAILABLE",
"message": "The Sinhala model is warming up or temporarily unavailable.",
}
def transcribe_audio_tuple(audio: tuple[int, np.ndarray] | None) -> str:
if audio is None:
return "Please record or upload a Sinhala voice clip."
try:
sample_rate, samples = audio
prepared = prepare_samples(np.asarray(samples), int(sample_rate))
transcript = normalize_transcript(engine.transcribe(prepared))
return transcript or "No Sinhala speech was recognized."
except ApiError as exc:
return exc.message
except Exception:
return "The model is warming up or temporarily unavailable. Please try again."
def decode_base64_audio(value: str) -> np.ndarray:
if not value or len(value) > (MAX_AUDIO_BYTES * 4 // 3) + 1024:
raise ApiError("INVALID_AUDIO", "The recording is empty or too large.")
if value.startswith("data:"):
_, _, value = value.partition(",")
try:
payload = base64.b64decode(value, validate=True)
except (binascii.Error, ValueError) as exc:
raise ApiError("INVALID_AUDIO", "The recording is not valid base64 audio.") from exc
if not payload or len(payload) > MAX_AUDIO_BYTES:
raise ApiError("AUDIO_TOO_LARGE", "The recording is empty or larger than 10 MB.")
try:
samples, sample_rate = sf.read(io.BytesIO(payload), dtype="float32", always_2d=True)
except Exception as exc:
raise ApiError("INVALID_AUDIO", "The recording must be valid WAV or FLAC audio.") from exc
return prepare_samples(samples, sample_rate)
def prepare_samples(samples: np.ndarray, sample_rate: int) -> np.ndarray:
if sample_rate <= 0 or samples.size == 0:
raise ApiError("EMPTY_AUDIO", "The recording contains no samples.")
if samples.ndim == 1:
mono = samples.astype(np.float32)
elif samples.ndim == 2:
mono = samples.astype(np.float32).mean(axis=1)
else:
raise ApiError("INVALID_AUDIO", "The recording has an unsupported channel layout.")
duration = mono.shape[0] / sample_rate
if duration < 0.15:
raise ApiError("AUDIO_TOO_SHORT", "Speak for at least a moment before stopping.")
if duration > MAX_DURATION_SECONDS + 0.05:
raise ApiError("AUDIO_TOO_LONG", "Recordings are limited to 30 seconds.")
if float(np.max(np.abs(mono))) < 0.0005:
raise ApiError("NO_SPEECH", "The recording is silent or the microphone level is too low.")
if sample_rate != TARGET_SAMPLE_RATE:
divisor = int(np.gcd(sample_rate, TARGET_SAMPLE_RATE))
mono = resample_poly(mono, TARGET_SAMPLE_RATE // divisor, sample_rate // divisor)
return np.asarray(mono, dtype=np.float32)
def parse_dictionary(value: str) -> dict[str, str]:
try:
parsed = json.loads(value or "{}")
except json.JSONDecodeError as exc:
raise ApiError("INVALID_DICTIONARY", "Dictionary must be a JSON object.") from exc
if not isinstance(parsed, dict) or len(parsed) > 100:
raise ApiError("INVALID_DICTIONARY", "Dictionary must contain at most 100 replacements.")
result: dict[str, str] = {}
for source, target in parsed.items():
if not isinstance(source, str) or not isinstance(target, str):
raise ApiError("INVALID_DICTIONARY", "Dictionary keys and values must be text.")
source, target = source.strip(), target.strip()
if not source or len(source) > 80 or len(target) > 80:
raise ApiError("INVALID_DICTIONARY", "Dictionary entries must be 1 to 80 characters.")
result[source] = target
return result
def normalize_transcript(text: str, replacements: dict[str, str] | None = None) -> str:
text = unicodedata.normalize("NFC", text or "")
text = re.sub(r"[\u200b\ufeff]", "", text)
text = re.sub(r"\s+", " ", text).strip()
text = _remove_repeated_phrases(text)
text = re.sub(r"\s+([,.!?;:।])", r"\1", text)
text = re.sub(r"([,.!?;:।])(?=[^\s,.!?;:।])", r"\1 ", text)
for source, target in sorted((replacements or {}).items(), key=lambda item: len(item[0]), reverse=True):
text = text.replace(source, target)
return unicodedata.normalize("NFC", text).strip()
def _remove_repeated_phrases(text: str) -> str:
tokens = text.split()
if len(tokens) < 4:
return text
output: list[str] = []
index = 0
while index < len(tokens):
removed = False
for size in range(min(8, (len(tokens) - index) // 2), 1, -1):
first = tokens[index : index + size]
second = tokens[index + size : index + size * 2]
if first == second:
output.extend(first)
index += size * 2
removed = True
break
if not removed:
output.append(tokens[index])
index += 1
return " ".join(output)