HealthAgent / speech.py
siyah1's picture
Rename Speech.py to speech.py
9d730c9 verified
Raw
History Blame Contribute Delete
6.34 kB
"""
speech.py — pluggable STT / TTS backends
------------------------------------------
FastRTC ships English-only built-ins out of the box:
- get_stt_model("moonshine/base") -> STT
- get_tts_model("kokoro") -> TTS
Neither supports Malayalam, so this module adds two small classes that
satisfy FastRTC's STTModel / TTSModel protocols (a `.stt(audio)` method
and `.tts(text, options)` / `.stream_tts_sync(text, options)` methods)
and can be dropped into `ReplyOnPause` exactly like the built-ins.
MalayalamSTT -> a Whisper checkpoint fine-tuned on Malayalam speech
MalayalamTTS -> Meta's MMS VITS Malayalam checkpoint
Swap either model id via env vars without touching any other code.
"""
from __future__ import annotations
import os
import re
from dataclasses import dataclass
import numpy as np
import torch
from transformers import (
AutoModelForSpeechSeq2Seq,
AutoProcessor,
VitsModel,
AutoTokenizer,
)
from fastrtc import KokoroTTSOptions, get_stt_model, get_tts_model
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
DTYPE = torch.float16 if DEVICE == "cuda" else torch.float32
# Fine-tuned Whisper checkpoint for Malayalam ASR. Swap for whichever
# checkpoint gives you the best WER on your audience's accent/domain --
# a few public options: vrclc/Whisper-medium-Malayalam,
# kavyamanohar/whisper-small-malayalam, Jithjacob123/whisper-small-Malayalam.
MALAYALAM_STT_MODEL_ID = os.environ.get(
"MEDGUIDE_ML_STT_MODEL_ID", "vrclc/Whisper-medium-Malayalam"
)
# Meta MMS single-speaker VITS checkpoint for Malayalam TTS. For higher
# quality/expressiveness (at the cost of a heavier model + prompt
# engineering) swap to ai4bharat/indic-parler-tts, which also covers
# Malayalam plus 20 other Indic languages.
MALAYALAM_TTS_MODEL_ID = os.environ.get(
"MEDGUIDE_ML_TTS_MODEL_ID", "facebook/mms-tts-mal"
)
TARGET_SR = 16000
def _to_mono_float32(audio: tuple[int, np.ndarray]) -> tuple[int, np.ndarray]:
"""Normalize a FastRTC (sample_rate, array) pair to mono float32 in [-1, 1]."""
sr, arr = audio
arr = np.asarray(arr)
if arr.ndim == 2:
arr = arr.mean(axis=0) if arr.shape[0] < arr.shape[1] else arr.mean(axis=1)
if np.issubdtype(arr.dtype, np.integer):
arr = arr.astype(np.float32) / 32768.0
else:
arr = arr.astype(np.float32)
return sr, arr
def _resample(arr: np.ndarray, sr: int, target_sr: int = TARGET_SR) -> np.ndarray:
if sr == target_sr:
return arr
from scipy import signal
n_target = int(round(len(arr) * target_sr / sr))
return signal.resample(arr, max(n_target, 1)).astype(np.float32)
class MalayalamSTT:
"""Implements FastRTC's STTModel protocol: `.stt(audio) -> str`."""
def __init__(self, model_id: str = MALAYALAM_STT_MODEL_ID):
self.processor = AutoProcessor.from_pretrained(model_id)
self.model = AutoModelForSpeechSeq2Seq.from_pretrained(
model_id, dtype=DTYPE
).to(DEVICE)
self.model.eval()
# Force Malayalam decoding regardless of Whisper's language auto-detect.
try:
self.forced_ids = self.processor.get_decoder_prompt_ids(
language="malayalam", task="transcribe"
)
except Exception:
self.forced_ids = None
def stt(self, audio: tuple[int, np.ndarray]) -> str:
sr, arr = _to_mono_float32(audio)
arr = _resample(arr, sr, TARGET_SR)
inputs = self.processor(
arr, sampling_rate=TARGET_SR, return_tensors="pt"
).to(DEVICE)
with torch.no_grad():
ids = self.model.generate(
**inputs,
forced_decoder_ids=self.forced_ids,
max_new_tokens=200,
)
return self.processor.batch_decode(ids, skip_special_tokens=True)[0].strip()
@dataclass
class SimpleTTSOptions:
speed: float = 1.0
_SENTENCE_SPLIT_RE = re.compile(r"(?<=[.!?।])\s+")
class MalayalamTTS:
"""Implements FastRTC's TTSModel protocol: `.tts()` and
`.stream_tts_sync()`. MMS-VITS is not autoregressive/streamable at the
token level, so "streaming" here means synthesizing sentence-by-sentence
and yielding each clip as soon as it's ready -- still cuts first-audio
latency noticeably versus waiting for the whole reply."""
def __init__(self, model_id: str = MALAYALAM_TTS_MODEL_ID):
self.tokenizer = AutoTokenizer.from_pretrained(model_id)
self.model = VitsModel.from_pretrained(model_id).to(DEVICE)
self.model.eval()
self.sample_rate = self.model.config.sampling_rate
def _synth_one(self, text: str) -> tuple[int, np.ndarray]:
inputs = self.tokenizer(text, return_tensors="pt").to(DEVICE)
with torch.no_grad():
waveform = self.model(**inputs).waveform
arr = waveform.squeeze().cpu().numpy().astype(np.float32)
return self.sample_rate, arr.reshape(1, -1)
def tts(self, text: str, options: SimpleTTSOptions | None = None):
return self._synth_one(text)
def stream_tts_sync(self, text: str, options: SimpleTTSOptions | None = None):
chunks = [c.strip() for c in _SENTENCE_SPLIT_RE.split(text) if c.strip()]
if not chunks:
return
for chunk in chunks:
yield self._synth_one(chunk)
# ---------------------------------------------------------------------
# Per-language registry the UI/handler pick from
# ---------------------------------------------------------------------
_english_stt = None
_english_tts = None
_malayalam_stt = None
_malayalam_tts = None
def get_stt(language: str):
global _english_stt, _malayalam_stt
if language == "Malayalam":
if _malayalam_stt is None:
_malayalam_stt = MalayalamSTT()
return _malayalam_stt
if _english_stt is None:
_english_stt = get_stt_model(model="moonshine/base")
return _english_stt
def get_tts(language: str):
global _english_tts, _malayalam_tts
if language == "Malayalam":
if _malayalam_tts is None:
_malayalam_tts = MalayalamTTS()
return _malayalam_tts
if _english_tts is None:
_english_tts = get_tts_model(model="kokoro")
return _english_tts
ENGLISH_TTS_OPTIONS = KokoroTTSOptions(voice="af_heart", speed=1.0, lang="en-us")