"""任务调度器: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