turn-detection / backend /streaming_semantic.py
amanetize's picture
Upload folder using huggingface_hub
dc98a34 verified
Raw
History Blame Contribute Delete
3.09 kB
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')