Files
audio2text/app/config.py
audio2text dev 7635e5e766 feat: 设置页调 batch_size + 删除任务 + /docs 去认证 + 离线模式
设置页(/settings):
- 新增 Setting DB 模型(键值存储,持久化运行时覆盖)
- config.py 加 _apply_overrides:get_settings() 合并 DB 覆盖值
  (_applying_overrides 标志防递归:DB 初始化回调 get_settings 时不重入)
- save_setting() 写 DB + 清 lru_cache,后续任务读到新值
- settings_router.py: GET/PUT /api/settings
- settings_html.py: 设置页表单(batch_size/beam_size/sort_by_length)
  + 设备只读信息 + 保存后 reload 确认
- 验证:改 asr_batch_size=16 beam_size=1 -> 任务 ASR 日志确认生效
- 验证:重启容器后设置从 DB 恢复(持久化)

删除任务:
- task_router.py: DELETE /api/tasks/{id},仅 done/failed 可删
  删字幕/音频/视频产物 + UploadSession + Task DB 记录
- home_html.py + history_html.py: done/failed 任务显示删除按钮
  + confirm 确认 + 调用 DELETE API + 刷新列表
- 验证:删除 task 15 成功,删除进行中 task 返回 409

/docs 去认证:
- 移除 require_docs_auth 依赖,/docs /redoc /openapi.json 直接公开

离线模式:
- Dockerfile dev/final 加 HF_HUB_OFFLINE=1 TRANSFORMERS_OFFLINE=1
  (模型已缓存在 ./models volume,无需联网验证)
- 验证:离线模式完整跑通 ASR+翻译(task 15 done)
2026-07-11 11:33:34 +08:00

239 lines
8.3 KiB
Python
Raw Permalink 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.

