Skip to content

Commit cbed336

Browse files
authored
fix(server): return speaker labels for spk requests (#3433)
Make spk=true functional across vLLM and fallback server paths, preserve speaker labels in verbose_json segments, and accept sentence_info sentence text.
1 parent 1b9926d commit cbed336

3 files changed

Lines changed: 370 additions & 19 deletions

File tree

funasr/bin/_server_app.py

Lines changed: 141 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -127,11 +127,95 @@ def prepare_audio_for_inference(audio_data, sr, target_sr=16000):
127127

128128
return audio_data.astype(np.float32), sr
129129

130+
131+
def attach_speaker_labels(audio_data, sr, segments, speaker_model, device):
132+
"""Run speaker diarization once and attach labels to timestamped segments."""
133+
if not segments:
134+
return segments
135+
136+
import torch
137+
from funasr.models.campplus.cluster_backend import ClusterBackend
138+
from funasr.models.campplus.utils import distribute_spk, postprocess, sv_chunk
139+
140+
audio_data, sr = prepare_audio_for_inference(audio_data, sr)
141+
duration = len(audio_data) / sr
142+
diarization_inputs = []
143+
segment_indexes = []
144+
for index, segment in enumerate(segments):
145+
start = max(float(segment.get("start", 0.0)), 0.0)
146+
end = min(max(float(segment.get("end", start)), start), duration)
147+
start_sample = int(start * sr)
148+
end_sample = int(end * sr)
149+
if end_sample <= start_sample:
150+
continue
151+
diarization_inputs.append([start, end, audio_data[start_sample:end_sample]])
152+
segment_indexes.append(index)
153+
154+
chunks = sv_chunk(diarization_inputs, fs=sr)
155+
if not chunks:
156+
return segments
157+
158+
speaker_results = speaker_model.generate(
159+
input=[chunk[2] for chunk in chunks], cache={}, is_final=True
160+
)
161+
embeddings = torch.cat(
162+
[speaker_result["spk_embedding"] for speaker_result in speaker_results], dim=0
163+
)
164+
labels = ClusterBackend(merge_thr=0.78).to(device)(embeddings.cpu(), oracle_num=None)
165+
if not isinstance(labels, np.ndarray):
166+
labels = np.asarray(labels)
167+
speaker_timeline = postprocess(
168+
sorted(chunks, key=lambda chunk: chunk[0]),
169+
None,
170+
labels,
171+
embeddings.detach().cpu().numpy(),
172+
)
173+
sentences = [
174+
{
175+
"text": segments[index]["text"],
176+
"start": int(float(segments[index]["start"]) * 1000),
177+
"end": int(float(segments[index]["end"]) * 1000),
178+
}
179+
for index in segment_indexes
180+
]
181+
distribute_spk(sentences, speaker_timeline)
182+
for index, sentence in zip(segment_indexes, sentences):
183+
speaker = sentence.get("spk")
184+
if speaker is not None:
185+
segments[index]["speaker"] = f"SPK{speaker}"
186+
return segments
187+
188+
189+
def build_openai_verbose_json(result, requested_language=None):
190+
"""Build OpenAI-compatible verbose JSON while preserving FunASR extensions."""
191+
segments = []
192+
for index, segment in enumerate(result.get("segments", [])):
193+
item = {
194+
"id": index,
195+
"start": segment["start"],
196+
"end": segment["end"],
197+
"text": segment["text"],
198+
"words": segment.get("words", []),
199+
}
200+
if segment.get("speaker") is not None:
201+
item["speaker"] = segment["speaker"]
202+
segments.append(item)
203+
204+
return {
205+
"task": "transcribe",
206+
"language": resolve_transcription_language(requested_language, result),
207+
"duration": result.get("duration", 0),
208+
"text": result["text"],
209+
"segments": segments,
210+
}
211+
212+
130213
def create_app(
131214
device: str = "cuda",
132215
preload_model: str = "auto",
133216
model_path: str = None,
134217
hub: str = "ms",
218+
spk_model: str = "cam++",
135219
cors_origins: Optional[Iterable[str]] = None,
136220
) -> FastAPI:
137221
if preload_model == "auto":
@@ -141,6 +225,8 @@ def create_app(
141225
app.state.device = device
142226
app.state.engine = None
143227
app.state.vad_model = None
228+
app.state.spk_model = None
229+
app.state.spk_model_name = spk_model
144230
app.state.fallback_models = {}
145231
app.state.model_path = model_path
146232
app.state.hub = hub
@@ -173,6 +259,24 @@ def create_app(
173259
},
174260
}
175261

262+
def _load_spk_model():
263+
"""Lazily load diarization only when a request opts in with spk=true."""
264+
if app.state.spk_model is not None:
265+
return app.state.spk_model
266+
if not app.state.spk_model_name:
267+
raise HTTPException(400, "Speaker diarization is disabled; configure --spk-model")
268+
269+
from funasr import AutoModel
270+
271+
logger.info(f"Loading speaker model: {app.state.spk_model_name}")
272+
app.state.spk_model = AutoModel(
273+
model=app.state.spk_model_name,
274+
device=device,
275+
disable_update=True,
276+
)
277+
logger.info("Speaker model ready.")
278+
return app.state.spk_model
279+
176280
def _load_vllm_engine():
177281
"""Load Fun-ASR-Nano vLLM engine. Falls back to AutoModel if vLLM unavailable."""
178282
if app.state.engine is not None or "fun-asr-nano" in app.state.fallback_models:
@@ -286,13 +390,22 @@ def _process_vllm(audio_data, sr, language=None, hotwords=None, use_spk=False):
286390
output_segments.append(seg_info)
287391
full_text_parts.append(text)
288392

393+
if use_spk:
394+
attach_speaker_labels(
395+
audio_data,
396+
sr,
397+
output_segments,
398+
_load_spk_model(),
399+
device,
400+
)
401+
289402
return {
290403
"text": "".join(full_text_parts),
291404
"segments": output_segments,
292405
"duration": len(audio_data) / sr,
293406
}
294407

295-
def _process_fallback(model_name, audio_path, language=None):
408+
def _process_fallback(model_name, audio_path, language=None, use_spk=False):
296409
"""Process with non-LLM model (SenseVoice/Paraformer)."""
297410
model = _load_fallback(model_name)
298411
try:
@@ -309,14 +422,27 @@ def _process_fallback(model_name, audio_path, language=None):
309422
segments = []
310423
if "sentence_info" in result[0]:
311424
for s in result[0]["sentence_info"]:
312-
segments.append({
425+
segment = {
313426
"start": s.get("start", 0)/1000,
314427
"end": s.get("end", 0)/1000,
315-
"text": re.sub(r'<\|[^|]*\|>', '', s.get("text", "")).strip(),
316-
"speaker": s.get("spk"),
317-
})
428+
"text": re.sub(
429+
r'<\|[^|]*\|>', '', s.get("text") or s.get("sentence", "")
430+
).strip(),
431+
}
432+
if s.get("spk") is not None:
433+
segment["speaker"] = s["spk"]
434+
segments.append(segment)
318435
if not segments and text:
319436
segments = build_openai_fallback_segments(text, duration)
437+
if use_spk and segments:
438+
audio_data, sr = sf.read(audio_path)
439+
attach_speaker_labels(
440+
audio_data,
441+
sr,
442+
segments,
443+
_load_spk_model(),
444+
device,
445+
)
320446
return {
321447
"text": text,
322448
"segments": segments,
@@ -356,7 +482,9 @@ async def transcribe(
356482
tmp.write(content)
357483
tmp_path = tmp.name
358484
try:
359-
result = _process_fallback("fun-asr-nano", tmp_path, language=language)
485+
result = _process_fallback(
486+
"fun-asr-nano", tmp_path, language=language, use_spk=spk
487+
)
360488
finally:
361489
os.unlink(tmp_path)
362490
elif model in FALLBACK_CONFIGS or model == "custom":
@@ -365,7 +493,9 @@ async def transcribe(
365493
tmp.write(content)
366494
tmp_path = tmp.name
367495
try:
368-
result = _process_fallback(model, tmp_path, language=language)
496+
result = _process_fallback(
497+
model, tmp_path, language=language, use_spk=spk
498+
)
369499
finally:
370500
os.unlink(tmp_path)
371501
else:
@@ -375,16 +505,7 @@ async def transcribe(
375505
t1 = time.perf_counter()
376506

377507
if response_format == "verbose_json":
378-
return JSONResponse({
379-
"task": "transcribe",
380-
"language": resolve_transcription_language(language, result),
381-
"duration": result.get("duration", 0),
382-
"text": result["text"],
383-
"segments": [
384-
{"id": i, "start": s["start"], "end": s["end"], "text": s["text"], "words": s.get("words", [])}
385-
for i, s in enumerate(result["segments"])
386-
],
387-
})
508+
return JSONResponse(build_openai_verbose_json(result, requested_language=language))
388509
elif response_format == "text":
389510
return JSONResponse(result["text"])
390511
else:
@@ -412,7 +533,9 @@ async def asr_endpoint(
412533
tmp.write(content)
413534
tmp_path = tmp.name
414535
try:
415-
result = _process_fallback("fun-asr-nano", tmp_path, language=language)
536+
result = _process_fallback(
537+
"fun-asr-nano", tmp_path, language=language, use_spk=spk
538+
)
416539
finally:
417540
os.unlink(tmp_path)
418541
t1 = time.perf_counter()

funasr/bin/server.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,11 @@ def build_parser():
4848
parser.add_argument("--model", default="auto", help="Pre-load model: auto (GPU=fun-asr-nano, CPU=sensevoice), sensevoice, paraformer, fun-asr-nano")
4949
parser.add_argument("--model-path", default=None, help="Local model path or model ID (overrides --model)")
5050
parser.add_argument("--hub", default="ms", help="Model hub: ms (ModelScope), hf (HuggingFace) (default: ms)")
51+
parser.add_argument(
52+
"--spk-model",
53+
default="cam++",
54+
help="Speaker model loaded on the first spk=true request (default: cam++)",
55+
)
5156
parser.add_argument(
5257
"--cors-origin",
5358
action="append",
@@ -81,6 +86,7 @@ def main():
8186
preload_model=args.model,
8287
model_path=args.model_path,
8388
hub=args.hub,
89+
spk_model=args.spk_model,
8490
cors_origins=args.cors_origin,
8591
)
8692

0 commit comments

Comments
 (0)