Initial commit: audio2text 双语字幕生成服务

- 音频/视频转双语(英/中)SRT 字幕,Docker 容器化,CPU 开发/GPU 生产同一份代码
- faster-whisper ASR(词级时间戳) + 断句时间戳重算 + NLLB 翻译(模型不共驻)
- 分片上传(断点续传) + SQLite 持久化 + 主页/历史/日志页面
- 历史页文件名搜索;缓存定时清理(默认保留7天,可配置)
- 双 Dockerfile(cpu/gpu) + setup/start/stop 脚本
This commit is contained in:
2026-07-06 06:54:19 +00:00
commit 00e2a95fb7
44 changed files with 4110 additions and 0 deletions

0
app/services/__init__.py Normal file
View File

View File

@@ -0,0 +1,65 @@
"""语音识别服务faster-whisper输出带词级时间戳的 segments。
CPU dev: tiny.en + int8GPU prod: large-v3-turbo + float16。同一份代码。
"""
from __future__ import annotations
import logging
from pathlib import Path
from ..config import get_settings
from .model_manager import get_model_manager
from .types import Segment, Word
logger = logging.getLogger("audio2text.asr")
def transcribe(wav_path: Path) -> list[Segment]:
"""转写 wav返回 segments含词级时间戳
Args:
wav_path: 16kHz mono PCM wav
Returns:
list[Segment],每个 Segment 带词级 words若 word_timestamps 启用)。
"""
s = get_settings().asr
if not wav_path.is_file():
raise FileNotFoundError(f"音频不存在:{wav_path}")
model = get_model_manager().get_asr()
logger.debug("开始转写 %smodel=%s language=%s", wav_path.name, s.model, s.language)
segments_gen, info = model.transcribe(
str(wav_path),
language=s.language,
word_timestamps=s.word_timestamps,
vad_filter=s.vad_filter,
beam_size=5,
)
logger.debug(
"音频时长 %.1fs检测语言=%s(置信度 %.2f",
info.duration, info.language, info.language_probability,
)
segments: list[Segment] = []
for seg in segments_gen:
words: list[Word] = []
if s.word_timestamps and getattr(seg, "words", None):
for w in seg.words:
words.append(Word(
text=w.word.strip(),
start=float(w.start),
end=float(w.end),
probability=float(getattr(w, "probability", 1.0)),
))
segments.append(Segment(
text=seg.text.strip(),
start=float(seg.start),
end=float(seg.end),
words=words,
))
logger.debug("转写完成:%d 段,%d 词。",
len(segments), sum(len(s.words) for s in segments))
return segments

View File