"""运行时配置:所有参数从 config.yaml 读取。
CPU dev / GPU prod 仅靠 device / model / compute_type 三项切换,代码完全不变。
默认值与 config.example.yaml 对齐,确保无 yaml 时也能用最小配置启动。
"""
from __future__ import annotations
import os
from functools import lru_cache
from pathlib import Path
import yaml
from pydantic import BaseModel, Field, field_validator
PROJECT_ROOT = Path(__file__).resolve().parent.parent
DEFAULT_CONFIG_PATH = PROJECT_ROOT / "config.yaml"
class ServerConfig(BaseModel):
host: str = "0.0.0.0"
port: int = 8000
workers: int = 1
class StorageConfig(BaseModel):
upload_dir: str = "/data/uploads"
work_dir: str = "/data/.work"
output_dir: str = "/data/outputs"
chunk_bytes: int = 1024 * 1024
chunk_session_ttl_seconds: int = 300
# 缓存保留:超期任务产物(输出字幕 / 中间音频 / 保留的原始视频)连同 DB 记录一并清理
cache_retention_days: int = 7
# 定时清理间隔(小时):启动时跑一次,之后按此间隔循环
cache_cleanup_interval_hours: int = 24
class ProcessingConfig(BaseModel):
delete_original_after_extract: bool = True
keep_audio: bool = False
class AsrConfig(BaseModel):
model: str = "tiny.en" # CPU: tiny.enGPU: large-v3-turbo
device: str = "cpu" # cpu | cuda
compute_type: str = "int8" # cpu: int8gpu: float16
language: str = "en"
word_timestamps: bool = True
vad_filter: bool = True
batch_size: int = 8 # BatchedInferencePipeline 的音频块批大小GPU 建议 16
beam_size: int = 5 # beam search 宽度GPU turbo 可降到 2 加速(候选数↓ 解码步数↓),质量损失小
class TranslationConfig(BaseModel):
model: str = "Helsinki-NLP/opus-mt-en-zh" # CPU: opus-mt轻量GPU: facebook/nllb-200-distilled-1.3B
device: str = "cpu" # cpu | cuda
src_lang: str = "eng_Latn" # NLLB 语言码:英语
tgt_lang: str = "zho_Hans" # NLLB 语言码:简体中文
batch_size: int = 8 # 排序后单批最大条数CPU: 8GPU: 32显存独占可用大 batch
max_length: int = 256 # 单条最大生成 token 数
sort_by_length: bool = True # 按句子长度排序后分批,减少批内 padding 浪费GPU 收益大)
class SegmentationConfig(BaseModel):
max_words_per_line: int = 14
max_duration_seconds: float = 7.0
min_duration_seconds: float = 1.0
max_chars_per_line: int = 42
class LoggingConfig(BaseModel):
"""日志配置:控制台 + 内存缓冲的最低级别,以及缓冲条数。
分层语义:
- debug进度详情任务 [status pct%]、转写/翻译逐批统计、ffmpeg 命令)
- info任务流转里程碑音频提取/ASR/翻译 的开始与完成、模型加载与卸载)
- error详细错误完整 traceback由 logger.exception 自带)
"""
level: str = "info" # debug | info | warning | error
buffer_size: int = 2000 # /logs 页面内存缓冲条数
@field_validator("level", mode="before")
@classmethod
def _normalize_level(cls, v):
return str(v).lower() if v else "info"
class DocsConfig(BaseModel):
"""/docs Basic Auth 凭据(明文,对齐 server"""
enabled: bool = True
username: str = "admin"
password: str = ""
realm: str = "audio2text docs"
@field_validator("password", mode="before")
@classmethod
def _coerce(cls, v):
return "" if v is None else str(v)
class Settings(BaseModel):
server: ServerConfig = ServerConfig()
storage: StorageConfig = StorageConfig()
processing: ProcessingConfig = ProcessingConfig()
asr: AsrConfig = AsrConfig()
translation: TranslationConfig = TranslationConfig()
segmentation: SegmentationConfig = SegmentationConfig()
logging: LoggingConfig = LoggingConfig()
docs: DocsConfig = DocsConfig()
def upload_dir(self) -> Path:
return Path(self.storage.upload_dir)
def work_dir(self) -> Path:
return Path(self.storage.work_dir)
def output_dir(self) -> Path:
return Path(self.storage.output_dir)
def _load_yaml(path: Path) -> dict:
if not path.exists():
raise FileNotFoundError(
f"未找到配置文件 {path};请先 cp config.example.yaml config.yaml 并填值。"
)
return yaml.safe_load(path.read_text(encoding="utf-8")) or {}
# ---------------- 运行时覆盖 ----------------
# 允许通过设置页修改的配置项(点分路径 -> 类型。config.yaml 是只读挂载,
# 改它需重启容器;运行时覆盖存 DB进程重启后自动加载无需重建镜像。
# 设置页保存时调 save_setting() 写 DB + 清 lru_cache下次 get_settings() 生效。
_applying_overrides = False # 防递归标志_apply_overrides 内部 DB 初始化会回调 get_settings()
_OVERIDEABLE_FIELDS: dict[str, type] = {
"asr.batch_size": int,
"asr.beam_size": int,
"translation.batch_size": int,
"translation.sort_by_length": bool,
}
def _apply_overrides(settings: Settings) -> Settings:
"""从 DB 读取覆盖值并应用到 Settings 对象。
在 lru_cache 的 get_settings() 内部调用,保证缓存的对象已含覆盖。
DB 还没初始化时(首次 import静默跳过用 YAML 原值。
注意get_session_local() -> get_engine() -> _db_path() -> get_settings()
会形成递归。用 _applying_overrides 标志阻断:递归调用直接返回当前 settings
(此时 DB 路径只需 work_dir无覆盖也无妨
"""
global _applying_overrides
if _applying_overrides:
return settings # 递归调用_db_path 触发),直接返回 YAML 原值
_applying_overrides = True
try:
from .database import get_session_local
from .models.setting import Setting
import json
db = get_session_local()()
try:
rows = db.query(Setting).all()
overrides = {r.key: r.value for r in rows}
finally:
db.close()
for key, type_ in _OVERIDEABLE_FIELDS.items():
if key not in overrides:
continue
try:
val = json.loads(overrides[key])
val = type_(val)
except (json.JSONDecodeError, ValueError, TypeError):
continue
_set_nested(settings, key, val)
except Exception:
# DB 未就绪(首次 import 时 database.py 可能还在初始化)-> 跳过,用 YAML 原值
pass
finally:
_applying_overrides = False
return settings
def _set_nested(settings: Settings, key: str, val) -> None:
"""按点分路径设置嵌套属性,如 'asr.batch_size' -> settings.asr.batch_size"""
parts = key.split(".")
obj = settings
for p in parts[:-1]:
obj = getattr(obj, p)
setattr(obj, parts[-1], val)
@lru_cache(maxsize=1)
def get_settings() -> Settings:
"""读取 config.yaml + 应用 DB 覆盖,返回完整 Settings。
结果被 lru_cache 缓存。修改设置后调 reload_settings() 清缓存,
下次调用返回含新值的 Settings。
"""
path = Path(os.getenv("CONFIG_PATH", str(DEFAULT_CONFIG_PATH)))
settings = Settings.model_validate(_load_yaml(path))
return _apply_overrides(settings)
def reload_settings() -> Settings:
"""清缓存并重新读取(含 DB 覆盖),供设置页保存后调用。"""
get_settings.cache_clear()
return get_settings()
def save_setting(key: str, value) -> None:
"""保存单个配置项覆盖到 DB + 清 lru_cache。
Args:
key: 点分路径,必须在 _OVERIDEABLE_FIELDS 中
value: 要保存的值(自动 JSON 编码)
"""
import json
if key not in _OVERIDEABLE_FIELDS:
raise ValueError(f"不允许修改的配置项:{key}")
from .database import get_session_local
from .models.setting import Setting
type_ = _OVERIDEABLE_FIELDS[key]
encoded = json.dumps(type_(value))
db = get_session_local()()
try:
row = db.get(Setting, key)
if row is None:
row = Setting(key=key, value=encoded)
db.add(row)
else:
row.value = encoded
db.commit()
finally:
db.close()
# 清缓存,让后续 get_settings() 读到新值
get_settings.cache_clear()