feat(asr): diarize uploaded audio

This commit is contained in:
Xulilong 2026-07-27 20:26:33 -07:00
parent 41baca051c
commit 0a369bd72c
3 changed files with 98 additions and 30 deletions

View File

@ -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):

View File

@ -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

View File

@ -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()
# 如果片段没有说话人,使用前一个片段的说话人