@@ -0,0 +1,194 @@
"""缓存清理:删除超期任务产物 + 孤儿目录,并删除对应 DB 记录。
「缓存」指每个任务落盘的产物:
<output_dir>/task_<id>/ 双语 / 英文 / 中文字幕
<work_dir>/task_<id>.wav 中间音频(若 keep_audio=true 未被管线删掉)
<upload_dir>/yyyy/mm/<uuid>.<ext> 保留的原始视频(若 delete_original_after_extract=false
超期 = `created_at` 早于 `now - cache_retention_days`。超期任务的全部产物连同
Task / UploadSession 行一起删除,避免历史页出现指向已删文件的死链接。
另做一次孤儿扫描output_dir / work_dir 下存在但无对应 Task 的目录(进程崩溃残留),
按目录 mtime 判超期后删除,让磁盘不被异常退出留下的碎片占满。
"""
from __future__ import annotations
import logging
import shutil
import time
from datetime import datetime, timedelta, timezone
from pathlib import Path
from sqlalchemy.orm import Session
from ..config import get_settings
from ..database import get_session_local
from ..models.task import Task
from ..models.upload_session import UploadSession
logger = logging.getLogger("audio2text.cache")
def _now_utc() -> datetime:
return datetime.now(timezone.utc)
def _aware(dt: datetime | None) -> datetime | None:
"""SQLite 存 naive datetime统一补 UTC。"""
if dt is None:
return None
return dt if dt.tzinfo is not None else dt.replace(tzinfo=timezone.utc)
def _rm_tree(path: Path) -> bool:
"""删目录或文件,失败仅警告不抛。返回是否实际删除。"""
try:
if path.is_dir():
shutil.rmtree(path, ignore_errors=False)
return True
if path.is_file():
path.unlink(missing_ok=True)
return True
except Exception as exc: # pragma: no cover
logger.warning("清理 %s 失败:%s", path, exc)
return False
def purge_expired_cache(db: Session | None = None) -> dict:
"""执行一次清理,返回统计 dict。
可传入已有 Session如复用请求会话不传则自建一个并关闭。
"""
own_db = db is None
if own_db:
db = get_session_local()()
try:
return _purge(db)
finally:
if own_db:
db.close()
def _purge(db: Session) -> dict:
s = get_settings()
retention = max(0, s.storage.cache_retention_days)
cutoff = _now_utc() - timedelta(days=retention)
upload_root = s.upload_dir()
work_root = s.work_dir()
output_root = s.output_dir()
stats = {"tasks": 0, "outputs": 0, "audio": 0, "videos": 0,
"orphans": 0, "retention_days": retention}
if retention <= 0:
logger.info("缓存清理已禁用cache_retention_days=%s)。", retention)
return stats
# ---------- 1. 超期任务 ----------
tasks = db.query(Task).all()
for task in tasks:
created = _aware(task.created_at)
if created is None or created >= cutoff:
continue
# 输出目录
out_dir = output_root / f"task_{task.id}"
if out_dir.is_dir() and _rm_tree(out_dir):
stats["outputs"] += 1
# 中间音频
wav = work_root / f"task_{task.id}.wav"
if wav.is_file() and _rm_tree(wav):
stats["audio"] += 1
# 保留的原始视频(若未在提取后删除)
if task.source_path:
src = upload_root / task.source_path
if src.is_file() and _rm_tree(src):
stats["videos"] += 1
# 删 DB 记录(先删关联的 UploadSession再删 Task
db.query(UploadSession).filter(UploadSession.task_id == task.id).delete()
db.delete(task)
stats["tasks"] += 1
logger.info(
"清理超期任务 id=%s file=%s created=%s",
task.id, task.filename, created.isoformat(),
)
if stats["tasks"]:
db.commit()
# ---------- 2. 孤儿目录扫描 ----------
existing_ids = {row[0] for row in db.query(Task.id).all()}
stats["orphans"] += _purge_orphans(output_root, "task_", cutoff, existing_ids)
stats["orphans"] += _purge_orphans(work_root, "task_", cutoff, existing_ids, suffix=".wav")
logger.info(
"缓存清理完成:任务 %d(输出 %d / 音频 %d / 视频 %d+ 孤儿 %d,保留期 %d",
stats["tasks"], stats["outputs"], stats["audio"], stats["videos"],
stats["orphans"], retention,
)
return stats
def _purge_orphans(root: Path, prefix: str, cutoff: datetime,
existing_ids: set[int], suffix: str | None = None) -> int:
"""删 root 下名为 `prefix<id>` 但无对应 Task 的孤儿条目。
output_dir 下是目录task_<id>work_dir 下可能是 task_<id>.wav 文件。
用 mtime 判超期,避免误删刚产生的中间产物。
"""
if not root.is_dir():
return 0
n = 0
for entry in root.iterdir():
name = entry.name
if not name.startswith(prefix):
continue
rest = name[len(prefix):]
if suffix:
if not rest.endswith(suffix):
continue
rest = rest[: -len(suffix)]
try:
tid = int(rest)
except ValueError:
continue # 不是 task_<id> 命名,跳过
if tid in existing_ids:
continue
# 孤儿:按 mtime 判超期
try:
mtime = datetime.fromtimestamp(entry.stat().st_mtime, tz=timezone.utc)
except OSError:
continue
if mtime >= cutoff:
continue
if _rm_tree(entry):
n += 1
logger.info("清理孤儿缓存 %smtime=%s", entry, mtime.isoformat())
return n
def run_forever(interval_hours: int) -> None:
"""后台线程入口:先跑一次,再按间隔循环。供 main.py lifespan 拉起。"""
interval = max(1, interval_hours) * 3600
logger.info("缓存清理调度启动:间隔 %d 小时,保留期 %d 天。",
interval_hours, get_settings().storage.cache_retention_days)
# 启动时先跑一次
try:
purge_expired_cache()
except Exception as exc: # pragma: no cover
logger.warning("启动缓存清理失败:%s", exc)
while True:
time.sleep(interval)
try:
purge_expired_cache()
except Exception as exc: # pragma: no cover
logger.warning("定时缓存清理失败:%s", exc)
if __name__ == "__main__": # pragma: no cover
# 容器内手动触发python -m app.services.cache_cleaner
import json
logging.basicConfig(level=logging.INFO,
format="%(asctime)s %(levelname)s %(name)s: %(message)s")
result = purge_expired_cache()
print(json.dumps(result, ensure_ascii=False))

View File

@@ -0,0 +1,64 @@
"""ffmpeg 服务:从视频/音频文件提取 16kHz 单声道 PCM wav。
16kHz mono PCM 正是 Whisper 的标准输入,省去模型内重采样。
"""
from __future__ import annotations
import logging
import shutil
import subprocess
from pathlib import Path
logger = logging.getLogger("audio2text.ffmpeg")
class FFmpegError(RuntimeError):
pass
def extract_audio(
src: Path,
out_wav: Path,
sample_rate: int = 16000,
) -> Path:
"""提取音频为 16kHz 单声道 PCM wav。
Args:
src: 输入视频/音频文件
out_wav: 输出 wav 路径
sample_rate: 采样率,默认 16000Whisper 标准)
Returns:
out_wav 路径
Raises:
FFmpegError: ffmpeg 不可用或提取失败
"""
if shutil.which("ffmpeg") is None:
raise FFmpegError("ffmpeg 未安装;容器内应通过 apt 装好。")
if not src.is_file():
raise FFmpegError(f"源文件不存在:{src}")
out_wav.parent.mkdir(parents=True, exist_ok=True)
# -vn 去视频;-ac 1 单声道;-ar 16k 采样率;-c:a pcm_s16le 16bit PCM
cmd = [
"ffmpeg", "-y", "-loglevel", "error",
"-i", str(src),
"-vn", "-ac", "1", "-ar", str(sample_rate),
"-c:a", "pcm_s16le",
str(out_wav),
]
logger.debug("提取音频:%s -> %s", src.name, out_wav.name)
logger.debug("ffmpeg 命令:%s", " ".join(cmd))
try:
result = subprocess.run(cmd, capture_output=True, text=True, timeout=3600)
except subprocess.TimeoutExpired as exc:
raise FFmpegError(f"ffmpeg 超时(>1h{src}") from exc
if result.returncode != 0:
raise FFmpegError(
f"ffmpeg 失败 (code={result.returncode}): {result.stderr.strip()[:500]}"
)
if not out_wav.is_file():
raise FFmpegError(f"ffmpeg 未生成输出文件:{out_wav}")
return out_wav

