Files
audio2text/app/services/pipeline.py
audio2text dev 78b87bfb24 feat: 进度条颗粒度优化 — ASR/翻译按批次细分进度
原来进度只在阶段结束时跳一次(ASR 5%→55%、翻译 60%→98%),长视频时
进度条卡住不动。现在每完成一个批次就更新进度。

ASR(asr_service.py):
- 用 info.duration_after_vad 算总 chunk 数(ceil(时长/30s))
- 消费生成器时按 seg.end 跨 30s chunk 边界回调 on_progress
- 30s 粒度自然节流,长视频约几十次更新

翻译(translate_service.py):
- _translate_sorted / _translate_sequential 每批完成后回调 on_progress(done, total)
- 批数循环前已知(len(batches)),每批都回调

pipeline.py:
- asr_phase / translate_phase 定义闭包回调,把 (current,total) 占比映射到
  对应进度区间(ASR 5%→55%、翻译 60%→98%),调 _set_status 写 DB(DEBUG 级)

验证(test/55.mp4, 640 条字幕):
- ASR: 5%→19.9%→23.8%→33.7%→48.6%→55% 平滑增长 
- 翻译: 20 批,65.3%→68.9%→74.2%→79.6%→84.9%→90.2%→93.8%→98% 
- 翻译卡在 60% 的 ~50s 是模型加载时间(卸载ASR+加载NLLB),属调度器层面
2026-07-06 23:18:30 +08:00

208 lines
8.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

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

"""转写管线各阶段:提取音频 → ASR → 断句 → 翻译 → 写 SRT。
阶段拆分供 scheduler 调度ffmpeg 阶段独立线程CPUGPU 阶段ASR+翻译)
由 scheduler 串行化并在切换模型前查队列复用已加载模型。
每个阶段函数接收 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 traceback
from dataclasses import asdict
from datetime import datetime, timezone
from pathlib import Path
from ..config import get_settings
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")
# 进度锚点(各阶段在 0-100 中的占比)
P_EXTRACT = 5.0
P_TRANSCRIBE_START = 5.0
P_TRANSCRIBE_END = 55.0
P_SEGMENT_START = 55.0
P_SEGMENT_END = 60.0
P_TRANSLATE_START = 60.0
P_TRANSLATE_END = 98.0
P_DONE = 100.0
# ---------------- 阶段 1提取音频CPU可并行----------------
def extract_phase(db, task) -> None:
"""ffmpeg 提取 16k mono wav写 task.wav_path状态置 transcribing。
由 scheduler 在独立线程调用(与 GPU 阶段并行)。提取完即可让 GPU 调度线程接管。
"""
s = get_settings()
src = s.upload_dir() / task.source_path
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)
# 按配置决定是否删原始视频(提取成功后)
if s.processing.delete_original_after_extract and src.is_file():
try:
src.unlink()
logger.info("已删除原始视频 %sdelete_original_after_extract=true", src.name)
except OSError as exc:
logger.warning("删除原始视频失败 %s: %s", src, exc)
# 记录 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)
# ---------------- 阶段 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)
# 进度回调:按 30s chunk 细分 ASR 进度5%→55% 区间)
# current=已处理 chunk 数, total=总 chunk 数
def on_asr_progress(current: int, total: int) -> None:
frac = current / total if total else 0.0
progress = P_TRANSCRIBE_START + (P_TRANSCRIBE_END - P_TRANSCRIBE_START) * frac
_set_status(db, task, STATUS_TRANSCRIBING, progress)
segments = asr_service.transcribe(wav_path, on_progress=on_asr_progress)
_set_status(db, task, STATUS_TRANSCRIBING, P_TRANSCRIBE_END,
note=f"识别出 {len(segments)}")
_set_status(db, task, STATUS_SEGMENTING, P_SEGMENT_START)
subs = segmenter.resegment(segments)
_set_status(db, task, STATUS_SEGMENTING, P_SEGMENT_END,
note=f"重组为 {len(subs)} 条字幕")
# 序列化断句结果供翻译阶段用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翻译 + 写 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)]
# 进度回调按批次细分翻译进度60%→98% 区间)
# done=已翻译条数, total=总条数
def on_translate_progress(done: int, total: int) -> None:
frac = done / total if total else 0.0
progress = P_TRANSLATE_START + (P_TRANSLATE_END - P_TRANSLATE_START) * frac
_set_status(db, task, STATUS_TRANSLATING, progress)
zh_texts = translate_service.translate(subs, on_progress=on_translate_progress)
_set_status(db, task, STATUS_TRANSLATING, P_TRANSLATE_END,
note=f"翻译 {len(zh_texts)}")
# 写 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)
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 = STATUS_DONE
task.progress = P_DONE
task.updated_at = datetime.now(timezone.utc)
db.commit()
# 清理中间音频
s = get_settings()
if not s.processing.keep_audio and task.wav_path:
try:
Path(task.wav_path).unlink()
except OSError:
pass
logger.info("任务 %d [完成] %s", task.id, bi_path.name)
# ---------------- 工具函数 ----------------
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)
db.commit()
if note:
logger.info("任务 %d [%s %.0f%%] %s", task.id, status, progress, note)
else:
logger.debug("任务 %d [%s %.0f%%]", task.id, status, progress)
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 = STATUS_FAILED
task.error = error[:2000]
task.updated_at = datetime.now(timezone.utc)
db.commit()
except Exception: # pragma: no cover
logger.error("写入失败状态时又失败:\n%s", traceback.format_exc())