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"))
|
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_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"))
|
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:
|
def _voiceprint_key(database_name: str, collection_name: str) -> str:
|
||||||
@ -224,6 +227,8 @@ def load_model():
|
|||||||
model="/app/models/seacomodel",
|
model="/app/models/seacomodel",
|
||||||
vad_model="/app/models/vadmodel",
|
vad_model="/app/models/vadmodel",
|
||||||
punc_model="/app/models/ctpuncmodel",
|
punc_model="/app/models/ctpuncmodel",
|
||||||
|
# Native diarization puts a `spk` cluster id on each VAD sentence.
|
||||||
|
spk_model=ASR_SPK_MODEL,
|
||||||
)
|
)
|
||||||
logger.info("模型加载完成")
|
logger.info("模型加载完成")
|
||||||
|
|
||||||
@ -370,20 +375,58 @@ async def websocket_asr_root(websocket: WebSocket):
|
|||||||
await run_asr_websocket(websocket)
|
await run_asr_websocket(websocket)
|
||||||
|
|
||||||
|
|
||||||
def make_speakr_response(text: str, language: str = "zh", speaker: str = "SPEAKER_00"):
|
def _speaker_label(value):
|
||||||
text = (text or "").strip()
|
"""Convert FunASR speaker ids to the labels expected by Speakr."""
|
||||||
return {
|
if value is None or value == "":
|
||||||
"text": text,
|
return "UNKNOWN_SPEAKER"
|
||||||
"language": language,
|
if isinstance(value, int) or (isinstance(value, str) and value.isdigit()):
|
||||||
"segments": [
|
return f"SPEAKER_{int(value):02d}"
|
||||||
{
|
return str(value)
|
||||||
"start": None,
|
|
||||||
"end": None,
|
|
||||||
"text": text,
|
def _speaker_segments(result):
|
||||||
"speaker": speaker or "SPEAKER_00",
|
"""Read native FunASR diarization output across supported versions."""
|
||||||
}
|
if not isinstance(result, dict):
|
||||||
] if text else [],
|
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(
|
async def run_asr(
|
||||||
@ -396,20 +439,22 @@ async def run_asr(
|
|||||||
result = model.generate(input=audio_path, hotword=hotword)
|
result = model.generate(input=audio_path, hotword=hotword)
|
||||||
|
|
||||||
if result and len(result) > 0:
|
if result and len(result) > 0:
|
||||||
text = result[0].get("text", "")
|
result_data = result[0]
|
||||||
|
text = result_data.get("text", "")
|
||||||
else:
|
else:
|
||||||
|
result_data = {}
|
||||||
text = ""
|
text = ""
|
||||||
|
|
||||||
text = cn_to_arabic(text)
|
text = cn_to_arabic(text)
|
||||||
|
if isinstance(result_data, dict):
|
||||||
|
result_data = dict(result_data)
|
||||||
|
result_data["text"] = text
|
||||||
logger.info(f"识别结果: {text}")
|
logger.info(f"识别结果: {text}")
|
||||||
|
|
||||||
if response_format == "text":
|
if response_format == "text":
|
||||||
return PlainTextResponse(content=text)
|
return PlainTextResponse(content=text)
|
||||||
|
|
||||||
return {
|
return {**result_data, "text": text, "language": language}
|
||||||
"text": text,
|
|
||||||
"language": language,
|
|
||||||
}
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"识别出错: {e}")
|
logger.error(f"识别出错: {e}")
|
||||||
@ -517,6 +562,10 @@ async def transcribe_speakr_compatible(
|
|||||||
language: Optional[str] = Form("zh"),
|
language: Optional[str] = Form("zh"),
|
||||||
database_name: Optional[str] = Form("NB"),
|
database_name: Optional[str] = Form("NB"),
|
||||||
collection_name: Optional[str] = Form("voice"),
|
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 兼容接口:
|
Speakr 兼容接口:
|
||||||
@ -560,11 +609,21 @@ async def transcribe_speakr_compatible(
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"声纹匹配失败,使用默认说话人: {e}")
|
logger.warning(f"声纹匹配失败,使用默认说话人: {e}")
|
||||||
|
|
||||||
return make_speakr_response(
|
response = make_speakr_response(
|
||||||
text=result.get("text", "") if isinstance(result, dict) else "",
|
result=result,
|
||||||
language=result.get("language", language or "zh") if isinstance(result, dict) else language or "zh",
|
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:
|
finally:
|
||||||
if tmp_path and os.path.exists(tmp_path):
|
if tmp_path and os.path.exists(tmp_path):
|
||||||
|
|||||||
@ -3,7 +3,8 @@
|
|||||||
# Speakr / ASR
|
# Speakr / ASR
|
||||||
USE_ASR_ENDPOINT=true
|
USE_ASR_ENDPOINT=true
|
||||||
ASR_BASE_URL=http://asr-test:59805
|
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_MIN_SPEAKERS=1
|
||||||
ASR_MAX_SPEAKERS=10
|
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:
|
with open(current_filepath, 'rb') as audio_file:
|
||||||
url = f"{ASR_BASE_URL}/asr"
|
url = f"{ASR_BASE_URL}/asr"
|
||||||
params = {
|
requested_speaker_count = 0
|
||||||
'batch_size_s':300,
|
if min_speakers and max_speakers and min_speakers == max_speakers:
|
||||||
'spk_diarization':True,
|
requested_speaker_count = min_speakers
|
||||||
'spk_num':0,
|
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_threshold':0.6,
|
||||||
'spk_cluster_method':'ahc',
|
'spk_cluster_method':'ahc',
|
||||||
'spk_smooth_window':3,
|
'spk_smooth_window':3,
|
||||||
@ -334,8 +339,9 @@ def transcribe_audio_asr(app_context, recording_id, filepath, original_filename,
|
|||||||
segments_without_speakers = 0
|
segments_without_speakers = 0
|
||||||
|
|
||||||
for segment in asr_response_data['segments']:
|
for segment in asr_response_data['segments']:
|
||||||
if 'speaker' in segment and segment['speaker'] is not None:
|
speaker = segment.get('speaker', segment.get('spk', segment.get('spk_name')))
|
||||||
all_speakers.add(segment['speaker'])
|
if speaker is not None:
|
||||||
|
all_speakers.add(speaker)
|
||||||
segments_with_speakers += 1
|
segments_with_speakers += 1
|
||||||
else:
|
else:
|
||||||
segments_without_speakers += 1
|
segments_without_speakers += 1
|
||||||
@ -355,7 +361,9 @@ def transcribe_audio_asr(app_context, recording_id, filepath, original_filename,
|
|||||||
last_known_speaker = None
|
last_known_speaker = None
|
||||||
|
|
||||||
for i, segment in enumerate(asr_response_data['segments']):
|
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()
|
text = segment.get('text', '').strip()
|
||||||
|
|
||||||
# 如果片段没有说话人,使用前一个片段的说话人
|
# 如果片段没有说话人,使用前一个片段的说话人
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user