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:
audio2text dev
2026-07-06 21:59:59 +08:00
parent 00e2a95fb7
commit 73110848f4
30 changed files with 1238 additions and 259 deletions

View File

@@ -8,3 +8,4 @@ logs/
*.pid *.pid
.git/ .git/
.gitignore .gitignore
test/

4
.gitignore vendored
View File

@@ -7,6 +7,7 @@ __pycache__/
# 运行时数据(走 docker volume不入库 # 运行时数据(走 docker volume不入库
data/ data/
data-gpu/
/models/ /models/
config.yaml config.yaml
*.pid *.pid
@@ -15,3 +16,6 @@ logs/
# 编辑器 # 编辑器
.vscode/ .vscode/
.idea/ .idea/
# 测试数据与脚本(本地测试用,不入库)
/test/

View File

@@ -1,44 +1,81 @@
# syntax=docker/dockerfile:1.7
# audio2text — 一份 DockerfileCPU(dev) / GPU(prod) 双形态。 # audio2text — 一份 DockerfileCPU(dev) / GPU(prod) 双形态。
# docker build --build-arg VARIANT=cpu -t audio2text:cpu . # docker build --build-arg VARIANT=cpu -t audio2text:cpu .
# docker build --build-arg VARIANT=gpu -t audio2text:gpu . # 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()=FalseCPU 预期)
# - GPU 版 torch.cuda.is_available()=TrueGPU 可调度)
# 因此 torch 必须按 VARIANT 分叉装不同 wheel绝不能跨 variant 共享依赖层。
# deps 阶段用 FROM base-${VARIANT}CPU/GPU 是两条独立构建链,各自装对应 torch。
ARG VARIANT=cpu ARG VARIANT=cpu
# ---------------------------------------------------------------------------
# 1. 基础镜像分叉CPU 用 slim PythonGPU 用 CUDA runtime + 手动装 python3.12
# ---------------------------------------------------------------------------
FROM python:3.12-slim AS base-cpu FROM python:3.12-slim AS base-cpu
# GPU 基础镜像带 CUDA 运行时torch 可装 CUDA 轮子
FROM nvidia/cuda:12.1.0-runtime-ubuntu22.04 AS base-gpu 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 \ 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/* \ && 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/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 ARG VARIANT
ENV VARIANT=${VARIANT} \ ENV PYTHONUNBUFFERED=1 \
PYTHONUNBUFFERED=1 \
PIP_NO_CACHE_DIR=1 \ PIP_NO_CACHE_DIR=1 \
HF_HOME=/models/huggingface \ HF_HOME=/models/huggingface \
CT2_CACHE=/models/ctranslate2 CT2_CACHE=/models/ctranslate2
# ffmpeg 是核心系统依赖,必须装 # ffmpeg 是核心系统依赖patchelf 用于修复 ctranslate2 可执行栈
RUN apt-get update -y && apt-get install -y --no-install-recommends \ 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 \ ffmpeg ca-certificates patchelf \
&& rm -rf /var/lib/apt/lists/* && rm -rf /var/lib/apt/lists/*
WORKDIR /app WORKDIR /app
COPY requirements.txt /app/requirements.txt COPY requirements.txt /app/requirements.txt
# 【顺序关键】必须先装 torch按 VARIANT 分叉),再装 requirements。
# CPU 装 CPU 版 torchGPU 走默认 index带 CUDA 的轮子) # 原因requirements 里的 transformers / accelerate 依赖 torch若先装 requirements
RUN if [ "$VARIANT" = "cpu" ]; then \ # 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()=FalseGPU 版 =True。不可共享。
RUN --mount=type=cache,target=/root/.cache/pip,sharing=locked \
pip install --upgrade pip && \ pip install --upgrade pip && \
if [ "$VARIANT" = "cpu" ]; then \
pip install torch --index-url https://download.pytorch.org/whl/cpu ; \ pip install torch --index-url https://download.pytorch.org/whl/cpu ; \
else \ else \
pip install --upgrade pip && \
pip install torch ; \ pip install torch ; \
fi 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 组合下 # ctranslate2 的 .so 带可执行栈标志PT_GNU_STACK X在某些内核 + Docker 组合下
# 会触发 "cannot enable executable stack as shared object requires"。用 patchelf # 会触发 "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; \ done; \
python -c "import ctranslate2; print('ctranslate2 stack fix verified', ctranslate2.__version__)" 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. finalprod从 deps 继承全部依赖,只加 app 代码。
# 放在最后 = `docker build` 默认 target。setup.sh / start.sh 依赖此默认行为。
# ---------------------------------------------------------------------------
FROM deps AS final
COPY app /app/app COPY app /app/app
COPY scripts /app/scripts
COPY config.example.yaml /app/config.example.yaml COPY config.example.yaml /app/config.example.yaml
# 运行时数据:上传 / 中间产物 / 输出字幕 / 模型缓存 # 运行时数据:上传 / 中间产物 / 输出字幕 / 模型缓存
# 全部走 volume镜像本身无状态、无敏感数据 # 全部走 volume镜像本身无状态、无敏感数据
VOLUME ["/data", "/models"] 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 EXPOSE 8000
CMD ["python", "-m", "uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"] CMD ["python", "-m", "uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"]

View File

@@ -54,8 +54,10 @@ Spring 风格分层HTTP 边界controllers与业务逻辑services
``` ```
audio2text/ audio2text/
├── Dockerfile # 一份 DockerfileARG VARIANT=cpu|gpu 出两个镜像 ├── Dockerfile # 一份 DockerfileARG VARIANT=cpu|gpu 出两个镜像
├── docker-compose.yml # cpu / gpu 个 profile ├── docker-compose.yml # cpu / gpu / dev 三个 profile
├── setup.sh / start.sh / stop.sh # 安装 / 启动 / 停止(包装 docker 命令) ├── setup.sh / start.sh / stop.sh # 安装 / 启动 / 停止(包装 docker 命令)
├── scripts/
│ └── prefetch_models.{sh,py} # 预拉模型权重到 ./models volume避免首次启动下载
├── requirements.txt ├── requirements.txt
├── config.example.yaml # 配置模板(复制为 config.yaml 后填值) ├── config.example.yaml # 配置模板(复制为 config.yaml 后填值)
├── README.md ├── README.md
@@ -149,6 +151,21 @@ cd /root/zikai/audio2text
首次启动会下载模型Whisper `tiny.en` ~39M + opus-mt ~300MB`./models` volume 首次启动会下载模型Whisper `tiny.en` ~39M + opus-mt ~300MB`./models` volume
之后秒起。启动后浏览器打开 `http://127.0.0.1:8000/`,拖入视频或音频文件即可。 之后秒起。启动后浏览器打开 `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 模型选型 ### CPU 模型选型
| 组件 | 模型 | 大小 | 说明 | | 组件 | 模型 | 大小 | 说明 |

View File

@@ -1,6 +1,7 @@
"""运行时配置:所有参数从 config.yaml 读取,对齐 server/config.py 的风格 """运行时配置:所有参数从 config.yaml 读取。
CPU dev / GPU prod 仅靠 device / model / compute_type 三项切换,代码完全不变。 CPU dev / GPU prod 仅靠 device / model / compute_type 三项切换,代码完全不变。
默认值与 config.example.yaml 对齐,确保无 yaml 时也能用最小配置启动。
""" """
from __future__ import annotations from __future__ import annotations
@@ -40,21 +41,23 @@ class ProcessingConfig(BaseModel):
class AsrConfig(BaseModel): class AsrConfig(BaseModel):
model: str = "small" model: str = "tiny.en" # CPU: tiny.enGPU: large-v3-turbo
device: str = "cpu" # cpu | cuda device: str = "cpu" # cpu | cuda
compute_type: str = "int8" # cpu: int8gpu: float16 compute_type: str = "int8" # cpu: int8gpu: float16
language: str = "en" language: str = "en"
word_timestamps: bool = True word_timestamps: bool = True
vad_filter: bool = True vad_filter: bool = True
batch_size: int = 8 # BatchedInferencePipeline 的音频块批大小GPU 建议 16
class TranslationConfig(BaseModel): class TranslationConfig(BaseModel):
model: str = "facebook/nllb-200-distilled-1.3B" model: str = "Helsinki-NLP/opus-mt-en-zh" # CPU: opus-mt轻量GPU: facebook/nllb-200-distilled-1.3B
device: str = "cpu" # cpu | cuda device: str = "cpu" # cpu | cuda
src_lang: str = "eng_Latn" src_lang: str = "eng_Latn" # NLLB 语言码:英语
tgt_lang: str = "zho_Hans" tgt_lang: str = "zho_Hans" # NLLB 语言码:简体中文
batch_size: int = 16 batch_size: int = 8 # 排序后单批最大条数CPU: 8GPU: 32显存独占可用大 batch
max_length: int = 256 max_length: int = 256 # 单条最大生成 token 数
sort_by_length: bool = True # 按句子长度排序后分批,减少批内 padding 浪费GPU 收益大)
class SegmentationConfig(BaseModel): class SegmentationConfig(BaseModel):
@@ -68,8 +71,8 @@ class LoggingConfig(BaseModel):
"""日志配置:控制台 + 内存缓冲的最低级别,以及缓冲条数。 """日志配置:控制台 + 内存缓冲的最低级别,以及缓冲条数。
分层语义: 分层语义:
- debug详细ffmpeg 命令、模型加载/卸载、转写逐段、翻译逐批进度 - debug进度详情(任务 [status pct%]、转写/翻译逐批统计、ffmpeg 命令
- info简略(仅任务阶段转换,如 "任务 N [transcribing 55%]" - info任务流转里程碑(音频提取/ASR/翻译 的开始与完成、模型加载与卸载
- error详细错误完整 traceback由 logger.exception 自带) - error详细错误完整 traceback由 logger.exception 自带)
""" """

View File

@@ -1,9 +1,9 @@
"""日志路由:查询内存日志缓冲,供 /logs 页面消费。 """日志路由:查询内存日志缓冲,供 /logs 页面消费。
分层语义: 分层语义:
- debug详细ffmpeg 命令、模型加载/卸载、转写逐段、翻译逐批进度) - debug进度详情(任务 [status pct%]、转写逐段统计、翻译逐批进度、ffmpeg 命令
- info简略(仅任务阶段转换,如 "任务 N [transcribing 55%]" - info任务流转里程碑(音频提取/ASR/翻译 的开始与完成、模型加载与卸载、批次汇总
- warning / error更详细error 含完整 traceback - warning / error异常与失败error 含完整 traceback
""" """
from __future__ import annotations from __future__ import annotations

View File

@@ -2,6 +2,7 @@
from __future__ import annotations from __future__ import annotations
from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
from fastapi import APIRouter, Depends, HTTPException, Query from fastapi import APIRouter, Depends, HTTPException, Query
@@ -10,7 +11,8 @@ from sqlalchemy.orm import Session
from ..config import get_settings from ..config import get_settings
from ..database import get_db 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 from ..schemas.task import TaskListResponse, TaskResponse
router = APIRouter(prefix="/api/tasks", tags=["task"]) 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="任务列表") @router.get("", response_model=TaskListResponse, summary="任务列表")
def list_tasks( def list_tasks(
limit: int = Query(100, ge=1, le=500), limit: int = Query(100, ge=1, le=500),
@@ -35,14 +62,30 @@ def list_tasks(
q: str = Query("", description="按文件名模糊搜索(大小写不敏感,匹配子串)"), q: str = Query("", description="按文件名模糊搜索(大小写不敏感,匹配子串)"),
db: Session = Depends(get_db), db: Session = Depends(get_db),
) -> TaskListResponse: ) -> TaskListResponse:
q_base = db.query(Task).order_by(Task.id.desc()) # 上传中的会话pending 状态)也作为虚拟任务返回,让前端能看到上传进度
upload_q = db.query(UploadSession).filter(UploadSession.status == "pending")
if q.strip(): if q.strip():
# SQLite 的 LIKE 默认大小写不敏感ASCIIilike 等价于 LIKE
like = f"%{q.strip()}%" like = f"%{q.strip()}%"
q_base = q_base.filter(Task.filename.ilike(like)) upload_q = upload_q.filter(UploadSession.filename.ilike(like))
total = q_base.count() uploads = upload_q.order_by(UploadSession.created_at.desc()).all()
tasks = q_base.offset(offset).limit(limit).all()
return TaskListResponse(tasks=[_to_resp(t) for t in tasks], total=total) # 已创建的 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="任务状态") @router.get("/{task_id}", response_model=TaskResponse, summary="任务状态")

View File

@@ -1,6 +1,6 @@
"""分片上传路由:建会话 / 查状态 / 传分片 / complete。 """分片上传路由:建会话 / 查状态 / 传分片 / complete。
协议与 server 完全一致,区别仅在 complete 后创建的是转写 Task 而非 UploadedFile complete 成功后创建转写 Task 并交由 scheduler 入队
""" """
from __future__ import annotations from __future__ import annotations
@@ -60,10 +60,11 @@ def complete_session(
"""拼接分片 + 创建转写任务 + 入队管线。 """拼接分片 + 创建转写任务 + 入队管线。
controller 负责编排service.complete 只管存储(拼接 + 建 Task controller 负责编排service.complete 只管存储(拼接 + 建 Task
管线触发由 controller 调用service 不依赖 pipeline避免循环依赖 管线触发由 controller 调用 schedulerffmpeg 异步 + GPU 串行调度)service 不依赖
scheduler避免循环依赖
""" """
resp = service.complete(upload_id) resp = service.complete(upload_id)
# 仅新建任务时入队(幂等 complete 返回的也是同一 task_idenqueue 幂等无副作用) # 仅新建任务时入队(幂等 complete 返回的也是同一 task_idenqueue 幂等无副作用)
from ..services.pipeline import enqueue_task from ..services.scheduler import enqueue_task
enqueue_task(resp.task_id) enqueue_task(resp.task_id)
return resp return resp

View File

@@ -51,12 +51,36 @@ class Base(DeclarativeBase):
def init_db_schema() -> None: def init_db_schema() -> None:
"""建表(幂等)。""" """建表(幂等)+ 旧库迁移(给 task 表补新字段)
SQLAlchemy 的 create_all 只建新表不改旧表。对已存在的 task 表,
需手动 ALTER TABLE ADD COLUMN 补 wav_path / segments_jsonnullable
"""
from .models.task import Task # noqa: F401 from .models.task import Task # noqa: F401
from .models.upload_session import UploadSession # noqa: F401 from .models.upload_session import UploadSession # noqa: F401
get_engine() engine = get_engine()
Base.metadata.create_all(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]: def get_db() -> Generator[Session, None, None]:

View File

@@ -76,6 +76,12 @@ async def lifespan(app: FastAPI):
await asyncio.to_thread(reap_stale_sessions) await asyncio.to_thread(reap_stale_sessions)
except Exception as exc: # pragma: no cover except Exception as exc: # pragma: no cover
logger.warning("启动 reaper 失败:%s", exc) 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() s = get_settings()
try: try:
@@ -144,7 +150,33 @@ def create_app() -> FastAPI:
@app.get("/health") @app.get("/health")
def health() -> dict: 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) @app.get("/history", response_class=HTMLResponse)
def history_page() -> HTMLResponse: def history_page() -> HTMLResponse:

View File

@@ -14,6 +14,17 @@ def _now() -> datetime:
return datetime.now(timezone.utc) 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): class Task(Base):
__tablename__ = "task" __tablename__ = "task"
@@ -23,7 +34,7 @@ class Task(Base):
# 视频在 upload_dir 下的相对路径(提取音频前后可能被删) # 视频在 upload_dir 下的相对路径(提取音频前后可能被删)
source_path: Mapped[str] = mapped_column(String(1024), nullable=False) source_path: Mapped[str] = mapped_column(String(1024), nullable=False)
# 状态机queued → extracting → transcribing → segmenting → translating → done | failed # 状态机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 进度 # 0-100 进度
progress: Mapped[float] = mapped_column(Float, nullable=False, default=0.0) 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) en_srt_path: Mapped[str | None] = mapped_column(String(1024), nullable=True)
zh_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 阶段写 JSONtranslate 阶段读。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) updated_at: Mapped[datetime] = mapped_column(DateTime, default=_now, onupdate=_now)
def __repr__(self) -> str: def __repr__(self) -> str:

View File

@@ -1,8 +1,9 @@
"""分片上传会话 ORM支撑断点续传。对齐 server 的 UploadSession 形态SQLite 版)。""" """分片上传会话 ORM支撑断点续传。"""
from __future__ import annotations from __future__ import annotations
import json import json
import logging
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Any from typing import Any
@@ -11,6 +12,8 @@ from sqlalchemy.orm import Mapped, mapped_column
from ..database import Base from ..database import Base
logger = logging.getLogger("audio2text.models")
def _now() -> datetime: def _now() -> datetime:
return datetime.now(timezone.utc) return datetime.now(timezone.utc)
@@ -26,7 +29,14 @@ class _IntList(TypeDecorator):
return json.dumps(value) if value is not None else None return json.dumps(value) if value is not None else None
def process_result_value(self, value: Any, dialect) -> list[int]: 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): class UploadSession(Base):
@@ -39,11 +49,13 @@ class UploadSession(Base):
total_chunks: Mapped[int] = mapped_column(Integer, nullable=False) total_chunks: Mapped[int] = mapped_column(Integer, nullable=False)
uploaded_chunks: Mapped[list[int]] = mapped_column(_IntList, default=list) uploaded_chunks: Mapped[list[int]] = mapped_column(_IntList, default=list)
# pending → completedcomplete 成功)| abandonedreaper 清理) # pending → completedcomplete 成功)| abandonedreaper 清理)
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 路径 # 拼接完成后的视频相对 upload_dir 路径
final_path: Mapped[str | None] = mapped_column(String(1024), nullable=True) final_path: Mapped[str | None] = mapped_column(String(1024), nullable=True)
# complete 后关联的 Task.id(直接引用,避免反向查找 source_path # complete 后关联的 Task.id。ondelete=SET NULLTask 被删时 session 保留task_id 置空
task_id: Mapped[int | None] = mapped_column(ForeignKey("task.id"), nullable=True) 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) updated_at: Mapped[datetime] = mapped_column(DateTime, default=_now, onupdate=_now)

View File

@@ -62,6 +62,11 @@ class TaskResponse(BaseModel):
error: str | None error: str | None
created_at: datetime created_at: datetime
updated_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): class TaskListResponse(BaseModel):

View File

@@ -1,4 +1,4 @@
"""/docs Basic Auth对齐 server/security.py。明文密码,常量时间比较。""" """/docs Basic Auth明文密码常量时间比较。"""
from __future__ import annotations from __future__ import annotations

View File

@@ -29,7 +29,8 @@ def transcribe(wav_path: Path) -> list[Segment]:
raise FileNotFoundError(f"音频不存在:{wav_path}") raise FileNotFoundError(f"音频不存在:{wav_path}")
model = get_model_manager().get_asr() model = get_model_manager().get_asr()
logger.debug("开始转写 %smodel=%s language=%s", wav_path.name, s.model, s.language) logger.debug("开始转写 %smodel=%s language=%s batch_size=%d",
wav_path.name, s.model, s.language, s.batch_size)
segments_gen, info = model.transcribe( segments_gen, info = model.transcribe(
str(wav_path), str(wav_path),
@@ -37,6 +38,8 @@ def transcribe(wav_path: Path) -> list[Segment]:
word_timestamps=s.word_timestamps, word_timestamps=s.word_timestamps,
vad_filter=s.vad_filter, vad_filter=s.vad_filter,
beam_size=5, beam_size=5,
batch_size=s.batch_size, # 批量解码:多音频块一次性送 GPU
without_timestamps=False, # BatchedInferencePipeline 默认 True需显式关闭以生成段级时间戳
) )
logger.debug( logger.debug(
"音频时长 %.1fs检测语言=%s(置信度 %.2f", "音频时长 %.1fs检测语言=%s(置信度 %.2f",

View File

@@ -54,15 +54,18 @@ class ModelManager:
if self._translator is not None: if self._translator is not None:
self._unload_translator_locked() self._unload_translator_locked()
s = get_settings().asr 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) 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 # 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, s.model, device=s.device, compute_type=s.compute_type,
) )
self._asr = BatchedInferencePipeline(model=whisper)
self._current = "asr" self._current = "asr"
logger.debug("ASR 模型已就绪") logger.info("ASR 模型已就绪batched, batch_size=%d)。", s.batch_size)
return self._asr return self._asr
def unload_asr(self) -> None: def unload_asr(self) -> None:
@@ -72,7 +75,7 @@ class ModelManager:
def _unload_asr_locked(self) -> None: def _unload_asr_locked(self) -> None:
if self._asr is None: if self._asr is None:
return return
logger.debug("卸载 ASR 模型(释放显存供翻译器独占)。") logger.info("卸载 ASR 模型(释放显存供翻译器独占)。")
# faster-whisper 模型无显式 closedel 即可 # faster-whisper 模型无显式 closedel 即可
del self._asr del self._asr
self._asr = None self._asr = None
@@ -89,7 +92,7 @@ class ModelManager:
if self._asr is not None: if self._asr is not None:
self._unload_asr_locked() self._unload_asr_locked()
s = get_settings().translation 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 from transformers import pipeline
self._translator = pipeline( self._translator = pipeline(
"translation", "translation",
@@ -97,9 +100,11 @@ class ModelManager:
device=s.device, device=s.device,
src_lang=s.src_lang, src_lang=s.src_lang,
tgt_lang=s.tgt_lang, tgt_lang=s.tgt_lang,
batch_size=s.batch_size, # pipeline 内部批大小,与 translate_service 分块对齐
) )
self._current = "translator" self._current = "translator"
logger.debug("翻译模型已就绪(独占显存,可用大 batch)。") logger.info("翻译模型已就绪(显存独占batch_size=%dsort_by_length=%s)。",
s.batch_size, s.sort_by_length)
return self._translator return self._translator
def unload_translator(self) -> None: def unload_translator(self) -> None:
@@ -109,7 +114,7 @@ class ModelManager:
def _unload_translator_locked(self) -> None: def _unload_translator_locked(self) -> None:
if self._translator is None: if self._translator is None:
return return
logger.debug("卸载翻译模型。") logger.info("卸载翻译模型。")
# 释放 pipeline 持有的 model + tokenizer # 释放 pipeline 持有的 model + tokenizer
mdl = getattr(self._translator, "model", None) mdl = getattr(self._translator, "model", None)
tok = getattr(self._translator, "tokenizer", None) tok = getattr(self._translator, "tokenizer", None)

View File

@@ -1,25 +1,31 @@
"""转写管线编排:提取音频 → ASR → 断句 → 翻译 → 写 SRT。 """转写管线各阶段:提取音频 → ASR → 断句 → 翻译 → 写 SRT。
任务状态机: 阶段拆分供 scheduler 调度ffmpeg 阶段独立线程CPUGPU 阶段ASR+翻译)
queued → extracting → transcribing → segmenting → translating → done 由 scheduler 串行化并在切换模型前查队列复用已加载模型。
任一步失败 → failed
模型不共驻ASR 与翻译分阶段加载,翻译时先卸载 Whisper 释放显存跑大 batch。 每个阶段函数接收 db Session + Task更新状态/进度,写入中间产物:
管线在后台线程跑(每个任务一个线程),通过 DB 更新状态与进度。 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 from __future__ import annotations
import json
import logging import logging
import threading
import traceback import traceback
from dataclasses import asdict
from datetime import datetime, timezone from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
from ..config import get_settings from ..config import get_settings
from ..database import get_session_local from ..models.task import (
from ..models.task import Task STATUS_EXTRACTING, STATUS_TRANSCRIBING, STATUS_SEGMENTING,
from . import ffmpeg_service, asr_service, segmenter, translate_service, srt_writer STATUS_TRANSLATING, STATUS_DONE, STATUS_FAILED,
)
from . import asr_service, ffmpeg_service, segmenter, srt_writer, translate_service
from .types import Subtitle from .types import Subtitle
logger = logging.getLogger("audio2text.pipeline") logger = logging.getLogger("audio2text.pipeline")
@@ -35,35 +41,17 @@ P_TRANSLATE_END = 98.0
P_DONE = 100.0 P_DONE = 100.0
def enqueue_task(task_id: int) -> None: # ---------------- 阶段 1提取音频CPU可并行----------------
"""把任务交给后台线程处理(非阻塞,供 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 extract_phase(db, task) -> None:
"""ffmpeg 提取 16k mono wav写 task.wav_path状态置 transcribing。
def _run_task(task_id: int) -> None: 由 scheduler 在独立线程调用(与 GPU 阶段并行)。提取完即可让 GPU 调度线程接管。
"""后台执行完整管线。所有异常都被捕获并写入 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() s = get_settings()
src = s.upload_dir() / task.source_path src = s.upload_dir() / task.source_path
logger.info("任务 %d [音频提取开始] %s", task.id, src.name)
# ---------- 1. 提取音频 ---------- _set_status(db, task, STATUS_EXTRACTING, P_EXTRACT)
_set_status(db, task, "extracting", P_EXTRACT)
wav = s.work_dir() / f"task_{task.id}.wav" wav = s.work_dir() / f"task_{task.id}.wav"
ffmpeg_service.extract_audio(src, wav) ffmpeg_service.extract_audio(src, wav)
@@ -75,61 +63,110 @@ def _pipeline(db, task: Task) -> None:
except OSError as exc: except OSError as exc:
logger.warning("删除原始视频失败 %s: %s", src, exc) logger.warning("删除原始视频失败 %s: %s", src, exc)
# ---------- 2. 语音识别 ---------- # 记录 wav 路径,状态置 transcribing待 GPU 调度线程接管 ASR
_set_status(db, task, "transcribing", P_TRANSCRIBE_START) task.wav_path = str(wav)
segments = asr_service.transcribe(wav) task.status = STATUS_TRANSCRIBING
_set_status(db, task, "transcribing", P_TRANSCRIBE_END, task.progress = P_TRANSCRIBE_START
task.updated_at = datetime.now(timezone.utc)
db.commit()
logger.info("任务 %d [音频提取完成] → 待 ASR%s", task.id, wav.name)
# ---------------- 阶段 2ASR + 断句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)}") note=f"识别出 {len(segments)}")
# ---------- 3. 断句 + 时间戳重算 ---------- _set_status(db, task, STATUS_SEGMENTING, P_SEGMENT_START)
_set_status(db, task, "segmenting", P_SEGMENT_START)
subs = segmenter.resegment(segments) 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)} 条字幕") note=f"重组为 {len(subs)} 条字幕")
# ---------- 4. 翻译 ---------- # 序列化断句结果供翻译阶段用dataclass → JSON
_set_status(db, task, "translating", P_TRANSLATE_START) task.segments_json = json.dumps([asdict(s) for s in subs], ensure_ascii=False)
# 翻译阶段model_manager 会自动卸载 ASR、加载翻译器独占显存 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翻译 + 写 SRTGPU----------------
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) 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)}") note=f"翻译 {len(zh_texts)}")
# ---------- 5. 写 SRT ---------- # 写 SRT
out_dir = s.output_dir() out_dir = s_output_dir()
stem = Path(task.filename).stem stem = Path(task.filename).stem
en_path = out_dir / f"task_{task.id}/{stem}.en.srt" en_path = out_dir / f"task_{task.id}/{stem}.en.srt"
zh_path = out_dir / f"task_{task.id}/{stem}.zh.srt" zh_path = out_dir / f"task_{task.id}/{stem}.zh.srt"
bi_path = out_dir / f"task_{task.id}/{stem}.srt" bi_path = out_dir / f"task_{task.id}/{stem}.srt"
srt_writer.write_srt(subs, en_path) srt_writer.write_srt(subs, en_path)
# 中文 SRT用译文 + 同时间戳)
zh_subs = [Subtitle(text=zh, start=sub.start, end=sub.end) zh_subs = [Subtitle(text=zh, start=sub.start, end=sub.end)
for zh, sub in zip(zh_texts, subs)] for zh, sub in zip(zh_texts, subs)]
srt_writer.write_srt(zh_subs, zh_path) srt_writer.write_srt(zh_subs, zh_path)
srt_writer.write_bilingual_srt(subs, zh_texts, bi_path) srt_writer.write_bilingual_srt(subs, zh_texts, bi_path)
# 记录相对路径
task.en_srt_path = str(en_path.relative_to(out_dir)) task.en_srt_path = str(en_path.relative_to(out_dir))
task.zh_srt_path = str(zh_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.bilingual_srt_path = str(bi_path.relative_to(out_dir))
task.status = "done" task.status = STATUS_DONE
task.progress = P_DONE task.progress = P_DONE
task.updated_at = datetime.now(timezone.utc) task.updated_at = datetime.now(timezone.utc)
db.commit() 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: try:
wav.unlink() Path(task.wav_path).unlink()
except OSError: except OSError:
pass 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.status = status
task.progress = progress task.progress = progress
task.updated_at = datetime.now(timezone.utc) 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: if note:
logger.info("任务 %d [%s %.0f%%] %s", task.id, status, progress, note) logger.info("任务 %d [%s %.0f%%] %s", task.id, status, progress, note)
else: 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: try:
task = db.get(Task, task_id) task = db.get(Task, task_id)
if task is None: if task is None:
return return
task.status = "failed" task.status = STATUS_FAILED
task.error = error[:2000] task.error = error[:2000]
task.updated_at = datetime.now(timezone.utc) task.updated_at = datetime.now(timezone.utc)
db.commit() db.commit()

236
app/services/scheduler.py Normal file
View 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

View File

@@ -1,12 +1,19 @@
"""翻译服务NLLB-200英译中。 """翻译服务NLLB-200英译中。
通过 model_manager 加载,确保 ASR 已卸载、翻译器独占显存,从而可用大 batch_size。 通过 model_manager 加载,确保 ASR 已卸载、翻译器独占显存,从而可用大 batch_size。
按字幕条目批量翻译,保留索引对应。
GPU 优化:按长度排序后分批翻译。
- 同一批内句子长度相近 → padding 浪费最小化 → GPU 有效计算占比提升
- 翻译完按原始下标散回,保证 zh_texts[i] 对应 subs[i](时间戳对齐不变)
- 批切分用 token 预算 + 条数上限双重约束:短句自动攒大批,长句自动拆小批
单条翻译失败时该位置回退为原英文。
""" """
from __future__ import annotations from __future__ import annotations
import logging import logging
import os
from ..config import get_settings from ..config import get_settings
from .model_manager import get_model_manager from .model_manager import get_model_manager
@@ -14,12 +21,15 @@ from .types import Subtitle
logger = logging.getLogger("audio2text.translate") logger = logging.getLogger("audio2text.translate")
# 估算每条字幕的 token 数:英文约 1 token/词,留 20% 余量覆盖标点/子词拆分
_TOKENS_PER_WORD = 1.2
def translate(subtitles: list[Subtitle]) -> list[str]: def translate(subtitles: list[Subtitle]) -> list[str]:
"""批量翻译英文字幕为中文。 """批量翻译英文字幕为中文。
Args: Args:
subtitles: 断句后的英文字幕条目 subtitles: 断句后的英文字幕条目(按时间顺序)
Returns: Returns:
list[str],与 subtitles 等长、顺序对应的中文译文。 list[str],与 subtitles 等长、顺序对应的中文译文。
@@ -30,31 +40,140 @@ def translate(subtitles: list[Subtitle]) -> list[str]:
s = get_settings().translation s = get_settings().translation
pipe = get_model_manager().get_translator() pipe = get_model_manager().get_translator()
batch = s.batch_size batch_size = s.batch_size
max_len = s.max_length max_len = s.max_length
# 取纯文本(去掉折行),避免翻译把换行符当语义 # 取纯文本(去掉折行),避免翻译把换行符当语义
texts = [sub.text.replace("\n", " ").strip() for sub in subtitles] 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] = [] # 环境变量覆盖:便于 A/B 基准对比test/bench_translate.py 用)
for i in range(0, len(texts), batch): if os.environ.get("TRANSLATE_NO_SORT") == "1":
chunk = texts[i:i + batch] sort_by_length = False
try:
out = pipe(chunk, max_length=max_len) if sort_by_length:
for item in out: results = _translate_sorted(pipe, texts, batch_size, max_len)
# pipeline 返回 [{"translation_text": "..."}] else:
results.append(item.get("translation_text", "").strip()) results = _translate_sequential(pipe, texts, batch_size, max_len)
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))
logger.debug("翻译完成:%d 条。", len(results)) logger.debug("翻译完成:%d 条。", len(results))
return 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

View File

@@ -7,8 +7,8 @@
... ...
<upload_dir>/<yyyy>/<mm>/<uuid>.<ext> complete 后的正式视频 <upload_dir>/<yyyy>/<mm>/<uuid>.<ext> complete 后的正式视频
与 server 的区别:视频无需 sha256 去重(每个视频都转写),complete 直接创建 Task complete 创建转写 Task,管线触发由 controller 调用 scheduler.enqueue_task
管线触发由 controller 调用 pipeline.enqueue_task本服务不依赖 pipeline 本服务不依赖 scheduler避免循环依赖
""" """
from __future__ import annotations 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: def complete(self, upload_id: str) -> CompleteResponse:
session = self._require_session(upload_id) session = self._require_session(upload_id)
@@ -135,6 +143,8 @@ class UploadService:
final_path = self._assemble(session) final_path = self._assemble(session)
rel = str(final_path.relative_to(self.upload_root)) rel = str(final_path.relative_to(self.upload_root))
# 单事务:建 Task + 更新 session 状态 + 关联 task_id 一次 commit
# 避免双 commit 之间崩溃产生孤儿 TaskTask 已建但 session.task_id 为空)
task = Task( task = Task(
filename=session.filename, filename=session.filename,
source_path=rel, source_path=rel,
@@ -144,14 +154,14 @@ class UploadService:
self.db.add(task) self.db.add(task)
session.status = "completed" session.status = "completed"
session.final_path = rel session.final_path = rel
session.task_id = None # 占位flush 后用 task.id 赋值
session.updated_at = datetime.now(timezone.utc) session.updated_at = datetime.now(timezone.utc)
self.db.commit() self.db.flush() # 拿到 task.id不 commit仍在事务内
self.db.refresh(task)
# 正向关联session → task替代旧的 source_path 反向查找)
session.task_id = task.id session.task_id = task.id
self.db.commit() self.db.commit()
self.db.refresh(task)
# 清理分片暂存 # 清理分片暂存commit 后,即使清理失败也不影响已建任务)
self._cleanup_session_dir(upload_id) self._cleanup_session_dir(upload_id)
logger.info("上传完成 task_id=%s file=%s size=%d", task.id, session.filename, session.size_bytes) 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: def _assemble(self, session: UploadSession) -> Path:
"""按 index 顺序拼接全部分片为正式视频文件。""" """按 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) now = datetime.now(timezone.utc)
sub = self.upload_root / f"{now:%Y}" / f"{now:%m}" sub = self.upload_root / f"{now:%Y}" / f"{now:%m}"
sub.mkdir(parents=True, exist_ok=True) sub.mkdir(parents=True, exist_ok=True)

View File

@@ -202,10 +202,11 @@ a { color: var(--accent); text-decoration: none; }
SHARED_JS = """ SHARED_JS = """
// 任务状态中文标签 // 任务状态中文标签
const STATUS_LABEL = { const STATUS_LABEL = {
queued: "排队中", extracting: "提取音频", transcribing: "语音识别", uploading: "上传中", queued: "排队中", extracting: "提取音频",
segmenting: "断句重算", translating: "翻译中", done: "完成", failed: "失败" transcribing: "语音识别", segmenting: "断句重算", translating: "翻译中",
done: "完成", failed: "失败"
}; };
const ACTIVE_STATES = ["queued","extracting","transcribing","segmenting","translating"]; const ACTIVE_STATES = ["uploading","queued","extracting","transcribing","segmenting","translating"];
// HTML 转义(防 XSS // HTML 转义(防 XSS
function escapeHtml(s) { function escapeHtml(s) {
@@ -222,6 +223,16 @@ function fmtBytes(n) {
return u === 0 ? x + " B" : x.toFixed(1) + " " + units[u]; 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 个并发 // 并发池indices 中的每个元素交给 worker最多 concurrency 个并发
async function runPool(indices, concurrency, worker) { async function runPool(indices, concurrency, worker) {
let cursor = 0; let cursor = 0;

View File

@@ -68,7 +68,11 @@ function renderTable(tasks) {{
const tr = document.createElement("tr"); const tr = document.createElement("tr");
const label = STATUS_LABEL[task.status] || task.status; const label = STATUS_LABEL[task.status] || task.status;
const stateClass = "st-" + (task.status === "done" ? "done" : task.status === "failed" ? "fail" : "running"); 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; let action;
if (task.status === "done") {{ if (task.status === "done") {{
@@ -84,10 +88,13 @@ function renderTable(tasks) {{
let progress; let progress;
if (task.status === "done") progress = "100%"; if (task.status === "done") progress = "100%";
else if (task.status === "failed") progress = ""; else if (task.status === "failed") progress = "";
else progress = `<div class="mini-bar"><div class="mini-fill" style="width:${{task.progress}}%"></div></div>${{task.progress.toFixed(0)}}%`; else {{
const pct = (task.progress == null) ? 0 : task.progress;
progress = `<div class="mini-bar"><div class="mini-fill" style="width:${{pct}}%"></div></div>${{pct.toFixed(0)}}%`;
}}
tr.innerHTML = ` tr.innerHTML = `
<td class="muted">#${{task.id}}</td> <td class="muted">${{idDisplay}}</td>
<td>${{escapeHtml(task.filename)}}</td> <td>${{escapeHtml(task.filename)}}</td>
<td><span class="status-tag ${{stateClass}}">${{label}}</span></td> <td><span class="status-tag ${{stateClass}}">${{label}}</span></td>
<td>${{progress}}</td> <td>${{progress}}</td>

View File

@@ -42,12 +42,13 @@ function renderTasks(tasks) {{
for (const task of tasks) {{ for (const task of tasks) {{
tasksEl.appendChild(makeTaskCard(task)); tasksEl.appendChild(makeTaskCard(task));
}} }}
// 对进行中的任务启动轮询 // 对进行中的任务启动轮询(上传会话用 upload_id 去重Task 用 id 去重)
for (const task of tasks) {{ for (const task of tasks) {{
if (ACTIVE_STATES.includes(task.status) && !pollingIds.has(task.id)) {{ if (!ACTIVE_STATES.includes(task.status)) continue;
pollingIds.add(task.id); const pollKey = task.is_upload ? task.upload_id : String(task.id);
pollTask(task.id); if (pollingIds.has(pollKey)) continue;
}} pollingIds.add(pollKey);
pollTask(task);
}} }}
}} }}
@@ -61,6 +62,7 @@ function makeTaskCard(task) {{
const el = document.createElement("div"); const el = document.createElement("div");
el.className = "card task server-task"; el.className = "card task server-task";
el.dataset.taskId = task.id; el.dataset.taskId = task.id;
if (task.is_upload) el.dataset.uploadId = task.upload_id;
el.innerHTML = renderTaskInner(task); el.innerHTML = renderTaskInner(task);
return el; return el;
}} }}
@@ -69,7 +71,11 @@ function renderTaskInner(task) {{
const label = STATUS_LABEL[task.status] || task.status; const label = STATUS_LABEL[task.status] || task.status;
const stateClass = task.status === "done" ? "state-done" const stateClass = task.status === "done" ? "state-done"
: task.status === "failed" ? "state-fail" : "state-running"; : 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; let body;
if (task.status === "done") {{ if (task.status === "done") {{
@@ -77,35 +83,72 @@ function renderTaskInner(task) {{
}} else if (task.status === "failed") {{ }} else if (task.status === "failed") {{
body = `<div class="task-meta fail-msg">${{escapeHtml(task.error || "未知错误")}}</div>`; body = `<div class="task-meta fail-msg">${{escapeHtml(task.error || "未知错误")}}</div>`;
}} else {{ }} else {{
body = `<div class="bar"><div class="fill" style="width:${{task.progress}}%"></div><span class="pct">${{task.progress.toFixed(0)}}%</span></div>`; const pct = (task.progress == null) ? 0 : task.progress;
const sizeInfo = task.size_bytes ? ` · ${{fmtBytes(task.size_bytes)}}` : "";
body = `<div class="bar"><div class="fill" style="width:${{pct}}%"></div><span class="pct">${{pct.toFixed(0)}}%${{sizeInfo}}</span></div>`;
}} }}
return ` return `
<div class="task-head"> <div class="task-head">
<span class="fname">#${{task.id}} ${{escapeHtml(task.filename)}}</span> <span class="fname">${{task.is_upload ? "" : "#" + task.id + " "}}${{escapeHtml(task.filename)}}</span>
<span class="fstate ${{stateClass}}">${{label}}</span> <span class="fstate ${{stateClass}}">${{label}}</span>
</div> </div>
${{body}} ${{body}}
<div class="task-time">${{created}}</div>`; <div class="task-time">${{created}}</div>`;
}} }}
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 () => {{ const tick = async () => {{
try {{ try {{
const r = await fetch("/api/tasks/" + taskId); const r = await fetch(url);
if (!r.ok) {{ pollingIds.delete(taskId); return; }} if (!r.ok) {{ pollingIds.delete(task.id); pollingIds.delete(task.upload_id); return; }}
const task = await r.json(); const data = await r.json();
const card = tasksEl.querySelector('.server-task[data-task-id="' + taskId + '"]');
if (!card) {{ pollingIds.delete(taskId); return; }}
card.innerHTML = renderTaskInner(task); if (isUpload) {{
// 上传会话complete 后 task_id 出现,切换为 Task 轮询
if (task.status === "done" || task.status === "failed") {{ if (data.completed && data.task_id) {{
pollingIds.delete(taskId); pollingIds.delete(task.upload_id);
pollingIds.add(data.task_id);
pollTask({{ id: data.task_id, is_upload: false }});
// 刷新列表让新 Task 卡片出现
refreshList();
return; return;
}} }}
setTimeout(tick, POLL_INTERVAL); // 更新上传进度
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);
}}
}} catch (e) {{ }} catch (e) {{
setTimeout(tick, POLL_INTERVAL); // 网络错误:指数退避,上限 30s
backoff = Math.min(backoff * 1.6, 30000);
setTimeout(tick, backoff);
}} }}
}}; }};
tick(); tick();

View File

@@ -2,25 +2,34 @@
共享 _shared.py 的 BASE_CSS / SHARED_JS。 共享 _shared.py 的 BASE_CSS / SHARED_JS。
页面专属级别过滤按钮、自动刷新开关、清空、traceback 折叠。 页面专属级别过滤按钮、自动刷新开关、清空、traceback 折叠。
显示策略:
- 默认 INFO仅阶段转换/模型加载卸载/任务流转开始结束),可切 DEBUG 看进度详情。
- 最新日志在顶部;用户向上滚动浏览历史时不会被自动刷新拉走(仅当停在顶部时跟随)。
- 轮询带指数退避,连续失败时拉长间隔,避免服务不可达时打爆。
""" """
from __future__ import annotations from __future__ import annotations
from ._shared import render_page from ._shared import render_page
POLL_INTERVAL_LOGS = 2000 POLL_INTERVAL_LOGS = 2000 # 正常轮询间隔ms
POLL_MAX_INTERVAL = 30000 # 退避上限ms
DEFAULT_TAIL = 200 DEFAULT_TAIL = 200
_PAGE_JS = f""" _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 DEFAULT_TAIL = {DEFAULT_TAIL};
const API = "/api/logs"; const API = "/api/logs";
let currentLevel = "debug"; let currentLevel = "info"; // 默认 INFO
let autoRefresh = true; let autoRefresh = true;
let timer = null; 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 emptyEl = document.getElementById("empty");
const statusEl = document.getElementById("status"); const statusEl = document.getElementById("status");
@@ -29,26 +38,34 @@ const LEVEL_CLASS = {{
CRITICAL: "st-fail" CRITICAL: "st-fail"
}}; }};
// 监听滚动:用户向下(往历史方向)浏览时暂停自动跟随,回到顶部则恢复
containerEl.addEventListener("scroll", () => {{
// scrollTop 越小越靠近顶部(最新)。接近顶部 = 用户在看最新
userScrolled = containerEl.scrollTop > 4;
}});
document.querySelectorAll(".filter").forEach(btn => {{ document.querySelectorAll(".filter").forEach(btn => {{
btn.addEventListener("click", () => {{ btn.addEventListener("click", () => {{
document.querySelectorAll(".filter").forEach(b => b.classList.remove("active")); document.querySelectorAll(".filter").forEach(b => b.classList.remove("active"));
btn.classList.add("active"); btn.classList.add("active");
currentLevel = btn.dataset.level; currentLevel = btn.dataset.level;
logsEl.innerHTML = ""; containerEl.innerHTML = "";
pollInterval = POLL_MIN_INTERVAL;
fetchLogs(); fetchLogs();
}}); }});
}}); }});
document.getElementById("autorefresh").addEventListener("change", e => {{ document.getElementById("autorefresh").addEventListener("change", e => {{
autoRefresh = e.target.checked; 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 () => {{ document.getElementById("clear-btn").addEventListener("click", async () => {{
if (!confirm("确定清空所有日志缓冲?")) return; if (!confirm("确定清空所有日志缓冲?")) return;
try {{ try {{
await fetch(API, {{ method: "DELETE" }}); await fetch(API, {{ method: "DELETE" }});
logsEl.innerHTML = ""; containerEl.innerHTML = "";
statusEl.textContent = "已清空"; statusEl.textContent = "已清空";
}} catch (e) {{ statusEl.textContent = "清空失败"; }} }} catch (e) {{ statusEl.textContent = "清空失败"; }}
}}); }});
@@ -56,28 +73,57 @@ document.getElementById("clear-btn").addEventListener("click", async () => {{
async function fetchLogs() {{ async function fetchLogs() {{
try {{ try {{
const r = await fetch(`${{API}}?level=${{currentLevel}}&tail=${{DEFAULT_TAIL}}`); 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(); const data = await r.json();
renderLogs(data.logs); renderLogs(data.logs);
statusEl.textContent = `${{data.count}} 条 · 更新 ${{new Date().toLocaleTimeString()}}`; statusEl.textContent = `${{data.count}} 条 · 更新 ${{fmtTime24(new Date())}}`;
scheduleNext(); pollInterval = POLL_MIN_INTERVAL; // 成功,重置间隔
scheduleNext(false);
}} catch (e) {{ }} catch (e) {{
statusEl.textContent = "获取失败"; statusEl.textContent = "获取失败";
scheduleNext(); scheduleNext(true); // 失败,退避
}} }}
}} }}
function renderLogs(logs) {{ function renderLogs(logs) {{
if (!logs || logs.length === 0) {{ if (!logs || logs.length === 0) {{
if (logsEl.children.length === 0) emptyEl.style.display = "block"; if (containerEl.children.length === 0) emptyEl.style.display = "block";
return; return;
}} }}
emptyEl.style.display = "none"; 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(); 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; if (existing.has(key)) continue;
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;
}}
}}
function buildRow(log, key) {{
const row = document.createElement("div"); const row = document.createElement("div");
row.className = "log-row " + (LEVEL_CLASS[log.level] || "st-running"); row.className = "log-row " + (LEVEL_CLASS[log.level] || "st-running");
row.dataset.key = key; row.dataset.key = key;
@@ -98,15 +144,16 @@ function renderLogs(logs) {{
pre.style.display = pre.style.display === "none" ? "block" : "none"; pre.style.display = pre.style.display === "none" ? "block" : "none";
}}); }});
}} }}
frag.appendChild(row); return row;
}}
logsEl.appendChild(frag);
while (logsEl.children.length > DEFAULT_TAIL) logsEl.removeChild(logsEl.firstChild);
logsEl.scrollTop = logsEl.scrollHeight;
}} }}
function scheduleNext() {{ function scheduleNext(failed) {{
if (autoRefresh) timer = setTimeout(fetchLogs, POLL_INTERVAL); if (!autoRefresh) return;
if (failed) {{
// 指数退避:每次失败 ×1.6,上限 30s
pollInterval = Math.min(pollInterval * 1.6, POLL_MAX_INTERVAL);
}}
timer = setTimeout(fetchLogs, pollInterval);
}} }}
fetchLogs(); fetchLogs();
@@ -136,12 +183,12 @@ _PAGE_CSS = """
_BODY = """ _BODY = """
<h1>日志</h1> <h1>日志</h1>
<p class="sub">实时查看服务日志。debug=详细子步骤info=仅阶段转换error=完整错误。自动刷新每 2 秒。</p> <p class="sub">默认显示 INFO任务流转/模型加载卸载)。切 DEBUG 看进度详情,切 警告+/仅错误 过滤问题。自动刷新每 2 秒,失败自动退避。最新日志在顶部,向下浏览历史时不会被拉走。</p>
<div class="toolbar"> <div class="toolbar">
<div class="filters"> <div class="filters">
<button class="filter active" data-level="debug">全部 (DEBUG)</button> <button class="filter" data-level="debug">全部 (DEBUG)</button>
<button class="filter" data-level="info">简略 (INFO)</button> <button class="filter active" data-level="info">简略 (INFO)</button>
<button class="filter" data-level="warning">警告+</button> <button class="filter" data-level="warning">警告+</button>
<button class="filter" data-level="error">仅错误</button> <button class="filter" data-level="error">仅错误</button>
</div> </div>

View File

@@ -30,6 +30,7 @@ asr:
language: en language: en
word_timestamps: true word_timestamps: true
vad_filter: true vad_filter: true
batch_size: 8 # CPU 无 GPU 并行收益,保持小批
translation: translation:
model: Helsinki-NLP/opus-mt-en-zh # 最轻量英译中(~300MBNLLB-600M 需 ~2.4GB2GB 机 OOM model: Helsinki-NLP/opus-mt-en-zh # 最轻量英译中(~300MBNLLB-600M 需 ~2.4GB2GB 机 OOM
@@ -38,6 +39,7 @@ translation:
tgt_lang: zho_Hans tgt_lang: zho_Hans
batch_size: 8 # opus-mt 轻量batch 适中 batch_size: 8 # opus-mt 轻量batch 适中
max_length: 256 max_length: 256
sort_by_length: true # 排序分批CPU 收益小,但无害)
segmentation: segmentation:
max_words_per_line: 14 max_words_per_line: 14
@@ -47,7 +49,7 @@ segmentation:
logging: logging:
level: info # debug=详细子步骤, info=仅阶段转换, error=完整 traceback level: info # debug | info | warning | error控制台最低级别
buffer_size: 2000 buffer_size: 2000
docs: docs:

View File

@@ -1,5 +1,9 @@
# audio2text 运行时配置。复制为 config.yaml 后填值。所有路径相对容器内 /app # audio2text 运行时配置模板
# CPU dev / GPU prod 仅靠 device + model + compute_type 三项切换,代码不变 # 复制为 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: server:
host: 0.0.0.0 # 容器内对外监听(由 docker -p 映射到宿主) host: 0.0.0.0 # 容器内对外监听(由 docker -p 映射到宿主)
@@ -10,9 +14,9 @@ storage:
upload_dir: /data/uploads # 上传视频落盘根目录 upload_dir: /data/uploads # 上传视频落盘根目录
work_dir: /data/.work # 分片会话暂存 + 中间音频 work_dir: /data/.work # 分片会话暂存 + 中间音频
output_dir: /data/outputs # 生成的 SRT 字幕输出 output_dir: /data/outputs # 生成的 SRT 字幕输出
chunk_bytes: 1048576 # 1 MiB 流式分片 chunk_bytes: 1048576 # 1 MiB 服务端流式读缓冲(前端分片大小由前端定)
chunk_session_ttl_seconds: 300 # 被放弃会话的存活秒数(后台 reaper 据此清理) chunk_session_ttl_seconds: 300 # 被放弃会话的存活秒数(后台 reaper 据此清理)
cache_retention_days: 7 # 任务产物保留天数,超期清理(字幕/中间音频/保留的原始视频+DB记录 cache_retention_days: 7 # 任务产物保留天数,超期清理(0=禁用
cache_cleanup_interval_hours: 24 # 定时清理间隔(启动时跑一次,之后循环) cache_cleanup_interval_hours: 24 # 定时清理间隔(启动时跑一次,之后循环)
processing: processing:
@@ -20,22 +24,27 @@ processing:
keep_audio: false # 完成后是否保留中间 wav默认删只留字幕 keep_audio: false # 完成后是否保留中间 wav默认删只留字幕
asr: asr:
# CPU devtiny.en + int8Whisper 同系列最小,英文专用) # CPU devtiny.en + int8Whisper 同系列最小,英文专用~39M
# GPU prodlarge-v3-turbo + float163090 上几 GB 视频几分钟出字幕 # GPU prodlarge-v3-turbo + float168x 速度,质量接近 large-v3~3GB 显存)
model: tiny.en model: tiny.en # CPU: tiny.en | GPU: large-v3-turbo
device: cpu # cpu | cuda device: cpu # cpu | cuda
compute_type: int8 # cpu: int8gpu: float16 compute_type: int8 # CPU: int8 | GPU: float16
language: en # 仅英语 language: en # 仅英语
word_timestamps: true # 词级时间戳:让断句精确而非纯匀速估算 word_timestamps: true # 词级时间戳:让断句精确而非纯匀速估算
vad_filter: true # 过滤静音段,提升质量与速度 vad_filter: true # 过滤静音段,提升质量与速度
batch_size: 8 # CPU: 8 | GPU: 16BatchedInferencePipeline 批量解码音频块)
translation: translation:
model: facebook/nllb-200-distilled-1.3B # CPU devHelsinki-NLP/opus-mt-en-zh~300MB2GB 内存可跑)
# GPU prodfacebook/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 device: cpu # cpu | cuda
src_lang: eng_Latn # NLLB 语言码:英语 src_lang: eng_Latn # NLLB 语言码:英语
tgt_lang: zho_Hans # NLLB 语言码:简体中文 tgt_lang: zho_Hans # NLLB 语言码:简体中文
batch_size: 16 # 不与 ASR 共驻:翻译时显存独占可用大 batch batch_size: 8 # CPU: 8 | GPU: 32(不共驻时显存独占可用大 batch
max_length: 256 # 单条翻译最大 token max_length: 256 # 单条翻译最大 token
sort_by_length: true # 按长度排序后分批,减少 padding 浪费GPU 收益大)
segmentation: segmentation:
max_words_per_line: 14 # 单行最多词数,超出按逗号拆 max_words_per_line: 14 # 单行最多词数,超出按逗号拆
@@ -43,13 +52,12 @@ segmentation:
min_duration_seconds: 1.0 # 单条字幕最短 1 秒(太短则合并) min_duration_seconds: 1.0 # 单条字幕最短 1 秒(太短则合并)
max_chars_per_line: 42 # SRT 规范:每行 ≤42 字符 max_chars_per_line: 42 # SRT 规范:每行 ≤42 字符
logging: logging:
level: info # debug | info | warning | error控制台 + 内存缓冲最低级别) level: info # debug | info | warning | error控制台最低级别
buffer_size: 2000 # /logs 页面内存缓冲条数 buffer_size: 2000 # /logs 页面内存缓冲条数
docs: docs:
enabled: true enabled: true
username: admin username: admin
password: "CHANGE_ME" # /docs Basic Auth明文root 持有,常量时间比较) password: "CHANGE_ME" # 部署前务必修改;留空则禁用 /docs
realm: "audio2text docs" realm: "audio2text docs"

View File

@@ -29,14 +29,16 @@ asr:
language: en language: en
word_timestamps: true word_timestamps: true
vad_filter: true vad_filter: true
batch_size: 16 # BatchedInferencePipeline每批解码 16 个 30s 音频块,填充 GPU
translation: translation:
model: facebook/nllb-200-distilled-1.3B # 质量最好 model: facebook/nllb-200-distilled-1.3B # 质量最好
device: cuda device: cuda
src_lang: eng_Latn src_lang: eng_Latn
tgt_lang: zho_Hans tgt_lang: zho_Hans
batch_size: 16 # 不共驻时显存独占,大 batch batch_size: 32 # 不共驻时显存独占,大 batch 填充 GPU
max_length: 256 max_length: 256
sort_by_length: true # 按句子长度排序后分批,减少批内 padding 浪费
segmentation: segmentation:
max_words_per_line: 14 max_words_per_line: 14
@@ -46,7 +48,7 @@ segmentation:
logging: logging:
level: info # debug=详细子步骤, info=仅阶段转换, error=完整 traceback level: info # debug | info | warning | error控制台最低级别
buffer_size: 2000 buffer_size: 2000
docs: docs:

View File

@@ -1,19 +1,46 @@
# audio2text — 双语字幕生成服务 # audio2text — 双语字幕生成服务
# CPU dev / GPU prod 两套配置,按需切换 service。 # CPU dev / GPU prod 两套配置,按需切换 service。
# #
# CPU 开发: # 开发CPU热重载改代码零重建
# docker compose up -d # 默认 cpu # docker compose --profile dev up -d --build
# GPU 生产: # CPU 生产:
# docker compose -f docker-compose.yml up -d # 本机有 nvidia runtime 时自动用 gpu # docker compose --profile cpu up -d --build
# GPU 生产(本机有 nvidia runtime 时):
# docker compose --profile gpu up -d --build
# #
# 模型缓存(/models跨容器复用首次启动下载、之后秒起。 # 模型缓存(/models跨容器复用首次启动下载、之后秒起。
# pip wheel 缓存由 BuildKit --mount=type=cache 管理,跨构建复用。
services: 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: audio2text-cpu:
build: build:
context: . context: .
args: args:
VARIANT: cpu VARIANT: cpu
target: final
image: audio2text:cpu image: audio2text:cpu
container_name: audio2text container_name: audio2text
ports: ports:
@@ -27,19 +54,21 @@ services:
restart: unless-stopped restart: unless-stopped
profiles: ["cpu", ""] profiles: ["cpu", ""]
# ---------------- GPU 生产 ----------------
audio2text-gpu: audio2text-gpu:
build: build:
context: . context: .
args: args:
VARIANT: gpu VARIANT: gpu
target: final
image: audio2text:gpu image: audio2text:gpu
container_name: audio2text container_name: audio2text-gpu
ports: ports:
- "8000:8000" - "8001:8000"
volumes: volumes:
- ./data:/data - ./data:/data
- ./models:/models - ./models:/models
- ./config.yaml:/app/config.yaml:ro - ./config.gpu.yaml:/app/config.yaml:ro
environment: environment:
CONFIG_PATH: /app/config.yaml CONFIG_PATH: /app/config.yaml
deploy: deploy:

137
scripts/prefetch_models.py Normal file
View File

@@ -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 缓存布局:<hf_home>/hub/models--<org>--<name>/snapshots/<hash>/
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())

View File

@@ -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 " 容器启动时会命中缓存,无需联网下载。"