201 lines
9.0 KiB
Python
201 lines
9.0 KiB
Python
"""Speech-to-text — AI4Bharat IndicConformer, with an English path and auto routing.
|
|
|
|
Why two models rather than one:
|
|
|
|
* **IndicConformer** decodes 22 Indian languages, and decodes *as* the language
|
|
you name — it does not detect. Handing it "ta" for English speech produces
|
|
Tamil-script nonsense. English is not one of its codes at all.
|
|
* **Whisper (multilingual)** covers English well and, usefully, can identify the
|
|
spoken language in a single decoder step.
|
|
|
|
So the default mode is `auto`: Whisper identifies the language from the audio,
|
|
English is transcribed by Whisper directly (the mel is already computed, so
|
|
this costs nothing extra), and anything Indic is routed to IndicConformer,
|
|
which is far stronger on those languages than Whisper is.
|
|
|
|
Code-mixed speech ("Tanglish") is the awkward case: language ID can land either
|
|
side of the fence on a short, mixed utterance. When Whisper is not confident,
|
|
the caller's preferred language wins rather than a coin toss — a Tamil-speaking
|
|
office gets Tamil, and the occasional English sentence still routes correctly
|
|
when it is clearly English.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
import time
|
|
|
|
import numpy as np
|
|
import torch
|
|
|
|
from .config import settings
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# The codes IndicConformer accepts.
|
|
INDIC_LANGS = {
|
|
"as", "bn", "brx", "doi", "gu", "hi", "kn", "kok", "ks", "mai", "ml",
|
|
"mni", "mr", "ne", "or", "pa", "sa", "sat", "sd", "ta", "te", "ur",
|
|
}
|
|
|
|
# Offered by the UI. "auto" first: most WeLe agents switch language mid-shift.
|
|
SUPPORTED = [
|
|
{"code": "auto", "label": "Auto-detect", "native": "Auto"},
|
|
{"code": "ta", "label": "Tamil", "native": "தமிழ்"},
|
|
{"code": "en", "label": "English", "native": "English"},
|
|
{"code": "hi", "label": "Hindi", "native": "हिन्दी"},
|
|
{"code": "te", "label": "Telugu", "native": "తెలుగు"},
|
|
{"code": "kn", "label": "Kannada", "native": "ಕನ್ನಡ"},
|
|
{"code": "ml", "label": "Malayalam", "native": "മലയാളം"},
|
|
{"code": "mr", "label": "Marathi", "native": "मराठी"},
|
|
{"code": "bn", "label": "Bengali", "native": "বাংলা"},
|
|
]
|
|
|
|
# Below this, trust the user's stated preference over the detector.
|
|
DETECT_CONFIDENCE = 0.60
|
|
|
|
|
|
class Transcriber:
|
|
def __init__(self) -> None:
|
|
self._indic = None
|
|
self._whisper = None
|
|
self._whisper_proc = None
|
|
|
|
# ── loading ──────────────────────────────────────────────────────────────
|
|
def load(self) -> None:
|
|
from transformers import AutoModel
|
|
|
|
t0 = time.perf_counter()
|
|
# float32: the checkpoint ships custom remote code that assumes fp32.
|
|
# ~2.4 GB at 600M, which still leaves room for Whisper and the TTS model.
|
|
self._indic = AutoModel.from_pretrained(settings.stt_model, trust_remote_code=True)
|
|
self._indic.to(settings.device).eval()
|
|
logger.info("STT loaded (%s) in %.1fs", settings.stt_model, time.perf_counter() - t0)
|
|
|
|
if settings.preload_english:
|
|
self._load_whisper()
|
|
|
|
def _load_whisper(self) -> None:
|
|
"""English + language ID. Multilingual on purpose — the .en checkpoint
|
|
cannot identify languages, which is the whole point of auto mode."""
|
|
if self._whisper is not None:
|
|
return
|
|
from transformers import WhisperForConditionalGeneration, WhisperProcessor
|
|
|
|
t0 = time.perf_counter()
|
|
self._whisper_proc = WhisperProcessor.from_pretrained(settings.english_model)
|
|
self._whisper = WhisperForConditionalGeneration.from_pretrained(
|
|
settings.english_model, torch_dtype=settings.dtype,
|
|
).to(settings.device).eval()
|
|
logger.info("English/ID model loaded (%s) in %.1fs", settings.english_model, time.perf_counter() - t0)
|
|
|
|
# ── inference ────────────────────────────────────────────────────────────
|
|
@torch.inference_mode()
|
|
def transcribe(self, audio: np.ndarray, lang: str = "auto") -> dict:
|
|
"""audio: float32 mono @16 kHz in [-1, 1].
|
|
|
|
`lang` may be an explicit code, or "auto" / "auto:ta" to detect with a
|
|
fallback preference.
|
|
"""
|
|
t0 = time.perf_counter()
|
|
if audio.size < settings.sample_rate_in // 5: # under 200 ms
|
|
return {"text": "", "lang": lang, "ms": 0, "note": "too short"}
|
|
|
|
detected = None
|
|
confidence = None
|
|
|
|
if lang.startswith("auto"):
|
|
prefer = lang.split(":", 1)[1] if ":" in lang else settings.default_lang
|
|
feats = self._features(audio)
|
|
detected, confidence = self._detect(feats)
|
|
|
|
if confidence is not None and confidence < DETECT_CONFIDENCE:
|
|
logger.info("language ID low confidence (%s @ %.2f) — using preferred %s",
|
|
detected, confidence, prefer)
|
|
use = prefer
|
|
elif detected == "en" or detected in INDIC_LANGS:
|
|
use = detected
|
|
else:
|
|
# Whisper reported something we cannot decode (e.g. "nn" on
|
|
# noise). Fall back rather than fail.
|
|
use = prefer
|
|
|
|
text = self._english(audio, feats=feats) if use == "en" else self._indic_decode(audio, use)
|
|
else:
|
|
use = lang
|
|
text = self._english(audio) if lang == "en" else self._indic_decode(audio, lang)
|
|
|
|
ms = int((time.perf_counter() - t0) * 1000)
|
|
audio_ms = int(1000 * audio.size / settings.sample_rate_in)
|
|
logger.info("STT %s%s: %dms audio → %dms → %r",
|
|
use, f" (detected {detected} {confidence:.2f})" if confidence is not None else "",
|
|
audio_ms, ms, text[:80])
|
|
|
|
return {
|
|
"text": text.strip(),
|
|
"lang": use,
|
|
"detected": detected,
|
|
"confidence": round(confidence, 3) if confidence is not None else None,
|
|
"ms": ms,
|
|
"audio_ms": audio_ms,
|
|
}
|
|
|
|
# ── internals ────────────────────────────────────────────────────────────
|
|
def _features(self, audio: np.ndarray):
|
|
self._load_whisper()
|
|
return self._whisper_proc(
|
|
audio, sampling_rate=settings.sample_rate_in, return_tensors="pt",
|
|
).input_features.to(settings.device, settings.dtype)
|
|
|
|
def _detect(self, feats) -> tuple[str | None, float | None]:
|
|
"""One decoder step: read the language-token distribution."""
|
|
try:
|
|
tok = self._whisper_proc.tokenizer
|
|
sot = tok.convert_tokens_to_ids("<|startoftranscript|>")
|
|
start = torch.tensor([[sot]], device=settings.device)
|
|
logits = self._whisper(feats, decoder_input_ids=start).logits[:, -1]
|
|
|
|
lang_ids, codes = [], []
|
|
for code in {*INDIC_LANGS, "en"}:
|
|
tid = tok.convert_tokens_to_ids(f"<|{code}|>")
|
|
# Unknown languages map to the unk id; skip those.
|
|
if tid is not None and tid != tok.unk_token_id:
|
|
lang_ids.append(tid)
|
|
codes.append(code)
|
|
if not lang_ids:
|
|
return None, None
|
|
|
|
probs = torch.softmax(logits[0, lang_ids].float(), dim=-1)
|
|
best = int(probs.argmax())
|
|
return codes[best], float(probs[best])
|
|
except Exception as e: # noqa: BLE001 — detection must never break a turn
|
|
logger.warning("language ID failed (%s) — falling back to preference", e)
|
|
return None, None
|
|
|
|
def _indic_decode(self, audio: np.ndarray, lang: str) -> str:
|
|
if lang not in INDIC_LANGS:
|
|
logger.warning("unsupported STT language %r — using %s", lang, settings.default_lang)
|
|
lang = settings.default_lang
|
|
wav = torch.from_numpy(audio).unsqueeze(0).to(settings.device) # (1, N)
|
|
out = self._indic(wav, lang, settings.stt_decoding)
|
|
if isinstance(out, (list, tuple)):
|
|
return str(out[0]) if out else ""
|
|
return str(out)
|
|
|
|
def _english(self, audio: np.ndarray, feats=None) -> str:
|
|
self._load_whisper()
|
|
if feats is None:
|
|
feats = self._features(audio)
|
|
ids = self._whisper.generate(feats, language="en", task="transcribe", max_new_tokens=180)
|
|
return self._whisper_proc.batch_decode(ids, skip_special_tokens=True)[0]
|
|
|
|
def warmup(self) -> None:
|
|
"""Silent pass so the first real utterance isn't paying for CUDA init."""
|
|
try:
|
|
silence = np.zeros(settings.sample_rate_in, dtype=np.float32)
|
|
self.transcribe(silence, settings.default_lang)
|
|
if settings.preload_english:
|
|
self.transcribe(silence, "en")
|
|
logger.info("STT warm")
|
|
except Exception as e: # noqa: BLE001
|
|
logger.warning("STT warmup skipped: %s", e)
|