test/app/services/asr/engines/base.py

786 lines
31 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

# -*- coding: utf-8 -*-
"""
ASR引擎基础模块
包含抽象基类和数据类定义
"""
import time
import logging
from typing import Optional, Dict, List, Any, Callable
from abc import ABC, abstractmethod
from dataclasses import dataclass
from app.core.config import settings
from app.core.exceptions import DefaultServerErrorException
from app.core.text_cleanup import deduplicate_asr_text
from app.core.text_cleanup import trim_segment_boundary_overlap
from app.core.hotword_resolver import apply_hotword_rules
from app.core.logging import log_inference_metrics
from app.utils.audio import get_audio_duration
logger = logging.getLogger(__name__)
def _log_asr_stage_timing(
stage: str,
duration_ms: float,
*,
task_id: Optional[str] = None,
model_id: Optional[str] = None,
audio_duration_sec: Optional[float] = None,
**extra: Any,
) -> None:
"""Record structured timing for one ASR processing stage."""
payload: Dict[str, Any] = {
"event": "asr_stage_timing",
"stage": stage,
"duration_ms": round(duration_ms, 2),
}
if task_id:
payload["task_id"] = task_id
if model_id:
payload["model_id"] = model_id
if audio_duration_sec is not None:
payload["audio_duration_sec"] = round(audio_duration_sec, 2)
if audio_duration_sec > 0:
payload["rtf"] = round((duration_ms / 1000) / audio_duration_sec, 4)
payload.update(extra)
logger.info("ASR阶段耗时", extra=payload)
@dataclass
class WordToken:
"""字词级时间戳信息"""
text: str # 字词文本
start_time: float # 开始时间(秒)
end_time: float # 结束时间(秒)
@dataclass
class ASRSegmentResult:
"""ASR 分段识别结果"""
text: str # 该段识别文本
start_time: float # 开始时间(秒)
end_time: float # 结束时间(秒)
speaker_id: Optional[str] = None # 说话人ID(多说话人模式)
speaker_name: Optional[str] = None # 已注册声纹命中后的显示名称
user_id: Optional[str] = None # 已注册声纹命中后的业务用户编号
speaker_embedding: Optional[Any] = None # 片段声纹向量,仅供服务端匹配使用
word_tokens: Optional[List[WordToken]] = None # 字词级时间戳(可选)
@dataclass
class ASRFullResult:
"""ASR 完整识别结果(支持长音频)"""
text: str # 完整识别文本
segments: List[ASRSegmentResult] # 分段结果
duration: float # 音频总时长(秒)
@dataclass
class ASRRawResult:
"""ASR 原始识别结果(包含时间戳)"""
text: str # 完整识别文本
segments: List[ASRSegmentResult] # 分段结果(从 VAD 时间戳解析)
class BaseASREngine(ABC):
"""基础ASR引擎抽象基类"""
@abstractmethod
def transcribe_file(
self,
audio_path: str,
hotwords: str = "",
enable_punctuation: bool = False,
enable_itn: bool = False,
enable_vad: bool = False,
sample_rate: int = 16000,
) -> str:
"""转录音频文件"""
pass
@abstractmethod
def transcribe_file_with_vad(
self,
audio_path: str,
hotwords: str = "",
enable_punctuation: bool = True,
enable_itn: bool = True,
sample_rate: int = 16000,
**kwargs,
) -> ASRRawResult:
"""使用 VAD 转录音频文件,返回带时间戳分段的结果
Args:
audio_path: 音频文件路径
hotwords: 热词/上下文提示
enable_punctuation: 是否启用标点
enable_itn: 是否启用 ITN
sample_rate: 采样率
**kwargs: 额外参数(如 word_timestamps 字词级时间戳)
Returns:
ASRRawResult 包含文本和分段信息
"""
pass
def transcribe_long_audio(
self,
audio_path: str,
hotwords: str = "",
enable_punctuation: bool = False,
enable_itn: bool = False,
sample_rate: int = 16000,
enable_speaker_diarization: bool = True,
enable_speaker_identification: bool = True,
enable_text_cleanup: bool = True,
word_timestamps: bool = False,
timestamp_scale: float = 1.0,
task_id: Optional[str] = None,
progress_callback: Optional[Callable[[str, str, int, Optional[dict[str, Any]]], None]] = None,
) -> ASRFullResult:
"""转录长音频文件(自动分段)
Args:
audio_path: 音频文件路径
hotwords: 热词
enable_punctuation: 是否启用标点
enable_itn: 是否启用 ITN
sample_rate: 采样率
enable_speaker_diarization: 是否启用说话人分离
enable_speaker_identification: 是否提取声纹向量用于匹配已注册声纹库
enable_text_cleanup: 是否启用文本去重和口头语清理
word_timestamps: 是否返回字词级时间戳(仅部分模型支持)
timestamp_scale: Timestamp correction factor from audio normalization.
task_id: 任务ID(用于日志追踪)
Returns:
ASRFullResult: 包含完整文本、分段结果和时长的结果
"""
from app.utils.audio_splitter import AudioSplitter
# 开始性能计时
start_time = time.time()
start_perf = time.perf_counter()
model_id = getattr(self, 'model_id', 'unknown')
stage_timings_ms: Dict[str, float] = {
"duration_probe_ms": 0.0,
"speaker_diarization_ms": 0.0,
"vad_audio_split_ms": 0.0,
"asr_inference_ms": 0.0,
"speaker_embedding_ms": 0.0,
"cleanup_ms": 0.0,
"text_cleanup_ms": 0.0,
"hotword_resolver_ms": 0.0,
}
task_prefix = f"[{task_id}] " if task_id else ""
def emit_progress(
stage: str,
message: str,
percentage: int,
detail: Optional[dict[str, Any]] = None,
) -> None:
if progress_callback is None:
return
try:
progress_callback(stage, message, percentage, detail)
except Exception as exc:
logger.warning("%s进度回调失败: %s", task_prefix, exc)
logger.info(
f"{task_prefix}[transcribe_long_audio] 音频: {audio_path}, "
f"speaker_diarization={enable_speaker_diarization}, "
f"speaker_identification={enable_speaker_identification}, "
f"text_cleanup={enable_text_cleanup}, "
f"word_level={word_timestamps}"
)
try:
# 获取音频时长
duration_started = time.perf_counter()
duration = get_audio_duration(audio_path)
stage_timings_ms["duration_probe_ms"] = (
time.perf_counter() - duration_started
) * 1000
_log_asr_stage_timing(
"duration_probe",
stage_timings_ms["duration_probe_ms"],
task_id=task_id,
model_id=model_id,
audio_duration_sec=duration,
audio_path=audio_path,
)
logger.info(f"{task_prefix}[transcribe_long_audio] 音频时长: {duration:.2f}秒")
emit_progress(
"analyzing",
f"音频时长 {duration:.1f} 秒,正在进行分割。",
15,
{"audio_duration_seconds": round(duration, 2)},
)
# 统一使用分段处理
speaker_segments = None
audio_segments = None
if enable_speaker_diarization:
# 多说话人:使用说话人分离
from app.utils.speaker_diarizer import SpeakerDiarizer
logger.info(f"{task_prefix}使用说话人分离模式")
emit_progress(
"diarizing",
"正在执行说话人分离。",
18,
{"audio_duration_seconds": round(duration, 2)},
)
diarizer = SpeakerDiarizer()
diarization_started = time.perf_counter()
speaker_segments = diarizer.split_audio_by_speakers(audio_path)
stage_timings_ms["speaker_diarization_ms"] = (
time.perf_counter() - diarization_started
) * 1000
speaker_audio_sec = sum(
float(getattr(seg, "duration_sec", 0.0))
for seg in (speaker_segments or [])
)
_log_asr_stage_timing(
"speaker_diarization",
stage_timings_ms["speaker_diarization_ms"],
task_id=task_id,
model_id=model_id,
audio_duration_sec=duration,
segment_count=len(speaker_segments or []),
segmented_audio_sec=round(speaker_audio_sec, 2),
)
if not speaker_segments:
logger.warning(f"{task_prefix}说话人分离未检测到片段,fallback 到 VAD 分割")
if not speaker_segments:
# 单说话人:使用 VAD 分割
logger.info(f"{task_prefix}使用 VAD 分割模式")
emit_progress(
"splitting",
"正在执行 VAD 音频分割。",
20,
{"audio_duration_seconds": round(duration, 2)},
)
splitter = AudioSplitter(device=self.device)
split_started = time.perf_counter()
audio_segments = splitter.split_audio_file(audio_path)
stage_timings_ms["vad_audio_split_ms"] = (
time.perf_counter() - split_started
) * 1000
split_audio_sec = sum(
float(getattr(seg, "duration_sec", 0.0))
for seg in (audio_segments or [])
)
_log_asr_stage_timing(
"vad_audio_split",
stage_timings_ms["vad_audio_split_ms"],
task_id=task_id,
model_id=model_id,
audio_duration_sec=duration,
segment_count=len(audio_segments or []),
segmented_audio_sec=round(split_audio_sec, 2),
)
# 选择要处理的片段
segments_to_process = speaker_segments if speaker_segments else audio_segments
if not segments_to_process:
raise DefaultServerErrorException("音频分割失败:未生成任何片段")
logger.info(f"{task_prefix}音频已分割为 {len(segments_to_process)} 段")
total_segments = len(segments_to_process)
emit_progress(
"transcribing",
"开始批量识别。",
30,
{
"segment_total": total_segments,
"segment_completed": 0,
"audio_duration_seconds": round(duration, 2),
},
)
results: List[ASRSegmentResult] = []
# 使用批处理推理
batch_size = settings.ASR_BATCH_SIZE
total_batches = (len(segments_to_process) + batch_size - 1) // batch_size
logger.info(
f"{task_prefix}使用批处理推理,batch_size={batch_size}, "
f"word_timestamps={word_timestamps}"
)
for batch_start in range(0, len(segments_to_process), batch_size):
batch_end = min(batch_start + batch_size, len(segments_to_process))
batch_segments = segments_to_process[batch_start:batch_end]
batch_index = batch_start // batch_size + 1
logger.info(
f"{task_prefix}推理批次 "
f"{batch_index}/{total_batches}: "
f"片段 {batch_start+1}-{batch_end}/{len(segments_to_process)}"
)
emit_progress(
"transcribing",
f"识别中:{batch_index}/{total_batches} 批,{batch_start}/{total_segments} 段",
min(88, 30 + int((batch_start / max(total_segments, 1)) * 58)),
{
"segment_total": total_segments,
"segment_completed": batch_start,
"batch_index": batch_index,
"batch_total": total_batches,
"audio_duration_seconds": round(duration, 2),
},
)
try:
# 批量推理,支持时间戳
batch_started = time.perf_counter()
batch_results = self._transcribe_batch(
segments=batch_segments,
hotwords=hotwords,
enable_punctuation=enable_punctuation,
enable_itn=enable_itn,
sample_rate=sample_rate,
word_timestamps=word_timestamps,
)
batch_inference_ms = (time.perf_counter() - batch_started) * 1000
stage_timings_ms["asr_inference_ms"] += batch_inference_ms
batch_audio_sec = sum(
float(getattr(seg, "duration_sec", 0.0))
for seg in batch_segments
)
valid_batch_results = 0
batch_embedding_ms = 0.0
for seg, result in zip(batch_segments, batch_results):
if result and result.text:
valid_batch_results += 1
start_sec = float(getattr(seg, "start_sec", 0.0))
end_sec = float(getattr(seg, "end_sec", start_sec))
speaker_embedding = None
if (
speaker_segments
and settings.SPEAKER_DB_ENABLED
and enable_speaker_identification
):
try:
from app.core.database import pg_speaker_db
from app.services.speaker_registry import (
get_speaker_registry_service,
)
audio_data = getattr(seg, "audio_data", None)
if pg_speaker_db.is_connected and audio_data is not None:
embedding_started = time.perf_counter()
speaker_embedding = (
get_speaker_registry_service()
.extract_embedding_from_audio(audio_data)
)
batch_embedding_ms += (
time.perf_counter() - embedding_started
) * 1000
except Exception as exc:
logger.warning(
"%s提取片段声纹向量失败,保留原说话人编号: %s",
task_prefix,
exc,
)
results.append(
ASRSegmentResult(
text=result.text,
start_time=start_sec,
end_time=end_sec,
speaker_id=getattr(seg, "speaker_id", None),
speaker_embedding=speaker_embedding,
word_tokens=result.word_tokens if word_timestamps else None,
)
)
stage_timings_ms["speaker_embedding_ms"] += batch_embedding_ms
_log_asr_stage_timing(
"asr_batch",
batch_inference_ms,
task_id=task_id,
model_id=model_id,
audio_duration_sec=batch_audio_sec,
batch_index=batch_index,
batch_total=total_batches,
segment_start=batch_start + 1,
segment_end=batch_end,
segment_total=total_segments,
batch_segment_count=len(batch_segments),
valid_segment_count=valid_batch_results,
speaker_embedding_ms=round(batch_embedding_ms, 2),
)
logger.info(
f"{task_prefix}批次推理完成,有效片段: "
f"{len([r for r in batch_results if r and r.text])}"
)
emit_progress(
"transcribing",
f"识别中:{batch_index}/{total_batches} 批,{batch_end}/{total_segments} 段",
min(90, 30 + int((batch_end / max(total_segments, 1)) * 60)),
{
"segment_total": total_segments,
"segment_completed": batch_end,
"batch_index": batch_index,
"batch_total": total_batches,
"audio_duration_seconds": round(duration, 2),
},
)
except Exception as e:
logger.error(f"{task_prefix}批次推理失败: {e}, 跳过该批次")
emit_progress(
"transcribing",
f"第 {batch_index}/{total_batches} 批识别失败,已跳过。",
min(90, 30 + int((batch_end / max(total_segments, 1)) * 60)),
{
"segment_total": total_segments,
"segment_completed": batch_end,
"batch_index": batch_index,
"batch_total": total_batches,
"error": str(e),
},
)
# 清理临时文件(独立清理,避免条件遗漏)
try:
cleanup_started = time.perf_counter()
if speaker_segments:
from app.utils.speaker_diarizer import SpeakerDiarizer
SpeakerDiarizer.cleanup_segments(speaker_segments)
if audio_segments:
AudioSplitter.cleanup_segments(audio_segments)
stage_timings_ms["cleanup_ms"] = (
time.perf_counter() - cleanup_started
) * 1000
_log_asr_stage_timing(
"temp_cleanup",
stage_timings_ms["cleanup_ms"],
task_id=task_id,
model_id=model_id,
segment_count=len(segments_to_process),
)
except Exception as e:
logger.warning(f"清理临时文件时出错: {e}")
overlap_trimmed_count = 0
if enable_text_cleanup:
text_cleanup_started = time.perf_counter()
cleaned_results: List[ASRSegmentResult] = []
previous_text = ""
for seg in results:
cleaned_text = deduplicate_asr_text(seg.text)
cleaned_text, was_trimmed = trim_segment_boundary_overlap(
previous_text,
cleaned_text,
)
if was_trimmed:
overlap_trimmed_count += 1
if not cleaned_text:
continue
seg.text = cleaned_text
cleaned_results.append(seg)
previous_text = cleaned_text
if len(cleaned_results) != len(results) or overlap_trimmed_count:
logger.info(
"%s文字去重完成:%s -> %s 段,边界重叠裁剪 %s 次",
task_prefix,
len(results),
len(cleaned_results),
overlap_trimmed_count,
)
results = cleaned_results
stage_timings_ms["text_cleanup_ms"] = (
time.perf_counter() - text_cleanup_started
) * 1000
_log_asr_stage_timing(
"text_cleanup",
stage_timings_ms["text_cleanup_ms"],
task_id=task_id,
model_id=model_id,
segment_count=len(results),
overlap_trimmed_count=overlap_trimmed_count,
cleanup_enabled=enable_text_cleanup,
)
hotword_resolver_started = time.perf_counter()
hotword_resolver_applied_count = 0
matched_hotwords: list[str] = []
if hotwords:
for seg in results:
resolved_text, was_applied, matched = apply_hotword_rules(
seg.text,
hotwords,
)
if resolved_text:
seg.text = resolved_text
if was_applied:
hotword_resolver_applied_count += 1
for hotword_text in matched:
if hotword_text not in matched_hotwords:
matched_hotwords.append(hotword_text)
if hotwords:
stage_timings_ms["hotword_resolver_ms"] = (
time.perf_counter() - hotword_resolver_started
) * 1000
_log_asr_stage_timing(
"hotword_resolver",
stage_timings_ms["hotword_resolver_ms"],
task_id=task_id,
model_id=model_id,
segment_count=len(results),
applied_segment_count=hotword_resolver_applied_count,
matched_hotwords=matched_hotwords,
)
if hotword_resolver_applied_count:
logger.info(
"%s热词纠偏完成:命中 %s 段,热词=%s",
task_prefix,
hotword_resolver_applied_count,
matched_hotwords,
)
all_texts = [seg.text for seg in results]
full_text = "\n".join(all_texts)
emit_progress(
"finalizing",
"分段识别完成,正在整理文本。",
92,
{
"segment_total": len(segments_to_process),
"segment_completed": len(segments_to_process),
"valid_segment_count": len(results),
"text_cleanup_enabled": enable_text_cleanup,
"overlap_trimmed_count": overlap_trimmed_count,
"audio_duration_seconds": round(duration, 2),
},
)
logger.info(
f"长音频识别完成,共 {len(results)} 个有效分段,"
f"总字符数: {len(full_text)}"
)
# 计算性能指标
total_duration_ms = (time.time() - start_time) * 1000
total_perf_ms = (time.perf_counter() - start_perf) * 1000
if timestamp_scale != 1.0:
for seg in results:
seg.start_time *= timestamp_scale
seg.end_time *= timestamp_scale
if seg.word_tokens:
for word_token in seg.word_tokens:
word_token.start_time *= timestamp_scale
word_token.end_time *= timestamp_scale
duration *= timestamp_scale
logger.info(
f"{task_prefix}Timestamp scaling applied: scale={timestamp_scale:.6f}"
)
_log_asr_stage_timing(
"asr_total",
total_perf_ms,
task_id=task_id,
model_id=model_id,
audio_duration_sec=duration,
segment_count=len(results),
total_segment_count=len(segments_to_process),
**{key: round(value, 2) for key, value in stage_timings_ms.items()},
)
log_inference_metrics(
logger=logger,
message="长音频识别完成",
task_id=task_id,
duration_ms=total_duration_ms,
audio_duration_sec=duration,
model_id=model_id,
status="success",
segments_count=len(results),
batch_size=settings.ASR_BATCH_SIZE,
event="asr_inference_metrics",
**{key: round(value, 2) for key, value in stage_timings_ms.items()},
enable_speaker_diarization=enable_speaker_diarization,
enable_speaker_identification=enable_speaker_identification,
enable_text_cleanup=enable_text_cleanup,
word_timestamps=word_timestamps,
)
return ASRFullResult(
text=full_text,
segments=results,
duration=duration,
)
except Exception as e:
# 计算失败时的性能指标
total_duration_ms = (time.time() - start_time) * 1000
try:
duration = get_audio_duration(audio_path)
except Exception:
duration = 0
log_inference_metrics(
logger=logger,
message="长音频识别失败",
task_id=task_id,
duration_ms=total_duration_ms,
audio_duration_sec=duration,
model_id=model_id,
status="error",
error=str(e),
)
logger.error(f"长音频识别失败: {e}")
raise DefaultServerErrorException(f"长音频识别失败: {str(e)}")
@abstractmethod
def is_model_loaded(self) -> bool:
"""检查模型是否已加载"""
pass
@property
@abstractmethod
def device(self) -> str:
"""获取设备信息"""
pass
@property
@abstractmethod
def supports_realtime(self) -> bool:
"""是否支持实时识别"""
pass
def _transcribe_batch(
self,
segments: List[Any],
hotwords: str = "",
enable_punctuation: bool = False,
enable_itn: bool = False,
sample_rate: int = 16000,
word_timestamps: bool = False,
) -> List[ASRSegmentResult]:
"""批量推理多个音频片段
Args:
segments: 音频片段列表(每个片段需要有 temp_file 属性)
hotwords: 热词
enable_punctuation: 是否启用标点
enable_itn: 是否启用 ITN
sample_rate: 采样率
word_timestamps: 是否返回字词级时间戳
Returns:
ASRSegmentResult 列表,与输入片段一一对应
"""
# 默认实现:逐个推理(子类可以重写实现真正的批处理)
results = []
for idx, seg in enumerate(segments):
try:
if not seg.temp_file:
logger.warning(f"批处理片段 {idx + 1} 临时文件不存在,跳过")
results.append(ASRSegmentResult(text="", start_time=0.0, end_time=0.0))
continue
if word_timestamps:
# 需要时间戳:使用 transcribe_file_with_vad
raw_result = self.transcribe_file_with_vad(
audio_path=seg.temp_file,
hotwords=hotwords,
enable_punctuation=enable_punctuation,
enable_itn=enable_itn,
sample_rate=sample_rate,
word_timestamps=True,
)
if raw_result.segments:
result_seg = raw_result.segments[0]
results.append(
ASRSegmentResult(
text=result_seg.text,
start_time=seg.start_sec,
end_time=seg.end_sec,
speaker_id=getattr(seg, 'speaker_id', None),
word_tokens=result_seg.word_tokens,
)
)
else:
results.append(
ASRSegmentResult(
text=raw_result.text,
start_time=seg.start_sec,
end_time=seg.end_sec,
speaker_id=getattr(seg, 'speaker_id', None),
)
)
else:
# 不需要时间戳:使用 transcribe_file
text = self.transcribe_file(
audio_path=seg.temp_file,
hotwords=hotwords,
enable_punctuation=enable_punctuation,
enable_itn=enable_itn,
enable_vad=False,
sample_rate=sample_rate,
)
results.append(
ASRSegmentResult(
text=text or "",
start_time=seg.start_sec,
end_time=seg.end_sec,
speaker_id=getattr(seg, 'speaker_id', None),
)
)
except Exception as e:
logger.error(f"批处理片段 {idx + 1} 推理失败: {e}")
results.append(
ASRSegmentResult(
text="",
start_time=getattr(seg, 'start_sec', 0.0),
end_time=getattr(seg, 'end_sec', 0.0),
speaker_id=getattr(seg, 'speaker_id', None),
)
)
return results
@staticmethod
def _detect_device(device: str = "auto") -> str:
"""检测可用设备"""
from app.core.device import detect_device
return detect_device(device)
class RealTimeASREngine(BaseASREngine):
"""实时ASR引擎抽象基类"""
@property
def supports_realtime(self) -> bool:
"""支持实时识别"""
return True
@abstractmethod
def transcribe_websocket(
self,
audio_chunk: bytes,
cache: Optional[Dict] = None,
is_final: bool = False,
**kwargs,
) -> str:
"""WebSocket流式语音识别"""
pass