WeLe Agentic AI: LangGraph multi-agent CRM assistant with voice
This commit is contained in:
@@ -0,0 +1,200 @@
|
||||
"""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)
|
||||
Reference in New Issue
Block a user