- scheduler: ffmpeg 异步线程 + GPU 串行调度 + 模型复用(2N→2 次加载) - pipeline: 阶段拆分(extract/asr/translate),中间数据存 Task 字段 - translate_service: 长度排序批处理,padding 浪费减少 91% - model_manager: ASR/翻译不共驻,BatchedInferencePipeline 批量解码 - 日志分级: INFO=任务流转里程碑,DEBUG=进度详情;默认 INFO - 前端: 日志最新在上+滚动感知+退避轮询;24h 时间;上传中状态显示 - /health: 返回完整 Whisper/NLLB 配置 - upload_service: 单事务 complete + 扩展名白名单 - task_router: 合并 UploadSession 虚拟任务到列表 - Dockerfile: CPU/GPU 独立构建链,deps 缓存稳定 - prefetch_models: 安装时预下载模型权重
193 lines
7.2 KiB
Python
193 lines
7.2 KiB
Python
"""转写管线各阶段:提取音频 → ASR → 断句 → 翻译 → 写 SRT。
|
||
|
||
阶段拆分供 scheduler 调度:ffmpeg 阶段独立线程(CPU),GPU 阶段(ASR+翻译)
|
||
由 scheduler 串行化并在切换模型前查队列复用已加载模型。
|
||
|
||
每个阶段函数接收 db Session + Task,更新状态/进度,写入中间产物:
|
||
extract_phase: queued → extracting → (写 wav_path, status 置 transcribing)
|
||
asr_phase: transcribing → segmenting → (写 segments_json, status 置 translating)
|
||
translate_phase: translating → done (写 SRT)
|
||
|
||
阶段间传递的中间数据存在 Task.wav_path / Task.segments_json,避免跨线程传对象。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import logging
|
||
import traceback
|
||
from dataclasses import asdict
|
||
from datetime import datetime, timezone
|
||
from pathlib import Path
|
||
|
||
from ..config import get_settings
|
||
from ..models.task import (
|
||
STATUS_EXTRACTING, STATUS_TRANSCRIBING, STATUS_SEGMENTING,
|
||
STATUS_TRANSLATING, STATUS_DONE, STATUS_FAILED,
|
||
)
|
||
from . import asr_service, ffmpeg_service, segmenter, srt_writer, translate_service
|
||
from .types import Subtitle
|
||
|
||
logger = logging.getLogger("audio2text.pipeline")
|
||
|
||
# 进度锚点(各阶段在 0-100 中的占比)
|
||
P_EXTRACT = 5.0
|
||
P_TRANSCRIBE_START = 5.0
|
||
P_TRANSCRIBE_END = 55.0
|
||
P_SEGMENT_START = 55.0
|
||
P_SEGMENT_END = 60.0
|
||
P_TRANSLATE_START = 60.0
|
||
P_TRANSLATE_END = 98.0
|
||
P_DONE = 100.0
|
||
|
||
|
||
# ---------------- 阶段 1:提取音频(CPU,可并行)----------------
|
||
|
||
def extract_phase(db, task) -> None:
|
||
"""ffmpeg 提取 16k mono wav,写 task.wav_path,状态置 transcribing。
|
||
|
||
由 scheduler 在独立线程调用(与 GPU 阶段并行)。提取完即可让 GPU 调度线程接管。
|
||
"""
|
||
s = get_settings()
|
||
src = s.upload_dir() / task.source_path
|
||
logger.info("任务 %d [音频提取开始] %s", task.id, src.name)
|
||
_set_status(db, task, STATUS_EXTRACTING, P_EXTRACT)
|
||
wav = s.work_dir() / f"task_{task.id}.wav"
|
||
ffmpeg_service.extract_audio(src, wav)
|
||
|
||
# 按配置决定是否删原始视频(提取成功后)
|
||
if s.processing.delete_original_after_extract and src.is_file():
|
||
try:
|
||
src.unlink()
|
||
logger.info("已删除原始视频 %s(delete_original_after_extract=true)。", src.name)
|
||
except OSError as exc:
|
||
logger.warning("删除原始视频失败 %s: %s", src, exc)
|
||
|
||
# 记录 wav 路径,状态置 transcribing(待 GPU 调度线程接管 ASR)
|
||
task.wav_path = str(wav)
|
||
task.status = STATUS_TRANSCRIBING
|
||
task.progress = P_TRANSCRIBE_START
|
||
task.updated_at = datetime.now(timezone.utc)
|
||
db.commit()
|
||
logger.info("任务 %d [音频提取完成] → 待 ASR:%s", task.id, wav.name)
|
||
|
||
|
||
# ---------------- 阶段 2:ASR + 断句(GPU)----------------
|
||
|
||
def asr_phase(db, task) -> None:
|
||
"""加载 Whisper 转写 + 断句,写 task.segments_json,状态置 translating。
|
||
|
||
由 GPU 调度线程调用。model_manager 保证 ASR 与翻译器不共驻。
|
||
"""
|
||
wav_path = Path(task.wav_path) if task.wav_path else None
|
||
if wav_path is None or not wav_path.is_file():
|
||
raise FileNotFoundError(f"音频不存在:{wav_path}(task {task.id})")
|
||
|
||
logger.info("任务 %d [ASR 开始] %s", task.id, wav_path.name)
|
||
_set_status(db, task, STATUS_TRANSCRIBING, P_TRANSCRIBE_START)
|
||
segments = asr_service.transcribe(wav_path)
|
||
_set_status(db, task, STATUS_TRANSCRIBING, P_TRANSCRIBE_END,
|
||
note=f"识别出 {len(segments)} 段")
|
||
|
||
_set_status(db, task, STATUS_SEGMENTING, P_SEGMENT_START)
|
||
subs = segmenter.resegment(segments)
|
||
_set_status(db, task, STATUS_SEGMENTING, P_SEGMENT_END,
|
||
note=f"重组为 {len(subs)} 条字幕")
|
||
|
||
# 序列化断句结果供翻译阶段用(dataclass → JSON)
|
||
task.segments_json = json.dumps([asdict(s) for s in subs], ensure_ascii=False)
|
||
task.status = STATUS_TRANSLATING
|
||
task.progress = P_TRANSLATE_START
|
||
task.updated_at = datetime.now(timezone.utc)
|
||
db.commit()
|
||
logger.info("任务 %d [ASR 完成] → 待翻译:%d 条字幕", task.id, len(subs))
|
||
|
||
|
||
# ---------------- 阶段 3:翻译 + 写 SRT(GPU)----------------
|
||
|
||
def translate_phase(db, task) -> None:
|
||
"""加载 NLLB 翻译 + 写 SRT,状态置 done。由 GPU 调度线程调用。"""
|
||
if not task.segments_json:
|
||
raise ValueError(f"任务 {task.id} 无 segments_json,无法翻译")
|
||
|
||
logger.info("任务 %d [翻译开始] %d 条字幕", task.id, len(json.loads(task.segments_json)))
|
||
_set_status(db, task, STATUS_TRANSLATING, P_TRANSLATE_START)
|
||
subs = [_dict_to_subtitle(d) for d in json.loads(task.segments_json)]
|
||
|
||
zh_texts = translate_service.translate(subs)
|
||
_set_status(db, task, STATUS_TRANSLATING, P_TRANSLATE_END,
|
||
note=f"翻译 {len(zh_texts)} 条")
|
||
|
||
# 写 SRT
|
||
out_dir = s_output_dir()
|
||
stem = Path(task.filename).stem
|
||
en_path = out_dir / f"task_{task.id}/{stem}.en.srt"
|
||
zh_path = out_dir / f"task_{task.id}/{stem}.zh.srt"
|
||
bi_path = out_dir / f"task_{task.id}/{stem}.srt"
|
||
|
||
srt_writer.write_srt(subs, en_path)
|
||
zh_subs = [Subtitle(text=zh, start=sub.start, end=sub.end)
|
||
for zh, sub in zip(zh_texts, subs)]
|
||
srt_writer.write_srt(zh_subs, zh_path)
|
||
srt_writer.write_bilingual_srt(subs, zh_texts, bi_path)
|
||
|
||
task.en_srt_path = str(en_path.relative_to(out_dir))
|
||
task.zh_srt_path = str(zh_path.relative_to(out_dir))
|
||
task.bilingual_srt_path = str(bi_path.relative_to(out_dir))
|
||
task.status = STATUS_DONE
|
||
task.progress = P_DONE
|
||
task.updated_at = datetime.now(timezone.utc)
|
||
db.commit()
|
||
|
||
# 清理中间音频
|
||
s = get_settings()
|
||
if not s.processing.keep_audio and task.wav_path:
|
||
try:
|
||
Path(task.wav_path).unlink()
|
||
except OSError:
|
||
pass
|
||
|
||
logger.info("任务 %d [完成] %s", task.id, bi_path.name)
|
||
|
||
|
||
# ---------------- 工具函数 ----------------
|
||
|
||
def s_output_dir() -> Path:
|
||
return get_settings().output_dir()
|
||
|
||
|
||
def _dict_to_subtitle(d: dict) -> Subtitle:
|
||
return Subtitle(text=d["text"], start=d["start"], end=d["end"])
|
||
|
||
|
||
def _set_status(db, task, status: str, progress: float, note: str = "") -> None:
|
||
"""更新任务状态/进度并落库。
|
||
|
||
日志级别策略:
|
||
- 带 note 的(如“识别出 N 段”)是阶段内里程碑 → INFO
|
||
- 仅进度百分比更新(同状态同阶段)→ DEBUG,避免 INFO 被进度刷屏
|
||
"""
|
||
task.status = status
|
||
task.progress = progress
|
||
task.updated_at = datetime.now(timezone.utc)
|
||
db.commit()
|
||
if note:
|
||
logger.info("任务 %d [%s %.0f%%] %s", task.id, status, progress, note)
|
||
else:
|
||
logger.debug("任务 %d [%s %.0f%%]", task.id, status, progress)
|
||
|
||
|
||
def mark_failed(db, task_id: int, error: str) -> None:
|
||
"""标记任务失败(对外公开,供 scheduler 调用)。"""
|
||
from ..models.task import Task
|
||
try:
|
||
task = db.get(Task, task_id)
|
||
if task is None:
|
||
return
|
||
task.status = STATUS_FAILED
|
||
task.error = error[:2000]
|
||
task.updated_at = datetime.now(timezone.utc)
|
||
db.commit()
|
||
except Exception: # pragma: no cover
|
||
logger.error("写入失败状态时又失败:\n%s", traceback.format_exc())
|