Support known speaker hints in REST API
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user