feat(asr): diarize uploaded audio
This commit is contained in:
parent
41baca051c
commit
0a369bd72c
103
asr/main.py
103
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):
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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()
|
||||
|
||||
# 如果片段没有说话人,使用前一个片段的说话人
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user