from __future__ import annotations import math import time import numpy as np import torch import config from backend.semantic import classify_transcript from backend.types import StageResult _device = 'mps' if torch.backends.mps.is_available() else 'cpu' _ctc_model = None _ctc_processor = None DEFAULT_CHUNK_MS = 1000 MAX_STEPS = 10 def _get_ctc(): global _ctc_model, _ctc_processor if _ctc_model is None: from transformers import Wav2Vec2ForCTC, Wav2Vec2Processor _ctc_processor = Wav2Vec2Processor.from_pretrained(config.HINDI_CTC_ID) _ctc_model = Wav2Vec2ForCTC.from_pretrained(config.HINDI_CTC_ID).to(_device).eval() return (_ctc_processor, _ctc_model) @torch.inference_mode() def transcribe_ctc_chunk(audio_prefix: np.ndarray, sample_rate: int=config.SAMPLE_RATE) -> str: processor, model = _get_ctc() inputs = processor(audio_prefix, sampling_rate=sample_rate, return_tensors='pt', padding=True) input_values = inputs.input_values.to(_device) logits = model(input_values).logits predicted_ids = torch.argmax(logits, dim=-1) text = processor.batch_decode(predicted_ids)[0] return text.strip() def run(audio: np.ndarray, sample_rate: int=config.SAMPLE_RATE, chunk_ms: int=DEFAULT_CHUNK_MS, temperature: float=0.2) -> StageResult: start = time.perf_counter() audio = np.asarray(audio, dtype=np.float32) total_audio_ms = len(audio) / sample_rate * 1000 chunk_samples = int(sample_rate * chunk_ms / 1000) history: list[dict] = [] first_decisive_audio_ms = None num_steps = min(MAX_STEPS, max(1, len(audio) // chunk_samples)) stride = max(chunk_samples, math.ceil(len(audio) / num_steps)) for step in range(1, num_steps + 1): prefix_len = min(len(audio), step * stride) prefix = audio[:prefix_len] elapsed_audio_ms = prefix_len / sample_rate * 1000 step_start = time.perf_counter() transcript_so_far = transcribe_ctc_chunk(prefix, sample_rate) verdict = 'incomplete' if transcript_so_far: verdict = classify_transcript(transcript_so_far, temperature)['verdict'] step_ms = (time.perf_counter() - step_start) * 1000 history.append({'elapsed_audio_ms': round(elapsed_audio_ms, 1), 'transcript_so_far': transcript_so_far, 'verdict': verdict, 'step_processing_ms': round(step_ms, 1)}) if first_decisive_audio_ms is None and verdict != 'incomplete': first_decisive_audio_ms = elapsed_audio_ms if prefix_len >= len(audio): break final = history[-1] timing_ms = (time.perf_counter() - start) * 1000 return StageResult(stage='semantic.qwen_local_streaming', timing_ms=timing_ms, output={'transcript': final['transcript_so_far'], 'verdict': final['verdict'], 'history': history, 'total_audio_ms': round(total_audio_ms, 1), 'first_decisive_audio_ms': first_decisive_audio_ms, 'eou_delay_saved_ms': round(total_audio_ms - first_decisive_audio_ms, 1) if first_decisive_audio_ms is not None else None}, available=True, provenance='architecture_reimplemented')