Files
audio2text/app/services/pipeline.py
zikai 00e2a95fb7 Initial commit: audio2text 双语字幕生成服务
- 音频/视频转双语(英/中)SRT 字幕,Docker 容器化,CPU 开发/GPU 生产同一份代码
- faster-whisper ASR(词级时间戳) + 断句时间戳重算 + NLLB 翻译(模型不共驻)
- 分片上传(断点续传) + SQLite 持久化 + 主页/历史/日志页面
- 历史页文件名搜索;缓存定时清理(默认保留7天,可配置)
- 双 Dockerfile(cpu/gpu) + setup/start/stop 脚本
2026-07-06 06:54:19 +00:00

154 lines
5.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""转写管线编排:提取音频 → ASR → 断句 → 翻译 → 写 SRT。
任务状态机:
queued → extracting → transcribing → segmenting → translating → done
任一步失败 → failed
模型不共驻ASR 与翻译分阶段加载,翻译时先卸载 Whisper 释放显存跑大 batch。
管线在后台线程跑(每个任务一个线程),通过 DB 更新状态与进度。
"""
from __future__ import annotations
import logging
import threading
import traceback
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 .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
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)
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:
s = get_settings()
src = s.upload_dir() / task.source_path
# ---------- 1. 提取音频 ----------
_set_status(db, task, "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("已删除原始视频 %sdelete_original_after_extract=true", src.name)
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,
note=f"识别出 {len(segments)}")
# ---------- 3. 断句 + 时间戳重算 ----------
_set_status(db, task, "segmenting", P_SEGMENT_START)
subs = segmenter.resegment(segments)
_set_status(db, task, "segmenting", P_SEGMENT_END,
note=f"重组为 {len(subs)} 条字幕")
# ---------- 4. 翻译 ----------
_set_status(db, task, "translating", P_TRANSLATE_START)
# 翻译阶段model_manager 会自动卸载 ASR、加载翻译器独占显存
zh_texts = translate_service.translate(subs)
_set_status(db, task, "translating", P_TRANSLATE_END,
note=f"翻译 {len(zh_texts)}")
# ---------- 5. 写 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.progress = P_DONE
task.updated_at = datetime.now(timezone.utc)
db.commit()
# 清理中间音频
if not s.processing.keep_audio and wav.is_file():
try:
wav.unlink()
except OSError:
pass
logger.info("任务 %d 完成:%s", task.id, bi_path.name)
# ---------------- DB 状态更新 ----------------
def _set_status(db, task: Task, status: str, progress: float, note: str = "") -> None:
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.info("任务 %d [%s %.0f%%]", task.id, status, progress)
def _mark_failed(db, task_id: int, error: str) -> None:
try:
task = db.get(Task, task_id)
if task is None:
return
task.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())