# -*- 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