|
|
|
|
@@ -0,0 +1,415 @@
|
|
|
|
|
"""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()
|