diff --git a/.dockerignore b/.dockerignore index 658500d..0e5132c 100644 --- a/.dockerignore +++ b/.dockerignore @@ -8,3 +8,4 @@ logs/ *.pid .git/ .gitignore +test/ diff --git a/.gitignore b/.gitignore index 9b98e73..dc29f78 100644 --- a/.gitignore +++ b/.gitignore @@ -7,6 +7,7 @@ __pycache__/ # 运行时数据(走 docker volume,不入库) data/ +data-gpu/ /models/ config.yaml *.pid @@ -15,3 +16,6 @@ logs/ # 编辑器 .vscode/ .idea/ + +# 测试数据与脚本(本地测试用,不入库) +/test/ diff --git a/Dockerfile b/Dockerfile index 989e246..e7e0bd3 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,44 +1,81 @@ +# syntax=docker/dockerfile:1.7 # audio2text — 一份 Dockerfile,CPU(dev) / GPU(prod) 双形态。 # docker build --build-arg VARIANT=cpu -t audio2text:cpu . # docker build --build-arg VARIANT=gpu -t audio2text:gpu . -# 区别仅在基础镜像与 torch 轮子;Python 依赖列表完全一致。 +# docker compose --profile dev up -d --build # 开发:dev target + 源码挂载 + reload +# +# 分层目标:依赖层(deps)稳定固化,代码层(final/dev)在最末,改代码不重装依赖。 +# BuildKit 缓存挂载(--mount=type=cache):pip wheel 跨构建复用,二次构建秒级。 +# +# 安全约束:PyTorch CPU wheel 与 GPU(CUDA) wheel 是两个不兼容二进制包。 +# - CPU 版 torch.cuda.is_available()=False(CPU 预期) +# - GPU 版 torch.cuda.is_available()=True(GPU 可调度) +# 因此 torch 必须按 VARIANT 分叉装不同 wheel,绝不能跨 variant 共享依赖层。 +# deps 阶段用 FROM base-${VARIANT},CPU/GPU 是两条独立构建链,各自装对应 torch。 ARG VARIANT=cpu + +# --------------------------------------------------------------------------- +# 1. 基础镜像分叉:CPU 用 slim Python,GPU 用 CUDA runtime + 手动装 python3.12 +# --------------------------------------------------------------------------- FROM python:3.12-slim AS base-cpu -# GPU 基础镜像带 CUDA 运行时,torch 可装 CUDA 轮子 + FROM nvidia/cuda:12.1.0-runtime-ubuntu22.04 AS base-gpu +# Ubuntu 22.04 默认源只有 python3.10/3.11,需加 deadsnakes PPA 才能装 python3.12。 +# 注意:不能用 python3-pip(那是 3.10 的系统 pip,会把包装到 3.10 site-packages)。 +# 用 python3.12 -m ensurepip 给 3.12 装 pip,确保所有包进 3.12 目录。 +# DEBIAN_FRONTEND=noninteractive 避免 tzdata 等包进入交互式配置卡住构建。 +ENV DEBIAN_FRONTEND=noninteractive RUN apt-get update -y && apt-get install -y --no-install-recommends \ - python3.12 python3.12-venv python3.12-dev python3-pip \ + software-properties-common gnupg ca-certificates \ + && add-apt-repository -y ppa:deadsnakes/ppa \ + && apt-get update -y && apt-get install -y --no-install-recommends \ + python3.12 python3.12-venv python3.12-dev \ && rm -rf /var/lib/apt/lists/* \ && ln -sf /usr/bin/python3.12 /usr/local/bin/python3 \ - && ln -sf /usr/bin/python3.12 /usr/local/bin/python + && ln -sf /usr/bin/python3.12 /usr/local/bin/python \ + && python3 -m ensurepip \ + && python3 -m pip install --upgrade pip -FROM base-${VARIANT} AS final +# --------------------------------------------------------------------------- +# 2. deps:系统依赖 + Python 依赖(含 torch)。CPU/GPU 各自一条独立链。 +# final / dev 都 FROM deps,继承对应 variant 的 torch,不串。 +# --------------------------------------------------------------------------- +FROM base-${VARIANT} AS deps ARG VARIANT -ENV VARIANT=${VARIANT} \ - PYTHONUNBUFFERED=1 \ +ENV PYTHONUNBUFFERED=1 \ PIP_NO_CACHE_DIR=1 \ HF_HOME=/models/huggingface \ CT2_CACHE=/models/ctranslate2 -# ffmpeg 是核心系统依赖,必须装 -RUN apt-get update -y && apt-get install -y --no-install-recommends \ +# ffmpeg 是核心系统依赖;patchelf 用于修复 ctranslate2 可执行栈 +RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \ + --mount=type=cache,target=/var/lib/apt,sharing=locked \ + apt-get update -y && apt-get install -y --no-install-recommends \ ffmpeg ca-certificates patchelf \ && rm -rf /var/lib/apt/lists/* WORKDIR /app COPY requirements.txt /app/requirements.txt - -# CPU 装 CPU 版 torch;GPU 走默认 index(带 CUDA 的轮子) -RUN if [ "$VARIANT" = "cpu" ]; then \ - pip install --upgrade pip && \ +# 【顺序关键】必须先装 torch(按 VARIANT 分叉),再装 requirements。 +# 原因:requirements 里的 transformers / accelerate 依赖 torch,若先装 requirements, +# pip 会从默认 PyPI 拉来 GPU 版 torch + nvidia-* 全家桶(~2GB),即使后续覆盖装 CPU torch, +# 那些无用的 nvidia 包仍残留在镜像里。先装 torch 让 pip 解析 requirements 时 torch 已满足。 +# +# CPU 走 cpu index(~200MB,无 nvidia 依赖);GPU 走默认 PyPI(带 CUDA,~2.5GB)。 +# 两套 wheel 二进制不兼容:CPU 版 cuda.is_available()=False,GPU 版 =True。不可共享。 +RUN --mount=type=cache,target=/root/.cache/pip,sharing=locked \ + pip install --upgrade pip && \ + if [ "$VARIANT" = "cpu" ]; then \ pip install torch --index-url https://download.pytorch.org/whl/cpu ; \ else \ - pip install --upgrade pip && \ pip install torch ; \ fi -RUN pip install -r /app/requirements.txt + +# torch 已就位,requirements 里的 transformers/accelerate 解析时复用已装 torch,不重复拉 +RUN --mount=type=cache,target=/root/.cache/pip,sharing=locked \ + pip install -r /app/requirements.txt # ctranslate2 的 .so 带可执行栈标志(PT_GNU_STACK X),在某些内核 + Docker 组合下 # 会触发 "cannot enable executable stack as shared object requires"。用 patchelf @@ -51,13 +88,37 @@ RUN for d in /usr/local/lib/python3.12/site-packages/ctranslate2.libs \ done; \ python -c "import ctranslate2; print('ctranslate2 stack fix verified', ctranslate2.__version__)" +# --------------------------------------------------------------------------- +# 3. dev:从 deps 继承依赖,不 COPY 代码(compose 用 volume 挂载源码)。 +# 开启 uvicorn --reload,改代码零重建、保存即生效。 +# 注意:dev 必须在 final 之前,保证 `docker build` 默认 target 是 final。 +# --------------------------------------------------------------------------- +FROM deps AS dev +VOLUME ["/data", "/models"] +# ctranslate2/faster-whisper 用系统动态链接器找 cuDNN(不走 torch 的库加载), +# 需把 torch wheel 自带的 nvidia 库目录加入 LD_LIBRARY_PATH,否则 GPU 推理报 +# "Unable to load libcudnn_ops.so.9"。放在 final/dev 而非 deps,避免 ENV 变化 +# 导致 deps 的 apt/pip 层缓存失效。CPU 镜像无此目录,路径被忽略不影响。 +ENV CONFIG_PATH=/app/config.yaml \ + LD_LIBRARY_PATH=/usr/local/lib/python3.12/site-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/site-packages/nvidia/cublas/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cublas/lib + +EXPOSE 8000 +CMD ["python", "-m", "uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000", "--reload"] + +# --------------------------------------------------------------------------- +# 4. final(prod):从 deps 继承全部依赖,只加 app 代码。 +# 放在最后 = `docker build` 默认 target。setup.sh / start.sh 依赖此默认行为。 +# --------------------------------------------------------------------------- +FROM deps AS final COPY app /app/app +COPY scripts /app/scripts COPY config.example.yaml /app/config.example.yaml # 运行时数据:上传 / 中间产物 / 输出字幕 / 模型缓存 # 全部走 volume,镜像本身无状态、无敏感数据 VOLUME ["/data", "/models"] -ENV CONFIG_PATH=/app/config.yaml +ENV CONFIG_PATH=/app/config.yaml \ + LD_LIBRARY_PATH=/usr/local/lib/python3.12/site-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cudnn/lib:/usr/local/lib/python3.12/site-packages/nvidia/cublas/lib:/usr/local/lib/python3.12/dist-packages/nvidia/cublas/lib EXPOSE 8000 CMD ["python", "-m", "uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"] diff --git a/README.md b/README.md index eea9ba5..3c5389e 100644 --- a/README.md +++ b/README.md @@ -54,8 +54,10 @@ Spring 风格分层,HTTP 边界(controllers)与业务逻辑(services) ``` audio2text/ ├── Dockerfile # 一份 Dockerfile,ARG VARIANT=cpu|gpu 出两个镜像 -├── docker-compose.yml # cpu / gpu 两个 profile +├── docker-compose.yml # cpu / gpu / dev 三个 profile ├── setup.sh / start.sh / stop.sh # 安装 / 启动 / 停止(包装 docker 命令) +├── scripts/ +│ └── prefetch_models.{sh,py} # 预拉模型权重到 ./models volume(避免首次启动下载) ├── requirements.txt ├── config.example.yaml # 配置模板(复制为 config.yaml 后填值) ├── README.md @@ -149,6 +151,21 @@ cd /root/zikai/audio2text 首次启动会下载模型(Whisper `tiny.en` ~39M + opus-mt ~300MB)到 `./models` volume, 之后秒起。启动后浏览器打开 `http://127.0.0.1:8000/`,拖入视频或音频文件即可。 +### 预拉模型(避免首次启动卡在下载) + +容器首次处理任务时会从 HuggingFace 下载模型,大模型(GPU 的 large-v3-turbo ~3GB + +NLLB-1.3B ~2.5GB)下载耗时较长。可用预拉脚本提前下好到 `./models` volume,之后容器启动 +即用、无需联网: + +```bash +./scripts/prefetch_models.sh # 读 config.yaml(当前激活配置) +./scripts/prefetch_models.sh config.gpu.yaml # 读指定配置(如切换到 GPU 前预拉大模型) +``` + +脚本用已构建的镜像跑一次性容器,读配置里的 `asr.model` / `translation.model`,下载到 +`./models/huggingface`(HF 标准缓存)。**幂等**:已下过的模型自动跳过。换 config 的模型 +名后重跑即可补下新模型,无需重建镜像。 + ### CPU 模型选型 | 组件 | 模型 | 大小 | 说明 | diff --git a/app/config.py b/app/config.py index 36cef0a..c632a10 100644 --- a/app/config.py +++ b/app/config.py @@ -1,6 +1,7 @@ -"""运行时配置:所有参数从 config.yaml 读取,对齐 server/config.py 的风格。 +"""运行时配置:所有参数从 config.yaml 读取。 CPU dev / GPU prod 仅靠 device / model / compute_type 三项切换,代码完全不变。 +默认值与 config.example.yaml 对齐,确保无 yaml 时也能用最小配置启动。 """ from __future__ import annotations @@ -40,21 +41,23 @@ class ProcessingConfig(BaseModel): class AsrConfig(BaseModel): - model: str = "small" - device: str = "cpu" # cpu | cuda - compute_type: str = "int8" # cpu: int8;gpu: float16 + model: str = "tiny.en" # CPU: tiny.en;GPU: large-v3-turbo + device: str = "cpu" # cpu | cuda + compute_type: str = "int8" # cpu: int8;gpu: float16 language: str = "en" word_timestamps: bool = True vad_filter: bool = True + batch_size: int = 8 # BatchedInferencePipeline 的音频块批大小;GPU 建议 16 class TranslationConfig(BaseModel): - model: str = "facebook/nllb-200-distilled-1.3B" - device: str = "cpu" # cpu | cuda - src_lang: str = "eng_Latn" - tgt_lang: str = "zho_Hans" - batch_size: int = 16 - max_length: int = 256 + model: str = "Helsinki-NLP/opus-mt-en-zh" # CPU: opus-mt(轻量);GPU: facebook/nllb-200-distilled-1.3B + device: str = "cpu" # cpu | cuda + src_lang: str = "eng_Latn" # NLLB 语言码:英语 + tgt_lang: str = "zho_Hans" # NLLB 语言码:简体中文 + batch_size: int = 8 # 排序后单批最大条数;CPU: 8,GPU: 32(显存独占可用大 batch) + max_length: int = 256 # 单条最大生成 token 数 + sort_by_length: bool = True # 按句子长度排序后分批,减少批内 padding 浪费(GPU 收益大) class SegmentationConfig(BaseModel): @@ -68,8 +71,8 @@ class LoggingConfig(BaseModel): """日志配置:控制台 + 内存缓冲的最低级别,以及缓冲条数。 分层语义: - - debug:详细(ffmpeg 命令、模型加载/卸载、转写逐段、翻译逐批进度) - - info:简略(仅任务阶段转换,如 "任务 N [transcribing 55%]") + - debug:进度详情(任务 [status pct%]、转写/翻译逐批统计、ffmpeg 命令) + - info:任务流转里程碑(音频提取/ASR/翻译 的开始与完成、模型加载与卸载) - error:详细错误(完整 traceback,由 logger.exception 自带) """ diff --git a/app/controllers/log_router.py b/app/controllers/log_router.py index 399c2b2..657bbe2 100644 --- a/app/controllers/log_router.py +++ b/app/controllers/log_router.py @@ -1,9 +1,9 @@ """日志路由:查询内存日志缓冲,供 /logs 页面消费。 分层语义: -- debug:详细(ffmpeg 命令、模型加载/卸载、转写逐段、翻译逐批进度) -- info:简略(仅任务阶段转换,如 "任务 N [transcribing 55%]") -- warning / error:更详细,error 含完整 traceback +- debug:进度详情(任务 [status pct%]、转写逐段统计、翻译逐批进度、ffmpeg 命令) +- info:任务流转里程碑(音频提取/ASR/翻译 的开始与完成、模型加载与卸载、批次汇总) +- warning / error:异常与失败,error 含完整 traceback """ from __future__ import annotations diff --git a/app/controllers/task_router.py b/app/controllers/task_router.py index eaec1bc..369f649 100644 --- a/app/controllers/task_router.py +++ b/app/controllers/task_router.py @@ -2,6 +2,7 @@ from __future__ import annotations +from datetime import datetime, timezone from pathlib import Path from fastapi import APIRouter, Depends, HTTPException, Query @@ -10,7 +11,8 @@ from sqlalchemy.orm import Session from ..config import get_settings from ..database import get_db -from ..models.task import Task +from ..models.task import Task, STATUS_UPLOADING +from ..models.upload_session import UploadSession from ..schemas.task import TaskListResponse, TaskResponse router = APIRouter(prefix="/api/tasks", tags=["task"]) @@ -28,6 +30,31 @@ def _to_resp(task: Task) -> TaskResponse: ) +def _upload_to_resp(session: UploadSession) -> TaskResponse: + """把上传中的 UploadSession 映射为虚拟 TaskResponse。 + + is_upload=true 让前端走上传状态轮询而非任务轮询。 + progress = 已传分片数 / 总分片数 × 100(映射到 0-5 区间,与 extract 阶段衔接)。 + """ + uploaded = len(session.uploaded_chunks or []) + total = session.total_chunks or 1 + # 上传进度映射到 0-4%(extract 从 5% 开始,留 1% 给 complete 拼接) + progress = min(4.0, uploaded / total * 4.0) + now = datetime.now(timezone.utc) + return TaskResponse( + id=0, # 虚拟 id,前端用 upload_id 轮询 + filename=session.filename, + status=STATUS_UPLOADING, + progress=progress, + error=None, + created_at=session.created_at, + updated_at=session.updated_at or now, + size_bytes=session.size_bytes, + is_upload=True, + upload_id=session.upload_id, + ) + + @router.get("", response_model=TaskListResponse, summary="任务列表") def list_tasks( limit: int = Query(100, ge=1, le=500), @@ -35,14 +62,30 @@ def list_tasks( q: str = Query("", description="按文件名模糊搜索(大小写不敏感,匹配子串)"), db: Session = Depends(get_db), ) -> TaskListResponse: - q_base = db.query(Task).order_by(Task.id.desc()) + # 上传中的会话(pending 状态)也作为虚拟任务返回,让前端能看到上传进度 + upload_q = db.query(UploadSession).filter(UploadSession.status == "pending") if q.strip(): - # SQLite 的 LIKE 默认大小写不敏感(ASCII),ilike 等价于 LIKE like = f"%{q.strip()}%" - q_base = q_base.filter(Task.filename.ilike(like)) - total = q_base.count() - tasks = q_base.offset(offset).limit(limit).all() - return TaskListResponse(tasks=[_to_resp(t) for t in tasks], total=total) + upload_q = upload_q.filter(UploadSession.filename.ilike(like)) + uploads = upload_q.order_by(UploadSession.created_at.desc()).all() + + # 已创建的 Task(含 extracting/transcribing/.../done/failed) + task_q = db.query(Task).order_by(Task.created_at.desc(), Task.id.desc()) + if q.strip(): + like = f"%{q.strip()}%" + task_q = task_q.filter(Task.filename.ilike(like)) + total_tasks = task_q.count() + tasks = task_q.offset(offset).limit(limit).all() + + # 合并:上传会话 + Task,按 created_at 倒序 + upload_resps = [_upload_to_resp(s) for s in uploads] + task_resps = [_to_resp(t) for t in tasks] + all_resps = upload_resps + task_resps + all_resps.sort(key=lambda r: r.created_at, reverse=True) + + # 分页:offset/limit 作用于合并后列表(上传会话通常很少,主要影响首页前几条) + paged = all_resps[offset:offset + limit] + return TaskListResponse(tasks=paged, total=len(all_resps)) @router.get("/{task_id}", response_model=TaskResponse, summary="任务状态") diff --git a/app/controllers/upload_router.py b/app/controllers/upload_router.py index 35116d4..fb0d159 100644 --- a/app/controllers/upload_router.py +++ b/app/controllers/upload_router.py @@ -1,6 +1,6 @@ """分片上传路由:建会话 / 查状态 / 传分片 / complete。 -协议与 server 完全一致,区别仅在 complete 后创建的是转写 Task 而非 UploadedFile。 +complete 成功后创建转写 Task 并交由 scheduler 入队。 """ from __future__ import annotations @@ -60,10 +60,11 @@ def complete_session( """拼接分片 + 创建转写任务 + 入队管线。 controller 负责编排:service.complete 只管存储(拼接 + 建 Task), - 管线触发由 controller 调用,service 不依赖 pipeline(避免循环依赖)。 + 管线触发由 controller 调用 scheduler(ffmpeg 异步 + GPU 串行调度),service 不依赖 + scheduler(避免循环依赖)。 """ resp = service.complete(upload_id) # 仅新建任务时入队(幂等 complete 返回的也是同一 task_id,enqueue 幂等无副作用) - from ..services.pipeline import enqueue_task + from ..services.scheduler import enqueue_task enqueue_task(resp.task_id) return resp diff --git a/app/database.py b/app/database.py index 3b74a97..4dae34e 100644 --- a/app/database.py +++ b/app/database.py @@ -51,12 +51,36 @@ class Base(DeclarativeBase): def init_db_schema() -> None: - """建表(幂等)。""" + """建表(幂等)+ 旧库迁移(给 task 表补新字段)。 + + SQLAlchemy 的 create_all 只建新表不改旧表。对已存在的 task 表, + 需手动 ALTER TABLE ADD COLUMN 补 wav_path / segments_json(nullable)。 + """ from .models.task import Task # noqa: F401 from .models.upload_session import UploadSession # noqa: F401 - get_engine() - Base.metadata.create_all(get_engine()) + engine = get_engine() + Base.metadata.create_all(engine) + _migrate_task_columns(engine) + + +def _migrate_task_columns(engine) -> None: + """检测 task 表缺失的列并 ALTER TABLE 补上(nullable,向后兼容)。""" + from sqlalchemy import inspect, text + + insp = inspect(engine) + if "task" not in insp.get_table_names(): + return # 新库,create_all 已建好完整表 + existing = {c["name"] for c in insp.get_columns("task")} + # 新增字段:(列名, 列定义) + additions = [ + ("wav_path", "VARCHAR(1024)"), + ("segments_json", "TEXT"), + ] + with engine.begin() as conn: + for col, coltype in additions: + if col not in existing: + conn.execute(text(f"ALTER TABLE task ADD COLUMN {col} {coltype}")) def get_db() -> Generator[Session, None, None]: diff --git a/app/main.py b/app/main.py index f8f8586..6e32ded 100644 --- a/app/main.py +++ b/app/main.py @@ -76,6 +76,12 @@ async def lifespan(app: FastAPI): await asyncio.to_thread(reap_stale_sessions) except Exception as exc: # pragma: no cover logger.warning("启动 reaper 失败:%s", exc) + # 启动 GPU 调度线程(常驻,串行处理 ASR+翻译,模型复用) + try: + from .services.scheduler import start_scheduler + start_scheduler() + except Exception as exc: # pragma: no cover + logger.warning("启动 GPU 调度线程失败:%s", exc) # 缓存清理:启动时跑一次 + 后台定时循环(守护线程,随进程退出) s = get_settings() try: @@ -144,7 +150,33 @@ def create_app() -> FastAPI: @app.get("/health") def health() -> dict: - return {"status": "ok"} + """存活探针 + 设备信息。 + + 返回 torch 版本、cuda 可用性、GPU 名称、配置的 device, + 便于一眼区分 CPU/GPU 容器是否正确调度到对应硬件。 + torch 导入失败时(理论上不会,因为镜像已装 torch)降级为仅 status。 + """ + info: dict = {"status": "ok"} + try: + import torch + info["torch"] = torch.__version__ + info["cuda_available"] = torch.cuda.is_available() + if torch.cuda.is_available(): + info["gpu"] = torch.cuda.get_device_name(0) + info["gpu_count"] = torch.cuda.device_count() + except Exception as e: # pragma: no cover + info["torch_error"] = str(e) + s = get_settings() + info["asr_device"] = s.asr.device + info["asr_model"] = s.asr.model + info["asr_compute_type"] = s.asr.compute_type + info["asr_batch_size"] = s.asr.batch_size + info["asr_language"] = s.asr.language + info["translation_device"] = s.translation.device + info["translation_model"] = s.translation.model + info["translation_batch_size"] = s.translation.batch_size + info["translation_sort_by_length"] = s.translation.sort_by_length + return info @app.get("/history", response_class=HTMLResponse) def history_page() -> HTMLResponse: diff --git a/app/models/task.py b/app/models/task.py index 17b8140..d95fdd4 100644 --- a/app/models/task.py +++ b/app/models/task.py @@ -14,6 +14,17 @@ def _now() -> datetime: return datetime.now(timezone.utc) +# 任务状态枚举(单一真源,供 scheduler / cache_cleaner / views 引用) +STATUS_UPLOADING = "uploading" # 虚拟状态:分片上传中(不在 Task 表,由 UploadSession 映射) +STATUS_QUEUED = "queued" +STATUS_EXTRACTING = "extracting" +STATUS_TRANSCRIBING = "transcribing" +STATUS_SEGMENTING = "segmenting" +STATUS_TRANSLATING = "translating" +STATUS_DONE = "done" +STATUS_FAILED = "failed" + + class Task(Base): __tablename__ = "task" @@ -23,7 +34,7 @@ class Task(Base): # 视频在 upload_dir 下的相对路径(提取音频前后可能被删) source_path: Mapped[str] = mapped_column(String(1024), nullable=False) # 状态机:queued → extracting → transcribing → segmenting → translating → done | failed - status: Mapped[str] = mapped_column(String(32), nullable=False, default="queued") + status: Mapped[str] = mapped_column(String(32), nullable=False, default=STATUS_QUEUED, index=True) # 0-100 进度 progress: Mapped[float] = mapped_column(Float, nullable=False, default=0.0) # 失败原因 @@ -33,7 +44,13 @@ class Task(Base): en_srt_path: Mapped[str | None] = mapped_column(String(1024), nullable=True) zh_srt_path: Mapped[str | None] = mapped_column(String(1024), nullable=True) - created_at: Mapped[datetime] = mapped_column(DateTime, default=_now) + # ---- 流水线阶段间传递的中间数据(scheduler 调度用)---- + # ffmpeg 提取出的 wav 绝对路径(extract 阶段写,asr 阶段读) + wav_path: Mapped[str | None] = mapped_column(String(1024), nullable=True) + # ASR + 断句结果(asr 阶段写 JSON,translate 阶段读)。JSON 序列化的 list[Subtitle] + segments_json: Mapped[str | None] = mapped_column(Text, nullable=True) + + created_at: Mapped[datetime] = mapped_column(DateTime, default=_now, index=True) updated_at: Mapped[datetime] = mapped_column(DateTime, default=_now, onupdate=_now) def __repr__(self) -> str: diff --git a/app/models/upload_session.py b/app/models/upload_session.py index 374bdac..9fb892e 100644 --- a/app/models/upload_session.py +++ b/app/models/upload_session.py @@ -1,8 +1,9 @@ -"""分片上传会话 ORM:支撑断点续传。对齐 server 的 UploadSession 形态(SQLite 版)。""" +"""分片上传会话 ORM:支撑断点续传。""" from __future__ import annotations import json +import logging from datetime import datetime, timezone from typing import Any @@ -11,6 +12,8 @@ from sqlalchemy.orm import Mapped, mapped_column from ..database import Base +logger = logging.getLogger("audio2text.models") + def _now() -> datetime: return datetime.now(timezone.utc) @@ -26,7 +29,14 @@ class _IntList(TypeDecorator): return json.dumps(value) if value is not None else None def process_result_value(self, value: Any, dialect) -> list[int]: - return json.loads(value) if value else [] + if not value: + return [] + try: + return json.loads(value) + except (json.JSONDecodeError, TypeError): + # DB 脏数据(手工修改/并发写截断)不应让整行读取失败 + logger.warning("uploaded_chunks JSON 解析失败,回退空列表:%r", value[:80]) + return [] class UploadSession(Base): @@ -39,11 +49,13 @@ class UploadSession(Base): total_chunks: Mapped[int] = mapped_column(Integer, nullable=False) uploaded_chunks: Mapped[list[int]] = mapped_column(_IntList, default=list) # pending → completed(complete 成功)| abandoned(reaper 清理) - status: Mapped[str] = mapped_column(String(32), nullable=False, default="pending") + status: Mapped[str] = mapped_column(String(32), nullable=False, default="pending", index=True) # 拼接完成后的视频相对 upload_dir 路径 final_path: Mapped[str | None] = mapped_column(String(1024), nullable=True) - # complete 后关联的 Task.id(直接引用,避免反向查找 source_path) - task_id: Mapped[int | None] = mapped_column(ForeignKey("task.id"), nullable=True) + # complete 后关联的 Task.id。ondelete=SET NULL:Task 被删时 session 保留(task_id 置空) + task_id: Mapped[int | None] = mapped_column( + ForeignKey("task.id", ondelete="SET NULL"), nullable=True + ) - created_at: Mapped[datetime] = mapped_column(DateTime, default=_now) + created_at: Mapped[datetime] = mapped_column(DateTime, default=_now, index=True) updated_at: Mapped[datetime] = mapped_column(DateTime, default=_now, onupdate=_now) diff --git a/app/schemas/task.py b/app/schemas/task.py index dfb1248..9d3666a 100644 --- a/app/schemas/task.py +++ b/app/schemas/task.py @@ -62,6 +62,11 @@ class TaskResponse(BaseModel): error: str | None created_at: datetime updated_at: datetime + # 上传会话虚拟任务用:is_upload=true 时 id 无意义(固定 0), + # 前端用 upload_id 轮询 /api/tasks/chunk-uploads/{upload_id}/status + size_bytes: int | None = None + is_upload: bool = False + upload_id: str | None = None class TaskListResponse(BaseModel): diff --git a/app/security.py b/app/security.py index c76c150..44b2ab5 100644 --- a/app/security.py +++ b/app/security.py @@ -1,4 +1,4 @@ -"""/docs Basic Auth:对齐 server/security.py。明文密码,常量时间比较。""" +"""/docs Basic Auth:明文密码,常量时间比较。""" from __future__ import annotations diff --git a/app/services/asr_service.py b/app/services/asr_service.py index ffca502..d526f50 100644 --- a/app/services/asr_service.py +++ b/app/services/asr_service.py @@ -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)", diff --git a/app/services/model_manager.py b/app/services/model_manager.py index 38bff3d..8b1cc49 100644 --- a/app/services/model_manager.py +++ b/app/services/model_manager.py @@ -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) diff --git a/app/services/pipeline.py b/app/services/pipeline.py index fc6f920..dd8b39d 100644 --- a/app/services/pipeline.py +++ b/app/services/pipeline.py @@ -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() diff --git a/app/services/scheduler.py b/app/services/scheduler.py new file mode 100644 index 0000000..3c000c3 --- /dev/null +++ b/app/services/scheduler.py @@ -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 diff --git a/app/services/translate_service.py b/app/services/translate_service.py index 04aad9f..d207f82 100644 --- a/app/services/translate_service.py +++ b/app/services/translate_service.py @@ -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 diff --git a/app/services/upload_service.py b/app/services/upload_service.py index 4a73b6c..8132963 100644 --- a/app/services/upload_service.py +++ b/app/services/upload_service.py @@ -7,8 +7,8 @@ ... ///. 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) diff --git a/app/views/_shared.py b/app/views/_shared.py index 73f2eb0..91bcc40 100644 --- a/app/views/_shared.py +++ b/app/views/_shared.py @@ -202,10 +202,11 @@ a { color: var(--accent); text-decoration: none; } SHARED_JS = """ // 任务状态中文标签 const STATUS_LABEL = { - queued: "排队中", extracting: "提取音频", transcribing: "语音识别", - segmenting: "断句重算", translating: "翻译中", done: "完成", failed: "失败" + uploading: "上传中", queued: "排队中", extracting: "提取音频", + transcribing: "语音识别", segmenting: "断句重算", translating: "翻译中", + done: "完成", failed: "失败" }; -const ACTIVE_STATES = ["queued","extracting","transcribing","segmenting","translating"]; +const ACTIVE_STATES = ["uploading","queued","extracting","transcribing","segmenting","translating"]; // HTML 转义(防 XSS) function escapeHtml(s) { @@ -222,6 +223,16 @@ function fmtBytes(n) { return u === 0 ? x + " B" : x.toFixed(1) + " " + units[u]; } +// 24 小时制时间格式化(不受浏览器 locale 影响,避免 am/pm 混淆) +function _pad2(n) { return n < 10 ? "0" + n : "" + n; } +function fmtTime24(d) { + return _pad2(d.getHours()) + ":" + _pad2(d.getMinutes()) + ":" + _pad2(d.getSeconds()); +} +function fmtDateTime24(d) { + return d.getFullYear() + "-" + _pad2(d.getMonth()+1) + "-" + _pad2(d.getDate()) + + " " + fmtTime24(d); +} + // 并发池:indices 中的每个元素交给 worker,最多 concurrency 个并发 async function runPool(indices, concurrency, worker) { let cursor = 0; diff --git a/app/views/history_html.py b/app/views/history_html.py index 1203787..1036692 100644 --- a/app/views/history_html.py +++ b/app/views/history_html.py @@ -68,7 +68,11 @@ function renderTable(tasks) {{ const tr = document.createElement("tr"); const label = STATUS_LABEL[task.status] || task.status; const stateClass = "st-" + (task.status === "done" ? "done" : task.status === "failed" ? "fail" : "running"); - const created = new Date(task.created_at).toLocaleString(); + let created; + try {{ created = fmtDateTime24(new Date(task.created_at + "Z")); }} + catch (e) {{ created = task.created_at; }} + + let idDisplay = task.is_upload ? "—" : "#" + task.id; let action; if (task.status === "done") {{ @@ -84,10 +88,13 @@ function renderTable(tasks) {{ let progress; if (task.status === "done") progress = "100%"; else if (task.status === "failed") progress = "—"; - else progress = `
${{task.progress.toFixed(0)}}%`; + else {{ + const pct = (task.progress == null) ? 0 : task.progress; + progress = `
${{pct.toFixed(0)}}%`; + }} tr.innerHTML = ` - #${{task.id}} + ${{idDisplay}} ${{escapeHtml(task.filename)}} ${{label}} ${{progress}} diff --git a/app/views/home_html.py b/app/views/home_html.py index 42e04a1..c743354 100644 --- a/app/views/home_html.py +++ b/app/views/home_html.py @@ -42,12 +42,13 @@ function renderTasks(tasks) {{ for (const task of tasks) {{ tasksEl.appendChild(makeTaskCard(task)); }} - // 对进行中的任务启动轮询 + // 对进行中的任务启动轮询(上传会话用 upload_id 去重,Task 用 id 去重) for (const task of tasks) {{ - if (ACTIVE_STATES.includes(task.status) && !pollingIds.has(task.id)) {{ - pollingIds.add(task.id); - pollTask(task.id); - }} + if (!ACTIVE_STATES.includes(task.status)) continue; + const pollKey = task.is_upload ? task.upload_id : String(task.id); + if (pollingIds.has(pollKey)) continue; + pollingIds.add(pollKey); + pollTask(task); }} }} @@ -61,6 +62,7 @@ function makeTaskCard(task) {{ const el = document.createElement("div"); el.className = "card task server-task"; el.dataset.taskId = task.id; + if (task.is_upload) el.dataset.uploadId = task.upload_id; el.innerHTML = renderTaskInner(task); return el; }} @@ -69,7 +71,11 @@ function renderTaskInner(task) {{ const label = STATUS_LABEL[task.status] || task.status; const stateClass = task.status === "done" ? "state-done" : task.status === "failed" ? "state-fail" : "state-running"; - const created = new Date(task.created_at).toLocaleString(); + + // 上传中:created_at 可能是 naive UTC,前端按 UTC 解析 + let created; + try {{ created = fmtDateTime24(new Date(task.created_at + "Z")); }} + catch (e) {{ created = task.created_at; }} let body; if (task.status === "done") {{ @@ -77,35 +83,72 @@ function renderTaskInner(task) {{ }} else if (task.status === "failed") {{ body = `
${{escapeHtml(task.error || "未知错误")}}
`; }} else {{ - body = `
${{task.progress.toFixed(0)}}%
`; + const pct = (task.progress == null) ? 0 : task.progress; + const sizeInfo = task.size_bytes ? ` · ${{fmtBytes(task.size_bytes)}}` : ""; + body = `
${{pct.toFixed(0)}}%${{sizeInfo}}
`; }} return `
- #${{task.id}} ${{escapeHtml(task.filename)}} + ${{task.is_upload ? "" : "#" + task.id + " "}}${{escapeHtml(task.filename)}} ${{label}}
${{body}}
${{created}}
`; }} -async function pollTask(taskId) {{ +async function pollTask(task) {{ + // 上传会话:轮询 upload status 接口;Task:轮询 task 接口 + const isUpload = task.is_upload === true && task.upload_id; + const url = isUpload + ? "/api/tasks/chunk-uploads/" + task.upload_id + "/status" + : "/api/tasks/" + task.id; + let backoff = POLL_INTERVAL; const tick = async () => {{ try {{ - const r = await fetch("/api/tasks/" + taskId); - if (!r.ok) {{ pollingIds.delete(taskId); return; }} - const task = await r.json(); - const card = tasksEl.querySelector('.server-task[data-task-id="' + taskId + '"]'); - if (!card) {{ pollingIds.delete(taskId); return; }} + const r = await fetch(url); + if (!r.ok) {{ pollingIds.delete(task.id); pollingIds.delete(task.upload_id); return; }} + const data = await r.json(); - card.innerHTML = renderTaskInner(task); - - if (task.status === "done" || task.status === "failed") {{ - pollingIds.delete(taskId); - return; + if (isUpload) {{ + // 上传会话:complete 后 task_id 出现,切换为 Task 轮询 + if (data.completed && data.task_id) {{ + pollingIds.delete(task.upload_id); + pollingIds.add(data.task_id); + pollTask({{ id: data.task_id, is_upload: false }}); + // 刷新列表让新 Task 卡片出现 + refreshList(); + return; + }} + // 更新上传进度 + const card = tasksEl.querySelector('.server-task[data-upload-id="' + task.upload_id + '"]'); + if (!card) {{ pollingIds.delete(task.upload_id); return; }} + const uploaded = (data.uploaded_chunks || []).length; + const total = data.total_chunks || 1; + const pct = Math.min(4, uploaded / total * 4); + const fakeTask = {{ + status: "uploading", progress: pct, is_upload: true, + upload_id: task.upload_id, filename: data.filename, + size_bytes: data.size_bytes, created_at: task.created_at, + }}; + card.innerHTML = renderTaskInner(fakeTask); + backoff = POLL_INTERVAL; + setTimeout(tick, backoff); + }} else {{ + // Task 轮询 + const card = tasksEl.querySelector('.server-task[data-task-id="' + task.id + '"]'); + if (!card) {{ pollingIds.delete(task.id); return; }} + card.innerHTML = renderTaskInner(data); + if (data.status === "done" || data.status === "failed") {{ + pollingIds.delete(task.id); + return; + }} + backoff = POLL_INTERVAL; + setTimeout(tick, backoff); }} - setTimeout(tick, POLL_INTERVAL); }} catch (e) {{ - setTimeout(tick, POLL_INTERVAL); + // 网络错误:指数退避,上限 30s + backoff = Math.min(backoff * 1.6, 30000); + setTimeout(tick, backoff); }} }}; tick(); diff --git a/app/views/logs_html.py b/app/views/logs_html.py index 4e1cba8..d057481 100644 --- a/app/views/logs_html.py +++ b/app/views/logs_html.py @@ -2,25 +2,34 @@ 共享 _shared.py 的 BASE_CSS / SHARED_JS。 页面专属:级别过滤按钮、自动刷新开关、清空、traceback 折叠。 + +显示策略: +- 默认 INFO(仅阶段转换/模型加载卸载/任务流转开始结束),可切 DEBUG 看进度详情。 +- 最新日志在顶部;用户向上滚动浏览历史时不会被自动刷新拉走(仅当停在顶部时跟随)。 +- 轮询带指数退避,连续失败时拉长间隔,避免服务不可达时打爆。 """ from __future__ import annotations from ._shared import render_page -POLL_INTERVAL_LOGS = 2000 +POLL_INTERVAL_LOGS = 2000 # 正常轮询间隔(ms) +POLL_MAX_INTERVAL = 30000 # 退避上限(ms) DEFAULT_TAIL = 200 _PAGE_JS = f""" -const POLL_INTERVAL = {POLL_INTERVAL_LOGS}; +const POLL_MIN_INTERVAL = {POLL_INTERVAL_LOGS}; +const POLL_MAX_INTERVAL = {POLL_MAX_INTERVAL}; const DEFAULT_TAIL = {DEFAULT_TAIL}; const API = "/api/logs"; -let currentLevel = "debug"; +let currentLevel = "info"; // 默认 INFO let autoRefresh = true; let timer = null; +let pollInterval = POLL_MIN_INTERVAL; // 动态退避 +let userScrolled = false; // 用户是否主动向下浏览历史 -const logsEl = document.getElementById("logs"); +const containerEl = document.getElementById("logs"); const emptyEl = document.getElementById("empty"); const statusEl = document.getElementById("status"); @@ -29,26 +38,34 @@ const LEVEL_CLASS = {{ CRITICAL: "st-fail" }}; +// 监听滚动:用户向下(往历史方向)浏览时暂停自动跟随,回到顶部则恢复 +containerEl.addEventListener("scroll", () => {{ + // scrollTop 越小越靠近顶部(最新)。接近顶部 = 用户在看最新 + userScrolled = containerEl.scrollTop > 4; +}}); + document.querySelectorAll(".filter").forEach(btn => {{ btn.addEventListener("click", () => {{ document.querySelectorAll(".filter").forEach(b => b.classList.remove("active")); btn.classList.add("active"); currentLevel = btn.dataset.level; - logsEl.innerHTML = ""; + containerEl.innerHTML = ""; + pollInterval = POLL_MIN_INTERVAL; fetchLogs(); }}); }}); document.getElementById("autorefresh").addEventListener("change", e => {{ autoRefresh = e.target.checked; - if (autoRefresh) fetchLogs(); else if (timer) {{ clearTimeout(timer); timer = null; }} + if (autoRefresh) {{ pollInterval = POLL_MIN_INTERVAL; fetchLogs(); }} + else if (timer) {{ clearTimeout(timer); timer = null; }} }}); document.getElementById("clear-btn").addEventListener("click", async () => {{ if (!confirm("确定清空所有日志缓冲?")) return; try {{ await fetch(API, {{ method: "DELETE" }}); - logsEl.innerHTML = ""; + containerEl.innerHTML = ""; statusEl.textContent = "已清空"; }} catch (e) {{ statusEl.textContent = "清空失败"; }} }}); @@ -56,57 +73,87 @@ document.getElementById("clear-btn").addEventListener("click", async () => {{ async function fetchLogs() {{ try {{ const r = await fetch(`${{API}}?level=${{currentLevel}}&tail=${{DEFAULT_TAIL}}`); - if (!r.ok) {{ statusEl.textContent = "HTTP " + r.status; scheduleNext(); return; }} + if (!r.ok) {{ statusEl.textContent = "HTTP " + r.status; scheduleNext(true); return; }} const data = await r.json(); renderLogs(data.logs); - statusEl.textContent = `${{data.count}} 条 · 更新 ${{new Date().toLocaleTimeString()}}`; - scheduleNext(); + statusEl.textContent = `${{data.count}} 条 · 更新 ${{fmtTime24(new Date())}}`; + pollInterval = POLL_MIN_INTERVAL; // 成功,重置间隔 + scheduleNext(false); }} catch (e) {{ statusEl.textContent = "获取失败"; - scheduleNext(); + scheduleNext(true); // 失败,退避 }} }} function renderLogs(logs) {{ if (!logs || logs.length === 0) {{ - if (logsEl.children.length === 0) emptyEl.style.display = "block"; + if (containerEl.children.length === 0) emptyEl.style.display = "block"; return; }} emptyEl.style.display = "none"; - const existing = new Set(Array.from(logsEl.children).map(el => el.dataset.key)); + + // 服务端已按时间升序返回(旧→新)。倒序渲染:最新在最上面。 + // 用 key=ts|level|logger|msg 去重(含 logger 名,避免不同模块同消息被误判重复) + const existing = new Set(Array.from(containerEl.children).map(el => el.dataset.key)); + const wasAtTop = containerEl.scrollTop <= 4; const frag = document.createDocumentFragment(); - for (const log of logs) {{ - const key = `${{log.ts}}|${{log.level}}|${{log.msg}}`; + + // 服务端按时间升序返回(旧→新)。倒序遍历 + appendChild → frag 内最新在前: + // logs[N-1](最新)先 append → frag 第一个;logs[0](最旧)最后 → frag 最后 + for (let i = logs.length - 1; i >= 0; i--) {{ + const log = logs[i]; + const key = `${{log.ts}}|${{log.level}}|${{log.logger}}|${{log.msg}}`; if (existing.has(key)) continue; - const row = document.createElement("div"); - row.className = "log-row " + (LEVEL_CLASS[log.level] || "st-running"); - row.dataset.key = key; - const hasTrace = !!log.traceback; - row.innerHTML = ` - ${{escapeHtml(log.ts)}} - ${{escapeHtml(log.level)}} - ${{escapeHtml(log.logger)}} - ${{escapeHtml(log.msg)}}${{hasTrace ? ' [traceback]' : ''}}`; - if (hasTrace) {{ - const pre = document.createElement("pre"); - pre.className = "trace"; - pre.textContent = log.traceback; - pre.style.display = "none"; - row.appendChild(pre); - row.querySelector(".trace-toggle").addEventListener("click", e => {{ - e.stopPropagation(); - pre.style.display = pre.style.display === "none" ? "block" : "none"; - }}); - }} - frag.appendChild(row); + frag.appendChild(buildRow(log, key)); + }} + + if (frag.children.length > 0) {{ + // 整块插入到容器顶部:最新日志在最上面 + containerEl.insertBefore(frag, containerEl.firstChild); + }} + + // 裁剪:超出上限删最旧(底部) + while (containerEl.children.length > DEFAULT_TAIL) {{ + containerEl.removeChild(containerEl.lastChild); + }} + + // 仅当用户停在顶部(看最新)时保持滚定在顶,否则不打扰浏览历史 + if (wasAtTop && !userScrolled) {{ + containerEl.scrollTop = 0; }} - logsEl.appendChild(frag); - while (logsEl.children.length > DEFAULT_TAIL) logsEl.removeChild(logsEl.firstChild); - logsEl.scrollTop = logsEl.scrollHeight; }} -function scheduleNext() {{ - if (autoRefresh) timer = setTimeout(fetchLogs, POLL_INTERVAL); +function buildRow(log, key) {{ + const row = document.createElement("div"); + row.className = "log-row " + (LEVEL_CLASS[log.level] || "st-running"); + row.dataset.key = key; + const hasTrace = !!log.traceback; + row.innerHTML = ` + ${{escapeHtml(log.ts)}} + ${{escapeHtml(log.level)}} + ${{escapeHtml(log.logger)}} + ${{escapeHtml(log.msg)}}${{hasTrace ? ' [traceback]' : ''}}`; + if (hasTrace) {{ + const pre = document.createElement("pre"); + pre.className = "trace"; + pre.textContent = log.traceback; + pre.style.display = "none"; + row.appendChild(pre); + row.querySelector(".trace-toggle").addEventListener("click", e => {{ + e.stopPropagation(); + pre.style.display = pre.style.display === "none" ? "block" : "none"; + }}); + }} + return row; +}} + +function scheduleNext(failed) {{ + if (!autoRefresh) return; + if (failed) {{ + // 指数退避:每次失败 ×1.6,上限 30s + pollInterval = Math.min(pollInterval * 1.6, POLL_MAX_INTERVAL); + }} + timer = setTimeout(fetchLogs, pollInterval); }} fetchLogs(); @@ -136,12 +183,12 @@ _PAGE_CSS = """ _BODY = """

日志

-

实时查看服务日志。debug=详细子步骤,info=仅阶段转换,error=完整错误。自动刷新每 2 秒。

+

默认显示 INFO(任务流转/模型加载卸载)。切 DEBUG 看进度详情,切 警告+/仅错误 过滤问题。自动刷新每 2 秒,失败自动退避。最新日志在顶部,向下浏览历史时不会被拉走。

- - + +
diff --git a/config.cpu.yaml b/config.cpu.yaml index 57bff19..3547ad2 100644 --- a/config.cpu.yaml +++ b/config.cpu.yaml @@ -30,6 +30,7 @@ asr: language: en word_timestamps: true vad_filter: true + batch_size: 8 # CPU 无 GPU 并行收益,保持小批 translation: model: Helsinki-NLP/opus-mt-en-zh # 最轻量英译中(~300MB;NLLB-600M 需 ~2.4GB,2GB 机 OOM) @@ -38,6 +39,7 @@ translation: tgt_lang: zho_Hans batch_size: 8 # opus-mt 轻量,batch 适中 max_length: 256 + sort_by_length: true # 排序分批(CPU 收益小,但无害) segmentation: max_words_per_line: 14 @@ -47,7 +49,7 @@ segmentation: logging: - level: info # debug=详细子步骤, info=仅阶段转换, error=完整 traceback + level: info # debug | info | warning | error(控制台最低级别) buffer_size: 2000 docs: diff --git a/config.example.yaml b/config.example.yaml index 775372c..f8cb167 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -1,5 +1,9 @@ -# audio2text 运行时配置。复制为 config.yaml 后填值。所有路径相对容器内 /app。 -# CPU dev / GPU prod 仅靠 device + model + compute_type 三项切换,代码不变。 +# audio2text 运行时配置模板。 +# 复制为 config.yaml 后填值。所有路径相对容器内 /app。 +# +# CPU dev / GPU prod 仅靠 asr + translation 两段的 device/model/compute_type/batch_size 切换。 +# 下方注释标注了 CPU(开发)与 GPU(生产)的推荐值对照。 +# 完整对照:CPU 用 ./setup.sh(默认),GPU 用 AUDIO2TEXT_VARIANT=gpu ./setup.sh。 server: host: 0.0.0.0 # 容器内对外监听(由 docker -p 映射到宿主) @@ -10,9 +14,9 @@ storage: upload_dir: /data/uploads # 上传视频落盘根目录 work_dir: /data/.work # 分片会话暂存 + 中间音频 output_dir: /data/outputs # 生成的 SRT 字幕输出 - chunk_bytes: 1048576 # 1 MiB 流式分片 + chunk_bytes: 1048576 # 1 MiB 服务端流式读缓冲(前端分片大小由前端定) chunk_session_ttl_seconds: 300 # 被放弃会话的存活秒数(后台 reaper 据此清理) - cache_retention_days: 7 # 任务产物保留天数,超期清理(字幕/中间音频/保留的原始视频+DB记录) + cache_retention_days: 7 # 任务产物保留天数,超期清理(0=禁用) cache_cleanup_interval_hours: 24 # 定时清理间隔(启动时跑一次,之后循环) processing: @@ -20,36 +24,40 @@ processing: keep_audio: false # 完成后是否保留中间 wav(默认删,只留字幕) asr: - # CPU dev:tiny.en + int8(Whisper 同系列最小,英文专用) - # GPU prod:large-v3-turbo + float16,3090 上几 GB 视频几分钟出字幕 - model: tiny.en - device: cpu # cpu | cuda - compute_type: int8 # cpu: int8;gpu: float16 - language: en # 仅英语 - word_timestamps: true # 词级时间戳:让断句精确而非纯匀速估算 - vad_filter: true # 过滤静音段,提升质量与速度 + # CPU dev:tiny.en + int8(Whisper 同系列最小,英文专用,~39M) + # GPU prod:large-v3-turbo + float16(8x 速度,质量接近 large-v3,~3GB 显存) + model: tiny.en # CPU: tiny.en | GPU: large-v3-turbo + device: cpu # cpu | cuda + compute_type: int8 # CPU: int8 | GPU: float16 + language: en # 仅英语 + word_timestamps: true # 词级时间戳:让断句精确而非纯匀速估算 + vad_filter: true # 过滤静音段,提升质量与速度 + batch_size: 8 # CPU: 8 | GPU: 16(BatchedInferencePipeline 批量解码音频块) translation: - model: facebook/nllb-200-distilled-1.3B - device: cpu # cpu | cuda - src_lang: eng_Latn # NLLB 语言码:英语 - tgt_lang: zho_Hans # NLLB 语言码:简体中文 - batch_size: 16 # 不与 ASR 共驻:翻译时显存独占,可用大 batch - max_length: 256 # 单条翻译最大 token + # CPU dev:Helsinki-NLP/opus-mt-en-zh(~300MB,2GB 内存可跑) + # GPU prod:facebook/nllb-200-distilled-1.3B(质量最好,~2.5GB 显存) + # 注意:model 与 device 必须配套——GPU 模型配 CPU device 会 OOM,反之浪费硬件。 + model: Helsinki-NLP/opus-mt-en-zh # CPU: opus-mt-en-zh | GPU: facebook/nllb-200-distilled-1.3B + device: cpu # cpu | cuda + src_lang: eng_Latn # NLLB 语言码:英语 + tgt_lang: zho_Hans # NLLB 语言码:简体中文 + batch_size: 8 # CPU: 8 | GPU: 32(不共驻时显存独占可用大 batch) + max_length: 256 # 单条翻译最大 token + sort_by_length: true # 按长度排序后分批,减少 padding 浪费(GPU 收益大) segmentation: - max_words_per_line: 14 # 单行最多词数,超出按逗号拆 - max_duration_seconds: 7.0 # 单条字幕最长 7 秒 - min_duration_seconds: 1.0 # 单条字幕最短 1 秒(太短则合并) - max_chars_per_line: 42 # SRT 规范:每行 ≤42 字符 - + max_words_per_line: 14 # 单行最多词数,超出按逗号拆 + max_duration_seconds: 7.0 # 单条字幕最长 7 秒 + min_duration_seconds: 1.0 # 单条字幕最短 1 秒(太短则合并) + max_chars_per_line: 42 # SRT 规范:每行 ≤42 字符 logging: - level: info # debug | info | warning | error(控制台 + 内存缓冲最低级别) - buffer_size: 2000 # /logs 页面内存缓冲条数 + level: info # debug | info | warning | error(控制台最低级别) + buffer_size: 2000 # /logs 页面内存缓冲条数 docs: enabled: true username: admin - password: "CHANGE_ME" # /docs Basic Auth,明文(root 持有,常量时间比较) + password: "CHANGE_ME" # 部署前务必修改;留空则禁用 /docs realm: "audio2text docs" diff --git a/config.gpu.yaml b/config.gpu.yaml index ded8daf..e793899 100644 --- a/config.gpu.yaml +++ b/config.gpu.yaml @@ -29,14 +29,16 @@ asr: language: en word_timestamps: true vad_filter: true + batch_size: 16 # BatchedInferencePipeline:每批解码 16 个 30s 音频块,填充 GPU translation: model: facebook/nllb-200-distilled-1.3B # 质量最好 device: cuda src_lang: eng_Latn tgt_lang: zho_Hans - batch_size: 16 # 不共驻时显存独占,大 batch + batch_size: 32 # 不共驻时显存独占,大 batch 填充 GPU max_length: 256 + sort_by_length: true # 按句子长度排序后分批,减少批内 padding 浪费 segmentation: max_words_per_line: 14 @@ -46,7 +48,7 @@ segmentation: logging: - level: info # debug=详细子步骤, info=仅阶段转换, error=完整 traceback + level: info # debug | info | warning | error(控制台最低级别) buffer_size: 2000 docs: diff --git a/docker-compose.yml b/docker-compose.yml index b53c6b9..c1319be 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,19 +1,46 @@ # audio2text — 双语字幕生成服务 # CPU dev / GPU prod 两套配置,按需切换 service。 # -# CPU 开发: -# docker compose up -d # 默认 cpu -# GPU 生产: -# docker compose -f docker-compose.yml up -d # 本机有 nvidia runtime 时自动用 gpu +# 开发(CPU,热重载,改代码零重建): +# docker compose --profile dev up -d --build +# CPU 生产: +# docker compose --profile cpu up -d --build +# GPU 生产(本机有 nvidia runtime 时): +# docker compose --profile gpu up -d --build # # 模型缓存(/models)跨容器复用,首次启动下载、之后秒起。 +# pip wheel 缓存由 BuildKit --mount=type=cache 管理,跨构建复用。 services: + # ---------------- 开发:CPU + 源码挂载 + uvicorn reload ---------------- + # 改 app/ 下代码无需重建镜像,保存即热重载。 + audio2text-cpu-dev: + build: + context: . + args: + VARIANT: cpu + target: dev + image: audio2text:cpu-dev + container_name: audio2text-dev + ports: + - "8000:8000" + volumes: + - ./app:/app/app + - ./data:/data + - ./models:/models + - ./config.yaml:/app/config.yaml:ro + environment: + CONFIG_PATH: /app/config.yaml + restart: unless-stopped + profiles: ["dev"] + + # ---------------- CPU 生产 ---------------- audio2text-cpu: build: context: . args: VARIANT: cpu + target: final image: audio2text:cpu container_name: audio2text ports: @@ -27,19 +54,21 @@ services: restart: unless-stopped profiles: ["cpu", ""] + # ---------------- GPU 生产 ---------------- audio2text-gpu: build: context: . args: VARIANT: gpu + target: final image: audio2text:gpu - container_name: audio2text + container_name: audio2text-gpu ports: - - "8000:8000" + - "8001:8000" volumes: - ./data:/data - ./models:/models - - ./config.yaml:/app/config.yaml:ro + - ./config.gpu.yaml:/app/config.yaml:ro environment: CONFIG_PATH: /app/config.yaml deploy: diff --git a/scripts/prefetch_models.py b/scripts/prefetch_models.py new file mode 100644 index 0000000..e60ada9 --- /dev/null +++ b/scripts/prefetch_models.py @@ -0,0 +1,137 @@ +"""预拉模型权重到本地缓存,避免容器启动时才下载(首次启动慢)。 + +在容器内执行(docker run + 挂载 ./models volume),复用镜像里的 huggingface_hub。 +读 config.yaml 拿 asr.model / translation.model,下载到 HF_HOME(= /models/huggingface)。 + +幂等:已下过的模型跳过(HF cache 命中检测)。 + +用法(经 scripts/prefetch_models.sh 包装): + ./scripts/prefetch_models.sh # 读 config.yaml + ./scripts/prefetch_models.sh config.gpu.yaml # 读指定配置 +""" + +from __future__ import annotations + +import os +import sys +import time +from pathlib import Path + + +def log(msg: str) -> None: + """带时间戳的日志(容器内无项目 logger,直接 print)。""" + ts = time.strftime("%H:%M:%S") + print(f"[{ts}] {msg}", flush=True) + + +def load_config(config_path: str) -> dict: + """读 YAML 配置,返回 asr.model / translation.model。""" + import yaml + with open(config_path, encoding="utf-8") as f: + return yaml.safe_load(f) + + +def resolve_asr_repo(model_name: str) -> str: + """faster-whisper 模型名 -> HF repo。 + + 预定义名(tiny.en 等)走 _MODELS 映射;已是 repo 路径(org/name)直接用。 + """ + from faster_whisper.utils import _MODELS + if "/" in model_name: + return model_name # 已是完整 repo 路径 + repo = _MODELS.get(model_name) + if repo is None: + raise ValueError(f"未知 faster-whisper 模型名:{model_name}(不在 _MODELS 里)") + return repo + + +def is_cached(repo_id: str, hf_home: str) -> bool: + """检测 HF cache 是否已有该模型(snapshot 目录存在且非空)。""" + # HF 缓存布局:/hub/models----/snapshots// + cache_dir = Path(hf_home) / "hub" / f"models--{repo_id.replace('/', '--')}" + snapshots = cache_dir / "snapshots" + if not snapshots.is_dir(): + return False + return any(snap.is_dir() and any(snap.iterdir()) for snap in snapshots.iterdir()) + + +def download_model(repo_id: str, hf_home: str) -> None: + """用 huggingface_hub 下载模型所有文件到 HF cache。""" + from huggingface_hub import snapshot_download + log(f" 下载 {repo_id} ...") + snapshot_download( + repo_id=repo_id, + local_dir=None, # 走标准 cache + cache_dir=Path(hf_home) / "hub", + ) + + +def prefetch_asr(model_name: str, hf_home: str) -> None: + """预拉 faster-whisper ASR 模型。""" + repo = resolve_asr_repo(model_name) + log(f"[ASR] model={model_name} repo={repo}") + if is_cached(repo, hf_home): + log(f" ✓ 已缓存,跳过") + return + download_model(repo, hf_home) + log(f" ✓ 完成") + + +def prefetch_translation(model_name: str, hf_home: str) -> None: + """预拉翻译模型(transformers pipeline 用的 HF repo)。""" + log(f"[翻译] model={model_name}") + if is_cached(model_name, hf_home): + log(f" ✓ 已缓存,跳过") + return + download_model(model_name, hf_home) + log(f" ✓ 完成") + + +def main() -> int: + config_path = os.environ.get("CONFIG_PATH", "/app/config.yaml") + if len(sys.argv) > 1: + config_path = sys.argv[1] + hf_home = os.environ.get("HF_HOME", "/models/huggingface") + + log(f"配置文件: {config_path}") + log(f"HF 缓存目录: {hf_home}") + + if not Path(config_path).is_file(): + log(f"✗ 配置文件不存在: {config_path}") + return 1 + + cfg = load_config(config_path) + asr_model = cfg.get("asr", {}).get("model", "tiny.en") + tr_model = cfg.get("translation", {}).get("model", "Helsinki-NLP/opus-mt-en-zh") + + log(f"ASR 模型: {asr_model}") + log(f"翻译模型: {tr_model}") + log("-" * 50) + + # 预拉 ASR + try: + prefetch_asr(asr_model, hf_home) + except Exception as e: + log(f"✗ ASR 模型预拉失败: {e}") + return 2 + + # 预拉翻译 + try: + prefetch_translation(tr_model, hf_home) + except Exception as e: + log(f"✗ 翻译模型预拉失败: {e}") + return 3 + + log("-" * 50) + log(f"全部完成。缓存位于: {hf_home}") + # 打印缓存大小 + try: + total = sum(f.stat().st_size for f in Path(hf_home).rglob("*") if f.is_file()) + log(f"缓存总大小: {total / 1024 / 1024 / 1024:.2f} GB") + except Exception: + pass + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/prefetch_models.sh b/scripts/prefetch_models.sh new file mode 100644 index 0000000..c853814 --- /dev/null +++ b/scripts/prefetch_models.sh @@ -0,0 +1,58 @@ +#!/usr/bin/env bash +# 预拉模型权重到 ./models volume,避免容器启动时才下载。 +# +# 用法: +# ./scripts/prefetch_models.sh # 读 config.yaml(当前激活配置) +# ./scripts/prefetch_models.sh config.gpu.yaml # 读指定配置文件 +# +# 原理:用已构建的 audio2text 镜像跑一次性容器,挂载 ./models volume, +# 执行 scripts/prefetch_models.py 把模型下到 HF cache。 +# 容器运行时 HF_HOME=/models/huggingface 命中缓存,秒级加载。 +# +# 幂等:已下过的模型跳过。换 config 的 model 后重跑即可补下新模型。 +set -euo pipefail + +cd "$(dirname "$0")/.." +ROOT="$(pwd)" + +CONFIG="${1:-config.yaml}" +VARIANT="${AUDIO2TEXT_VARIANT:-cpu}" + +# 选镜像:优先用 gpu 镜像(GPU 模型大,可能要 cuda 才能 correctly 下载部分文件), +# 否则 cpu。两者都能下 HF 模型(下载是纯网络 IO,不依赖 GPU)。 +IMAGE="audio2text:cpu" +if docker image inspect audio2text:gpu >/dev/null 2>&1 && [ "$VARIANT" = "gpu" ]; then + IMAGE="audio2text:gpu" +fi + +if ! docker image inspect "$IMAGE" >/dev/null 2>&1; then + echo "✗ 镜像 $IMAGE 不存在,请先 ./setup.sh 构建。" >&2 + exit 1 +fi + +if [ ! -f "$ROOT/$CONFIG" ]; then + echo "✗ 配置文件不存在: $ROOT/$CONFIG" >&2 + exit 1 +fi + +mkdir -p "$ROOT/models" + +echo "==> 预拉模型" +echo " 镜像: $IMAGE" +echo " 配置: $CONFIG" +echo " 缓存: $ROOT/models (volume)" +echo + +# 注意:CONFIG_PATH 指向容器内路径,挂载配置文件为只读 +MSYS_NO_PATHCONV=1 docker run --rm \ + -v "$ROOT/models:/models" \ + -v "$ROOT/$CONFIG:/app/config.yaml:ro" \ + -e HF_HOME=/models/huggingface \ + -e CT2_CACHE=/models/ctranslate2 \ + -e CONFIG_PATH=/app/config.yaml \ + "$IMAGE" \ + python /app/scripts/prefetch_models.py /app/config.yaml + +echo +echo "==> 完成。模型已缓存到 $ROOT/models/huggingface" +echo " 容器启动时会命中缓存,无需联网下载。"