@@ -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+
130213def 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 ()
0 commit comments