Files
MeetVault/Wireframes/Desktop/MeetVault-Desktop-Demo/speech-engine/server.py
2026-09-30 00:50:44 +07:00

416 lines
15 KiB
Python

"""MeetVault local live speech engine.
Protocol
--------
WebSocket: ws://127.0.0.1:8765/ws/live
1. Client sends a JSON config message.
2. Client streams signed 16-bit little-endian mono PCM at 16 kHz.
3. Server returns JSON events: ready, partial, final, health, error.
Real mode combines faster-whisper with diart speaker diarization. Mock mode is
available so the desktop UI can be exercised without ML dependencies/models.
"""
from __future__ import annotations
import argparse
import asyncio
import json
import os
import queue
import threading
import time
import uuid
from dataclasses import dataclass
from typing import Optional
import numpy as np
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.middleware.cors import CORSMiddleware
import uvicorn
SAMPLE_RATE = 16000
@dataclass
class EngineConfig:
model: str = "small"
compute: str = "auto"
language: str = "auto"
diarization: bool = True
sample_rate: int = SAMPLE_RATE
class SpeakerTimeline:
"""Thread-safe speaker timeline built from diart streaming predictions."""
def __init__(self) -> None:
self._segments: list[tuple[float, float, str]] = []
self._label_map: dict[str, str] = {}
self._lock = threading.Lock()
def _stable_label(self, raw: str) -> str:
with self._lock:
if raw not in self._label_map:
self._label_map[raw] = f"speaker_{len(self._label_map) + 1}"
return self._label_map[raw]
def update_from_annotation(self, annotation) -> None:
fresh: list[tuple[float, float, str]] = []
try:
for segment, _, label in annotation.itertracks(yield_label=True):
fresh.append((float(segment.start), float(segment.end), self._stable_label(str(label))))
except Exception:
return
if not fresh:
return
latest_end = max(end for _, end, _ in fresh)
with self._lock:
# Keep a rolling two-minute history, then merge the latest prediction.
self._segments = [s for s in self._segments if s[1] >= latest_end - 120]
self._segments.extend(fresh)
def resolve(self, start: float, end: float) -> str:
with self._lock:
candidates = list(self._segments)
best_label = "speaker_1"
best_overlap = 0.0
for s, e, label in candidates:
overlap = max(0.0, min(end, e) - max(start, s))
if overlap > best_overlap:
best_overlap = overlap
best_label = label
return best_label
class QueueAudioSource:
"""diart AudioSource fed by PCM pushed from the WebSocket thread."""
def __init__(self, sample_rate: int):
from diart.sources import AudioSource
class _Source(AudioSource):
def __init__(self, sr: int):
super().__init__("meetvault-live", sr)
self.q: queue.Queue[Optional[np.ndarray]] = queue.Queue()
self.cursor = 0
@property
def is_regular(self):
return False
@property
def duration(self):
return None
def push(self, samples: np.ndarray):
self.q.put(samples.astype(np.float32, copy=False))
def close(self):
self.q.put(None)
def read(self):
from pyannote.core import SlidingWindow, SlidingWindowFeature
while True:
item = self.q.get()
if item is None:
self.stream.on_completed()
return
start = self.cursor / self.sample_rate
window = SlidingWindow(start=start, duration=1 / self.sample_rate, step=1 / self.sample_rate)
feature = SlidingWindowFeature(item.reshape(-1, 1), window)
self.cursor += len(item)
self.stream.on_next(feature)
self.source = _Source(sample_rate)
def push(self, samples: np.ndarray) -> None:
self.source.push(samples)
def close(self) -> None:
self.source.close()
class DiartWorker:
def __init__(self, enabled: bool, sample_rate: int = SAMPLE_RATE):
self.enabled = enabled
self.timeline = SpeakerTimeline()
self.audio_source: Optional[QueueAudioSource] = None
self.thread: Optional[threading.Thread] = None
self.error: Optional[str] = None
if enabled:
self._start(sample_rate)
def _start(self, sample_rate: int) -> None:
try:
from diart import SpeakerDiarization
from diart.inference import StreamingInference
pipeline = SpeakerDiarization()
expected_rate = int(getattr(pipeline.config, "sample_rate", sample_rate))
if expected_rate != sample_rate:
raise RuntimeError(f"diart expects {expected_rate} Hz but MeetVault sends {sample_rate} Hz")
self.audio_source = QueueAudioSource(sample_rate)
inference = StreamingInference(pipeline, self.audio_source.source, do_plot=False)
def hook(annotation_and_waveform):
try:
self.timeline.update_from_annotation(annotation_and_waveform[0])
except Exception:
pass
inference.attach_hooks(hook)
self.thread = threading.Thread(target=inference, name="meetvault-diart", daemon=True)
self.thread.start()
except Exception as exc:
self.error = str(exc)
self.enabled = False
def push(self, samples: np.ndarray) -> None:
if self.enabled and self.audio_source:
self.audio_source.push(samples)
def close(self) -> None:
if self.audio_source:
self.audio_source.close()
def resolve(self, start: float, end: float) -> str:
return self.timeline.resolve(start, end) if self.enabled else "speaker_1"
class LiveSpeechSession:
def __init__(self, config: EngineConfig):
from faster_whisper import WhisperModel
self.config = config
device = config.compute
if device == "auto":
device = "cpu"
compute_type = "float16" if device == "cuda" else "int8"
try:
self.model = WhisperModel(config.model, device=device, compute_type=compute_type)
except Exception:
# Useful fallback when CUDA was selected but is not available.
self.model = WhisperModel(config.model, device="cpu", compute_type="int8")
self.diarizer = DiartWorker(config.diarization, config.sample_rate)
self.samples = np.empty(0, dtype=np.float32)
self.lock = threading.Lock()
self.emitted_until = 0.0
self.last_process_samples = 0
self.closed = False
def push_pcm(self, raw: bytes) -> None:
pcm = np.frombuffer(raw, dtype="<i2").astype(np.float32) / 32768.0
if pcm.size == 0:
return
self.diarizer.push(pcm)
with self.lock:
self.samples = np.concatenate([self.samples, pcm])
# Keep at most 20 minutes in memory. Long meetings are still recorded by the desktop app.
max_samples = self.config.sample_rate * 60 * 20
if self.samples.size > max_samples:
dropped = self.samples.size - max_samples
self.samples = self.samples[dropped:]
self.emitted_until = max(0.0, self.emitted_until - dropped / self.config.sample_rate)
def total_seconds(self) -> float:
with self.lock:
return self.samples.size / self.config.sample_rate
def should_process(self) -> bool:
with self.lock:
new_samples = self.samples.size - self.last_process_samples
return new_samples >= int(self.config.sample_rate * 1.4)
def transcribe_once(self) -> tuple[list[dict], Optional[str]]:
with self.lock:
total = self.samples.size
if total < self.config.sample_rate:
return [], None
window_size = min(total, self.config.sample_rate * 10)
audio = self.samples[-window_size:].copy()
self.last_process_samples = total
window_start = max(0.0, total / self.config.sample_rate - len(audio) / self.config.sample_rate)
segments_iter, _ = self.model.transcribe(
audio,
language=None if self.config.language in ("", "auto") else self.config.language,
beam_size=1,
best_of=1,
vad_filter=True,
word_timestamps=True,
condition_on_previous_text=False,
temperature=0.0,
)
segments = list(segments_iter)
results: list[dict] = []
partial: Optional[str] = None
safe_cutoff = total / self.config.sample_rate - 0.65
for seg in segments:
text = seg.text.strip()
if not text:
continue
abs_start = window_start + float(seg.start)
abs_end = window_start + float(seg.end)
if abs_end <= self.emitted_until + 0.08:
continue
if abs_end > safe_cutoff:
partial = text
continue
speaker_id = self.diarizer.resolve(abs_start, abs_end)
results.append({
"id": str(uuid.uuid4()),
"speakerId": speaker_id,
"start": abs_start,
"end": abs_end,
"text": text,
"final": True,
})
self.emitted_until = max(self.emitted_until, abs_end)
return results, partial
def close(self) -> None:
self.closed = True
self.diarizer.close()
class MockSession:
def __init__(self, config: EngineConfig):
self.config = config
self.sample_count = 0
self.index = 0
self.lines = [
("speaker_1", "We should finalize the upload flow before Friday."),
("speaker_2", "I can take the local recording and video playback work."),
("speaker_1", "Great. Keep the backend in processing state for this demo."),
("speaker_3", "The overlay should keep updating while another application is open."),
]
def push_pcm(self, raw: bytes) -> Optional[dict]:
self.sample_count += len(raw) // 2
threshold = self.config.sample_rate * 2.4
if self.sample_count < threshold:
return None
self.sample_count = 0
speaker, text = self.lines[self.index % len(self.lines)]
self.index += 1
end = self.index * 2.4
return {
"id": str(uuid.uuid4()), "speakerId": speaker,
"start": end - 2.1, "end": end, "text": text, "final": True,
}
def close(self) -> None:
pass
app = FastAPI(title="MeetVault Local Speech Engine")
app.add_middleware(CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"])
MOCK_ONLY = False
@app.get("/health")
async def health():
return {"ok": True, "mode": "mock" if MOCK_ONLY else "real-capable", "sample_rate": SAMPLE_RATE}
@app.websocket("/ws/live")
async def live(websocket: WebSocket):
await websocket.accept()
session = None
process_task = None
try:
first = await websocket.receive_text()
payload = json.loads(first)
config = EngineConfig(
model=str(payload.get("model", "small")),
compute=str(payload.get("compute", "auto")),
language=str(payload.get("language", "auto")),
diarization=bool(payload.get("diarization", True)),
sample_rate=int(payload.get("sample_rate", SAMPLE_RATE)),
)
if config.sample_rate != SAMPLE_RATE:
await websocket.send_json({"type": "error", "message": "MeetVault speech engine currently requires 16 kHz mono PCM."})
return
if MOCK_ONLY:
session = MockSession(config)
await websocket.send_json({"type": "ready", "engine": "mock-whisper", "diarization": "mock-diart"})
else:
try:
session = await asyncio.to_thread(LiveSpeechSession, config)
except Exception as exc:
await websocket.send_json({
"type": "error",
"message": (
f"Could not initialize Whisper/diart: {exc}. "
"Use --mock for UI testing. Real diarization may require Hugging Face authentication and accepted pyannote model terms."
),
})
return
diarization_state = "diart ready" if session.diarizer.enabled else f"diart unavailable: {session.diarizer.error}"
await websocket.send_json({"type": "ready", "engine": f"Whisper {config.model}", "diarization": diarization_state})
async def processor():
while not session.closed:
await asyncio.sleep(0.35)
if not session.should_process():
continue
try:
finals, partial = await asyncio.to_thread(session.transcribe_once)
for segment in finals:
await websocket.send_json({"type": "final", "segment": segment})
if partial:
await websocket.send_json({"type": "partial", "text": partial})
except Exception as exc:
await websocket.send_json({"type": "error", "message": f"Live transcription error: {exc}"})
await asyncio.sleep(1)
process_task = asyncio.create_task(processor())
while True:
message = await websocket.receive()
if message.get("bytes") is not None:
raw = message["bytes"]
if isinstance(session, MockSession):
segment = session.push_pcm(raw)
if segment:
await websocket.send_json({"type": "final", "segment": segment})
else:
session.push_pcm(raw)
elif message.get("text"):
command = json.loads(message["text"])
if command.get("type") == "stop":
break
except WebSocketDisconnect:
pass
except Exception as exc:
try:
await websocket.send_json({"type": "error", "message": str(exc)})
except Exception:
pass
finally:
if session:
session.close()
if process_task:
process_task.cancel()
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, default=8765)
parser.add_argument("--mock", action="store_true", help="Run without Whisper/diart models")
args = parser.parse_args()
global MOCK_ONLY
MOCK_ONLY = args.mock
uvicorn.run(app, host=args.host, port=args.port, log_level="info")
if __name__ == "__main__":
main()