Files
ARIA-AGENT/xtts/voxtral/bridge.py
T
duffyduckandClaude Opus 4.8 36f04f83ff fix(voxtral): willkuerliche Abbrueche — semantischen Endpoint + Live-Partials raus
Ursache: Voxtral-3B transkribiert den ganzen wachsenden Buffer (~5-6s bei langen Aufnahmen). Diese Partial-Latenz war groesser als der semantische Endpoint-Timeout (4.8s) → 'Text waechst nicht mehr' feuerte faelschlich → Abbruch nach 20-40s. Fix: keine Live-Partials mehr, kein semantischer Endpoint — Turn-Ende rein akustisch (Stille-VAD), transkribiert wird nur EINMAL im _finalize. max_new_tokens 512->4096 (512 schnitt lange Diktate ab).

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-08-15 12:02:44 +02:00

408 lines
16 KiB
Python

#!/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 tempfile
import time
from dataclasses import dataclass, field
from typing import Optional
import numpy as np
import soundfile as sf
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:
import torch
proc, model = self.processor, self.model
if proc is None or model is None or audio_f32.size == 0:
return ""
# VoxtralProcessor verlangt bei rohen Arrays ein 'format'. Robuster:
# in ein temp-WAV schreiben und den PFAD uebergeben — der Processor liest
# Format + Samplerate selbst, kein 'format'-Argument noetig.
wav_path = None
try:
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tf:
wav_path = tf.name
sf.write(wav_path, audio_f32, 16000, subtype="PCM_16")
inputs = proc.apply_transcription_request(
language=language, audio=wav_path, model_id=VOXTRAL_MODEL,
)
inputs = inputs.to(VOXTRAL_DEVICE, dtype=torch.bfloat16)
with torch.no_grad():
# hoch genug fuer lange Diktate (stoppt eh am EOS); 512 hat
# mehrminutige Aufnahmen abgeschnitten.
outputs = model.generate(**inputs, max_new_tokens=4096)
trimmed = outputs[:, inputs.input_ids.shape[1]:]
text = proc.batch_decode(trimmed, skip_special_tokens=True)
return (text[0] if text else "").strip()
finally:
if wav_path:
try:
os.unlink(wav_path)
except Exception:
pass
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 (M0.1). KEINE Live-Partials mehr:
# Voxtral-3B transkribiert den ganzen WACHSENDEN Buffer und braucht dafuer
# bei langen Aufnahmen 5-6 s — zu langsam fuer Live-Text, UND diese Latenz
# hat den semantischen Endpoint faelschlich ausgeloest (Partial-Latenz >
# Timeout → willkuerliche Abbrueche nach 20-40 s). Deshalb: Turn-Ende rein
# AKUSTISCH, transkribiert wird nur EINMAL im _finalize.
rms = self._tail_rms(sess)
if rms >= self._voice_threshold(sess):
sess.last_voice_at = now
else:
self._update_noise_floor(sess, rms)
# Endpoint: hat der User schon gesprochen UND ist es seit endpoint_ms still?
if sess.last_voice_at > 0 and (now - sess.last_voice_at) * 1000.0 >= sess.endpoint_ms:
await self._finalize(sess, "endpoint")
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