109
app/services/log_buffer.py Normal file
View File

@@ -0,0 +1,109 @@
"""内存日志缓冲:捕获所有 audio2text.* 日志到有界 deque供 /logs 页面查询。
设计:
- MemoryLogHandler 挂到 `audio2text` logger子 logger 的记录经传播自动汇入。
- deque(maxlen=N) 有界,旧记录自动淘汰,避免内存无限增长。
- emit() 在 logging 内部锁下执行deque.append 线程安全。
- 查询时按级别过滤(≥指定级别)、取最近 N 条,返回结构化 dict。
"""
from __future__ import annotations
import logging
import threading
from collections import deque
from datetime import datetime, timezone
# 级别名 → 数值,用于过滤
_LEVELS = {
"debug": logging.DEBUG,
"info": logging.INFO,
"warning": logging.WARNING,
"error": logging.ERROR,
}
class _BoundedLogRecord:
"""从 LogRecord 提取显示所需字段,避免持有完整对象引用。"""
__slots__ = ("ts", "level_no", "level_name", "logger_name", "message", "traceback")
def __init__(self, record: logging.LogRecord) -> None:
self.ts: float = record.created
self.level_no: int = record.levelno
self.level_name: str = record.levelname
self.logger_name: str = record.name
# getMessage() 应用 %-style 懒格式化
self.message: str = record.getMessage()
# exc_info 存的是 (type, value, tb) 三元组,格式化为字符串
self.traceback: str | None = None
if record.exc_info:
import traceback as _tb
self.traceback = "".join(_tb.format_exception(*record.exc_info))
class MemoryLogHandler(logging.Handler):
"""把日志记录存入有界 deque供 /api/logs 查询。"""
def __init__(self, buffer_size: int = 2000) -> None:
super().__init__()
self._buffer: deque[_BoundedLogRecord] = deque(maxlen=buffer_size)
self._lock = threading.Lock()
def emit(self, record: logging.LogRecord) -> None:
try:
entry = _BoundedLogRecord(record)
with self._lock:
self._buffer.append(entry)
except Exception: # pragma: no cover —— logging 不应抛
self.handleError(record)
def get_records(
self, level: str = "info", tail: int = 200,
) -> list[dict]:
"""返回最近 tail 条、≥指定级别的日志JSON 友好的 dict 列表)。
Args:
level: debug | info | warning | error最小级别
tail: 最多返回条数
"""
min_level = _LEVELS.get(level.lower(), logging.INFO)
with self._lock:
snapshot = list(self._buffer)
# 过滤 + 取最近 tail 条
filtered = [r for r in snapshot if r.level_no >= min_level]
result = []
for r in filtered[-tail:]:
ts_str = datetime.fromtimestamp(r.ts, tz=timezone.utc).strftime(
"%Y-%m-%d %H:%M:%S"
)
result.append({
"ts": ts_str,
"level": r.level_name,
"logger": r.logger_name,
"msg": r.message,
"traceback": r.traceback,
})
return result
def clear(self) -> None:
with self._lock:
self._buffer.clear()
# 进程级单例
_buffer: MemoryLogHandler | None = None
def get_log_buffer() -> MemoryLogHandler:
global _buffer
if _buffer is None:
_buffer = MemoryLogHandler()
return _buffer
def init_log_buffer(buffer_size: int) -> MemoryLogHandler:
"""(重新)创建缓冲并返回,供 main.py 启动时按配置初始化。"""
global _buffer
_buffer = MemoryLogHandler(buffer_size=buffer_size)
return _buffer

View File

