diff --git a/app/services/asr_service.py b/app/services/asr_service.py index 9d2e40b..c1c0d63 100644 --- a/app/services/asr_service.py +++ b/app/services/asr_service.py @@ -6,7 +6,9 @@ CPU dev: tiny.en + int8;GPU 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_length(秒):BatchedInferencePipeline 按 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("开始转写 %s(model=%s language=%s batch_size=%d)", - wav_path.name, s.model, s.language, s.batch_size) + logger.debug("开始转写 %s(model=%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, + "音频时长 %.1fs(VAD 后 %.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 + diff --git a/app/services/pipeline.py b/app/services/pipeline.py index dd8b39d..8f0e0fc 100644 --- a/app/services/pipeline.py +++ b/app/services/pipeline.py @@ -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)} 条") diff --git a/app/services/translate_service.py b/app/services/translate_service.py index d207f82..8453b21 100644 --- a/app/services/translate_service.py +++ b/app/services/translate_service.py @@ -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