"""MuseTalk Flask HTTP 服务 — 反向轮询架构的服务端部分. 部署在 RTX2060 本地,接收 gpu_worker.py 的推理请求,调用 MuseTalk 生成口型同步视频。 本文件修复了原 worker.py 的 8 个工程 bug,并新增 /cancel 端点。 环境变量: 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_$$ 接口: GET /health 健康检查 + GPU 显存信息 POST /inference 推理请求(multipart: video + audio) POST /cancel 终止当前推理任务 """ from __future__ import annotations import atexit import logging import os import shutil import signal import subprocess 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()}") # ── 全局状态 ────────────────────────────────────────────────────────── inference_lock = threading.Lock() current_task: dict = {"task_id": None, "process": None, "start_time": 0.0} shutdown_event = threading.Event() # ── 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 _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 def _run_inference(video_path: Path, audio_path: Path, output_path: Path) -> None: """执行 MuseTalk 推理(可被子线程和测试独立调用). 实际部署时替换为 MuseTalk 真实推理逻辑。 此处为示例实现:提取帧 → 合并音视频。 """ fps = _get_video_fps(video_path) logger.info("视频 fps: %.2f", fps) frames_dir = video_path.parent / "frames" frames_dir.mkdir(parents=True, exist_ok=True) _run_ffmpeg( [ "ffmpeg", "-y", "-i", str(video_path), "-r", str(fps), str(frames_dir / "frame_%05d.png"), ], timeout=120, ) frame_files = sorted(frames_dir.glob("*.png")) if not frame_files: raise RuntimeError("未从视频中提取到帧") # TODO: 替换为 MuseTalk 实际推理逻辑 logger.warning("使用示例推理逻辑,未实际调用 MuseTalk 模型") _run_ffmpeg( [ "ffmpeg", "-y", "-i", str(video_path), "-i", str(audio_path), "-c:v", "libx264", "-c:a", "aac", "-shortest", str(output_path), ], timeout=300, ) if not output_path.exists() or output_path.stat().st_size < 1024: raise RuntimeError("推理产物不存在或过小") # ── 路由 ────────────────────────────────────────────────────────────── @app.route("/health", methods=["GET"]) def health(): """健康检查 + GPU 显存信息.""" 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, "timestamp": time.time(), } ) @app.route("/inference", methods=["POST"]) def inference(): """推理请求:multipart form 包含 video 和 audio 文件.""" # 并发控制:检查锁 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())}") # 文件大小检查 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.wav" 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", task_id, video_path.name, audio_path.name) # 更新当前任务信息 current_task["task_id"] = task_id current_task["start_time"] = time.time() # 启动推理进程(用 subprocess 包装,便于超时终止) # 此处直接调用推理函数,实际可改为 subprocess 调用外部脚本 current_task["process"] = "inference_thread" # 标记为运行中 # 在线程中运行推理(支持超时) result_container = {"error": None} def inference_thread(): try: _run_inference(video_path, audio_path, output_path) except Exception as 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(): """启动 Flask 服务.""" # 创建临时目录 Path(Config.temp_dir).mkdir(parents=True, exist_ok=True) logger.info("临时目录: %s", Config.temp_dir) # 打印配置 logger.info("=" * 60) logger.info("MuseTalk Flask Server 启动") logger.info(" 端口: %d", Config.port) logger.info(" 最大并发: %d", Config.max_concurrent) logger.info(" 推理超时: %.0fs", Config.inference_timeout) logger.info(" 视频大小限制: %dMB", Config.video_max_mb) logger.info(" 音频大小限制: %dMB", Config.audio_max_mb) logger.info(" 默认 fps: %.1f", Config.default_fps) logger.info("=" * 60) # 检查 GPU gpu_info = _get_gpu_info() logger.info("GPU 信息: %s", gpu_info) # 启动 Flask(threaded=True 处理并发请求) app.run(host="0.0.0.0", port=Config.port, threaded=True) if __name__ == "__main__": main()