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:
@@ -29,7 +29,8 @@ def transcribe(wav_path: Path) -> list[Segment]:
|
||||
raise FileNotFoundError(f"音频不存在:{wav_path}")
|
||||
|
||||
model = get_model_manager().get_asr()
|
||||
logger.debug("开始转写 %s(model=%s language=%s)", wav_path.name, s.model, s.language)
|
||||
logger.debug("开始转写 %s(model=%s language=%s batch_size=%d)",
|
||||
wav_path.name, s.model, s.language, s.batch_size)
|
||||
|
||||
segments_gen, info = model.transcribe(
|
||||
str(wav_path),
|
||||
@@ -37,6 +38,8 @@ def transcribe(wav_path: Path) -> list[Segment]:
|
||||
word_timestamps=s.word_timestamps,
|
||||
vad_filter=s.vad_filter,
|
||||
beam_size=5,
|
||||
batch_size=s.batch_size, # 批量解码:多音频块一次性送 GPU
|
||||
without_timestamps=False, # BatchedInferencePipeline 默认 True,需显式关闭以生成段级时间戳
|
||||
)
|
||||
logger.debug(
|
||||
"音频时长 %.1fs,检测语言=%s(置信度 %.2f)",
|
||||
|
||||
@@ -54,15 +54,18 @@ class ModelManager:
|
||||
if self._translator is not None:
|
||||
self._unload_translator_locked()
|
||||
s = get_settings().asr
|
||||
logger.debug("加载 ASR 模型 model=%s device=%s compute_type=%s",
|
||||
logger.info("加载 ASR 模型 model=%s device=%s compute_type=%s",
|
||||
s.model, s.device, s.compute_type)
|
||||
from faster_whisper import WhisperModel
|
||||
from faster_whisper import WhisperModel, BatchedInferencePipeline
|
||||
# device/compute_type 组合:cpu+int8 / cuda+float16
|
||||
self._asr = WhisperModel(
|
||||
# BatchedInferencePipeline 包装 WhisperModel,使 transcribe() 支持 batch_size,
|
||||
# 多个音频块(chunk_length=30s)一次性送 GPU 解码,配合内部 prefill 提高利用率。
|
||||
whisper = WhisperModel(
|
||||
s.model, device=s.device, compute_type=s.compute_type,
|
||||
)
|
||||
self._asr = BatchedInferencePipeline(model=whisper)
|
||||
self._current = "asr"
|
||||
logger.debug("ASR 模型已就绪。")
|
||||
logger.info("ASR 模型已就绪(batched, batch_size=%d)。", s.batch_size)
|
||||
return self._asr
|
||||
|
||||
def unload_asr(self) -> None:
|
||||
@@ -72,7 +75,7 @@ class ModelManager:
|
||||
def _unload_asr_locked(self) -> None:
|
||||
if self._asr is None:
|
||||
return
|
||||
logger.debug("卸载 ASR 模型(释放显存供翻译器独占)。")
|
||||
logger.info("卸载 ASR 模型(释放显存供翻译器独占)。")
|
||||
# faster-whisper 模型无显式 close,del 即可
|
||||
del self._asr
|
||||
self._asr = None
|
||||
@@ -89,7 +92,7 @@ class ModelManager:
|
||||
if self._asr is not None:
|
||||
self._unload_asr_locked()
|
||||
s = get_settings().translation
|
||||
logger.debug("加载翻译模型 model=%s device=%s", s.model, s.device)
|
||||
logger.info("加载翻译模型 model=%s device=%s", s.model, s.device)
|
||||
from transformers import pipeline
|
||||
self._translator = pipeline(
|
||||
"translation",
|
||||
@@ -97,9 +100,11 @@ class ModelManager:
|
||||
device=s.device,
|
||||
src_lang=s.src_lang,
|
||||
tgt_lang=s.tgt_lang,
|
||||
batch_size=s.batch_size, # pipeline 内部批大小,与 translate_service 分块对齐
|
||||
)
|
||||
self._current = "translator"
|
||||
logger.debug("翻译模型已就绪(独占显存,可用大 batch)。")
|
||||
logger.info("翻译模型已就绪(显存独占,batch_size=%d,sort_by_length=%s)。",
|
||||
s.batch_size, s.sort_by_length)
|
||||
return self._translator
|
||||
|
||||
def unload_translator(self) -> None:
|
||||
@@ -109,7 +114,7 @@ class ModelManager:
|
||||
def _unload_translator_locked(self) -> None:
|
||||
if self._translator is None:
|
||||
return
|
||||
logger.debug("卸载翻译模型。")
|
||||
logger.info("卸载翻译模型。")
|
||||
# 释放 pipeline 持有的 model + tokenizer
|
||||
mdl = getattr(self._translator, "model", None)
|
||||
tok = getattr(self._translator, "tokenizer", None)
|
||||
|
||||
@@ -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()
|
||||
|
||||
236
app/services/scheduler.py
Normal file
236
app/services/scheduler.py
Normal file
@@ -0,0 +1,236 @@
|
||||
"""任务调度器:ffmpeg 异步提取 + GPU 阶段串行 + 模型复用。
|
||||
|
||||
设计动机:多任务时不应串行等一个任务全跑完才下一个。ffmpeg 是纯 CPU,可与 GPU 阶段
|
||||
并行;GPU 阶段(ASR + 翻译)串行化(共享显存),但卸载模型前查队列,有同类待处理
|
||||
任务就继续用当前模型,减少重复加载/卸载。
|
||||
|
||||
数据流:
|
||||
enqueue_task ──► ffmpeg 线程(每任务一个,异步,CPU)
|
||||
│ 提取音频 → task.wav_path → status=transcribing
|
||||
▼(唤醒 GPU 线程)
|
||||
GPU 调度线程(单线程,常驻)
|
||||
① 取 status=transcribing 的任务,get_asr()
|
||||
while 还有 transcribing 任务: asr_phase → status=translating
|
||||
(ASR 队列空,切翻译)
|
||||
② 取 status=translating 的任务,get_translator()
|
||||
while 还有 translating 任务: translate_phase → status=done
|
||||
(翻译队列空,回到 ① 等待)
|
||||
|
||||
模型复用:N 个任务的模型切换次数从 2N 降到最优 2 次(一批 ASR 全做完 → 切翻译 → 一批翻译全做完)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from ..database import get_session_local
|
||||
from ..models.task import (
|
||||
Task, STATUS_EXTRACTING, STATUS_TRANSCRIBING, STATUS_SEGMENTING,
|
||||
STATUS_TRANSLATING, STATUS_QUEUED, STATUS_FAILED,
|
||||
)
|
||||
from . import pipeline
|
||||
from .model_manager import get_model_manager
|
||||
|
||||
logger = logging.getLogger("audio2text.scheduler")
|
||||
|
||||
# 唤醒 GPU 调度线程的事件(新任务入队或 ffmpeg 完成时 set)
|
||||
_wake_event = threading.Event()
|
||||
# GPU 调度线程单例
|
||||
_scheduler_thread: threading.Thread | None = None
|
||||
_scheduler_started = False
|
||||
|
||||
|
||||
def enqueue_task(task_id: int) -> None:
|
||||
"""任务入队:起 ffmpeg 线程提取音频 + 唤醒 GPU 调度线程。
|
||||
|
||||
替代旧 pipeline.enqueue_task(每任务一个线程跑完整管线)。
|
||||
ffmpeg 在独立线程跑(CPU,与 GPU 并行),完成后 GPU 调度线程接管 ASR+翻译。
|
||||
"""
|
||||
_ensure_scheduler_running()
|
||||
t = threading.Thread(
|
||||
target=_extract_audio_async, args=(task_id,),
|
||||
name=f"ffmpeg-{task_id}", daemon=True,
|
||||
)
|
||||
t.start()
|
||||
logger.info("任务 %d 已入队,开始音频提取。", task_id)
|
||||
|
||||
|
||||
def start_scheduler() -> None:
|
||||
"""启动 GPU 调度线程(应用启动时调一次,幂等)。"""
|
||||
global _scheduler_thread, _scheduler_started
|
||||
if _scheduler_started:
|
||||
return
|
||||
_scheduler_started = True
|
||||
_reset_stuck_tasks()
|
||||
_scheduler_thread = threading.Thread(
|
||||
target=_gpu_scheduler, name="gpu-scheduler", daemon=True,
|
||||
)
|
||||
_scheduler_thread.start()
|
||||
logger.info("GPU 调度线程已启动。")
|
||||
|
||||
|
||||
def _reset_stuck_tasks() -> None:
|
||||
"""启动时清理卡在中间状态的任务(进程上次崩溃残留)。
|
||||
|
||||
transcribing 但无 wav_path、translating 但无 segments_json 的任务,
|
||||
是上次进程异常退出留下的孤儿。标记为 failed 避免调度线程反复尝试。
|
||||
"""
|
||||
db = get_session_local()()
|
||||
try:
|
||||
stuck = (
|
||||
db.query(Task)
|
||||
.filter(Task.status.in_([
|
||||
STATUS_EXTRACTING, STATUS_TRANSCRIBING, STATUS_SEGMENTING, STATUS_TRANSLATING,
|
||||
]))
|
||||
.all()
|
||||
)
|
||||
n = 0
|
||||
for task in stuck:
|
||||
reason = ""
|
||||
if task.status in (STATUS_TRANSCRIBING, STATUS_SEGMENTING) and not task.wav_path:
|
||||
reason = f"重启时发现 {task.status} 状态但无 wav_path"
|
||||
elif task.status == STATUS_TRANSLATING and not task.segments_json:
|
||||
reason = f"重启时发现 translating 状态但无 segments_json"
|
||||
elif task.status == STATUS_EXTRACTING:
|
||||
reason = "重启时发现 extracting 状态(ffmpeg 未完成)"
|
||||
if reason:
|
||||
task.status = STATUS_FAILED
|
||||
task.error = reason[:2000]
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
n += 1
|
||||
logger.warning("清理卡住的任务 %d:%s", task.id, reason)
|
||||
if n:
|
||||
db.commit()
|
||||
logger.info("共清理 %d 个卡住的任务。", n)
|
||||
except Exception as exc: # pragma: no cover
|
||||
logger.error("清理卡住任务时出错:%s", exc)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _ensure_scheduler_running() -> None:
|
||||
"""确保 GPU 调度线程在跑(enqueue 时调,防止 lifespan 未启动的边界情况)。"""
|
||||
if not _scheduler_started:
|
||||
start_scheduler()
|
||||
|
||||
|
||||
# ---------------- ffmpeg 异步提取 ----------------
|
||||
|
||||
def _extract_audio_async(task_id: int) -> None:
|
||||
"""独立线程跑 ffmpeg 提取(CPU),完成后唤醒 GPU 调度线程。"""
|
||||
db = get_session_local()()
|
||||
try:
|
||||
task = db.get(Task, task_id)
|
||||
if task is None:
|
||||
logger.error("任务 %d 不存在,ffmpeg 线程退出。", task_id)
|
||||
return
|
||||
pipeline.extract_phase(db, task)
|
||||
_wake_event.set() # 通知 GPU 线程有新任务
|
||||
except Exception as exc:
|
||||
logger.exception("任务 %d ffmpeg 提取失败:%s", task_id, exc)
|
||||
pipeline.mark_failed(db, task_id, str(exc))
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
# ---------------- GPU 调度线程 ----------------
|
||||
|
||||
def _gpu_scheduler() -> None:
|
||||
"""常驻 GPU 调度线程:串行处理 ASR + 翻译,切换模型前查队列复用。
|
||||
|
||||
循环逻辑:
|
||||
1. 处理所有待 ASR 任务(Whisper 只加载一次)
|
||||
2. 处理所有待翻译任务(NLLB 只加载一次)
|
||||
3. 都空了 → 等待唤醒
|
||||
每个阶段失败的任务标记 failed,不影响其他任务。
|
||||
"""
|
||||
logger.info("GPU 调度线程开始运行。")
|
||||
while True:
|
||||
try:
|
||||
# 优先处理 ASR 队列:把所有待 ASR 的任务一次性做完(模型复用)
|
||||
asr_count = _drain_asr_queue()
|
||||
# 再处理翻译队列:把所有待翻译的任务一次性做完(模型复用)
|
||||
trans_count = _drain_translate_queue()
|
||||
|
||||
if asr_count == 0 and trans_count == 0:
|
||||
# 两队列都空,等待新任务唤醒
|
||||
_wake_event.wait(timeout=60)
|
||||
_wake_event.clear()
|
||||
except Exception as exc: # pragma: no cover
|
||||
# 调度线程不能死,任何异常都捕获后继续
|
||||
logger.exception("GPU 调度线程异常(已恢复):%s", exc)
|
||||
|
||||
|
||||
def _drain_asr_queue() -> int:
|
||||
"""连续处理所有 status=transcribing 的任务,Whisper 只加载一次。
|
||||
|
||||
Returns: 本轮处理的任务数
|
||||
"""
|
||||
n = 0
|
||||
mm = get_model_manager()
|
||||
while True:
|
||||
db = get_session_local()()
|
||||
task = None # 预定义,防 query 抛异常时 except 引用未定义变量
|
||||
try:
|
||||
# 取最早一个待 ASR 的任务(按 id 升序,FIFO)
|
||||
task = (
|
||||
db.query(Task)
|
||||
.filter(Task.status == STATUS_TRANSCRIBING)
|
||||
.order_by(Task.id.asc())
|
||||
.first()
|
||||
)
|
||||
if task is None:
|
||||
break # ASR 队列空
|
||||
# 加载 ASR 模型(若翻译器在内存,model_manager 自动卸载它,并记 INFO)
|
||||
mm.get_asr()
|
||||
# 执行 ASR + 断句
|
||||
pipeline.asr_phase(db, task)
|
||||
n += 1
|
||||
except Exception as exc:
|
||||
tid = task.id if task is not None else -1
|
||||
logger.exception("任务 %d ASR 阶段失败:%s", tid, exc)
|
||||
if task is not None:
|
||||
pipeline.mark_failed(db, task.id, str(exc))
|
||||
finally:
|
||||
db.close()
|
||||
if n > 0:
|
||||
logger.info("ASR 批次完成:处理 %d 个任务。", n)
|
||||
return n
|
||||
|
||||
|
||||
def _drain_translate_queue() -> int:
|
||||
"""连续处理所有 status=translating 的任务,NLLB 只加载一次。
|
||||
|
||||
Returns: 本轮处理的任务数
|
||||
"""
|
||||
n = 0
|
||||
mm = get_model_manager()
|
||||
while True:
|
||||
db = get_session_local()()
|
||||
task = None # 预定义,防 query 抛异常时 except 引用未定义变量
|
||||
try:
|
||||
task = (
|
||||
db.query(Task)
|
||||
.filter(Task.status == STATUS_TRANSLATING)
|
||||
.order_by(Task.id.asc())
|
||||
.first()
|
||||
)
|
||||
if task is None:
|
||||
break # 翻译队列空
|
||||
# 加载翻译模型(若 ASR 在内存,model_manager 自动卸载它,并记 INFO)
|
||||
mm.get_translator()
|
||||
# 执行翻译 + 写 SRT
|
||||
pipeline.translate_phase(db, task)
|
||||
n += 1
|
||||
except Exception as exc:
|
||||
tid = task.id if task is not None else -1
|
||||
logger.exception("任务 %d 翻译阶段失败:%s", tid, exc)
|
||||
if task is not None:
|
||||
pipeline.mark_failed(db, task.id, str(exc))
|
||||
finally:
|
||||
db.close()
|
||||
if n > 0:
|
||||
logger.info("翻译批次完成:处理 %d 个任务。", n)
|
||||
return n
|
||||
@@ -1,12 +1,19 @@
|
||||
"""翻译服务:NLLB-200,英译中。
|
||||
|
||||
通过 model_manager 加载,确保 ASR 已卸载、翻译器独占显存,从而可用大 batch_size。
|
||||
按字幕条目批量翻译,保留索引对应。
|
||||
|
||||
GPU 优化:按长度排序后分批翻译。
|
||||
- 同一批内句子长度相近 → padding 浪费最小化 → GPU 有效计算占比提升
|
||||
- 翻译完按原始下标散回,保证 zh_texts[i] 对应 subs[i](时间戳对齐不变)
|
||||
- 批切分用 token 预算 + 条数上限双重约束:短句自动攒大批,长句自动拆小批
|
||||
|
||||
单条翻译失败时该位置回退为原英文。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
|
||||
from ..config import get_settings
|
||||
from .model_manager import get_model_manager
|
||||
@@ -14,12 +21,15 @@ from .types import Subtitle
|
||||
|
||||
logger = logging.getLogger("audio2text.translate")
|
||||
|
||||
# 估算每条字幕的 token 数:英文约 1 token/词,留 20% 余量覆盖标点/子词拆分
|
||||
_TOKENS_PER_WORD = 1.2
|
||||
|
||||
|
||||
def translate(subtitles: list[Subtitle]) -> list[str]:
|
||||
"""批量翻译英文字幕为中文。
|
||||
|
||||
Args:
|
||||
subtitles: 断句后的英文字幕条目
|
||||
subtitles: 断句后的英文字幕条目(按时间顺序)
|
||||
|
||||
Returns:
|
||||
list[str],与 subtitles 等长、顺序对应的中文译文。
|
||||
@@ -30,31 +40,140 @@ def translate(subtitles: list[Subtitle]) -> list[str]:
|
||||
|
||||
s = get_settings().translation
|
||||
pipe = get_model_manager().get_translator()
|
||||
batch = s.batch_size
|
||||
batch_size = s.batch_size
|
||||
max_len = s.max_length
|
||||
|
||||
# 取纯文本(去掉折行),避免翻译把换行符当语义
|
||||
texts = [sub.text.replace("\n", " ").strip() for sub in subtitles]
|
||||
logger.debug("开始翻译 %d 条字幕(batch_size=%d)...", len(texts), batch)
|
||||
sort_by_length = s.sort_by_length
|
||||
|
||||
results: list[str] = []
|
||||
for i in range(0, len(texts), batch):
|
||||
chunk = texts[i:i + batch]
|
||||
try:
|
||||
out = pipe(chunk, max_length=max_len)
|
||||
for item in out:
|
||||
# pipeline 返回 [{"translation_text": "..."}]
|
||||
results.append(item.get("translation_text", "").strip())
|
||||
except Exception as exc: # pragma: no cover
|
||||
logger.warning("第 %d-%d 批翻译失败,逐条重试:%s", i, i + len(chunk), exc)
|
||||
for t in chunk:
|
||||
try:
|
||||
out = pipe([t], max_length=max_len)
|
||||
results.append(out[0].get("translation_text", "").strip())
|
||||
except Exception:
|
||||
results.append(t) # 回退原文
|
||||
if (i // batch + 1) % 5 == 0:
|
||||
logger.debug("已翻译 %d/%d 条。", min(i + len(chunk), len(texts)), len(texts))
|
||||
# 环境变量覆盖:便于 A/B 基准对比(test/bench_translate.py 用)
|
||||
if os.environ.get("TRANSLATE_NO_SORT") == "1":
|
||||
sort_by_length = False
|
||||
|
||||
if sort_by_length:
|
||||
results = _translate_sorted(pipe, texts, batch_size, max_len)
|
||||
else:
|
||||
results = _translate_sequential(pipe, texts, batch_size, max_len)
|
||||
|
||||
logger.debug("翻译完成:%d 条。", len(results))
|
||||
return results
|
||||
|
||||
|
||||
# ---------------- 长度排序批处理(默认)----------------
|
||||
|
||||
def _translate_sorted(
|
||||
pipe, texts: list[str], batch_size: int, max_len: int,
|
||||
) -> list[str]:
|
||||
"""按长度排序后分批翻译,翻译完按原序散回。
|
||||
|
||||
1. 记录 (orig_idx, text, est_tokens)
|
||||
2. 按 est_tokens 升序排序 → 相近长度的聚到同一批
|
||||
3. token 预算 + 条数上限双重约束切批:短句攒大批,长句拆小批
|
||||
4. 逐批翻译,按 orig_idx 把译文放回 results[orig_idx]
|
||||
"""
|
||||
n = len(texts)
|
||||
# 估算每条 token 数(用词数 × 1.2,至少 1 避免除零)
|
||||
items = [
|
||||
(i, texts[i], max(1, int(len(texts[i].split()) * _TOKENS_PER_WORD)))
|
||||
for i in range(n)
|
||||
]
|
||||
# 按 token 长度升序:短句在前,长句在后
|
||||
items.sort(key=lambda x: x[2])
|
||||
|
||||
# token 预算上限:一批的总 token 不超过 batch_size * max_len
|
||||
# 短句(每条 ~10 token)可攒到 batch_size 条;长句(~200 token)自动拆成更小批
|
||||
token_budget = batch_size * max_len
|
||||
|
||||
batches: list[list[tuple[int, str, int]]] = [] # [(orig_idx, text, tok), ...]
|
||||
cur_batch: list[tuple[int, str, int]] = []
|
||||
cur_max = 0 # 当前批内最长句的 token 数
|
||||
|
||||
for orig_idx, text, tok in items:
|
||||
new_max = max(cur_max, tok)
|
||||
new_tokens = (len(cur_batch) + 1) * new_max # 批内所有句都 pad 到 new_max
|
||||
if cur_batch and (len(cur_batch) >= batch_size or new_tokens > token_budget):
|
||||
batches.append(cur_batch)
|
||||
cur_batch = []
|
||||
cur_max = 0
|
||||
new_max = tok
|
||||
cur_batch.append((orig_idx, text, tok))
|
||||
cur_max = new_max
|
||||
if cur_batch:
|
||||
batches.append(cur_batch)
|
||||
|
||||
# padding 浪费对比(DEBUG 日志量化收益)
|
||||
pad_sorted = sum(len(b) * max(t for _, _, t in b) - sum(t for _, _, t in b) for b in batches)
|
||||
pad_seq = _estimate_sequential_padding(texts, batch_size)
|
||||
saving = (1 - pad_sorted / pad_seq) * 100 if pad_seq else 0
|
||||
logger.debug(
|
||||
"翻译分批:%d 条 → %d 批(长度排序)。padding 浪费:顺序 %d → 排序 %d token(节省 %.0f%%)",
|
||||
n, len(batches), pad_seq, pad_sorted, saving,
|
||||
)
|
||||
|
||||
results: list[str | None] = [None] * n
|
||||
done = 0
|
||||
for batch in batches:
|
||||
orig_indices = [b[0] for b in batch]
|
||||
batch_texts = [b[1] for b in batch]
|
||||
translated = _translate_batch(pipe, batch_texts, max_len)
|
||||
for idx, zh in zip(orig_indices, translated):
|
||||
results[idx] = zh
|
||||
done += len(batch)
|
||||
if (done // batch_size + 1) % 5 == 0:
|
||||
logger.debug("已翻译 %d/%d 条。", done, n)
|
||||
|
||||
# None(理论不会发生,_translate_batch 保证返回等长)→ 回退原文
|
||||
return [results[i] or texts[i] for i in range(n)]
|
||||
|
||||
|
||||
def _estimate_sequential_padding(texts: list[str], batch_size: int) -> int:
|
||||
"""估算按原序分批的 padding 浪费(token 数)。"""
|
||||
total = 0
|
||||
for i in range(0, len(texts), batch_size):
|
||||
chunk = texts[i:i + batch_size]
|
||||
toks = [max(1, int(len(t.split()) * _TOKENS_PER_WORD)) for t in chunk]
|
||||
batch_max = max(toks)
|
||||
total += batch_max * len(chunk) - sum(toks)
|
||||
return total
|
||||
|
||||
|
||||
# ---------------- 顺序批处理(A/B 对比用 / sort_by_length=false)----------------
|
||||
|
||||
def _translate_sequential(
|
||||
pipe, texts: list[str], batch_size: int, max_len: int,
|
||||
) -> list[str]:
|
||||
"""按原序分批翻译(旧行为,便于 A/B 对比)。"""
|
||||
results: list[str] = []
|
||||
for i in range(0, len(texts), batch_size):
|
||||
chunk = texts[i:i + batch_size]
|
||||
translated = _translate_batch(pipe, chunk, max_len)
|
||||
results.extend(translated)
|
||||
if (i // batch_size + 1) % 5 == 0:
|
||||
logger.debug("已翻译 %d/%d 条。", min(i + len(chunk), len(texts)), len(texts))
|
||||
return results
|
||||
|
||||
|
||||
# ---------------- 单批翻译 + 逐条重试回退 ----------------
|
||||
|
||||
def _translate_batch(pipe, chunk: list[str], max_len: int) -> list[str]:
|
||||
"""翻译一个批次,失败时降级到逐条重试。
|
||||
|
||||
Args:
|
||||
chunk: 本批的文本列表
|
||||
max_len: 单条最大生成长度
|
||||
"""
|
||||
try:
|
||||
out = pipe(chunk, max_length=max_len, truncation=True)
|
||||
return [item.get("translation_text", "").strip() for item in out]
|
||||
except Exception as exc: # pragma: no cover
|
||||
logger.warning("批次翻译失败(%d 条),逐条重试:%s", len(chunk), exc)
|
||||
results: list[str] = []
|
||||
for t in chunk:
|
||||
try:
|
||||
out = pipe([t], max_length=max_len, truncation=True)
|
||||
results.append(out[0].get("translation_text", "").strip())
|
||||
except Exception as exc2:
|
||||
logger.warning("单条翻译失败,回退原文:%s", exc2)
|
||||
results.append(t) # 回退原文
|
||||
return results
|
||||
|
||||
@@ -7,8 +7,8 @@
|
||||
...
|
||||
<upload_dir>/<yyyy>/<mm>/<uuid>.<ext> complete 后的正式视频
|
||||
|
||||
与 server 的区别:视频无需 sha256 去重(每个视频都转写),complete 直接创建 Task。
|
||||
管线触发由 controller 调用 pipeline.enqueue_task,本服务不依赖 pipeline。
|
||||
complete 创建转写 Task,管线触发由 controller 调用 scheduler.enqueue_task,
|
||||
本服务不依赖 scheduler(避免循环依赖)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -112,6 +112,14 @@ class UploadService:
|
||||
|
||||
# ---------------- 拼接 + 创建任务 ----------------
|
||||
|
||||
# 允许的音视频扩展名白名单(防可执行文件落盘到上传目录)
|
||||
_ALLOWED_EXTS = frozenset({
|
||||
".mp4", ".mkv", ".avi", ".mov", ".webm", ".flv",
|
||||
".mp3", ".wav", ".flac", ".aac", ".m4a", ".ogg", ".wma",
|
||||
})
|
||||
|
||||
# ---------------- 拼接 + 创建任务 ----------------
|
||||
|
||||
def complete(self, upload_id: str) -> CompleteResponse:
|
||||
session = self._require_session(upload_id)
|
||||
|
||||
@@ -135,6 +143,8 @@ class UploadService:
|
||||
final_path = self._assemble(session)
|
||||
rel = str(final_path.relative_to(self.upload_root))
|
||||
|
||||
# 单事务:建 Task + 更新 session 状态 + 关联 task_id 一次 commit
|
||||
# 避免双 commit 之间崩溃产生孤儿 Task(Task 已建但 session.task_id 为空)
|
||||
task = Task(
|
||||
filename=session.filename,
|
||||
source_path=rel,
|
||||
@@ -144,14 +154,14 @@ class UploadService:
|
||||
self.db.add(task)
|
||||
session.status = "completed"
|
||||
session.final_path = rel
|
||||
session.task_id = None # 占位,flush 后用 task.id 赋值
|
||||
session.updated_at = datetime.now(timezone.utc)
|
||||
self.db.commit()
|
||||
self.db.refresh(task)
|
||||
# 正向关联:session → task(替代旧的 source_path 反向查找)
|
||||
self.db.flush() # 拿到 task.id(不 commit,仍在事务内)
|
||||
session.task_id = task.id
|
||||
self.db.commit()
|
||||
self.db.refresh(task)
|
||||
|
||||
# 清理分片暂存
|
||||
# 清理分片暂存(commit 后,即使清理失败也不影响已建任务)
|
||||
self._cleanup_session_dir(upload_id)
|
||||
|
||||
logger.info("上传完成 task_id=%s file=%s size=%d", task.id, session.filename, session.size_bytes)
|
||||
@@ -215,7 +225,10 @@ class UploadService:
|
||||
|
||||
def _assemble(self, session: UploadSession) -> Path:
|
||||
"""按 index 顺序拼接全部分片为正式视频文件。"""
|
||||
ext = Path(session.filename).suffix or ".mp4"
|
||||
# 扩展名取自客户端 filename,但做白名单净化:不在允许列表内则回退 .bin
|
||||
ext = Path(session.filename).suffix.lower()
|
||||
if ext not in self._ALLOWED_EXTS:
|
||||
ext = ".bin"
|
||||
now = datetime.now(timezone.utc)
|
||||
sub = self.upload_root / f"{now:%Y}" / f"{now:%m}"
|
||||
sub.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
Reference in New Issue
Block a user