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),属调度器层面
This commit is contained in:
audio2text dev
2026-07-06 23:18:30 +08:00
parent 5f6a242114
commit 78b87bfb24
3 changed files with 69 additions and 15 deletions

View File

@@ -6,7 +6,9 @@ CPU dev: tiny.en + int8GPU prod: large-v3-turbo + float16。同一份代码
from __future__ import annotations
import logging
import math
from pathlib import Path
from typing import Callable
from ..config import get_settings
from .model_manager import get_model_manager
@@ -14,12 +16,20 @@ from .types import Segment, Word
logger = logging.getLogger("audio2text.asr")
# Whisper 默认 chunk_lengthBatchedInferencePipeline 按 30s 窗口切音频
_CHUNK_SECONDS = 30.0
def transcribe(wav_path: Path) -> list[Segment]:
def transcribe(
wav_path: Path,
on_progress: Callable[[int, int], None] | None = None,
) -> list[Segment]:
"""转写 wav返回 segments含词级时间戳
Args:
wav_path: 16kHz mono PCM wav
on_progress: 可选进度回调 (current_chunk, total_chunks)。
每 transcribe 完一个 30s chunk 调一次,用于细分进度条。
Returns:
list[Segment],每个 Segment 带词级 words若 word_timestamps 启用)。
@@ -29,8 +39,8 @@ def transcribe(wav_path: Path) -> list[Segment]:
raise FileNotFoundError(f"音频不存在:{wav_path}")
model = get_model_manager().get_asr()
logger.debug("开始转写 %smodel=%s language=%s batch_size=%d",
wav_path.name, s.model, s.language, s.batch_size)
logger.debug("开始转写 %smodel=%s language=%s batch_size=%d beam_size=%d",
wav_path.name, s.model, s.language, s.batch_size, s.beam_size)
segments_gen, info = model.transcribe(
str(wav_path),
@@ -41,12 +51,19 @@ def transcribe(wav_path: Path) -> list[Segment]:
batch_size=s.batch_size, # 批量解码:多音频块一次性送 GPU
without_timestamps=False, # BatchedInferencePipeline 默认 True需显式关闭以生成段级时间戳
)
# VAD 过滤后的实际语音时长 → 算总 chunk 数(进度颗粒度细分用)
duration = info.duration_after_vad or info.duration
total_chunks = max(1, math.ceil(duration / _CHUNK_SECONDS))
logger.debug(
"音频时长 %.1fs检测语言=%s(置信度 %.2f",
info.duration, info.language, info.language_probability,
"音频时长 %.1fsVAD 后 %.1fs,检测语言=%s(置信度 %.2f,约 %d 个 chunk",
info.duration, duration, info.language, info.language_probability, total_chunks,
)
if on_progress is not None:
on_progress(0, total_chunks)
segments: list[Segment] = []
last_chunk = 0 # 已报进度的 chunk 序号(避免同 chunk 内多个 segment 重复回调)
for seg in segments_gen:
words: list[Word] = []
if s.word_timestamps and getattr(seg, "words", None):
@@ -63,6 +80,15 @@ def transcribe(wav_path: Path) -> list[Segment]:
end=float(seg.end),
words=words,
))
# 按 30s chunk 边界报进度seg.end 跨过 chunk 边界时回调
if on_progress is not None:
cur_chunk = min(total_chunks, int(seg.end / _CHUNK_SECONDS) + 1)
if cur_chunk > last_chunk:
last_chunk = cur_chunk
on_progress(cur_chunk, total_chunks)
logger.debug("转写完成:%d 段,%d 词。",
len(segments), sum(len(s.words) for s in segments))
len(segments), sum(len(seg.words) for seg in segments))
return segments

View File

@@ -85,7 +85,15 @@ def asr_phase(db, task) -> None:
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)
# 进度回调:按 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)}")
@@ -114,7 +122,14 @@ def translate_phase(db, task) -> None:
_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)
# 进度回调按批次细分翻译进度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)}")

View File

@@ -14,6 +14,7 @@ from __future__ import annotations
import logging
import os
from typing import Callable
from ..config import get_settings
from .model_manager import get_model_manager
@@ -25,11 +26,15 @@ logger = logging.getLogger("audio2text.translate")
_TOKENS_PER_WORD = 1.2
def translate(subtitles: list[Subtitle]) -> list[str]:
def translate(
subtitles: list[Subtitle],
on_progress: Callable[[int, int], None] | None = None,
) -> list[str]:
"""批量翻译英文字幕为中文。
Args:
subtitles: 断句后的英文字幕条目(按时间顺序)
on_progress: 可选进度回调 (done_count, total_count),每批完成时调一次。
Returns:
list[str],与 subtitles 等长、顺序对应的中文译文。
@@ -52,9 +57,9 @@ def translate(subtitles: list[Subtitle]) -> list[str]:
sort_by_length = False
if sort_by_length:
results = _translate_sorted(pipe, texts, batch_size, max_len)
results = _translate_sorted(pipe, texts, batch_size, max_len, on_progress)
else:
results = _translate_sequential(pipe, texts, batch_size, max_len)
results = _translate_sequential(pipe, texts, batch_size, max_len, on_progress)
logger.debug("翻译完成:%d 条。", len(results))
return results
@@ -64,6 +69,7 @@ def translate(subtitles: list[Subtitle]) -> list[str]:
def _translate_sorted(
pipe, texts: list[str], batch_size: int, max_len: int,
on_progress: Callable[[int, int], None] | None = None,
) -> list[str]:
"""按长度排序后分批翻译,翻译完按原序散回。
@@ -120,7 +126,9 @@ def _translate_sorted(
for idx, zh in zip(orig_indices, translated):
results[idx] = zh
done += len(batch)
if (done // batch_size + 1) % 5 == 0:
if on_progress is not None:
on_progress(done, n)
elif (done // batch_size + 1) % 5 == 0:
logger.debug("已翻译 %d/%d 条。", done, n)
# None理论不会发生_translate_batch 保证返回等长)→ 回退原文
@@ -142,15 +150,20 @@ def _estimate_sequential_padding(texts: list[str], batch_size: int) -> int:
def _translate_sequential(
pipe, texts: list[str], batch_size: int, max_len: int,
on_progress: Callable[[int, int], None] | None = None,
) -> list[str]:
"""按原序分批翻译(旧行为,便于 A/B 对比)。"""
results: list[str] = []
for i in range(0, len(texts), batch_size):
n = len(texts)
for i in range(0, n, 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))
done = min(i + len(chunk), n)
if on_progress is not None:
on_progress(done, n)
elif (i // batch_size + 1) % 5 == 0:
logger.debug("已翻译 %d/%d 条。", done, n)
return results