From 0a369bd72c2372f1f964bb9ae694f6b6b62bac4e Mon Sep 17 00:00:00 2001 From: Xulilong Date: Mon, 27 Jul 2026 20:26:33 -0700 Subject: [PATCH] feat(asr): diarize uploaded audio --- asr/main.py | 103 +++++++++++++++---- config.env | 3 +- speakr/src/services/transcription_service.py | 22 ++-- 3 files changed, 98 insertions(+), 30 deletions(-) diff --git a/asr/main.py b/asr/main.py index ac18d4c..ccc3b1e 100644 --- a/asr/main.py +++ b/asr/main.py @@ -103,6 +103,9 @@ VOICEPRINT_STORE = os.environ.get("VOICEPRINT_STORE", "/tmp/asr_voiceprints.json VOICEPRINT_MATCH_THRESHOLD = float(os.environ.get("VOICEPRINT_MATCH_THRESHOLD", "0.45")) REALTIME_PARTIAL_INTERVAL_SECONDS = float(os.environ.get("REALTIME_PARTIAL_INTERVAL_SECONDS", "2.0")) REALTIME_PARTIAL_MIN_SECONDS = float(os.environ.get("REALTIME_PARTIAL_MIN_SECONDS", "1.5")) +ASR_SPK_MODEL = os.environ.get( + "ASR_SPK_MODEL", "iic/speech_campplus_sv_zh-cn_16k-common" +) def _voiceprint_key(database_name: str, collection_name: str) -> str: @@ -224,6 +227,8 @@ def load_model(): model="/app/models/seacomodel", vad_model="/app/models/vadmodel", punc_model="/app/models/ctpuncmodel", + # Native diarization puts a `spk` cluster id on each VAD sentence. + spk_model=ASR_SPK_MODEL, ) logger.info("模型加载完成") @@ -370,20 +375,58 @@ async def websocket_asr_root(websocket: WebSocket): await run_asr_websocket(websocket) -def make_speakr_response(text: str, language: str = "zh", speaker: str = "SPEAKER_00"): - text = (text or "").strip() - return { - "text": text, - "language": language, - "segments": [ - { - "start": None, - "end": None, - "text": text, - "speaker": speaker or "SPEAKER_00", - } - ] if text else [], - } +def _speaker_label(value): + """Convert FunASR speaker ids to the labels expected by Speakr.""" + if value is None or value == "": + return "UNKNOWN_SPEAKER" + if isinstance(value, int) or (isinstance(value, str) and value.isdigit()): + return f"SPEAKER_{int(value):02d}" + return str(value) + + +def _speaker_segments(result): + """Read native FunASR diarization output across supported versions.""" + if not isinstance(result, dict): + return [] + + sentence_info = result.get("sentence_info") or result.get("segments") or [] + if not isinstance(sentence_info, list): + return [] + + segments = [] + for sentence in sentence_info: + if not isinstance(sentence, dict): + continue + text = str(sentence.get("text") or "").strip() + if not text: + continue + segments.append({ + "start": sentence.get("start"), + "end": sentence.get("end"), + "text": text, + "speaker": _speaker_label(sentence.get("speaker", sentence.get("spk"))), + }) + return segments + + +def make_speakr_response( + result, + language: str = "zh", + diarize: bool = True, + fallback_speaker: str = "SPEAKER_00", +): + """Return diarized sentences in Speakr's upload-transcription contract.""" + result = result if isinstance(result, dict) else {} + text = str(result.get("text") or "").strip() + segments = _speaker_segments(result) if diarize else [] + if not segments and text: + segments = [{ + "start": None, + "end": None, + "text": text, + "speaker": fallback_speaker or "SPEAKER_00", + }] + return {"text": text, "language": language, "segments": segments} async def run_asr( @@ -396,20 +439,22 @@ async def run_asr( result = model.generate(input=audio_path, hotword=hotword) if result and len(result) > 0: - text = result[0].get("text", "") + result_data = result[0] + text = result_data.get("text", "") else: + result_data = {} text = "" text = cn_to_arabic(text) + if isinstance(result_data, dict): + result_data = dict(result_data) + result_data["text"] = text logger.info(f"识别结果: {text}") if response_format == "text": return PlainTextResponse(content=text) - return { - "text": text, - "language": language, - } + return {**result_data, "text": text, "language": language} except Exception as e: logger.error(f"识别出错: {e}") @@ -517,6 +562,10 @@ async def transcribe_speakr_compatible( language: Optional[str] = Form("zh"), database_name: Optional[str] = Form("NB"), collection_name: Optional[str] = Form("voice"), + spk_diarization: bool = True, + spk_num: int = 0, + spk_min_speakers: Optional[int] = None, + spk_max_speakers: Optional[int] = None, ): """ Speakr 兼容接口: @@ -560,11 +609,21 @@ async def transcribe_speakr_compatible( except Exception as e: logger.warning(f"声纹匹配失败,使用默认说话人: {e}") - return make_speakr_response( - text=result.get("text", "") if isinstance(result, dict) else "", + response = make_speakr_response( + result=result, language=result.get("language", language or "zh") if isinstance(result, dict) else language or "zh", - speaker=speaker_name, + diarize=spk_diarization, + fallback_speaker=speaker_name, ) + logger.info( + "Speakr diarization: enabled=%s requested_speakers=%s range=%s-%s detected=%s", + spk_diarization, + spk_num, + spk_min_speakers, + spk_max_speakers, + len({segment["speaker"] for segment in response["segments"]}), + ) + return response finally: if tmp_path and os.path.exists(tmp_path): diff --git a/config.env b/config.env index 9b3e589..defbecb 100644 --- a/config.env +++ b/config.env @@ -3,7 +3,8 @@ # Speakr / ASR USE_ASR_ENDPOINT=true ASR_BASE_URL=http://asr-test:59805 -ASR_DIARIZE=true +ASR_DIARIZE=true +ASR_SPK_MODEL=iic/speech_campplus_sv_zh-cn_16k-common ASR_MIN_SPEAKERS=1 ASR_MAX_SPEAKERS=10 diff --git a/speakr/src/services/transcription_service.py b/speakr/src/services/transcription_service.py index d21b5c8..b0b5d91 100644 --- a/speakr/src/services/transcription_service.py +++ b/speakr/src/services/transcription_service.py @@ -221,10 +221,15 @@ def transcribe_audio_asr(app_context, recording_id, filepath, original_filename, with open(current_filepath, 'rb') as audio_file: url = f"{ASR_BASE_URL}/asr" - params = { - 'batch_size_s':300, - 'spk_diarization':True, - 'spk_num':0, + requested_speaker_count = 0 + if min_speakers and max_speakers and min_speakers == max_speakers: + requested_speaker_count = min_speakers + params = { + 'batch_size_s':300, + 'spk_diarization':bool(diarize), + 'spk_num':requested_speaker_count, + 'spk_min_speakers':min_speakers, + 'spk_max_speakers':max_speakers, 'spk_threshold':0.6, 'spk_cluster_method':'ahc', 'spk_smooth_window':3, @@ -334,8 +339,9 @@ def transcribe_audio_asr(app_context, recording_id, filepath, original_filename, segments_without_speakers = 0 for segment in asr_response_data['segments']: - if 'speaker' in segment and segment['speaker'] is not None: - all_speakers.add(segment['speaker']) + speaker = segment.get('speaker', segment.get('spk', segment.get('spk_name'))) + if speaker is not None: + all_speakers.add(speaker) segments_with_speakers += 1 else: segments_without_speakers += 1 @@ -355,7 +361,9 @@ def transcribe_audio_asr(app_context, recording_id, filepath, original_filename, last_known_speaker = None for i, segment in enumerate(asr_response_data['segments']): - speaker = segment.get('speaker') + # Native FunASR uses `spk`; the realtime bridge exposes + # `spk_name`. Normalize both to Speakr's `speaker` field. + speaker = segment.get('speaker', segment.get('spk', segment.get('spk_name'))) text = segment.get('text', '').strip() # 如果片段没有说话人,使用前一个片段的说话人