turn-detection / backend /semantic.py
amanetize's picture
Upload folder using huggingface_hub
dc98a34 verified
Raw
History Blame Contribute Delete
5.35 kB
from __future__ import annotations
import re
import time
import unicodedata
import numpy as np
import torch
import config
from backend.types import StageResult
_device = 'mps' if torch.backends.mps.is_available() else 'cpu'
_asr_pipeline = None
_qwen_model = None
_qwen_tokenizer = None
_livekit_model = None
_livekit_tokenizer = None
_livekit_im_end_id = None
PROMPT_TEMPLATE = 'You are analyzing a snippet of transcribed speech from a Hindi-English (Hinglish) conversation with a voice assistant. Code-switching between Hindi and English mid-sentence is normal and not a sign of incompleteness. Filler words like "matlab", "toh", "haan", "wo kya bolte hain", "um", "uh" indicate the speaker is still thinking and has NOT completed their turn.\n\nTranscript: "{transcript}"\n\nClassify whether the speaker\'s turn is:\n- complete: the utterance is a complete thought, the speaker is done\n- incomplete: the utterance is grammatically or semantically incomplete, more is coming\n- wait: the speaker is explicitly asking for a pause (e.g. "ek second", "hold on", "wait")\n\nRespond with exactly one word: complete, incomplete, or wait.'
_VALID_LABELS = ('complete', 'incomplete', 'wait')
def _get_asr_pipeline():
global _asr_pipeline
if _asr_pipeline is None:
from transformers import pipeline
_asr_pipeline = pipeline('automatic-speech-recognition', model=config.WHISPER_TINY_ID, device=_device if _device != 'mps' else -1)
return _asr_pipeline
def _get_qwen():
global _qwen_model, _qwen_tokenizer
if _qwen_model is None:
from transformers import AutoModelForCausalLM, AutoTokenizer
_qwen_tokenizer = AutoTokenizer.from_pretrained(config.QWEN_LOCAL_ID)
_qwen_model = AutoModelForCausalLM.from_pretrained(config.QWEN_LOCAL_ID).to(_device).eval()
return (_qwen_model, _qwen_tokenizer)
def _get_livekit():
global _livekit_model, _livekit_tokenizer, _livekit_im_end_id
if _livekit_model is None:
from transformers import AutoModelForCausalLM, AutoTokenizer
_livekit_tokenizer = AutoTokenizer.from_pretrained(config.LIVEKIT_TURN_DETECTOR_ID)
_livekit_model = AutoModelForCausalLM.from_pretrained(config.LIVEKIT_TURN_DETECTOR_ID).to(_device).eval()
_livekit_im_end_id = _livekit_tokenizer.convert_tokens_to_ids('<|im_end|>')
return (_livekit_model, _livekit_tokenizer, _livekit_im_end_id)
_PUNCTUATION_RE = re.compile("[^\\w\\s'-]", re.UNICODE)
_WHITESPACE_RE = re.compile('\\s+')
def _normalize_for_livekit(transcript: str) -> str:
text = unicodedata.normalize('NFKC', transcript).lower()
text = _PUNCTUATION_RE.sub(' ', text)
return _WHITESPACE_RE.sub(' ', text).strip()
def transcribe(audio: np.ndarray, sample_rate: int=config.SAMPLE_RATE) -> str:
asr = _get_asr_pipeline()
result = asr({'raw': np.asarray(audio, dtype=np.float32), 'sampling_rate': sample_rate})
return result['text'].strip()
def _parse_label(raw_text: str) -> str:
lowered = raw_text.lower()
for label in _VALID_LABELS:
if re.search(f'\\b{label}\\b', lowered):
return label
return 'incomplete'
@torch.inference_mode()
def classify_transcript(transcript: str, temperature: float=0.2) -> dict:
model, tokenizer = _get_qwen()
messages = [{'role': 'user', 'content': PROMPT_TEMPLATE.format(transcript=transcript)}]
prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
inputs = tokenizer(prompt, return_tensors='pt').to(_device)
do_sample = temperature > 0
output_ids = model.generate(**inputs, max_new_tokens=8, do_sample=do_sample, temperature=temperature if do_sample else None, pad_token_id=tokenizer.eos_token_id)
generated = output_ids[0][inputs['input_ids'].shape[1]:]
raw_text = tokenizer.decode(generated, skip_special_tokens=True)
label = _parse_label(raw_text)
return {'verdict': label, 'raw_response': raw_text.strip()}
def run(audio: np.ndarray, sample_rate: int=config.SAMPLE_RATE, temperature: float=0.2) -> StageResult:
start = time.perf_counter()
transcript = transcribe(audio, sample_rate)
result = classify_transcript(transcript, temperature)
timing_ms = (time.perf_counter() - start) * 1000
return StageResult(stage='semantic.qwen_local', timing_ms=timing_ms, output={'transcript': transcript, 'verdict': result['verdict'], 'raw_response': result['raw_response']}, available=True, provenance='real_checkpoint')
@torch.inference_mode()
def classify_transcript_livekit(transcript: str) -> float:
model, tokenizer, im_end_id = _get_livekit()
normalized = _normalize_for_livekit(transcript)
prompt = f'<|im_start|><|user|>{normalized}'
inputs = tokenizer(prompt, return_tensors='pt').to(_device)
logits = model(**inputs).logits[0, -1, :]
probs = torch.softmax(logits, dim=-1)
return probs[im_end_id].item()
def run_livekit(audio: np.ndarray, sample_rate: int=config.SAMPLE_RATE) -> StageResult:
start = time.perf_counter()
transcript = transcribe(audio, sample_rate)
probability = classify_transcript_livekit(transcript)
timing_ms = (time.perf_counter() - start) * 1000
return StageResult(stage='semantic.livekit_eou', timing_ms=timing_ms, output={'transcript': transcript, 'probability': probability}, available=True, provenance='real_checkpoint')