"""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)