@@ -0,0 +1,136 @@
"""模型管理器ASR 与翻译模型不共驻,任一时刻 GPU 上只有一个模型。
设计:
- 单例 ModelManager 跟踪当前加载的模型类型none / asr / translator
- get_asr(): 若翻译器在内存 → 先卸载del + gc + empty_cache→ 加载 faster-whisper
- get_translator(): 若 ASR 在内存 → 先卸载 → 加载 NLLB
- 翻译阶段独占显存,可用大 batch_sizeASR 阶段同理
这样在 24G 3090 上无需担心显存叠加CPU dev 时也省内存。
"""
from __future__ import annotations
import gc
import logging
import threading
from typing import Any
from ..config import get_settings
logger = logging.getLogger("audio2text.models")
# 全局单例 + 锁:模型加载/卸载必须串行
_lock = threading.Lock()
def _free_memory() -> None:
"""释放 Python 对象与 GPU 缓存。"""
gc.collect()
try:
import torch
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.synchronize()
except Exception: # pragma: no cover
pass
class ModelManager:
"""ASR / 翻译模型的不共驻管理器。"""
def __init__(self) -> None:
self._asr: Any = None # faster_whisper.WhisperModel
self._translator: Any = None # transformers pipeline
self._current: str = "none" # none | asr | translator
# ---------------- ASR ----------------
def get_asr(self) -> Any:
"""返回已加载的 faster-whisper 模型;必要时先卸载翻译器。"""
with _lock:
if self._asr is not None:
return self._asr
if self._translator is not None:
self._unload_translator_locked()
s = get_settings().asr
logger.debug("加载 ASR 模型 model=%s device=%s compute_type=%s",
s.model, s.device, s.compute_type)
from faster_whisper import WhisperModel
# device/compute_type 组合cpu+int8 / cuda+float16
self._asr = WhisperModel(
s.model, device=s.device, compute_type=s.compute_type,
)
self._current = "asr"
logger.debug("ASR 模型已就绪。")
return self._asr
def unload_asr(self) -> None:
with _lock:
self._unload_asr_locked()
def _unload_asr_locked(self) -> None:
if self._asr is None:
return
logger.debug("卸载 ASR 模型(释放显存供翻译器独占)。")
# faster-whisper 模型无显式 closedel 即可
del self._asr
self._asr = None
self._current = "none" if self._translator is None else "translator"
_free_memory()
# ---------------- 翻译器 ----------------
def get_translator(self) -> Any:
"""返回已加载的 NLLB 翻译 pipeline必要时先卸载 ASR。"""
with _lock:
if self._translator is not None:
return self._translator
if self._asr is not None:
self._unload_asr_locked()
s = get_settings().translation
logger.debug("加载翻译模型 model=%s device=%s", s.model, s.device)
from transformers import pipeline
self._translator = pipeline(
"translation",
model=s.model,
device=s.device,
src_lang=s.src_lang,
tgt_lang=s.tgt_lang,
)
self._current = "translator"
logger.debug("翻译模型已就绪(独占显存,可用大 batch")
return self._translator
def unload_translator(self) -> None:
with _lock:
self._unload_translator_locked()
def _unload_translator_locked(self) -> None:
if self._translator is None:
return
logger.debug("卸载翻译模型。")
# 释放 pipeline 持有的 model + tokenizer
mdl = getattr(self._translator, "model", None)
tok = getattr(self._translator, "tokenizer", None)
del self._translator, mdl, tok
self._translator = None
self._current = "none" if self._asr is None else "asr"
_free_memory()
# ---------------- 状态 ----------------
@property
def current(self) -> str:
return self._current
# 进程级单例
_manager: ModelManager | None = None
def get_model_manager() -> ModelManager:
global _manager
if _manager is None:
_manager = ModelManager()
return _manager

153
app/services/pipeline.py Normal file
View File

@@ -0,0 +1,153 @@
"""转写管线编排:提取音频 → ASR → 断句 → 翻译 → 写 SRT。
任务状态机:
queued → extracting → transcribing → segmenting → translating → done
任一步失败 → failed
模型不共驻ASR 与翻译分阶段加载,翻译时先卸载 Whisper 释放显存跑大 batch。
管线在后台线程跑(每个任务一个线程),通过 DB 更新状态与进度。
"""
from __future__ import annotations
import logging
import threading
import traceback
from datetime import datetime, timezone
from pathlib import Path
from ..config import get_settings
from ..database import get_session_local
from ..models.task import Task
from . import ffmpeg_service, asr_service, segmenter, translate_service, srt_writer
from .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
def enqueue_task(task_id: int) -> None:
"""把任务交给后台线程处理(非阻塞,供 upload_service.complete 调用)。"""
t = threading.Thread(target=_run_task, args=(task_id,), daemon=True)
t.start()
logger.info("任务 %d 已入队(后台线程 %s)。", task_id, t.name)
def _run_task(task_id: int) -> None:
"""后台执行完整管线。所有异常都被捕获并写入 task.error。"""
db = get_session_local()()
try:
task = db.get(Task, task_id)
if task is None:
logger.error("任务 %d 不存在。", task_id)
return
_pipeline(db, task)
except Exception as exc:
logger.exception("任务 %d 失败:%s", task_id, exc)
_mark_failed(db, task_id, str(exc))
finally:
db.close()
def _pipeline(db, task: Task) -> None:
s = get_settings()
src = s.upload_dir() / task.source_path
# ---------- 1. 提取音频 ----------
_set_status(db, task, "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)
# ---------- 2. 语音识别 ----------
_set_status(db, task, "transcribing", P_TRANSCRIBE_START)
segments = asr_service.transcribe(wav)
_set_status(db, task, "transcribing", P_TRANSCRIBE_END,
note=f"识别出 {len(segments)}")
# ---------- 3. 断句 + 时间戳重算 ----------
_set_status(db, task, "segmenting", P_SEGMENT_START)
subs = segmenter.resegment(segments)
_set_status(db, task, "segmenting", P_SEGMENT_END,
note=f"重组为 {len(subs)} 条字幕")
# ---------- 4. 翻译 ----------
_set_status(db, task, "translating", P_TRANSLATE_START)
# 翻译阶段model_manager 会自动卸载 ASR、加载翻译器独占显存
zh_texts = translate_service.translate(subs)
_set_status(db, task, "translating", P_TRANSLATE_END,
note=f"翻译 {len(zh_texts)}")
# ---------- 5. 写 SRT ----------
out_dir = s.output_dir()
stem = Path(task.filename).stem
en_path = out_dir / f"task_{task.id}/{stem}.en.srt"
zh_path = out_dir / f"task_{task.id}/{stem}.zh.srt"
bi_path = out_dir / f"task_{task.id}/{stem}.srt"
srt_writer.write_srt(subs, en_path)
# 中文 SRT用译文 + 同时间戳)
zh_subs = [Subtitle(text=zh, start=sub.start, end=sub.end)
for zh, sub in zip(zh_texts, subs)]
srt_writer.write_srt(zh_subs, zh_path)
srt_writer.write_bilingual_srt(subs, zh_texts, bi_path)
# 记录相对路径
task.en_srt_path = str(en_path.relative_to(out_dir))
task.zh_srt_path = str(zh_path.relative_to(out_dir))
task.bilingual_srt_path = str(bi_path.relative_to(out_dir))
task.status = "done"
task.progress = P_DONE
task.updated_at = datetime.now(timezone.utc)
db.commit()
# 清理中间音频
if not s.processing.keep_audio and wav.is_file():
try:
wav.unlink()
except OSError:
pass
logger.info("任务 %d 完成:%s", task.id, bi_path.name)
# ---------------- DB 状态更新 ----------------
def _set_status(db, task: Task, status: str, progress: float, note: str = "") -> None:
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.info("任务 %d [%s %.0f%%]", task.id, status, progress)
def _mark_failed(db, task_id: int, error: str) -> None:
try:
task = db.get(Task, task_id)
if task is None:
return
task.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())

