feat: 调度器+并发管线+GPU优化+日志分级+前端修复
- 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: 安装时预下载模型权重
This commit is contained in:
@@ -1,25 +1,31 @@
|
||||
"""转写管线编排:提取音频 → ASR → 断句 → 翻译 → 写 SRT。
|
||||
"""转写管线各阶段:提取音频 → ASR → 断句 → 翻译 → 写 SRT。
|
||||
|
||||
任务状态机:
|
||||
queued → extracting → transcribing → segmenting → translating → done
|
||||
任一步失败 → failed
|
||||
阶段拆分供 scheduler 调度:ffmpeg 阶段独立线程(CPU),GPU 阶段(ASR+翻译)
|
||||
由 scheduler 串行化并在切换模型前查队列复用已加载模型。
|
||||
|
||||
模型不共驻:ASR 与翻译分阶段加载,翻译时先卸载 Whisper 释放显存跑大 batch。
|
||||
管线在后台线程跑(每个任务一个线程),通过 DB 更新状态与进度。
|
||||
每个阶段函数接收 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 threading
|
||||
import traceback
|
||||
from dataclasses import asdict
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
|
||||
from ..config import get_settings
|
||||
from ..database import get_session_local
|
||||
from ..models.task import Task
|
||||
from . import ffmpeg_service, asr_service, segmenter, translate_service, srt_writer
|
||||
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")
|
||||
@@ -35,35 +41,17 @@ P_TRANSLATE_END = 98.0
|
||||
P_DONE = 100.0
|
||||
|
||||
|
||||
def enqueue_task(task_id: int) -> None:
|
||||
"""把任务交给后台线程处理(非阻塞,供 upload_service.complete 调用)。"""
|
||||
t = threading.Thread(target=_run_task, args=(task_id,), daemon=True)
|
||||
t.start()
|
||||
logger.info("任务 %d 已入队(后台线程 %s)。", task_id, t.name)
|
||||
# ---------------- 阶段 1:提取音频(CPU,可并行)----------------
|
||||
|
||||
def extract_phase(db, task) -> None:
|
||||
"""ffmpeg 提取 16k mono wav,写 task.wav_path,状态置 transcribing。
|
||||
|
||||
def _run_task(task_id: int) -> None:
|
||||
"""后台执行完整管线。所有异常都被捕获并写入 task.error。"""
|
||||
db = get_session_local()()
|
||||
try:
|
||||
task = db.get(Task, task_id)
|
||||
if task is None:
|
||||
logger.error("任务 %d 不存在。", task_id)
|
||||
return
|
||||
_pipeline(db, task)
|
||||
except Exception as exc:
|
||||
logger.exception("任务 %d 失败:%s", task_id, exc)
|
||||
_mark_failed(db, task_id, str(exc))
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _pipeline(db, task: Task) -> None:
|
||||
由 scheduler 在独立线程调用(与 GPU 阶段并行)。提取完即可让 GPU 调度线程接管。
|
||||
"""
|
||||
s = get_settings()
|
||||
src = s.upload_dir() / task.source_path
|
||||
|
||||
# ---------- 1. 提取音频 ----------
|
||||
_set_status(db, task, "extracting", P_EXTRACT)
|
||||
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)
|
||||
|
||||
@@ -75,61 +63,110 @@ def _pipeline(db, task: Task) -> None:
|
||||
except OSError as exc:
|
||||
logger.warning("删除原始视频失败 %s: %s", src, exc)
|
||||
|
||||
# ---------- 2. 语音识别 ----------
|
||||
_set_status(db, task, "transcribing", P_TRANSCRIBE_START)
|
||||
segments = asr_service.transcribe(wav)
|
||||
_set_status(db, task, "transcribing", P_TRANSCRIBE_END,
|
||||
# 记录 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)} 段")
|
||||
|
||||
# ---------- 3. 断句 + 时间戳重算 ----------
|
||||
_set_status(db, task, "segmenting", P_SEGMENT_START)
|
||||
_set_status(db, task, STATUS_SEGMENTING, P_SEGMENT_START)
|
||||
subs = segmenter.resegment(segments)
|
||||
_set_status(db, task, "segmenting", P_SEGMENT_END,
|
||||
_set_status(db, task, STATUS_SEGMENTING, P_SEGMENT_END,
|
||||
note=f"重组为 {len(subs)} 条字幕")
|
||||
|
||||
# ---------- 4. 翻译 ----------
|
||||
_set_status(db, task, "translating", P_TRANSLATE_START)
|
||||
# 翻译阶段:model_manager 会自动卸载 ASR、加载翻译器(独占显存)
|
||||
# 序列化断句结果供翻译阶段用(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, "translating", P_TRANSLATE_END,
|
||||
_set_status(db, task, STATUS_TRANSLATING, P_TRANSLATE_END,
|
||||
note=f"翻译 {len(zh_texts)} 条")
|
||||
|
||||
# ---------- 5. 写 SRT ----------
|
||||
out_dir = s.output_dir()
|
||||
# 写 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)
|
||||
# 中文 SRT(用译文 + 同时间戳)
|
||||
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 = "done"
|
||||
task.status = STATUS_DONE
|
||||
task.progress = P_DONE
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
|
||||
# 清理中间音频
|
||||
if not s.processing.keep_audio and wav.is_file():
|
||||
s = get_settings()
|
||||
if not s.processing.keep_audio and task.wav_path:
|
||||
try:
|
||||
wav.unlink()
|
||||
Path(task.wav_path).unlink()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
logger.info("任务 %d 完成:%s", task.id, bi_path.name)
|
||||
logger.info("任务 %d [完成] %s", task.id, bi_path.name)
|
||||
|
||||
|
||||
# ---------------- DB 状态更新 ----------------
|
||||
# ---------------- 工具函数 ----------------
|
||||
|
||||
def _set_status(db, task: Task, status: str, progress: float, note: str = "") -> None:
|
||||
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)
|
||||
@@ -137,15 +174,17 @@ def _set_status(db, task: Task, status: str, progress: float, note: str = "") ->
|
||||
if note:
|
||||
logger.info("任务 %d [%s %.0f%%] %s", task.id, status, progress, note)
|
||||
else:
|
||||
logger.info("任务 %d [%s %.0f%%]", task.id, status, progress)
|
||||
logger.debug("任务 %d [%s %.0f%%]", task.id, status, progress)
|
||||
|
||||
|
||||
def _mark_failed(db, task_id: int, error: str) -> None:
|
||||
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 = "failed"
|
||||
task.status = STATUS_FAILED
|
||||
task.error = error[:2000]
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
|
||||
Reference in New Issue
Block a user