Update agent
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user