diff --git a/deploy/gpu_worker/gpu_worker.py b/deploy/gpu_worker/gpu_worker.py index ecdeac35b..d9d53f9b3 100644 --- a/deploy/gpu_worker/gpu_worker.py +++ b/deploy/gpu_worker/gpu_worker.py @@ -9,9 +9,12 @@ WORKER_ID 本机唯一 ID(默认 hostname+网卡MAC 后4位) MUSE_TALK_URL 本地 MuseTalk 地址,默认 http://127.0.0.1:7861 POLL_INTERVAL 轮询间隔秒,默认 5 - HEARTBEAT_INTERVAL 心跳间隔秒,默认 15 - REQUEST_TIMEOUT HTTP 请求超时秒,默认 60 - TASK_MAX_RETRY 单个任务最大重试次数(在 Worker 本地的重试),默认 2 + HEARTBEAT_INTERVAL 空闲心跳间隔秒,默认 15 + REQUEST_TIMEOUT HTTP 请求超时秒(下载/推理/上传统一使用),默认 900 + 需与服务端 GPU_TASK_TIMEOUT_SECONDS(默认 900)对齐 + TASK_MAX_RETRY 单任务本地最大重试次数(仅对瞬时错误重试),默认 1 + TASK_HEARTBEAT_INTERVAL 推理期间任务心跳间隔秒,默认 30 + MIN_VIDEO_DURATION_SECONDS 最短输入视频时长秒,小于则直接上报失败,默认 3 用法: python gpu_worker.py @@ -19,13 +22,13 @@ from __future__ import annotations -import json import logging import os import platform import socket import sys import tempfile +import threading import time import uuid from pathlib import Path @@ -54,8 +57,17 @@ class Config: muse_talk_url: str = _env("MUSE_TALK_URL", "http://127.0.0.1:7861").rstrip("/") poll_interval: float = float(_env("POLL_INTERVAL", "5")) heartbeat_interval: float = float(_env("HEARTBEAT_INTERVAL", "15")) - request_timeout: float = float(_env("REQUEST_TIMEOUT", "300")) - task_max_retry: int = int(_env("TASK_MAX_RETRY", "2")) + # #1970:RTX2060 6G 处理 720p 长视频可能 >5min;与服务端 + # GPU_TASK_TIMEOUT_SECONDS 默认值对齐为 900,避免推理被本地/服务端先掐断。 + request_timeout: float = float(_env("REQUEST_TIMEOUT", "900")) + # 本地只在网络/MuseTalk 瞬时错误时重试 1 次;服务端 MAX_ATTEMPTS=3 + # 负责跨 worker/真正超时后的重派发,总尝试次数不再相乘放大。 + task_max_retry: int = int(_env("TASK_MAX_RETRY", "1")) + # 推理期间任务心跳间隔(独立线程 POST /gpu/register 带 task_id) + task_heartbeat_interval: float = float(_env("TASK_HEARTBEAT_INTERVAL", "30")) + # 输入视频最短时长(秒):过短(如 1s)MuseTalk 会 division by zero, + # 本地前置拦截,直接上报 failed,不浪费 GPU 时间 + min_video_duration_seconds: float = float(_env("MIN_VIDEO_DURATION_SECONDS", "3")) worker_id: str = _env("WORKER_ID", "") @classmethod @@ -96,11 +108,20 @@ def _check_musetalk_health() -> tuple[bool, dict]: return False, {"error": str(exc)} -def _register() -> bool: - """向服务端注册 / 心跳,附带 GPU 信息.""" +def _register(task_id: Optional[str] = None) -> bool: + """向服务端注册 / 心跳,附带 GPU 信息。 + + 推理期间的心跳线程传 task_id:服务端会同步刷新该 processing 任务的 + last_heartbeat_at,防止长推理被误判超时回收。 + """ ok, info = _check_musetalk_health() - free_vram = int(info.get("free_vram_mb", 0) or 0) if isinstance(info, dict) else 0 - gpu_name = info.get("gpu_name", "") if isinstance(info, dict) else "" + if isinstance(info, dict): + gpu_info = info.get("gpu", info) + free_vram = int(gpu_info.get("free_vram_mb", gpu_info.get("memory_free_mb", 0)) or 0) + gpu_name = gpu_info.get("gpu_name", info.get("gpu_name", "")) + else: + free_vram = 0 + gpu_name = "" if not gpu_name: # 尝试在 Windows 上读 nvidia-smi gpu_name = _probe_gpu_name() @@ -111,6 +132,8 @@ def _register() -> bool: "free_vram_mb": free_vram, "capabilities": "musetalk", } + if task_id: + payload["task_id"] = task_id try: r = requests.post( f"{Config.api_base_url}/api/v1/gpu/register", @@ -181,11 +204,13 @@ def _download(url: str, path: Path) -> bool: return False -def _call_musetalk(video_path: Path, audio_path: Path, out_path: Path) -> tuple[bool, float, str]: +def _call_musetalk(video_path: Path, audio_path: Path, out_path: Path) -> tuple[bool, float, str, bool]: """调用本地 MuseTalk /inference. - 返回 (success, duration_seconds, error_msg). + 返回 (success, duration_seconds, error_msg, retryable)。 duration 用 ffprobe 读结果视频,失败填 0。 + retryable 仅对瞬时错误(连接失败/超时/5xx)为 True;HTTP 4xx、结果过小 + 等确定性失败不重试,直接上报服务端(服务端 MAX_ATTEMPTS 再决定是否重派发)。 """ try: with open(video_path, "rb") as vf, open(audio_path, "rb") as af: @@ -199,17 +224,34 @@ def _call_musetalk(video_path: Path, audio_path: Path, out_path: Path) -> tuple[ timeout=Config.request_timeout, ) if r.status_code != 200: - return False, 0.0, f"MuseTalk HTTP {r.status_code}: {r.text[:500]}" + retryable = r.status_code >= 500 + return False, 0.0, f"MuseTalk HTTP {r.status_code}: {r.text[:500]}", retryable out_path.parent.mkdir(parents=True, exist_ok=True) out_path.write_bytes(r.content) if out_path.stat().st_size < 1024: - return False, 0.0, f"MuseTalk 返回结果过小 ({out_path.stat().st_size} bytes)" + # 确定性失败(推理产物异常),本地重试大概率还是坏的,不重试 + return False, 0.0, f"MuseTalk 返回结果过小 ({out_path.stat().st_size} bytes)", False duration = _probe_duration(out_path) - return True, duration, "" - except requests.exceptions.Timeout: - return False, 0.0, f"MuseTalk 推理超时(>{Config.request_timeout}s)" + return True, duration, "", False + except (requests.exceptions.Timeout, requests.exceptions.ConnectionError): + # 瞬时网络/超时错误,允许本地重试 1 次;同时调 /cancel 让服务端终止僵尸推理 + _cancel_musetalk() + return False, 0.0, f"MuseTalk 推理超时或连接失败(>{Config.request_timeout}s)", True except Exception as exc: - return False, 0.0, f"MuseTalk 调用异常: {exc}" + return False, 0.0, f"MuseTalk 调用异常: {exc}", False + + +def _cancel_musetalk() -> None: + """调 MuseTalk /cancel 端点终止服务端僵尸推理进程,避免超时后任务还在跑占显存.""" + try: + r = requests.post(f"{Config.muse_talk_url}/cancel", timeout=10) + if r.status_code == 200: + logger.info("已调 MuseTalk /cancel,服务端终止推理") + else: + logger.warning("MuseTalk /cancel 返回 %d: %s", r.status_code, r.text[:200]) + except Exception as exc: + # /cancel 失败不应影响主流程上报 + logger.warning("调 MuseTalk /cancel 异常(忽略): %s", exc) def _probe_duration(path: Path) -> float: @@ -219,9 +261,13 @@ def _probe_duration(path: Path) -> float: out = subprocess.check_output( [ - "ffprobe", "-v", "error", - "-show_entries", "format=duration", - "-of", "default=noprint_wrappers=1:nokey=1", + "ffprobe", + "-v", + "error", + "-show_entries", + "format=duration", + "-of", + "default=noprint_wrappers=1:nokey=1", str(path), ], stderr=subprocess.DEVNULL, @@ -276,42 +322,91 @@ def _report_result(task_id: str, success: bool, duration: float = 0.0, error_msg return False +class TaskHeartbeat(threading.Thread): + """推理期间的任务心跳线程。 + + 主循环的空闲心跳在 ``_handle_task`` 同步阻塞(下载/推理/上传最长 900s) + 期间无法发送,服务端会因任务 last_heartbeat_at 停滞而误判超时回退 pending。 + 本线程每 task_heartbeat_interval 秒(默认 30s)POST /gpu/register 并 + 携带当前 task_id,让服务端持续续期任务心跳;任务处理结束 stop()。 + """ + + def __init__(self, task_id: str, interval: float): + super().__init__(daemon=True, name=f"hb-{task_id[:8]}") + self.task_id = task_id + self.interval = max(5.0, interval) + self._stop_event = threading.Event() + + def run(self) -> None: + # 先立即发一次,再按间隔循环(首次心跳失败不影响主流程) + while not self._stop_event.is_set(): + try: + if _register(self.task_id): + logger.debug("任务 %s 心跳已发送", self.task_id) + except Exception as exc: # noqa: BLE001 + logger.warning("任务 %s 心跳异常(忽略): %s", self.task_id, exc) + self._stop_event.wait(self.interval) + + def stop(self) -> None: + self._stop_event.set() + + def _handle_task(task: dict) -> None: - """处理一条任务(整个串行流程:下载→推理→上传→上报).""" + """处理一条任务(整个串行流程:下载→时长校验→推理→上传→上报)。""" task_id = task["task_id"] logger.info("开始处理任务 %s", task_id) - with tempfile.TemporaryDirectory(prefix="musetalk_") as tmpdir: - tmp = Path(tmpdir) - video_path = tmp / "input.mp4" - audio_path = tmp / "input_audio.bin" - out_path = tmp / "output.mp4" + # 领取任务后立即启动任务级心跳线程,覆盖下载/推理/上报全过程 + hb = TaskHeartbeat(task_id, Config.task_heartbeat_interval) + hb.start() + try: + with tempfile.TemporaryDirectory(prefix="musetalk_") as tmpdir: + tmp = Path(tmpdir) + video_path = tmp / "input.mp4" + audio_path = tmp / "input_audio.bin" + out_path = tmp / "output.mp4" - # 1. 下载 - if not _download(task["video_url"], video_path): - _report_result(task_id, False, 0.0, "下载人物视频失败") - return - if not _download(task["audio_url"], audio_path): - _report_result(task_id, False, 0.0, "下载驱动音频失败") - return + # 1. 下载 + if not _download(task["video_url"], video_path): + _report_result(task_id, False, 0.0, "下载人物视频失败") + return + if not _download(task["audio_url"], audio_path): + _report_result(task_id, False, 0.0, "下载驱动音频失败") + return - # 2. 推理(本地重试) - success = False - duration = 0.0 - err = "" - for attempt in range(Config.task_max_retry + 1): - if attempt > 0: - logger.info("任务 %s 第 %d 次重试...", task_id, attempt + 1) - time.sleep(2) - success, duration, err = _call_musetalk(video_path, audio_path, out_path) - if success: - break - if not success: - logger.error("任务 %s 推理失败: %s", task_id, err) - _report_result(task_id, False, 0.0, err) - return + # 2. 输入时长前置校验:短视频 MuseTalk 会 division by zero, + # 直接上报 failed,不浪费 GPU 时间。ffprobe 不可用/读失败(0.0) + # 时不拦截,交给 MuseTalk 处理,避免误杀。 + video_duration = _probe_duration(video_path) + if video_duration and video_duration < Config.min_video_duration_seconds: + msg = ( + f"视频过短({video_duration:.2f}s < {Config.min_video_duration_seconds:.0f}s)," + "MuseTalk 无法处理" + ) + logger.error("任务 %s %s", task_id, msg) + _report_result(task_id, False, 0.0, msg) + return - # 3. 上报结果(multipart 同时上传文件 → API 代为 PUT 到 OSS,逻辑最稳) - _report_success_with_file(task_id, duration, out_path) + # 3. 推理(本地仅对瞬时错误重试) + success = False + duration = 0.0 + err = "" + retryable = False + for attempt in range(Config.task_max_retry + 1): + if attempt > 0: + logger.info("任务 %s 第 %d 次重试(瞬时错误)...", task_id, attempt + 1) + time.sleep(2) + success, duration, err, retryable = _call_musetalk(video_path, audio_path, out_path) + if success or not retryable: + break + if not success: + logger.error("任务 %s 推理失败: %s", task_id, err) + _report_result(task_id, False, 0.0, err) + return + + # 4. 上报结果(multipart 同时上传文件 → API 代为 PUT 到 OSS,逻辑最稳) + _report_success_with_file(task_id, duration, out_path) + finally: + hb.stop() def _report_success_with_file(task_id: str, duration: float, file_path: Path) -> None: diff --git a/deploy/gpu_worker/musetalk_server.py b/deploy/gpu_worker/musetalk_server.py new file mode 100644 index 000000000..c040e725d --- /dev/null +++ b/deploy/gpu_worker/musetalk_server.py @@ -0,0 +1,1041 @@ +"""MuseTalk Flask HTTP 服务 — 反向轮询架构的服务端部分. + +部署在 RTX2060 本地,接收 gpu_worker.py 的推理请求,调用 MuseTalk 生成口型同步视频。 + +#2000 关键修复: + - 集成真实 MuseTalk 推理(替换原有 stub 代码) + - 音频预处理:22050Hz MP3 → 16kHz mono 16bit WAV(MuseTalk 要求) + - 模型懒加载:首次推理时加载,后续复用,避免重复加载 + - 视频帧循环使用 mirror indexing(乒乓模式),消除循环边界跳变 + - bbox_shift 可通过请求参数配置 + +#1978 性能修复(v2 架构): + MuseTalk 原生支持长音频输入(内部循环视频帧),不需要我们先 loop 视频。 + 正确流程:原视频 + 全量音频 → MuseTalk 推理 → 输出时长=音频时长的无声画面 + → ffmpeg 快速 -c:v copy 替换音轨。推理时间不变(~14s),后处理几秒。 + +环境变量: + MUSE_PORT 监听端口,默认 7861 + MUSE_MAX_CONCURRENT 最大并发推理数,默认 1(GPU 一次只能处理一个) + MUSE_INFERENCE_TIMEOUT 推理超时秒数,默认 600 + MUSE_VIDEO_MAX_MB 视频上传大小限制 MB,默认 100 + MUSE_AUDIO_MAX_MB 音频上传大小限制 MB,默认 20 + MUSE_DEFAULT_FPS 视频 fps 兜底值,默认 25.0 + MUSE_TEMP_DIR 临时文件目录,默认 /tmp/musetalk_$$ + MUSE_VIDEO_ENCODER 循环视频时的编码器(仅兜底):auto(默认)/h264_nvenc/libx264 + MUSE_DIR MuseTalk 仓库路径,默认 /home/ying/projects/MuseTalk + MUSE_MODEL_DIR 模型目录(相对 MUSE_DIR),默认 models/musetalk + MUSE_USE_FLOAT16 使用 FP16 推理,默认 1(开启) + MUSE_BATCH_SIZE 推理批次大小,默认 8 + +接口: + GET /health 健康检查 + GPU 显存信息 + POST /inference 推理请求(multipart: video + audio, form: bbox_shift) + POST /cancel 终止当前推理任务 +""" + +from __future__ import annotations + +import atexit +import copy +import logging +import os +import shutil +import signal +import subprocess +import sys +import threading +import time +from pathlib import Path +from typing import Optional + +from flask import Flask, jsonify, request, send_file + +# ── 日志 ────────────────────────────────────────────────────────────── +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s [%(levelname)s] %(message)s", + datefmt="%Y-%m-%d %H:%M:%S", +) +logger = logging.getLogger("musetalk-server") + +# ── 配置 ────────────────────────────────────────────────────────────── + + +def _env(name: str, default: str = "") -> str: + v = os.environ.get(name, default) + return v.strip() if isinstance(v, str) else default + + +class Config: + port: int = int(_env("MUSE_PORT", "7861")) + max_concurrent: int = int(_env("MUSE_MAX_CONCURRENT", "1")) + inference_timeout: float = float(_env("MUSE_INFERENCE_TIMEOUT", "600")) + video_max_mb: int = int(_env("MUSE_VIDEO_MAX_MB", "100")) + audio_max_mb: int = int(_env("MUSE_AUDIO_MAX_MB", "20")) + default_fps: float = float(_env("MUSE_DEFAULT_FPS", "25.0")) + temp_dir: str = _env("MUSE_TEMP_DIR", f"/tmp/musetalk_{os.getpid()}") + # 循环视频时的编码器(仅当 MuseTalk 输出画面短于音频时的兜底) + video_encoder: str = _env("MUSE_VIDEO_ENCODER", "auto") or "auto" + # 判定音视频时长差异的容差(秒) + duration_epsilon: float = 0.25 + # MuseTalk 仓库路径 + muse_dir: str = _env("MUSE_DIR", "/home/ying/projects/MuseTalk") + # 模型目录(相对 MUSE_DIR) + muse_model_dir: str = _env("MUSE_MODEL_DIR", "models/musetalk") + # 是否使用 FP16(节省显存,RTX2060 建议开启) + use_float16: bool = _env("MUSE_USE_FLOAT16", "1") == "1" + # 推理批次大小(RTX2060 6G 显存建议 4-8) + batch_size: int = int(_env("MUSE_BATCH_SIZE", "8")) + + +# ── 全局状态 ────────────────────────────────────────────────────────── +inference_lock = threading.Lock() +current_task: dict = {"task_id": None, "process": None, "start_time": 0.0} +shutdown_event = threading.Event() + +# ── MuseTalk 模型懒加载 ───────────────────────────────────────────── +_muse_models = None +_muse_models_lock = threading.Lock() +_muse_models_loaded = False +_muse_load_error = None + + +# ── Flask App ───────────────────────────────────────────────────────── +app = Flask(__name__) + + +def _cleanup_temp_dir(): + """退出时清理临时目录.""" + if os.path.exists(Config.temp_dir): + try: + shutil.rmtree(Config.temp_dir) + logger.info("已清理临时目录: %s", Config.temp_dir) + except Exception as exc: + logger.warning("清理临时目录失败: %s", exc) + + +atexit.register(_cleanup_temp_dir) + + +def _signal_handler(signum, frame): + """优雅退出.""" + logger.info("收到信号 %s,准备退出...", signum) + shutdown_event.set() + if current_task["process"]: + logger.info("终止正在进行的推理进程...") + try: + current_task["process"].terminate() + current_task["process"].wait(timeout=5) + except Exception: + pass + _cleanup_temp_dir() + exit(0) + + +signal.signal(signal.SIGTERM, _signal_handler) +signal.signal(signal.SIGINT, _signal_handler) + + +# ── 工具函数 ────────────────────────────────────────────────────────── + + +def _get_gpu_info() -> dict: + """获取 GPU 显存信息(通过 nvidia-smi).""" + try: + out = subprocess.check_output( + [ + "nvidia-smi", + "--query-gpu=name,memory.total,memory.used,memory.free", + "--format=csv,noheader,nounits", + ], + stderr=subprocess.DEVNULL, + timeout=5, + ) + parts = out.decode().strip().split(",") + if len(parts) >= 4: + return { + "gpu_name": parts[0].strip(), + "memory_total_mb": int(parts[1].strip()), + "memory_used_mb": int(parts[2].strip()), + "memory_free_mb": int(parts[3].strip()), + } + except Exception as exc: + logger.warning("nvidia-smi 失败: %s", exc) + return {"gpu_name": "unknown", "memory_total_mb": 0, "memory_used_mb": 0, "memory_free_mb": 0} + + +def _get_video_fps(video_path: Path) -> float: + """用 ffprobe 读视频帧率,失败或为 0 时返回 default_fps.""" + try: + out = subprocess.check_output( + [ + "ffprobe", + "-v", + "error", + "-select_streams", + "v:0", + "-show_entries", + "stream=r_frame_rate", + "-of", + "default=noprint_wrappers=1:nokey=1", + str(video_path), + ], + stderr=subprocess.DEVNULL, + timeout=10, + ) + fps_str = out.decode().strip() + if "/" in fps_str: + num, den = fps_str.split("/") + fps = float(num) / float(den) if float(den) != 0 else 0.0 + else: + fps = float(fps_str) if fps_str else 0.0 + return fps if fps > 0 else Config.default_fps + except Exception as exc: + logger.warning("ffprobe 读 fps 失败: %s,使用默认 %.1f", exc, Config.default_fps) + return Config.default_fps + + +def _get_media_duration(path: Path) -> float: + """用 ffprobe 读媒体时长(秒),失败返回 0.0.""" + try: + out = subprocess.check_output( + [ + "ffprobe", + "-v", + "error", + "-show_entries", + "format=duration", + "-of", + "default=noprint_wrappers=1:nokey=1", + str(path), + ], + stderr=subprocess.DEVNULL, + timeout=10, + ) + duration = float(out.decode().strip()) + return duration if duration > 0 else 0.0 + except Exception as exc: + logger.warning("ffprobe 读时长失败 %s: %s", path, exc) + return 0.0 + + +def _pick_video_encoder() -> str: + """选择视频编码器:配置指定则用指定值;auto 时探测 NVENC 是否可用,不可用回退 libx264.""" + configured = Config.video_encoder.strip() + if configured in ("h264_nvenc", "libx264"): + return configured + # auto:探测本机 ffmpeg 是否编译了 h264_nvenc + try: + result = subprocess.run( + ["ffmpeg", "-hide_banner", "-encoders"], + stdout=subprocess.PIPE, + stderr=subprocess.DEVNULL, + timeout=10, + check=False, + ) + if b"h264_nvenc" in result.stdout: + return "h264_nvenc" + except Exception as exc: + logger.warning("探测 ffmpeg 编码器失败,回退 libx264: %s", exc) + return "libx264" + + +def _preprocess_audio(input_path: Path, output_path: Path, target_sr: int = 16000) -> None: + """将输入音频转换为 MuseTalk 要求的格式:16kHz mono 16bit WAV. + + MuseTalk 的 whisper audio2feature 要求 16kHz 采样率的单声道音频。 + 当前 TTS 输出为 22050Hz MP3,不转换会导致 mel 频谱错位、 + 音素特征提取错误,口型只跟能量不跟音素。 + """ + cmd = [ + "ffmpeg", "-y", "-v", "warning", + "-i", str(input_path), + "-ar", str(target_sr), # 重采样到 16kHz + "-ac", "1", # 单声道 + "-sample_fmt", "s16", # 16bit PCM + str(output_path), + ] + _run_ffmpeg(cmd, timeout=60) + + if not output_path.exists() or output_path.stat().st_size < 100: + raise RuntimeError(f"音频预处理失败: {output_path}") + + logger.info("音频预处理完成: %s → 16kHz mono WAV", input_path.name) + + +def _mux_video_with_audio( + video_path: Path, + audio_path: Path, + output_path: Path, + timeout: float = 300, +) -> None: + """把无声画面视频与驱动音频封装为最终结果. + + #1978 v2 架构:MuseTalk 已处理全量音频,输出视频时长=音频时长。 + 此处仅做快速封装:-map 0:v:0 -map 1:a:0 强制取画面+驱动音频, + -c:v copy 无损秒级封装(不重编码),-shortest 以较短流为准。 + + 仅当 MuseTalk 输出画面短于音频时(极端兜底),才启用 -stream_loop + NVENC + 循环视频到音频长度。正常情况下走 copy 快速路径。 + """ + video_duration = _get_media_duration(video_path) + audio_duration = _get_media_duration(audio_path) + + # 判断是否需要兜底循环(正常情况下 MuseTalk 输出已 >= 音频时长) + need_loop_fallback = bool( + audio_duration > 0 and video_duration > 0 and video_duration < audio_duration - Config.duration_epsilon + ) + + if need_loop_fallback: + # 兜底:MuseTalk 输出画面不足,循环补齐 + encoder = _pick_video_encoder() + preset = "p4" if encoder == "h264_nvenc" else "veryfast" + logger.warning( + "MuseTalk 输出(%.2fs)短于音频(%.2fs),兜底循环视频以 %s 重编码", + video_duration, + audio_duration, + encoder, + ) + + def build_cmd(enc: str, pre: str) -> list: + return [ + "ffmpeg", + "-y", + "-stream_loop", + "-1", + "-i", + str(video_path), + "-i", + str(audio_path), + "-map", + "0:v:0", + "-map", + "1:a:0", + "-c:v", + enc, + "-preset", + pre, + "-c:a", + "aac", + "-b:a", + "128k", + "-t", + f"{audio_duration:.3f}", + str(output_path), + ] + + try: + _run_ffmpeg(build_cmd(encoder, preset), timeout=timeout) + except RuntimeError: + if encoder == "h264_nvenc": + logger.warning("h264_nvenc 兜底失败,回退 libx264 重试") + _run_ffmpeg(build_cmd("libx264", "veryfast"), timeout=timeout) + else: + raise + else: + # 正常快速路径:-c:v copy 无损封装,仅替换音轨为驱动音频 + cmd = [ + "ffmpeg", + "-y", + "-i", + str(video_path), + "-i", + str(audio_path), + "-map", + "0:v:0", + "-map", + "1:a:0", + "-c:v", + "copy", + "-c:a", + "aac", + "-b:a", + "128k", + "-shortest", + str(output_path), + ] + _run_ffmpeg(cmd, timeout=timeout) + + +def _check_file_size(file, max_mb: int, label: str) -> Optional[str]: + """检查文件大小,超限返回错误信息,否则返回 None.""" + file.seek(0, 2) + size = file.tell() + file.seek(0) + max_bytes = max_mb * 1024 * 1024 + if size > max_bytes: + return f"{label} 文件大小 {size / (1024*1024):.1f}MB 超过限制 {max_mb}MB" + if size == 0: + return f"{label} 文件为空" + return None + + +def _run_ffmpeg(cmd: list, timeout: float = 120) -> subprocess.CompletedProcess: + """运行 ffmpeg 命令,检查返回码和超时.""" + try: + result = subprocess.run( + cmd, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + timeout=timeout, + check=True, + ) + return result + except subprocess.CalledProcessError as exc: + stderr = exc.stderr.decode(errors="ignore") if exc.stderr else "" + raise RuntimeError(f"ffmpeg 失败 (code={exc.returncode}): {stderr[:500]}") from exc + except subprocess.TimeoutExpired as exc: + raise RuntimeError(f"ffmpeg 超时(>{timeout}s)") from exc + + +# ── MuseTalk 模型加载 ──────────────────────────────────────────────── + + +def _load_musetalk_models(): + """懒加载 MuseTalk 模型(全局单例,首次调用时加载). + + 加载 VAE、UNet、PositionalEncoder 三个核心组件。 + 加载到 GPU 后转为 FP16(如果配置开启)以节省显存。 + RTX2060 6G 显存,FP16 大约需要 3-4GB。 + """ + global _muse_models, _muse_models_loaded, _muse_load_error + + if _muse_models_loaded: + return _muse_models + if _muse_load_error is not None: + raise _muse_load_error + + with _muse_models_lock: + if _muse_models_loaded: + return _muse_models + + try: + muse_dir = Path(Config.muse_dir) + if not muse_dir.exists(): + raise FileNotFoundError( + f"MuseTalk 目录不存在: {muse_dir}\n" + f"请设置 MUSE_DIR 环境变量指向 MuseTalk 仓库路径" + ) + + # 将 MuseTalk 加入 sys.path(只在首次加载时) + muse_str = str(muse_dir) + if muse_str not in sys.path: + sys.path.insert(0, muse_str) + + import torch + from musetalk.utils.utils import load_all_model + + device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") + logger.info("MuseTalk 使用设备: %s", device) + + # 自动检测模型路径 + # 优先检测 v1.5 模型,然后回退到 v1 + v15_unet = muse_dir / "models" / "musetalkV15" / "unet.pth" + v1_unet = muse_dir / "models" / "musetalk" / "pytorch_model.bin" + + if v15_unet.exists(): + unet_model_path = str(v15_unet) + unet_config = str(muse_dir / "models" / "musetalkV15" / "musetalk.json") + model_version = "v15" + elif v1_unet.exists(): + unet_model_path = str(v1_unet) + unet_config = str(muse_dir / "models" / "musetalk" / "config.json") + model_version = "v1" + else: + raise FileNotFoundError( + f"未找到 MuseTalk 模型权重。\n" + f"检查路径: {v15_unet} 或 {v1_unet}\n" + f"请确认模型已下载到 MuseTalk 仓库的 models/ 目录下" + ) + + logger.info("加载 MuseTalk %s 模型: %s", model_version, unet_model_path) + + vae, unet, pe = load_all_model( + unet_model_path=unet_model_path, + vae_type="sd-vae", + unet_config=unet_config, + device=device, + ) + timesteps = torch.tensor([0], device=device) + + # FP16 转换(节省 ~50% 显存) + if Config.use_float16: + pe = pe.half() + vae.vae = vae.vae.half() + unet.model = unet.model.half() + logger.info("已启用 FP16 推理") + + pe = pe.to(device) + vae.vae = vae.vae.to(device) + unet.model = unet.model.to(device) + + # 加载 AudioProcessor 和 face parsing + from musetalk.utils.audio_processor import AudioProcessor + from musetalk.utils.face_parsing import FaceParsing + + audio_processor = AudioProcessor() + face_parsing = FaceParsing() + + _muse_models = { + "vae": vae, + "unet": unet, + "pe": pe, + "timesteps": timesteps, + "audio_processor": audio_processor, + "face_parsing": face_parsing, + "device": device, + "model_version": model_version, + } + _muse_models_loaded = True + logger.info("MuseTalk 模型加载完成 (版本=%s, 设备=%s, fp16=%s)", + model_version, device, Config.use_float16) + + # 打印显存使用情况 + if torch.cuda.is_available(): + allocated = torch.cuda.memory_allocated() / 1024**2 + reserved = torch.cuda.memory_reserved() / 1024**2 + logger.info("GPU 显存: 已分配 %.0fMB, 已预留 %.0fMB", allocated, reserved) + + return _muse_models + + except Exception as exc: + _muse_load_error = exc + logger.error("MuseTalk 模型加载失败: %s", exc) + raise + + +# ── MuseTalk 推理核心 ───────────────────────────────────────────────── + + +def _mirror_index(size: int, index: int) -> int: + """乒乓式循环索引,避免循环边界硬切跳变. + + 效果: 0→1→2→...→N→N-1→...→1→0→1→... + 比简单的 index % size 在边界处更平滑。 + """ + if size == 0: + return 0 + turn = index // size + res = index % size + if turn % 2 == 0: + return res + else: + return size - res - 1 + + +def _run_inference( + video_path: Path, + audio_path: Path, + output_path: Path, + bbox_shift: int = 0, +) -> None: + """执行 MuseTalk 真实推理. + + 流程: + 1. 音频预处理:任意格式 → 16kHz mono 16bit WAV + 2. 加载/复用 MuseTalk 模型(VAE + UNet + PE + Whisper) + 3. 视频预处理:提取帧 → 人脸检测 → 获取 bbox → VAE 编码 latent + 4. 音频特征提取:whisper 提取 audio features (50×384 per chunk) + 5. 批量推理:UNet 去噪 → VAE 解码 → 得到口型同步的人脸帧 + 6. 帧合成:将生成的人脸贴回原帧(使用 face parsing 做边缘融合) + 7. 输出无声视频(后续由 _mux_video_with_audio 封装 TTS 音频) + + Args: + video_path: 输入视频路径 + audio_path: 输入音频路径(任意格式,会被预处理为 16kHz WAV) + output_path: 输出无声视频路径 + bbox_shift: 口型区域垂直偏移量,默认 0,范围 [-5, 5] + """ + import cv2 + import numpy as np + import torch + from tqdm import tqdm + + from musetalk.utils.preprocessing import get_landmark_and_bbox as _orig_get_landmark_and_bbox + from musetalk.utils.blending import get_image + import tempfile as _tempfile, math as _math, shutil as _shutil + from einops import rearrange as _rearrange + + # read_imgs: 读取视频帧(支持视频文件路径),返回 numpy BGR 帧列表 + def read_imgs(path): + import cv2 as _cv2 + cap = _cv2.VideoCapture(str(path)) + frames = [] + while True: + ret, frame = cap.read() + if not ret: + break + frames.append(frame) + cap.release() + return frames + + # get_landmark_and_bbox 适配:旧版签名(img_list, upperbondrange=0),且 img_list 是文件路径列表 + def get_landmark_and_bbox(frames, vid_pts=0, bbox_shift=0): + import cv2 as _cv2 + _tmpdir = _tempfile.mkdtemp(prefix="muse_frames_") + frame_paths = [] + for _i, _frm in enumerate(frames): + _fp = f"{_tmpdir}/{_i:08d}.png" + _cv2.imwrite(_fp, _frm) + frame_paths.append(_fp) + coords_list, _ = _orig_get_landmark_and_bbox(frame_paths, upperbondrange=bbox_shift) + _shutil.rmtree(_tmpdir, ignore_errors=True) + _sentinel = object() + coords_list = [c if c is not None else _sentinel for c in coords_list] + return coords_list, _sentinel + + # 加载模型(首次调用时加载,后续复用) + models = _load_musetalk_models() + vae = models["vae"] + unet = models["unet"] + pe = models["pe"] + timesteps = models["timesteps"] + audio_processor = models["audio_processor"] + device = models["device"] + model_version = models["model_version"] + + # 给旧版 AudioProcessor 动态添加 feature2chunks 方法 + import types as _types + def _feature2chunks(self, feature_array, fps=25, weight_dtype=None, + batch_size=8, audio_padding_length_left=2, + audio_padding_length_right=2): + import torch + sr = 16000 + audio_fps = 50 + chunk_len = 2 * (audio_padding_length_left + audio_padding_length_right + 1) + whisper_idx_multiplier = audio_fps / fps + num_frames = int(_math.floor((len(feature_array) / sr) * fps)) + actual_length = int(_math.floor((len(feature_array) / sr) * audio_fps)) + inputs = self.feature_extractor( + feature_array, return_tensors="pt", sampling_rate=sr + ).input_features.to(device) + if weight_dtype is not None: + inputs = inputs.to(dtype=weight_dtype) + global _whisper_enc_model + if "_whisper_enc_model" not in globals() or _whisper_enc_model is None: + from transformers import WhisperModel + _wp = str(Path(Config.muse_dir) / "models" / "whisper") + _whisper_enc_model = WhisperModel.from_pretrained(_wp).to(device) + _whisper_enc_model.eval() + if Config.use_float16: + _whisper_enc_model = _whisper_enc_model.half() + with torch.no_grad(): + _af = _whisper_enc_model.encoder(inputs, output_hidden_states=True).hidden_states + _af = torch.stack(_af, dim=2) + _af = _af[0, :actual_length, ...] + _pn = int(_math.ceil(whisper_idx_multiplier)) + _af = torch.cat([ + torch.zeros_like(_af[:_pn * audio_padding_length_left]), + _af, + torch.zeros_like(_af[:_pn * 3 * audio_padding_length_right]), + ], dim=0) + _all = [] + for _fi in range(num_frames): + _ai = int(_math.floor(_fi * whisper_idx_multiplier)) + _clip = _af[_ai:_ai + chunk_len] + if _clip.shape[0] < chunk_len: + _pad = torch.zeros(chunk_len - _clip.shape[0], *_clip.shape[1:], + device=device, dtype=_clip.dtype) + _clip = torch.cat([_clip, _pad], dim=0) + _all.append(_clip) + _prompts = torch.stack(_all, dim=0) + _prompts = _rearrange(_prompts, "b c h w -> b (c h) w") + return _prompts + audio_processor.feature2chunks = _types.MethodType(_feature2chunks, audio_processor) + + fps = _get_video_fps(video_path) + audio_duration = _get_media_duration(audio_path) + video_duration = _get_media_duration(video_path) + logger.info( + "MuseTalk 推理开始: video=%.2fs, audio=%.2fs, fps=%.1f, bbox_shift=%d", + video_duration, audio_duration, fps, bbox_shift, + ) + + # ── Step 1: 音频预处理(关键修复:22050Hz MP3 → 16kHz mono WAV)── + audio_wav_path = video_path.parent / "audio_16k_mono.wav" + _preprocess_audio(audio_path, audio_wav_path, target_sr=16000) + + # ── Step 2: 视频帧提取 ── + input_frames = read_imgs(str(video_path)) + total_frames = len(input_frames) + if total_frames == 0: + raise RuntimeError("未能从视频中提取到任何帧") + logger.info("提取到 %d 帧视频画面", total_frames) + + # ── Step 3: 人脸检测 & bbox 计算 ── + coord_list, coord_placeholder = get_landmark_and_bbox( + input_frames, vid_pts=0, bbox_shift=bbox_shift + ) + logger.info("人脸检测完成,有效 bbox: %d/%d", sum(1 for c in coord_list if c is not coord_placeholder), total_frames) + + # 使用 mirror indexing 循环帧和坐标(避免硬切跳变) + num_output_frames = int(audio_duration * fps) + if num_output_frames <= 0: + num_output_frames = total_frames + + # ── Step 4: 音频特征提取 ── + # 使用 librosa 加载预处理后的 16kHz 音频 + import librosa + audio_array, _ = librosa.load(str(audio_wav_path), sr=16000, mono=True) + whisper_features = audio_processor.feature2chunks( + feature_array=audio_array, + fps=fps, + weight_dtype=(torch.float16 if Config.use_float16 else torch.float32), + batch_size=Config.batch_size, + ) + if isinstance(whisper_features, torch.Tensor): + whisper_features = whisper_features.detach().cpu() + torch.cuda.empty_cache() + logger.info("音频特征提取完成: %d 个 chunk", len(whisper_features)) + + # ── Step 5: 逐帧裁剪人脸并编码为 8ch latent (masked+ref) ── + face_parsing = models.get("face_parsing", None) + input_latent_list = [] + valid_frame_indices = [] # 记录成功编码的帧索引(跳过无脸帧) + + with torch.no_grad(): + for idx, (frame, bbox) in enumerate(zip(input_frames, coord_list)): + if bbox is coord_placeholder: + continue + x1, y1, x2, y2 = bbox + # v1.5 额外扩展下边界(下巴区域),与 Step 7 保持一致 + extra_y2 = 10 if model_version == "v15" else 0 + y2_eff = min(y2 + extra_y2, frame.shape[0]) + if y2_eff <= y1 or x2 <= x1: + continue + # 裁剪人脸区域 → resize 256×256 + crop = frame[y1:y2_eff, x1:x2] + if crop.size == 0: + continue + crop_rgb = cv2.cvtColor(crop, cv2.COLOR_BGR2RGB) + crop_resized = cv2.resize(crop_rgb, (256, 256), interpolation=cv2.INTER_LANCZOS4) + # 使用 VAE 的 get_latents_for_unet 得到 8 通道输入 + # get_latents_for_unet 内部: preprocess(half_mask=True) encode + preprocess(half_mask=False) encode → cat → [1,8,32,32] + latents = vae.get_latents_for_unet(crop_resized).detach().cpu() + input_latent_list.append(latents) + # 保存此帧的实际bbox(含extra_y2)和原帧索引供 Step 7 使用 + valid_frame_indices.append((idx, x1, y1, x2, y2_eff)) + + # 构建循环列表:正序+倒序,实现旧版的平滑首尾帧循环 + frame_list_cycle = input_frames + input_frames[::-1] + coord_cycle = [] + for _i, _x1, _y1, _x2, _y2 in valid_frame_indices: + coord_cycle.append((_x1, _y1, _x2, _y2)) + coord_cycle = coord_cycle + coord_cycle[::-1] + latent_cycle = input_latent_list + input_latent_list[::-1] + valid_cycle = valid_frame_indices + [(i, x1, y1, x2, y2) for (i, x1, y1, x2, y2) in reversed(valid_frame_indices)] + + torch.cuda.empty_cache() + logger.info("人脸裁剪+VAE 编码完成: %d 个有效latent", len(input_latent_list)) + + # ── Step 6: 批量推理(仿旧版 datagen 循环)── + res_frame_list = [] + video_num = len(whisper_features) + bs = min(Config.batch_size, 2) # RTX2060 6G 限制batch=2防OOM + total_batches = (video_num + bs - 1) // bs + + for bi in tqdm(range(total_batches), desc="MuseTalk 推理"): + whisper_batch = whisper_features[bi*bs:(bi+1)*bs] + if len(whisper_batch) == 0: + break + # 对应 latent 索引(循环取 latent_cycle) + latent_batch_parts = [] + for j in range(len(whisper_batch)): + global_idx = bi*bs + j + lat_idx = global_idx % len(latent_cycle) + latent_batch_parts.append(latent_cycle[lat_idx]) + + # whisper_batch 是 [bs,50,384] tensor slice (feature2chunks 已返回 stacked tensor) + if isinstance(whisper_batch, list): + whisper_batch_t = torch.stack(whisper_batch).to(device) + else: + whisper_batch_t = whisper_batch.to(device) + latent_batch_t = torch.cat(latent_batch_parts, dim=0).to(device) + if Config.use_float16: + latent_batch_t = latent_batch_t.to(dtype=unet.model.dtype) + whisper_batch_t = whisper_batch_t.to(dtype=unet.model.dtype) + + audio_feature_batch = pe(whisper_batch_t) + + with torch.no_grad(): + pred_latents = unet.model( + latent_batch_t, + timesteps, + encoder_hidden_states=audio_feature_batch, + ).sample + + recon_frames = vae.decode_latents(pred_latents) + for rf in recon_frames: + res_frame_list.append(rf) + del pred_latents, recon_frames, latent_batch_t, whisper_batch_t + if "audio_feature_batch" in dir(): + try: del audio_feature_batch + except: pass + torch.cuda.empty_cache() + + logger.info("推理完成,生成 %d 帧", len(res_frame_list)) + + # ── Step 7: 合成最终帧并写PNG ── + output_frames_dir = video_path.parent / "output_frames" + output_frames_dir.mkdir(parents=True, exist_ok=True) + + for i, res_frame in enumerate(tqdm(res_frame_list[:num_output_frames], desc="合成帧")): + cyc_i = i % len(coord_cycle) + bbox = coord_cycle[cyc_i] + ori_frame = copy.deepcopy(frame_list_cycle[cyc_i]) + if bbox is None: + combined = ori_frame + else: + x1, y1, x2, y2 = bbox + try: + res_frame_resized = cv2.resize( + res_frame.astype(np.uint8), (x2-x1, y2-y1) + ) + except Exception: + combined = ori_frame + cv2.imwrite(str(output_frames_dir / f"{i:08d}.png"), combined) + continue + # face parsing 融合(fp 实例存在时传入) + try: + if face_parsing is not None: + combined = get_image(ori_frame, res_frame_resized, + [x1, y1, x2, y2], fp=face_parsing) + else: + combined = get_image(ori_frame, res_frame_resized, + [x1, y1, x2, y2]) + except Exception as e: + logger.warning("face parsing 融合失败,使用简单粘贴: %s", e) + # 简单粘贴 fallback + combined = ori_frame.copy() + try: + combined[y1:y2, x1:x2] = res_frame_resized + except Exception: + combined = ori_frame + cv2.imwrite(str(output_frames_dir / f"{i:08d}.png"), combined) + + # 帧序列 → 无声视频 + silent_video_path = video_path.parent / "silent_output.mp4" + _run_ffmpeg( + [ + "ffmpeg", "-y", "-v", "warning", + "-r", str(fps), + "-f", "image2", + "-i", str(output_frames_dir / "%08d.png"), + "-vcodec", "libx264", + "-vf", "format=yuv420p", + "-crf", "18", + str(silent_video_path), + ], + timeout=300, + ) + + # 将无声视频复制到输出路径 + shutil.copy2(str(silent_video_path), str(output_path)) + + # 清理中间文件 + try: + shutil.rmtree(str(output_frames_dir)) + if audio_wav_path.exists(): + audio_wav_path.unlink() + if silent_video_path.exists() and str(silent_video_path) != str(output_path): + silent_video_path.unlink() + except Exception as e: + logger.warning("清理中间文件失败: %s", e) + + logger.info("MuseTalk 推理完成: output=%s, duration=%.2fs", + output_path.name, _get_media_duration(output_path)) + + +# ── 路由 ────────────────────────────────────────────────────────────── + + +@app.route("/health", methods=["GET"]) +def health(): + """健康检查 + GPU 显存信息 + MuseTalk 模型状态.""" + gpu_info = _get_gpu_info() + task_info = { + "task_id": current_task["task_id"], + "running": current_task["process"] is not None, + "elapsed_seconds": time.time() - current_task["start_time"] if current_task["start_time"] else 0.0, + } + return jsonify( + { + "status": "healthy", + "gpu": gpu_info, + "current_task": task_info, + "musetalk_loaded": _muse_models_loaded, + "musetalk_load_error": str(_muse_load_error) if _muse_load_error else None, + "timestamp": time.time(), + } + ) + + +@app.route("/inference", methods=["POST"]) +def inference(): + """推理请求:multipart form 包含 video 和 audio 文件. + + 可选 form 参数: + bbox_shift: 口型区域垂直偏移量,默认 0,范围 [-5, 5] + """ + # 并发控制:检查锁 + if not inference_lock.acquire(blocking=False): + return jsonify({"error": "GPU 正在处理其他任务,请稍后重试", "status": "busy"}), 503 + + task_id = None + video_path = None + audio_path = None + output_path = None + + try: + # 解析参数 + if "video" not in request.files or "audio" not in request.files: + return jsonify({"error": "缺少 video 或 audio 文件"}), 400 + + video_file = request.files["video"] + audio_file = request.files["audio"] + task_id = request.form.get("task_id", f"task_{int(time.time())}") + bbox_shift = int(request.form.get("bbox_shift", "0")) + bbox_shift = max(-5, min(5, bbox_shift)) # 限制范围 + + # 文件大小检查 + err = _check_file_size(video_file, Config.video_max_mb, "视频") + if err: + return jsonify({"error": err}), 413 + err = _check_file_size(audio_file, Config.audio_max_mb, "音频") + if err: + return jsonify({"error": err}), 413 + + # 保存到临时目录 + task_dir = Path(Config.temp_dir) / task_id + task_dir.mkdir(parents=True, exist_ok=True) + video_path = task_dir / "input.mp4" + audio_path = task_dir / "input_audio.bin" + output_path = task_dir / "output.mp4" + + video_file.save(str(video_path)) + audio_file.save(str(audio_path)) + + logger.info("开始推理 task_id=%s, video=%s, audio=%s, bbox_shift=%d", + task_id, video_path.name, audio_path.name, bbox_shift) + + # 更新当前任务信息 + current_task["task_id"] = task_id + current_task["start_time"] = time.time() + current_task["process"] = "inference_thread" # 标记为运行中 + + # 在线程中运行推理(支持超时) + result_container = {"error": None} + + def inference_thread(): + try: + _run_inference(video_path, audio_path, output_path, bbox_shift=bbox_shift) + except Exception as exc: + logger.exception("推理异常: %s", exc) + result_container["error"] = str(exc) + + thread = threading.Thread(target=inference_thread) + thread.start() + thread.join(timeout=Config.inference_timeout) + + if thread.is_alive(): + # 超时,终止 + logger.error("推理超时 (>%ds),终止任务 %s", Config.inference_timeout, task_id) + return jsonify({"error": f"推理超时(>{Config.inference_timeout}s)", "task_id": task_id}), 504 + + if result_container["error"]: + logger.error("推理失败 task_id=%s: %s", task_id, result_container["error"]) + return jsonify({"error": result_container["error"], "task_id": task_id}), 500 + + # 返回结果文件 + logger.info("推理完成 task_id=%s, output=%s", task_id, output_path) + return send_file(str(output_path), mimetype="video/mp4", as_attachment=True, download_name=f"{task_id}.mp4") + + except Exception as exc: + logger.exception("推理异常: %s", exc) + return jsonify({"error": str(exc)}), 500 + + finally: + # 释放锁,清理当前任务信息 + inference_lock.release() + current_task["task_id"] = None + current_task["process"] = None + current_task["start_time"] = 0.0 + + # 清理临时文件 + if video_path and video_path.parent.exists(): + try: + shutil.rmtree(video_path.parent) + logger.info("已清理临时目录: %s", video_path.parent) + except Exception as exc: + logger.warning("清理临时目录失败: %s", exc) + + +@app.route("/cancel", methods=["POST"]) +def cancel(): + """终止当前正在进行的推理任务.""" + if current_task["task_id"] is None: + return jsonify({"message": "当前无正在运行的任务"}) + + task_id = current_task["task_id"] + logger.info("收到取消请求,终止任务 %s", task_id) + + # 终止推理进程(如果是 subprocess) + if current_task["process"] and current_task["process"] != "inference_thread": + try: + current_task["process"].terminate() + current_task["process"].wait(timeout=5) + logger.info("已终止推理进程") + except Exception as exc: + logger.warning("终止进程失败: %s", exc) + + # 清理临时文件 + task_dir = Path(Config.temp_dir) / task_id + if task_dir.exists(): + try: + shutil.rmtree(task_dir) + logger.info("已清理临时目录: %s", task_dir) + except Exception as exc: + logger.warning("清理临时目录失败: %s", exc) + + # 重置当前任务 + current_task["task_id"] = None + current_task["process"] = None + current_task["start_time"] = 0.0 + + return jsonify({"message": f"已取消任务 {task_id}"}) + + +# ── 主入口 ──────────────────────────────────────────────────────────── + + +def main(): + import os as _os + _os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "max_split_size_mb:128") + """启动 Flask 服务.""" + # 创建临时目录 + Path(Config.temp_dir).mkdir(parents=True, exist_ok=True) + logger.info("临时目录: %s", Config.temp_dir) + + gpu_info = _get_gpu_info() + logger.info( + "GPU: %s (显存 %dMB / %dMB)", + gpu_info["gpu_name"], + gpu_info["memory_used_mb"], + gpu_info["memory_total_mb"], + ) + logger.info("MuseTalk 仓库路径: %s", Config.muse_dir) + logger.info( + "启动 MuseTalk Server: port=%d, timeout=%.0fs, max_concurrent=%d, fp16=%s, batch_size=%d", + Config.port, + Config.inference_timeout, + Config.max_concurrent, + Config.use_float16, + Config.batch_size, + ) + + app.run(host="0.0.0.0", port=Config.port, threaded=True) + + +if __name__ == "__main__": + main()