#!/usr/bin/env python3 """ ARIA Voxtral-STT-3B Bridge (Transformers) — Ersatz fuer whisper. Laeuft auf Treiber 550/CUDA 12.4 via torch cu124 (kein Treiber-Upgrade noetig). Modell: Voxtral-Mini-3B-2507 (bf16, ~9 GB) → GPU 1 (12 GB, per Compose gepinnt). Arbeitsweise = Zwilling der whisper-Bridge: App schickt live PCM-Chunks; wir transkribieren alle ~STREAM_TRANSCRIBE_INTERVAL_MS auf dem Ringbuffer (Partials) und feuern stt_endpoint, sobald der ADAPTIVE Endpointer (Rausch-Boden-VAD + semantische Stagnation, aus M0.1) "fertig" sagt. RVS-Wire-Protokoll identisch zu whisper → drop-in (die App merkt nur bessere Genauigkeit). ⚠️ VERIFY-ON-FIRST-RUN: Die exakte Transformers-Transkriptions-API von Voxtral (apply_transcription_request / generate / decode) ist unten in EINER Methode (VoxtralRunner._transcribe_blocking) gekapselt und nach dem HF-Modelcard-Muster modelliert. Beim ersten echten Lauf gegen die Voxtral-Modelcard pruefen und dort anpassen. Alles andere (RVS, Endpointer) ist bewaehrt. Env: RVS_HOST, RVS_PORT, RVS_TLS, RVS_TLS_FALLBACK, RVS_TOKEN VOXTRAL_MODEL Default: mistralai/Voxtral-Mini-3B-2507 VOXTRAL_LANGUAGE Default: de VOXTRAL_DEVICE Default: cuda STREAM_TRANSCRIBE_INTERVAL_MS Default 1000 (3B ist schwerer als whisper-small) """ import asyncio import base64 import json import logging import os import time from dataclasses import dataclass, field from typing import Optional import numpy as np import websockets logging.basicConfig( level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s", datefmt="%H:%M:%S", ) logger = logging.getLogger("voxtral-bridge") RVS_HOST = os.getenv("RVS_HOST", "").strip() RVS_PORT = int(os.getenv("RVS_PORT", "443")) RVS_TLS = os.getenv("RVS_TLS", "true").lower() == "true" RVS_TLS_FALLBACK = os.getenv("RVS_TLS_FALLBACK", "true").lower() == "true" RVS_TOKEN = os.getenv("RVS_TOKEN", "").strip() VOXTRAL_MODEL = os.getenv("VOXTRAL_MODEL", "mistralai/Voxtral-Mini-3B-2507") VOXTRAL_LANGUAGE = os.getenv("VOXTRAL_LANGUAGE", "de") VOXTRAL_DEVICE = os.getenv("VOXTRAL_DEVICE", "cuda") STREAM_TRANSCRIBE_INTERVAL_MS = int(os.getenv("STREAM_TRANSCRIBE_INTERVAL_MS", "1000")) STREAM_DEFAULT_ENDPOINT_MS = 2400 STREAM_DEFAULT_HARD_CAP_MS = 300000 STREAM_MIN_AUDIO_MS = 600 STREAM_SESSION_TTL_S = 120 STREAM_ENERGY_WINDOW_MS = 300 STREAM_SEMANTIC_BACKUP_FACTOR = 2.0 # Adaptiver Voice-Schwellwert (M0.1): relativ zum gemessenen Rausch-Boden. STREAM_VOICE_FACTOR = 2.5 STREAM_VOICE_RMS_MIN = 0.005 STREAM_VOICE_RMS_MAX = 0.020 def pcm_s16le_to_float32(data: bytes) -> np.ndarray: if not data: return np.zeros(0, dtype=np.float32) return np.frombuffer(data, dtype=np.int16).astype(np.float32) / 32768.0 async def _send(ws, mtype: str, payload: dict) -> None: try: await ws.send(json.dumps({ "type": mtype, "payload": payload, "timestamp": int(time.time() * 1000), })) except Exception as e: logger.warning("RVS-Send fehlgeschlagen (%s): %s", mtype, e) class VoxtralRunner: """Haelt das Voxtral-Modell (Transformers). transcribe() blockiert → aus dem Event-Loop via run_in_executor aufrufen. Ein Lock serialisiert GPU-Zugriffe.""" def __init__(self) -> None: self.model = None self.processor = None self._lock = asyncio.Lock() def load(self) -> None: import torch from transformers import AutoProcessor, VoxtralForConditionalGeneration t0 = time.time() logger.info("Lade Voxtral '%s' (device=%s, bf16)…", VOXTRAL_MODEL, VOXTRAL_DEVICE) self.processor = AutoProcessor.from_pretrained(VOXTRAL_MODEL) self.model = VoxtralForConditionalGeneration.from_pretrained( VOXTRAL_MODEL, torch_dtype=torch.bfloat16, device_map=VOXTRAL_DEVICE, ) logger.info("Voxtral geladen in %.1fs", time.time() - t0) def _transcribe_blocking(self, audio_f32: np.ndarray, language: str) -> str: # ⚠️ VERIFY: exakte Voxtral-Transformers-API gegen die HF-Modelcard. import torch proc, model = self.processor, self.model if proc is None or model is None or audio_f32.size == 0: return "" inputs = proc.apply_transcription_request( language=language, audio=audio_f32, model_id=VOXTRAL_MODEL, sampling_rate=16000, ) inputs = inputs.to(VOXTRAL_DEVICE, dtype=torch.bfloat16) with torch.no_grad(): outputs = model.generate(**inputs, max_new_tokens=512) trimmed = outputs[:, inputs.input_ids.shape[1]:] text = proc.batch_decode(trimmed, skip_special_tokens=True) return (text[0] if text else "").strip() async def transcribe(self, audio_f32: np.ndarray, language: str) -> str: loop = asyncio.get_running_loop() async with self._lock: return await loop.run_in_executor(None, self._transcribe_blocking, audio_f32, language) @dataclass class StreamSession: request_id: str audio_request_id: str language: str endpoint_ms: int hard_cap_ms: int voice: str = "" speed: float = 1.0 interrupted: bool = False location: Optional[dict] = None sample_rate: int = 16000 voice_factor: float = STREAM_VOICE_FACTOR voice_rms_min: float = STREAM_VOICE_RMS_MIN voice_rms_max: float = STREAM_VOICE_RMS_MAX pcm_buffer: bytearray = field(default_factory=bytearray) started_at: float = field(default_factory=time.time) last_chunk_at: float = field(default_factory=time.time) last_partial: str = "" last_growth_at: float = 0.0 last_transcribe_at: float = 0.0 last_voice_at: float = 0.0 noise_floor: float = 0.0 closed: bool = False endpoint_sent: bool = False class SessionManager: def __init__(self, runner: VoxtralRunner) -> None: self.runner = runner self._sessions: dict[str, StreamSession] = {} self._ws = None def attach_ws(self, ws) -> None: self._ws = ws def start_session(self, payload: dict) -> None: rid = (payload.get("requestId") or "").strip() if not rid: return try: endpoint_ms = int(payload.get("endpointMs") or STREAM_DEFAULT_ENDPOINT_MS) except (TypeError, ValueError): endpoint_ms = STREAM_DEFAULT_ENDPOINT_MS try: hard_cap_ms = int(payload.get("hardCapMs") or STREAM_DEFAULT_HARD_CAP_MS) except (TypeError, ValueError): hard_cap_ms = STREAM_DEFAULT_HARD_CAP_MS try: voice_factor = float(payload.get("voiceFactor") or STREAM_VOICE_FACTOR) except (TypeError, ValueError): voice_factor = STREAM_VOICE_FACTOR self._sessions[rid] = StreamSession( request_id=rid, audio_request_id=payload.get("audioRequestId", "") or "", language=payload.get("language") or VOXTRAL_LANGUAGE, endpoint_ms=endpoint_ms, hard_cap_ms=hard_cap_ms, voice=payload.get("voice", "") or "", speed=float(payload.get("speed") or 1.0), voice_factor=voice_factor, interrupted=bool(payload.get("interrupted", False)), location=payload.get("location") or None, sample_rate=int(payload.get("sampleRate") or 16000), ) logger.info("Voxtral-Session offen: id=%s lang=%s endpointMs=%d", rid[:8], self._sessions[rid].language, endpoint_ms) def feed_chunk(self, payload: dict) -> bool: sess = self._sessions.get(payload.get("requestId", "")) if sess is None or sess.closed: return False pcm_b64 = payload.get("pcm", "") if pcm_b64: try: sess.pcm_buffer.extend(base64.b64decode(pcm_b64)) except Exception: pass sess.last_chunk_at = time.time() return True def end_session(self, request_id: str) -> None: sess = self._sessions.get(request_id) if sess is not None: sess.closed = True def drop(self, request_id: str) -> None: self._sessions.pop(request_id, None) # ── Endpointer (adaptiv, M0.1) ── def _buffer_ms(self, sess: StreamSession) -> float: samples = len(sess.pcm_buffer) // 2 return (samples / sess.sample_rate) * 1000.0 if samples else 0.0 def _tail_rms(self, sess: StreamSession) -> float: win = int(sess.sample_rate * STREAM_ENERGY_WINDOW_MS / 1000) * 2 if win <= 0: return 0.0 tail = sess.pcm_buffer[-win:] if len(tail) < 2: return 0.0 arr = pcm_s16le_to_float32(bytes(tail)) return float(np.sqrt(np.mean(arr * arr))) if arr.size else 0.0 def _voice_threshold(self, sess: StreamSession) -> float: nf = sess.noise_floor if nf <= 0.0: return sess.voice_rms_min return min(max(nf * sess.voice_factor, sess.voice_rms_min), sess.voice_rms_max) def _update_noise_floor(self, sess: StreamSession, rms: float) -> None: nf = sess.noise_floor if nf <= 0.0: sess.noise_floor = rms elif rms < nf: sess.noise_floor = 0.90 * nf + 0.10 * rms else: sess.noise_floor = 0.98 * nf + 0.02 * rms async def run_endpointer(self) -> None: logger.info("Voxtral-Endpointer gestartet (adaptiver VAD, interval=%dms)", STREAM_TRANSCRIBE_INTERVAL_MS) while True: await asyncio.sleep(0.2) now = time.time() for sid, sess in list(self._sessions.items()): try: await self._tick(sess, now) except Exception: logger.exception("Tick crashed (session=%s)", sid[:8]) for sid, sess in list(self._sessions.items()): if now - sess.last_chunk_at > STREAM_SESSION_TTL_S: logger.info("Stream %s: TTL — drop", sid[:8]) self.drop(sid) async def _tick(self, sess: StreamSession, now: float) -> None: if sess.endpoint_sent: return if (now - sess.started_at) * 1000.0 > sess.hard_cap_ms and not sess.closed: await self._finalize(sess, "hardcap") return if sess.closed: await self._finalize(sess, "stream_end") return if self._buffer_ms(sess) < STREAM_MIN_AUDIO_MS: return # adaptive akustische Sprach-Aktivitaet rms = self._tail_rms(sess) if rms >= self._voice_threshold(sess): sess.last_voice_at = now else: self._update_noise_floor(sess, rms) # Endpoint-Entscheidung, sobald Text erkannt wurde if sess.last_growth_at > 0.0: ac_sil = (now - sess.last_voice_at) * 1000.0 if sess.last_voice_at > 0 else 0.0 se_sil = (now - sess.last_growth_at) * 1000.0 ac_done = sess.last_voice_at > 0 and ac_sil >= sess.endpoint_ms se_done = se_sil >= sess.endpoint_ms * STREAM_SEMANTIC_BACKUP_FACTOR if ac_done or se_done: await self._finalize(sess, "endpoint" if ac_done else "endpoint_semantic") return # Partial-Transkription (throttled) if (now - sess.last_transcribe_at) * 1000.0 < STREAM_TRANSCRIBE_INTERVAL_MS: return sess.last_transcribe_at = now audio = pcm_s16le_to_float32(bytes(sess.pcm_buffer)) try: text = (await self.runner.transcribe(audio, sess.language)).strip() except Exception: logger.exception("Stream %s: Partial-Transcribe crashed", sess.request_id[:8]) return if text and text != sess.last_partial: sess.last_partial = text sess.last_growth_at = now if self._ws is not None: await _send(self._ws, "stt_partial", { "requestId": sess.request_id, "audioRequestId": sess.audio_request_id, "text": text, }) async def _finalize(self, sess: StreamSession, reason: str) -> None: if sess.endpoint_sent: return sess.endpoint_sent = True audio = pcm_s16le_to_float32(bytes(sess.pcm_buffer)) t0 = time.time() try: final_text = (await self.runner.transcribe(audio, sess.language)).strip() except Exception: logger.exception("Stream %s: Final-Transcribe crashed", sess.request_id[:8]) final_text = sess.last_partial stt_ms = int((time.time() - t0) * 1000) duration_s = audio.size / 16000.0 logger.info("Stream %s: FINAL (reason=%s, %.1fs, %dms): %r", sess.request_id[:8], reason, duration_s, stt_ms, final_text[:120]) if self._ws is not None: payload = { "requestId": sess.request_id, "audioRequestId": sess.audio_request_id, "text": final_text, "reason": reason, "durationS": duration_s, "sttMs": stt_ms, "voice": sess.voice, "speed": sess.speed, "interrupted": sess.interrupted, } if sess.location: payload["location"] = sess.location await _send(self._ws, "stt_endpoint", payload) await _send(self._ws, "stt_stream_done", { "requestId": sess.request_id, "audioRequestId": sess.audio_request_id, "text": final_text, "reason": reason, }) self.drop(sess.request_id) async def _broadcast_status(ws, state: str, **extra) -> None: payload = {"service": "voxtral", "state": state} payload.update(extra) await _send(ws, "service_status", payload) async def run_loop(sessions: SessionManager) -> None: use_tls = RVS_TLS retry_s = 2 tls_fallback_tried = False while True: scheme = "wss" if use_tls else "ws" url = f"{scheme}://{RVS_HOST}:{RVS_PORT}/ws?token={RVS_TOKEN}" masked = url.replace(RVS_TOKEN, "***") if RVS_TOKEN else url try: logger.info("Verbinde zu RVS: %s", masked) async with websockets.connect(url, ping_interval=20, ping_timeout=10, max_size=50 * 1024 * 1024) as ws: logger.info("RVS verbunden") retry_s = 2 tls_fallback_tried = False sessions.attach_ws(ws) await _broadcast_status(ws, "ready", model=VOXTRAL_MODEL) await _send(ws, "config_request", {"service": "voxtral"}) async for raw in ws: try: msg = json.loads(raw) except Exception: continue mtype = msg.get("type", "") payload = msg.get("payload", {}) or {} if mtype == "stt_stream_start": sessions.start_session(payload) elif mtype == "stt_audio_chunk": sessions.feed_chunk(payload) elif mtype == "stt_stream_end": sessions.end_session(payload.get("requestId", "")) except Exception as e: logger.warning("RVS-Verbindung verloren: %s — retry in %ds", e, retry_s) if use_tls and RVS_TLS_FALLBACK and not tls_fallback_tried: use_tls = False tls_fallback_tried = True continue await asyncio.sleep(retry_s) retry_s = min(retry_s * 2, 30) use_tls = RVS_TLS async def main() -> None: if not RVS_HOST or not RVS_TOKEN: logger.error("RVS_HOST/RVS_TOKEN fehlen — .env pruefen. Abbruch.") return runner = VoxtralRunner() loop = asyncio.get_running_loop() await loop.run_in_executor(None, runner.load) # Modell laden (blockierend) sessions = SessionManager(runner) logger.info("Voxtral-Bridge startet — Modell=%s", VOXTRAL_MODEL) await asyncio.gather(run_loop(sessions), sessions.run_endpointer()) if __name__ == "__main__": try: asyncio.run(main()) except KeyboardInterrupt: pass