22
app/services/reaper.py Normal file
View File

@@ -0,0 +1,22 @@
"""后台 reaper清理被放弃的分片上传会话与临时文件。供 main.py 启动时调用。"""
from __future__ import annotations
import logging
from ..database import get_session_local
from .upload_service import UploadService
logger = logging.getLogger("audio2text.reaper")
def reap_stale_sessions() -> int:
"""执行一次过期会话清理。"""
db = get_session_local()()
try:
return UploadService(db).reap_stale_sessions()
except Exception as exc: # pragma: no cover
logger.warning("reaper 执行失败:%s", exc)
return 0
finally:
db.close()

233
app/services/segmenter.py Normal file
View File

@@ -0,0 +1,233 @@
"""断句 + 时间戳重算(纯算法,零模型开销)。
Whisper 原始 segment 的断句通常很混乱:每段不一定是完整句子,时间戳也不对齐句界。
本模块基于词级时间戳重新断句,得到规范的字幕条目。
两路策略:
1. 精确路(有 word_timestamps按句末标点. ! ? ;)切句,超长句再按逗号拆,
每条字幕的时间戳直接取首词.start ~ 末词.end精确无误。
2. 匀速估算路(无 word_timestamps段内按字符数比例分配时间——
句start = 段start + (前缀字符数 / 段总字符数) × 段时长。
即「短时匀速」假设,无需大模型。
最后做 SRT 规范化:单条 17 秒、≤2 行、每行 ≤42 字符。
"""
from __future__ import annotations
import logging
import re
from ..config import get_settings
from .types import Segment, Subtitle, Word
logger = logging.getLogger("audio2text.segmenter")
# 句末标点:句号、感叹号、问号、分号
_SENTENCE_END = re.compile(r"[.!?;]+")
# 句内停顿:逗号、冒号、破折号
_CLAUSE_BREAK = re.compile(r"[,:\-—]+")
def resegment(segments: list[Segment]) -> list[Subtitle]:
"""把 ASR segments 重组为规范字幕条目。
Args:
segments: asr_service.transcribe 的输出(可能含词级时间戳)
Returns:
list[Subtitle],已做时长/字数规范化。
"""
cfg = get_settings().segmentation
max_words = cfg.max_words_per_line
max_dur = cfg.max_duration_seconds
min_dur = cfg.min_duration_seconds
max_chars = cfg.max_chars_per_line
# 第一步:把所有词串成大列表(精确路)或退化为段级(估算路)
all_words: list[Word] = []
has_word_ts = True
for seg in segments:
if not seg.words:
has_word_ts = False
break
all_words.extend(seg.words)
if has_word_ts and all_words:
subs = _resegment_by_words(all_words, max_words, max_dur)
else:
logger.warning("无词级时间戳,退化为匀速估算路。")
subs = _resegment_by_estimate(segments, max_words, max_dur)
# 规范化:合并过短条目、拆分行宽
subs = _normalize(subs, min_dur, max_chars)
logger.debug("断句完成:%d 条字幕。", len(subs))
return subs
# ---------------- 精确路:按词级时间戳 ----------------
def _resegment_by_words(
words: list[Word], max_words: int, max_dur: float,
) -> list[Subtitle]:
"""按句末标点切句,超长句按逗号拆,时间戳取首末词。"""
subs: list[Subtitle] = []
# 当前句的词
current: list[Word] = []
def flush(wlist: list[Word]) -> None:
if not wlist:
return
text = " ".join(w.text for w in wlist).strip()
if not text:
return
subs.append(Subtitle(
text=text,
start=wlist[0].start,
end=wlist[-1].end,
))
for w in words:
current.append(w)
# 句末标点 → 收尾
if _SENTENCE_END.search(w.text):
_maybe_split_and_flush(current, max_words, max_dur, flush)
current = []
continue
# 超长(词数或时长)→ 优先在最近的逗号处断
cur_dur = (current[-1].end - current[0].start) if len(current) > 1 else 0
if len(current) >= max_words or cur_dur >= max_dur:
_maybe_split_and_flush(current, max_words, max_dur, flush)
current = []
flush(current)
return subs
def _maybe_split_and_flush(
wlist: list[Word], max_words: int, max_dur: float, flush,
) -> None:
"""若 wlist 过长,在逗号处再拆;否则整条 flush。"""
if len(wlist) <= max_words and (len(wlist) <= 1 or
wlist[-1].end - wlist[0].start < max_dur):
flush(wlist)
return
# 找逗号断点
parts: list[list[Word]] = []
cur: list[Word] = []
for w in wlist:
cur.append(w)
if _CLAUSE_BREAK.search(w.text) and len(cur) >= max_words // 2:
parts.append(cur)
cur = []
if cur:
parts.append(cur)
# 若逗号拆不开(无逗号),强制按 max_words 等分
if len(parts) == 1 and len(parts[0]) > max_words:
parts = [parts[0][i:i + max_words] for i in range(0, len(parts[0]), max_words)]
for p in parts:
flush(p)
# ---------------- 匀速估算路:段内按字符比例 ----------------
def _resegment_by_estimate(
segments: list[Segment], max_words: int, max_dur: float,
) -> list[Subtitle]:
"""无词级时间戳时:先按文本断句,再按字符数比例估算时间戳。"""
subs: list[Subtitle] = []
for seg in segments:
text = seg.text.strip()
if not text:
continue
dur = seg.end - seg.start
# 按句末标点切
sentences = _split_sentences(text)
if not sentences:
sentences = [text]
# 段内按字符数比例分配时间
total_chars = sum(len(s) for s in sentences) or 1
cursor = seg.start
for sent in sentences:
frac = len(sent) / total_chars
est_end = cursor + dur * frac
# 超长句再按逗号拆(时间按字符比例再分)
if len(sent.split()) > max_words or (est_end - cursor) > max_dur:
for clause in _split_clauses(sent):
cfrac = len(clause) / len(sent) if len(sent) else 1
c_end = cursor + (est_end - cursor) * cfrac
subs.append(Subtitle(text=clause.strip(), start=cursor, end=c_end))
cursor = c_end
else:
subs.append(Subtitle(text=sent.strip(), start=cursor, end=est_end))
cursor = est_end
return subs
def _split_sentences(text: str) -> list[str]:
"""按句末标点切句,保留标点。"""
parts = _SENTENCE_END.split(text)
marks = _SENTENCE_END.findall(text)
out = []
for i, p in enumerate(parts):
p = p.strip()
if not p:
continue
out.append(p + (marks[i] if i < len(marks) else ""))
return out
def _split_clauses(sentence: str) -> list[str]:
"""按逗号/冒号拆子句,保留标点。"""
parts = _CLAUSE_BREAK.split(sentence)
marks = _CLAUSE_BREAK.findall(sentence)
out = []
for i, p in enumerate(parts):
p = p.strip()
if not p:
continue
out.append(p + (marks[i - 1] if 0 < i <= len(marks) else ""))
return out
# ---------------- 规范化 ----------------
def _normalize(subs: list[Subtitle], min_dur: float, max_chars: int) -> list[Subtitle]:
"""合并过短条目、拆分过宽行。"""
# 1. 合并过短(< min_dur 且非末尾)
merged: list[Subtitle] = []
for s in subs:
if merged and (s.end - s.start) < min_dur:
prev = merged[-1]
prev.text = (prev.text + " " + s.text).strip()
prev.end = s.end
else:
merged.append(Subtitle(text=s.text, start=s.start, end=s.end))
# 2. 拆分超过 max_chars 的行(按词折行,不改时间戳)
out: list[Subtitle] = []
for s in merged:
if len(s.text) <= max_chars:
out.append(s)
continue
lines = _wrap_text(s.text, max_chars)
out.append(Subtitle(text=lines, start=s.start, end=s.end))
return out
def _wrap_text(text: str, max_chars: int) -> str:
"""按词折成 ≤2 行,每行 ≤max_chars 字符SRT 规范)。"""
words = text.split()
lines: list[str] = []
cur = ""
for idx, w in enumerate(words):
if cur and len(cur) + 1 + len(w) > max_chars:
lines.append(cur)
# 已有一行 + 当前行,合并剩余为一行
cur = " ".join([w] + words[idx + 1:])
break
else:
cur = (cur + " " + w).strip() if cur else w
if cur:
lines.append(cur)
return "\n".join(lines[:2])

