786 lines
31 KiB
Python
786 lines
31 KiB
Python
# -*- 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
|