Support known speaker hints in REST API

This commit is contained in:
Aaron Boxer
2026-06-01 19:46:41 -04:00
committed by Aaron Boxer
parent 32c1b18c9f
commit 44940e2834
5 changed files with 180 additions and 28 deletions
+2
View File
@@ -245,6 +245,8 @@ When enabled, completed segments include a `speaker` field:
``` ```
Diarization uses online cosine-similarity clustering of speaker embeddings. If `pyannote.audio` is not installed, the server logs a warning and continues without diarization. Diarization uses online cosine-similarity clustering of speaker embeddings. If `pyannote.audio` is not installed, the server logs a warning and continues without diarization.
The OpenAI-compatible REST endpoint also accepts `known_speaker_names` and uploaded `known_speaker_references` multipart fields. When speaker fields are supplied with `response_format="verbose_json"`, segments include a `speaker` field.
#### Batch Inference #### Batch Inference
Batch multiple client sessions into single GPU calls for higher throughput: Batch multiple client sessions into single GPU calls for higher throughput:
```bash ```bash
+49
View File
@@ -8,6 +8,7 @@ class TestSpeakerDiarizer(unittest.TestCase):
def _make_diarizer(self, **kwargs): def _make_diarizer(self, **kwargs):
from whisper_live.diarization import SpeakerDiarizer from whisper_live.diarization import SpeakerDiarizer
d = SpeakerDiarizer(**kwargs) d = SpeakerDiarizer(**kwargs)
# Mock the embedding model to return deterministic embeddings # Mock the embedding model to return deterministic embeddings
d._model = MagicMock() d._model = MagicMock()
@@ -82,8 +83,26 @@ class TestSpeakerDiarizer(unittest.TestCase):
self.assertEqual(len(d.speakers), 0) self.assertEqual(len(d.speakers), 0)
self.assertEqual(d._speaker_count, 0) self.assertEqual(d._speaker_count, 0)
def test_enroll_speaker_uses_known_name(self):
d = self._make_diarizer(similarity_threshold=0.8)
self._set_embedding(d, [1.0, 0.0, 0.0])
audio = np.zeros(16000, dtype=np.float32)
self.assertTrue(d.enroll_speaker("Alice", audio))
self._set_embedding(d, [0.99, 0.01, 0.0])
speaker = d.identify_speaker(audio)
self.assertEqual(speaker, "Alice")
def test_speaker_names_label_new_speakers(self):
d = self._make_diarizer(speaker_names=["Alice"])
self._set_embedding(d, [1.0, 0.0, 0.0])
audio = np.zeros(16000, dtype=np.float32)
speaker = d.identify_speaker(audio)
self.assertEqual(speaker, "Alice")
def test_import_error_without_pyannote(self): def test_import_error_without_pyannote(self):
from whisper_live.diarization import SpeakerDiarizer from whisper_live.diarization import SpeakerDiarizer
d = SpeakerDiarizer() d = SpeakerDiarizer()
with patch.dict("sys.modules", {"pyannote": None, "pyannote.audio": None}): with patch.dict("sys.modules", {"pyannote": None, "pyannote.audio": None}):
with self.assertRaises(ImportError): with self.assertRaises(ImportError):
@@ -100,8 +119,10 @@ class TestDiarizationInBase(unittest.TestCase):
def __init__(self, **kwargs): def __init__(self, **kwargs):
super().__init__(**kwargs) super().__init__(**kwargs)
self.language = "en" self.language = "en"
def transcribe_audio(self, input_sample): def transcribe_audio(self, input_sample):
return None return None
def handle_transcription_output(self, result, duration): def handle_transcription_output(self, result, duration):
pass pass
@@ -148,5 +169,33 @@ class TestDiarizationInBase(unittest.TestCase):
mock_diarizer.identify_speaker.assert_called_once() mock_diarizer.identify_speaker.assert_called_once()
class TestRestDiarizationHelpers(unittest.TestCase):
def test_normalize_form_list_accepts_repeated_or_comma_separated_values(self):
from whisper_live.server import TranscriptionServer
self.assertEqual(
TranscriptionServer._normalize_form_list(["Alice,Bob", "Carol"]),
["Alice", "Bob", "Carol"],
)
def test_speaker_labels_for_segments(self):
from whisper_live.server import TranscriptionServer
segment = MagicMock()
segment.start = 0.0
segment.end = 1.0
diarizer = MagicMock()
diarizer.identify_speaker.return_value = "Alice"
speakers = TranscriptionServer._speaker_labels_for_segments(
[segment],
np.zeros(16000, dtype=np.float32),
diarizer,
)
self.assertEqual(speakers, {0: "Alice"})
diarizer.identify_speaker.assert_called_once()
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
+5 -15
View File
@@ -343,11 +343,6 @@ class TestStreamTranscription(unittest.TestCase):
def _make_app(self): def _make_app(self):
"""Create a FastAPI app with the transcribe endpoint that has streaming support.""" """Create a FastAPI app with the transcribe endpoint that has streaming support."""
from fastapi import FastAPI, UploadFile, Form from fastapi import FastAPI, UploadFile, Form
from fastapi.testclient import TestClient
from starlette.responses import StreamingResponse
import os
import tempfile
import shutil
app = FastAPI() app = FastAPI()
server = TranscriptionServer() server = TranscriptionServer()
@@ -504,10 +499,9 @@ class TestRESTAPIParamWarnings(unittest.TestCase):
def setUpClass(cls): def setUpClass(cls):
"""Build a FastAPI test app by extracting the endpoint definition.""" """Build a FastAPI test app by extracting the endpoint definition."""
import logging import logging
from fastapi import FastAPI, UploadFile, Form from fastapi import FastAPI, UploadFile, Form, File
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
from typing import Optional, List from typing import Optional, List
from starlette.responses import PlainTextResponse, JSONResponse
app = FastAPI() app = FastAPI()
@@ -523,16 +517,12 @@ class TestRESTAPIParamWarnings(unittest.TestCase):
chunking_strategy: Optional[str] = Form(default=None), chunking_strategy: Optional[str] = Form(default=None),
include: Optional[List[str]] = Form(default=None), include: Optional[List[str]] = Form(default=None),
known_speaker_names: Optional[List[str]] = Form(default=None), known_speaker_names: Optional[List[str]] = Form(default=None),
known_speaker_references: Optional[List[str]] = Form(default=None), known_speaker_references: Optional[List[UploadFile]] = File(default=None),
stream: bool = Form(default=False), stream: bool = Form(default=False),
): ):
ignored_params = [] ignored_params = []
if chunking_strategy: if chunking_strategy:
ignored_params.append(f"chunking_strategy='{chunking_strategy}'") ignored_params.append(f"chunking_strategy='{chunking_strategy}'")
if known_speaker_names:
ignored_params.append("known_speaker_names")
if known_speaker_references:
ignored_params.append("known_speaker_references")
if include: if include:
ignored_params.append(f"include={include}") ignored_params.append(f"include={include}")
if ignored_params: if ignored_params:
@@ -565,17 +555,17 @@ class TestRESTAPIParamWarnings(unittest.TestCase):
ignored = resp.json()["ignored"] ignored = resp.json()["ignored"]
self.assertTrue(any("include" in p for p in ignored)) self.assertTrue(any("include" in p for p in ignored))
def test_known_speaker_names_warning(self): def test_known_speaker_names_supported(self):
resp = self._post(known_speaker_names="alice") resp = self._post(known_speaker_names="alice")
self.assertEqual(resp.status_code, 200) self.assertEqual(resp.status_code, 200)
ignored = resp.json()["ignored"] ignored = resp.json()["ignored"]
self.assertTrue(any("known_speaker_names" in p for p in ignored)) self.assertFalse(any("known_speaker_names" in p for p in ignored))
def test_multiple_ignored_params(self): def test_multiple_ignored_params(self):
resp = self._post(chunking_strategy="auto", known_speaker_names="bob") resp = self._post(chunking_strategy="auto", known_speaker_names="bob")
self.assertEqual(resp.status_code, 200) self.assertEqual(resp.status_code, 200)
ignored = resp.json()["ignored"] ignored = resp.json()["ignored"]
self.assertGreaterEqual(len(ignored), 2) self.assertEqual(len(ignored), 1)
class TestAPIKeyAuth(unittest.TestCase): class TestAPIKeyAuth(unittest.TestCase):
+41 -2
View File
@@ -12,6 +12,28 @@ import logging
import numpy as np import numpy as np
def load_audio(file_path, sample_rate=16000):
"""Load an audio file as mono float32 PCM at the requested sample rate."""
import av
container = av.open(file_path)
resampler = av.AudioResampler(format="flt", layout="mono", rate=sample_rate)
chunks = []
try:
for frame in container.decode(audio=0):
for resampled_frame in resampler.resample(frame):
chunks.append(
resampled_frame.to_ndarray().reshape(-1).astype(np.float32)
)
finally:
container.close()
if not chunks:
return np.array([], dtype=np.float32)
return np.concatenate(chunks)
class SpeakerDiarizer: class SpeakerDiarizer:
"""Real-time speaker diarization using speaker embeddings and online clustering. """Real-time speaker diarization using speaker embeddings and online clustering.
@@ -38,15 +60,22 @@ class SpeakerDiarizer:
max_speakers=10, max_speakers=10,
embedding_model="pyannote/wespeaker-voxceleb-resnet34-LM", embedding_model="pyannote/wespeaker-voxceleb-resnet34-LM",
hf_token=None, hf_token=None,
speaker_names=None,
): ):
self.similarity_threshold = similarity_threshold self.similarity_threshold = similarity_threshold
self.max_speakers = max_speakers self.max_speakers = max_speakers
self.speaker_names = list(speaker_names or [])
self.speakers = {} # speaker_id -> embedding (averaged) self.speakers = {} # speaker_id -> embedding (averaged)
self._speaker_count = 0 self._speaker_count = 0
self._model = None self._model = None
self._embedding_model_name = embedding_model self._embedding_model_name = embedding_model
self._hf_token = hf_token self._hf_token = hf_token
def _next_speaker_id(self):
if self._speaker_count < len(self.speaker_names):
return self.speaker_names[self._speaker_count]
return f"SPEAKER_{self._speaker_count:02d}"
def _load_model(self): def _load_model(self):
"""Lazy-load the embedding model on first use.""" """Lazy-load the embedding model on first use."""
if self._model is not None: if self._model is not None:
@@ -128,14 +157,24 @@ class SpeakerDiarizer:
if len(self.speakers) >= self.max_speakers: if len(self.speakers) >= self.max_speakers:
# Assign to closest speaker # Assign to closest speaker
return best_speaker if best_speaker else f"SPEAKER_{self._speaker_count:02d}" return (
best_speaker if best_speaker else f"SPEAKER_{self._speaker_count:02d}"
)
# Create a new speaker # Create a new speaker
speaker_id = f"SPEAKER_{self._speaker_count:02d}" speaker_id = self._next_speaker_id()
self._speaker_count += 1 self._speaker_count += 1
self.speakers[speaker_id] = embedding self.speakers[speaker_id] = embedding
return speaker_id return speaker_id
def enroll_speaker(self, speaker_name, audio_np, sample_rate=16000):
"""Enroll a known speaker from reference audio."""
embedding = self._compute_embedding(audio_np, sample_rate)
if embedding is None:
return False
self.speakers[speaker_name] = embedding
return True
def reset(self): def reset(self):
"""Reset all speaker state.""" """Reset all speaker state."""
self.speakers.clear() self.speakers.clear()
+83 -11
View File
@@ -9,19 +9,18 @@ import logging
import shutil import shutil
import tempfile import tempfile
from typing import Optional, List from typing import Optional, List
from fastapi import FastAPI, UploadFile, Form, Request from fastapi import FastAPI, UploadFile, Form, Request, File
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from starlette.responses import PlainTextResponse, JSONResponse, StreamingResponse from starlette.responses import PlainTextResponse, StreamingResponse
import uvicorn import uvicorn
from faster_whisper import WhisperModel from faster_whisper import WhisperModel
import torch import torch
from enum import Enum from enum import Enum
from whisper_live import metrics as wl_metrics
from typing import List, Optional
import numpy as np import numpy as np
from whisper_live import metrics as wl_metrics
from websockets.sync.server import serve from websockets.sync.server import serve
from websockets.exceptions import ConnectionClosed from websockets.exceptions import ConnectionClosed
from whisper_live.vad import VoiceActivityDetector from whisper_live.vad import VoiceActivityDetector
@@ -523,6 +522,67 @@ class TranscriptionServer:
return StreamingResponse(_sse_generator(), media_type="text/event-stream") return StreamingResponse(_sse_generator(), media_type="text/event-stream")
@staticmethod
def _normalize_form_list(values):
"""Normalize repeated or comma-separated multipart form fields."""
if not values:
return []
normalized = []
for value in values:
if isinstance(value, str):
normalized.extend(item.strip() for item in value.split(",") if item.strip())
return normalized
async def _create_rest_diarizer(self, known_speaker_names, known_speaker_references):
"""Create a diarizer from OpenAI-compatible known speaker fields."""
speaker_names = self._normalize_form_list(known_speaker_names)
speaker_references = known_speaker_references or []
if speaker_references and not speaker_names:
raise ValueError("known_speaker_references requires matching known_speaker_names")
if speaker_names and speaker_references and len(speaker_names) != len(speaker_references):
raise ValueError("known_speaker_names and known_speaker_references must have the same length")
if not speaker_names and not speaker_references:
return None
from whisper_live.diarization import SpeakerDiarizer, load_audio
diarizer = SpeakerDiarizer(
max_speakers=max(10, len(speaker_names)),
speaker_names=speaker_names,
)
for speaker_name, reference in zip(speaker_names, speaker_references):
suffix = os.path.splitext(reference.filename or "")[1] or ".wav"
reference_path = None
try:
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
tmp.write(await reference.read())
reference_path = tmp.name
audio_np = load_audio(reference_path)
if not diarizer.enroll_speaker(speaker_name, audio_np):
raise ValueError(f"known_speaker_references for '{speaker_name}' is too short")
finally:
if reference_path and os.path.exists(reference_path):
os.unlink(reference_path)
return diarizer
@staticmethod
def _speaker_labels_for_segments(segments, audio_np, diarizer, sample_rate=16000):
if diarizer is None or audio_np is None:
return {}
labels = {}
for index, segment in enumerate(segments):
start = max(0, int(segment.start * sample_rate))
end = min(len(audio_np), int(segment.end * sample_rate))
if end <= start:
continue
speaker = diarizer.identify_speaker(audio_np[start:end], sample_rate)
if speaker:
labels[index] = speaker
return labels
def run(self, def run(self,
host, host,
port=9090, port=9090,
@@ -666,7 +726,7 @@ class TranscriptionServer:
chunking_strategy: Optional[str] = Form(default=None), chunking_strategy: Optional[str] = Form(default=None),
include: Optional[List[str]] = Form(default=None), include: Optional[List[str]] = Form(default=None),
known_speaker_names: Optional[List[str]] = Form(default=None), known_speaker_names: Optional[List[str]] = Form(default=None),
known_speaker_references: Optional[List[str]] = Form(default=None), known_speaker_references: Optional[List[UploadFile]] = File(default=None),
stream: bool = Form(default=False), stream: bool = Form(default=False),
hotwords: Optional[str] = Form(default=None), hotwords: Optional[str] = Form(default=None),
): ):
@@ -680,10 +740,6 @@ class TranscriptionServer:
ignored_params = [] ignored_params = []
if chunking_strategy: if chunking_strategy:
ignored_params.append(f"chunking_strategy='{chunking_strategy}'") ignored_params.append(f"chunking_strategy='{chunking_strategy}'")
if known_speaker_names:
ignored_params.append("known_speaker_names")
if known_speaker_references:
ignored_params.append("known_speaker_references")
if include: if include:
ignored_params.append(f"include={include}") ignored_params.append(f"include={include}")
if ignored_params: if ignored_params:
@@ -698,6 +754,7 @@ class TranscriptionServer:
logging.warning(f"Model '{model}' requested; using 'small' as fallback.") logging.warning(f"Model '{model}' requested; using 'small' as fallback.")
model_name = faster_whisper_custom_model_path or "small" model_name = faster_whisper_custom_model_path or "small"
tmp_path = None
try: try:
suffix = os.path.splitext(file.filename)[1] or ".wav" suffix = os.path.splitext(file.filename)[1] or ".wav"
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp: with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
@@ -717,9 +774,9 @@ class TranscriptionServer:
word_timestamps=(timestamp_granularities and "word" in timestamp_granularities), word_timestamps=(timestamp_granularities and "word" in timestamp_granularities),
hotwords=hotwords, hotwords=hotwords,
) )
segments = list(segments)
text = " ".join([s.text.strip() for s in segments]) text = " ".join([s.text.strip() for s in segments])
os.unlink(tmp_path)
if response_format == "text": if response_format == "text":
wl_metrics.track_rest_request(endpoint="transcriptions", status=200) wl_metrics.track_rest_request(endpoint="transcriptions", status=200)
@@ -735,7 +792,17 @@ class TranscriptionServer:
"text": text, "text": text,
"segments": [] "segments": []
} }
for seg in segments: speaker_labels = {}
try:
rest_diarizer = await self._create_rest_diarizer(known_speaker_names, known_speaker_references)
except ValueError as e:
wl_metrics.track_rest_request(endpoint="transcriptions", status=400)
return JSONResponse({"error": str(e)}, status_code=400)
if rest_diarizer is not None:
from whisper_live.diarization import load_audio
audio_np = load_audio(tmp_path)
speaker_labels = self._speaker_labels_for_segments(segments, audio_np, rest_diarizer)
for index, seg in enumerate(segments):
seg_dict = { seg_dict = {
"id": seg.id, "id": seg.id,
"seek": seg.seek, "seek": seg.seek,
@@ -748,6 +815,8 @@ class TranscriptionServer:
"compression_ratio": seg.compression_ratio, "compression_ratio": seg.compression_ratio,
"no_speech_prob": seg.no_speech_prob "no_speech_prob": seg.no_speech_prob
} }
if index in speaker_labels:
seg_dict["speaker"] = speaker_labels[index]
if timestamp_granularities and "word" in timestamp_granularities: if timestamp_granularities and "word" in timestamp_granularities:
seg_dict["words"] = [{"word": w.word, "start": w.start, "end": w.end, "probability": w.probability} for w in seg.words] seg_dict["words"] = [{"word": w.word, "start": w.start, "end": w.end, "probability": w.probability} for w in seg.words]
verbose["segments"].append(seg_dict) verbose["segments"].append(seg_dict)
@@ -768,6 +837,9 @@ class TranscriptionServer:
wl_metrics.track_rest_request(endpoint="transcriptions", status=500) wl_metrics.track_rest_request(endpoint="transcriptions", status=500)
wl_metrics.track_error("rest_transcription") wl_metrics.track_error("rest_transcription")
return JSONResponse({"error": str(e)}, status_code=500) return JSONResponse({"error": str(e)}, status_code=500)
finally:
if tmp_path and os.path.exists(tmp_path):
os.unlink(tmp_path)
threading.Thread( threading.Thread(
target=uvicorn.run, target=uvicorn.run,