View File

@@ -0,0 +1,68 @@
"""SRT 字幕文件写入与双语合并。"""
from __future__ import annotations
from pathlib import Path
from .types import Subtitle
def _format_ts(seconds: float) -> str:
"""秒 → SRT 时间戳 HH:MM:SS,mmm。"""
if seconds < 0:
seconds = 0.0
ms = int(round((seconds - int(seconds)) * 1000))
s = int(seconds) % 60
m = (int(seconds) // 60) % 60
h = int(seconds) // 3600
if ms == 1000: # 四舍五入进位
ms = 0
s += 1
if s == 60:
s = 0
m += 1
if m == 60:
m = 0
h += 1
return f"{h:02d}:{m:02d}:{s:02d},{ms:03d}"
def write_srt(subtitles: list[Subtitle], path: Path) -> Path:
"""写单语 SRT。"""
path.parent.mkdir(parents=True, exist_ok=True)
lines: list[str] = []
for i, sub in enumerate(subtitles, 1):
lines.append(str(i))
lines.append(f"{_format_ts(sub.start)} --> {_format_ts(sub.end)}")
lines.append(sub.text)
lines.append("")
path.write_text("\n".join(lines), encoding="utf-8")
return path
def write_bilingual_srt(
en_subs: list[Subtitle],
zh_texts: list[str],
path: Path,
) -> Path:
"""写双语合并 SRT英文在上、中文在下同一时间戳。
Args:
en_subs: 英文字幕条目
zh_texts: 与 en_subs 等长、顺序对应的中文译文
path: 输出路径
"""
if len(en_subs) != len(zh_texts):
raise ValueError(
f"英文字幕数({len(en_subs)}) 与中文译文数({len(zh_texts)}) 不一致"
)
path.parent.mkdir(parents=True, exist_ok=True)
lines: list[str] = []
for i, (sub, zh) in enumerate(zip(en_subs, zh_texts), 1):
lines.append(str(i))
lines.append(f"{_format_ts(sub.start)} --> {_format_ts(sub.end)}")
lines.append(sub.text)
lines.append(zh)
lines.append("")
path.write_text("\n".join(lines), encoding="utf-8")
return path

View File

@@ -0,0 +1,60 @@
"""翻译服务NLLB-200英译中。
通过 model_manager 加载,确保 ASR 已卸载、翻译器独占显存,从而可用大 batch_size。
按字幕条目批量翻译,保留索引对应。
"""
from __future__ import annotations
import logging
from ..config import get_settings
from .model_manager import get_model_manager
from .types import Subtitle
logger = logging.getLogger("audio2text.translate")
def translate(subtitles: list[Subtitle]) -> list[str]:
"""批量翻译英文字幕为中文。
Args:
subtitles: 断句后的英文字幕条目
Returns:
list[str],与 subtitles 等长、顺序对应的中文译文。
单条翻译失败时该位置回退为原英文。
"""
if not subtitles:
return []
s = get_settings().translation
pipe = get_model_manager().get_translator()
batch = s.batch_size
max_len = s.max_length
# 取纯文本(去掉折行),避免翻译把换行符当语义
texts = [sub.text.replace("\n", " ").strip() for sub in subtitles]
logger.debug("开始翻译 %d 条字幕batch_size=%d...", len(texts), batch)
results: list[str] = []
for i in range(0, len(texts), batch):
chunk = texts[i:i + batch]
try:
out = pipe(chunk, max_length=max_len)
for item in out:
# pipeline 返回 [{"translation_text": "..."}]
results.append(item.get("translation_text", "").strip())
except Exception as exc: # pragma: no cover
logger.warning("%d-%d 批翻译失败,逐条重试:%s", i, i + len(chunk), exc)
for t in chunk:
try:
out = pipe([t], max_length=max_len)
results.append(out[0].get("translation_text", "").strip())
except Exception:
results.append(t) # 回退原文
if (i // batch + 1) % 5 == 0:
logger.debug("已翻译 %d/%d 条。", min(i + len(chunk), len(texts)), len(texts))
logger.debug("翻译完成:%d 条。", len(results))
return results

38
app/services/types.py Normal file
View File

@@ -0,0 +1,38 @@
"""管线各阶段共享的数据传输对象DTO
独立于任何模型加载逻辑——segmenter纯算法和 srt_writer纯 IO可以只依赖
本模块,不拉入 faster_whisper / torch。
"""
from __future__ import annotations
from dataclasses import dataclass
@dataclass
class Word:
"""ASR 识别出的单个词,带时间戳。"""
text: str
start: float # 秒
end: float
probability: float = 1.0
@dataclass
class Segment:
"""ASR 输出的一段文本,可能含词级时间戳。"""
text: str
start: float
end: float
words: list[Word]
@dataclass
class Subtitle:
"""断句后的一条字幕:英文文本 + 起止时间戳。"""
text: str
start: float
end: float

View File

@@ -0,0 +1,249 @@
"""分片上传服务:会话管理 + 分片落盘 + 拼接 + 创建转写任务。
存储布局::
<work_dir>/<upload_id>/0.part 分片暂存
<work_dir>/<upload_id>/1.part
...
<upload_dir>/<yyyy>/<mm>/<uuid>.<ext> complete 后的正式视频
与 server 的区别:视频无需 sha256 去重每个视频都转写complete 直接创建 Task。
管线触发由 controller 调用 pipeline.enqueue_task本服务不依赖 pipeline。
"""
from __future__ import annotations
import logging
import os
import shutil
import uuid
from datetime import datetime, timedelta, timezone
from pathlib import Path
from fastapi import HTTPException
from sqlalchemy.orm import Session
from ..config import get_settings
from ..models.task import Task
from ..models.upload_session import UploadSession
from ..schemas.task import (
ChunkUploadResponse,
CompleteResponse,
CreateSessionRequest,
CreateSessionResponse,
SessionStatusResponse,
)
logger = logging.getLogger("audio2text.upload")
class UploadService:
def __init__(self, db: Session) -> None:
self.db = db
s = get_settings()
self.upload_root = s.upload_dir()
self.work_root = s.work_dir()
self.chunk_bytes = s.storage.chunk_bytes
self.session_ttl = s.storage.chunk_session_ttl_seconds
# ---------------- 会话生命周期 ----------------
def create_session(self, body: CreateSessionRequest) -> CreateSessionResponse:
upload_id = uuid.uuid4().hex
session = UploadSession(
upload_id=upload_id,
filename=body.filename,
size_bytes=body.size_bytes,
chunk_size=body.chunk_size,
total_chunks=body.total_chunks,
uploaded_chunks=[],
status="pending",
)
self.db.add(session)
self.db.commit()
self._session_dir(upload_id).mkdir(parents=True, exist_ok=True)
return CreateSessionResponse(
upload_id=upload_id,
filename=body.filename,
size_bytes=body.size_bytes,
chunk_size=body.chunk_size,
total_chunks=body.total_chunks,
)
def get_status(self, upload_id: str) -> SessionStatusResponse:
session = self._require_session(upload_id)
return SessionStatusResponse(
upload_id=session.upload_id,
filename=session.filename,
size_bytes=session.size_bytes,
chunk_size=session.chunk_size,
total_chunks=session.total_chunks,
uploaded_chunks=list(session.uploaded_chunks or []),
completed=(session.status == "completed"),
task_id=session.task_id,
)
# ---------------- 分片写入 ----------------
def write_chunk(self, upload_id: str, index: int, data: bytes) -> list[int]:
session = self._require_session(upload_id)
self._validate_index(session, index)
session_dir = self._session_dir(upload_id)
session_dir.mkdir(parents=True, exist_ok=True)
chunk_path = session_dir / f"{index}.part"
try:
with chunk_path.open("wb") as out:
out.write(data)
out.flush()
os.fsync(out.fileno())
except Exception:
chunk_path.unlink(missing_ok=True)
raise
uploaded = list(session.uploaded_chunks or [])
if index not in uploaded:
uploaded.append(index)
session.uploaded_chunks = uploaded
session.updated_at = datetime.now(timezone.utc)
self.db.commit()
return sorted(uploaded)
# ---------------- 拼接 + 创建任务 ----------------
def complete(self, upload_id: str) -> CompleteResponse:
session = self._require_session(upload_id)
# 幂等:已 complete 则返回已建任务
if session.status == "completed" and session.task_id is not None:
task = self.db.get(Task, session.task_id)
if task is not None:
return CompleteResponse(
task_id=task.id, filename=task.filename,
size_bytes=session.size_bytes, status=task.status,
)
uploaded = set(session.uploaded_chunks or [])
missing = [i for i in range(session.total_chunks) if i not in uploaded]
if missing:
raise HTTPException(
409,
f"分片不齐全:缺失 {len(missing)} 个,例如 {sorted(missing)[:10]}",
)
final_path = self._assemble(session)
rel = str(final_path.relative_to(self.upload_root))
task = Task(
filename=session.filename,
source_path=rel,
status="queued",
progress=0.0,
)
self.db.add(task)
session.status = "completed"
session.final_path = rel
session.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(task)
# 正向关联session → task替代旧的 source_path 反向查找)
session.task_id = task.id
self.db.commit()
# 清理分片暂存
self._cleanup_session_dir(upload_id)
logger.info("上传完成 task_id=%s file=%s size=%d", task.id, session.filename, session.size_bytes)
return CompleteResponse(
task_id=task.id, filename=task.filename,
size_bytes=session.size_bytes, status=task.status,
)
# ---------------- 过期会话清理 ----------------
def reap_stale_sessions(self) -> int:
"""清理被放弃的会话pending 且 updated_at 超过 ttl。
时间比较统一用 aware UTC datetime避免 naive datetime 的 .timestamp()
按本地时区算导致的偏移。
"""
cutoff = datetime.now(timezone.utc) - timedelta(seconds=self.session_ttl)
sessions = (
self.db.query(UploadSession)
.filter(UploadSession.status == "pending")
.all()
)
n = 0
for session in sessions:
updated = session.updated_at
if updated is None:
continue
# SQLite 存 naive datetime统一补 UTC 后比较
if updated.tzinfo is None:
updated = updated.replace(tzinfo=timezone.utc)
if updated < cutoff:
self._cleanup_session_dir(session.upload_id)
self.db.delete(session)
n += 1
logger.info("清理过期分片会话 upload_id=%s file=%s", session.upload_id, session.filename)
if n:
self.db.commit()
return n
# ---------------- 内部 ----------------
def _require_session(self, upload_id: str) -> UploadSession:
if not upload_id:
raise HTTPException(400, "upload_id 不能为空")
session = (
self.db.query(UploadSession)
.filter(UploadSession.upload_id == upload_id)
.first()
)
if session is None:
raise HTTPException(404, f"会话不存在或已过期:{upload_id}")
return session
@staticmethod
def _validate_index(session: UploadSession, index: int) -> None:
if index < 0 or index >= session.total_chunks:
raise HTTPException(400, f"分片下标越界:{index} 不在 [0, {session.total_chunks})")
def _session_dir(self, upload_id: str) -> Path:
return self.work_root / upload_id
def _assemble(self, session: UploadSession) -> Path:
"""按 index 顺序拼接全部分片为正式视频文件。"""
ext = Path(session.filename).suffix or ".mp4"
now = datetime.now(timezone.utc)
sub = self.upload_root / f"{now:%Y}" / f"{now:%m}"
sub.mkdir(parents=True, exist_ok=True)
final = sub / f"{uuid.uuid4().hex}{ext}"
part_path = final.with_suffix(final.suffix + ".part")
session_dir = self._session_dir(session.upload_id)
try:
with part_path.open("wb") as out:
for index in range(session.total_chunks):
chunk_path = session_dir / f"{index}.part"
if not chunk_path.is_file():
raise HTTPException(409, f"拼接时发现分片缺失:{index}.part")
with chunk_path.open("rb") as src:
while buf := src.read(self.chunk_bytes):
out.write(buf)
out.flush()
os.fsync(out.fileno())
os.replace(part_path, final)
except Exception:
part_path.unlink(missing_ok=True)
raise
return final
def _cleanup_session_dir(self, upload_id: str) -> None:
session_dir = self._session_dir(upload_id)
try:
if session_dir.exists():
shutil.rmtree(session_dir, ignore_errors=True)
except Exception as exc: # pragma: no cover
logger.warning("清理会话目录失败 upload_id=%s: %s", upload_id, exc)