Files
xiaoxia-saas/deploy/gpu_worker/musetalk_server.py
T
xiaoxia 1855dc54dc
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (push) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (push) Has been skipped
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
CI/CD Pipeline / Check push changed paths (push) Successful in 5s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 3s
CI/CD Pipeline / Frontend Lint (push) Has been skipped
CI/CD Pipeline / PR Build API Image (push) Has been skipped
CI/CD Pipeline / PR Build Web Image (push) Has been skipped
CI/CD Pipeline / PR Build Worker Image (push) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Validate - Style (pull_request) Has been skipped
CI/CD Pipeline / Validate - Security (pull_request) Has been skipped
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been skipped
CI/CD Pipeline / Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 42s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m8s
CI/CD Pipeline / Build Staging API Image (push) Successful in 1m1s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (push) Successful in 27s
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m44s
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
CI/CD Pipeline / Build Staging Web Image (push) Successful in 1m9s
CI/CD Pipeline / Retag skipped Staging API Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (push) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (push) Failing after 2m45s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Successful in 1m0s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m16s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 4m55s
CI/CD Pipeline / ACR Image Cleanup (push) Successful in 1m47s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 4m52s
CI/CD Pipeline / Validate - Style (push) Successful in 5m32s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 2m57s
CI/CD Pipeline / Validate - Python (mypy + alembic) (push) Successful in 6m10s
AI Code Review / AI Code Review (pull_request) Successful in 6m53s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 4m3s
CI/CD Pipeline / Unit Tests (push) Successful in 11m50s
CI/CD Pipeline / Validate - Security (push) Successful in 15m34s
CI/CD Pipeline / Build Production API Image (push) Has been skipped
CI/CD Pipeline / Build Production Web Image (push) Has been skipped
CI/CD Pipeline / Build Production Worker Image (push) Has been skipped
CI/CD Pipeline / CI Gate (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Canary Release to Production (push) Has been skipped
fix(gpu): 修复MuseTalk蓝色块 — Step5删除多余BGR2RGB转换,直接传BGR给VAE (#2015)
Co-authored-by: xiaoxia <dev@xiaoxiajianji.com>
Co-committed-by: xiaoxia <dev@xiaoxiajianji.com>
2026-09-23 00:08:08 +08:00

1251 lines
51 KiB
Python
Raw 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.
"""MuseTalk Flask HTTP 服务 — 反向轮询架构的服务端部分.
部署在 RTX2060 本地,接收 gpu_worker.py 的推理请求,调用 MuseTalk 生成口型同步视频。
Bug 修复(2026-09-21):
- Bug1: 60fps降帧逻辑 — 输入视频 >30fps 时先降帧至 25fps 推理,直接输出 25fps 结果
(MCI 运动补偿插帧已移除,CPU密集且口型场景 25fps 足够)
- Bug2: 超时终止机制 — 推理线程改为 daemon + abort_event 机制,超时时 set event
让推理循环检测退出,同时 kill 所有活跃 ffmpeg 子进程,等线程退出后再释放锁和清理目录
- Bug3: /health 接口增加 gfpgan_loaded 和 gfpgan_load_error 字段
- Bonus: 动态超时计算(基础120s + 帧数*0.15s,上限1800s)
- Bonus: GFPGAN 增强异常时 log warning 而非静默跳过
#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"))
# GFPGAN 人脸超分增强(提升生成人脸清晰度,+~170MB VRAM, +60ms/帧)
use_gfpgan: bool = _env("MUSE_USE_GFPGAN", "1") == "1"
# ── 全局状态 ──────────────────────────────────────────────────────────
inference_lock = threading.Lock()
current_task: dict = {"task_id": None, "process": None, "start_time": 0.0}
shutdown_event = threading.Event()
# 当前推理线程的终止信号和线程引用
current_abort_event: Optional[threading.Event] = None
current_thread: Optional[threading.Thread] = None
# 全局追踪正在运行的 ffmpeg 子进程(用于超时终止)
_active_ffmpeg_procs: list[subprocess.Popen] = []
_active_ffmpeg_lock = threading.Lock()
# GFPGAN 加载状态(供 /health 接口查询)
_gfpgan_loaded = False
_gfpgan_load_error: Optional[str] = None
# ── 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 命令,检查返回码和超时。进程会被注册到全局列表以便外部终止."""
proc = None
try:
proc = subprocess.Popen(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
)
with _active_ffmpeg_lock:
_active_ffmpeg_procs.append(proc)
try:
stdout, stderr = proc.communicate(timeout=timeout)
except subprocess.TimeoutExpired:
proc.kill()
proc.wait(timeout=5)
raise RuntimeError(f"ffmpeg 超时(>{timeout}s)")
if proc.returncode != 0:
stderr_text = stderr.decode(errors="ignore") if stderr else ""
raise RuntimeError(f"ffmpeg 失败 (code={proc.returncode}): {stderr_text[:500]}")
return subprocess.CompletedProcess(cmd, proc.returncode, stdout, stderr)
finally:
if proc is not None:
with _active_ffmpeg_lock:
try:
_active_ffmpeg_procs.remove(proc)
except ValueError:
pass
# ── 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, _gfpgan_loaded, _gfpgan_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()
# 加载 GFPGAN 人脸超分模型(FP16,仅 ~170MB VRAM)
gfpgan_model = None
if Config.use_gfpgan:
try:
from gfpgan.archs.gfpganv1_clean_arch import GFPGANv1Clean
gfpgan_path = muse_dir / "models" / "GFPGAN" / "GFPGANv1.4.pth"
if gfpgan_path.exists():
logger.info("加载 GFPGANv1.4 人脸超分模型: %s", gfpgan_path)
gfpgan_ckpt = torch.load(str(gfpgan_path), map_location="cpu")
gfpgan_model = GFPGANv1Clean(
out_size=512, num_style_feat=512, channel_multiplier=2,
decoder_load_path=None, fix_decoder=False, num_mlp=8,
input_is_latent=True, different_w=True, narrow=1, sft_half=True,
)
gfpgan_key = "params_ema" if "params_ema" in gfpgan_ckpt else "params"
gfpgan_model.load_state_dict(gfpgan_ckpt[gfpgan_key], strict=True)
gfpgan_model.eval()
# GFPGAN 始终使用 FP32 推理,避免 FP16 色偏导致紫/灰色块
gfpgan_model = gfpgan_model.to(device)
del gfpgan_ckpt
_gfpgan_loaded = True
_gfpgan_load_error = None
logger.info("GFPGAN 加载完成 (FP32,避免色偏)")
else:
_gfpgan_loaded = False
_gfpgan_load_error = f"模型文件不存在: {gfpgan_path}"
logger.warning("GFPGAN 模型不存在: %s,跳过人脸增强", gfpgan_path)
except Exception as e:
_gfpgan_loaded = False
_gfpgan_load_error = str(e)
logger.warning("GFPGAN 加载失败,跳过人脸增强: %s", e)
gfpgan_model = None
else:
_gfpgan_loaded = False
_gfpgan_load_error = "已通过环境变量禁用 (MUSE_USE_GFPGAN=0)"
logger.info("GFPGAN 已禁用 (MUSE_USE_GFPGAN=0)")
_muse_models = {
"vae": vae,
"unet": unet,
"pe": pe,
"timesteps": timesteps,
"audio_processor": audio_processor,
"face_parsing": face_parsing,
"gfpgan": gfpgan_model,
"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,
abort_event: threading.Event = None,
) -> 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)
original_fps = fps
audio_duration = _get_media_duration(audio_path)
video_duration = _get_media_duration(video_path)
logger.info(
"MuseTalk 推理开始: video=%.2fs, audio=%.2fs, input_fps=%.1f, bbox_shift=%d",
video_duration, audio_duration, fps, bbox_shift,
)
# ── Step 0: 高帧率视频降帧(>30fps → 25fps 推理,推理后再插帧回原帧率)──
inference_fps = fps
video_downsampled = False
if fps > 30:
inference_fps = 25.0
downsampled_video = video_path.parent / "input_25fps.mp4"
logger.info("检测到高帧率视频 %.1f fps,降帧至 %.1f fps 进行推理", fps, inference_fps)
try:
_run_ffmpeg([
"ffmpeg", "-y", "-v", "warning",
"-i", str(video_path),
"-r", str(inference_fps),
"-c:v", "libx264", "-preset", "veryfast",
"-crf", "18", "-pix_fmt", "yuv420p",
str(downsampled_video),
], timeout=120)
video_path = downsampled_video
video_downsampled = True
# 更新帧数和时长
total_input_frames_approx = int(audio_duration * inference_fps)
logger.info("降帧完成: input_fps=%.1f → inference_fps=%.1f", original_fps, inference_fps)
except Exception as e:
logger.warning("视频降帧失败,使用原始帧率 %.1f fps 推理: %s", fps, e)
inference_fps = fps
# ── 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 * inference_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=inference_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
# 直接传 BGR 给 VAE:VAE 内部 preprocess_img 已做 BGR→RGB 转换,
# 此处再 cvtColor 会造成双重转换、R/B 通道互换(蓝色块根因)。
crop_resized = cv2.resize(crop, (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, 8) # RTX3060 12G 显存,FP16+GFPGAN batch=8 约用 7-8GB,留足余量
total_batches = (video_num + bs - 1) // bs
for bi in tqdm(range(total_batches), desc="MuseTalk 推理"):
# 检查是否收到终止信号
if abort_event is not None and abort_event.is_set():
logger.warning("收到终止信号,中止推理 (batch %d/%d)", bi, total_batches)
raise InterruptedError("推理被外部终止")
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: 合成最终帧 → ffmpeg pipe 编码(零磁盘IO) ──
gfpgan_enhancer = models.get("gfpgan")
silent_video_path = video_path.parent / "silent_output.mp4"
frame_h, frame_w = frame_list_cycle[0].shape[:2]
# 启动 ffmpeg:stdin 接收 raw BGR24 帧,直接编码 H.264(省去PNG落盘+回读)
_ff_cmd = [
"ffmpeg", "-y", "-v", "warning",
"-f", "rawvideo", "-pix_fmt", "bgr24",
"-s", f"{frame_w}x{frame_h}", "-r", str(inference_fps),
"-i", "-",
"-vcodec", "libx264", "-preset", "veryfast",
"-vf", "format=yuv420p", "-crf", "18",
str(silent_video_path),
]
import subprocess as _sp
_ff_proc = _sp.Popen(_ff_cmd, stdin=_sp.PIPE, stdout=_sp.DEVNULL, stderr=_sp.PIPE)
# 注册 ffmpeg 进程到全局列表,以便超时终止
with _active_ffmpeg_lock:
_active_ffmpeg_procs.append(_ff_proc)
n_out = min(len(res_frame_list), num_output_frames)
try:
for i in tqdm(range(n_out), desc="合成帧"):
# 检查终止信号
if abort_event is not None and abort_event.is_set():
logger.warning("收到终止信号,中止合成 (frame %d/%d)", i, n_out)
raise InterruptedError("推理被外部终止")
cyc_i = i % len(coord_cycle)
x1, y1, x2, y2 = coord_cycle[cyc_i]
ori_frame = copy.deepcopy(frame_list_cycle[cyc_i])
res_frame = res_frame_list[i]
try:
res_frame_resized = cv2.resize(
res_frame.astype(np.uint8), (x2-x1, y2-y1),
interpolation=cv2.INTER_LANCZOS4
)
except Exception:
_ff_proc.stdin.write(ori_frame.tobytes())
continue
# GFPGAN 人脸超分增强(FP32 推理,避免 FP16 色偏)
# 色彩通道约定:ori_frame / res_frame / _face_up 均为 BGR(OpenCV 默认);
# GFPGAN 输出用 return_rgb=True 拿到 RGB,再转 BGR,与后续 face_parsing 融合保持一致。
if gfpgan_enhancer is not None:
try:
_fh, _fw = res_frame_resized.shape[:2]
_face_up = cv2.resize(res_frame_resized, (512, 512),
interpolation=cv2.INTER_LANCZOS4)
_face_rgb = cv2.cvtColor(_face_up, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
_face_t = torch.from_numpy(_face_rgb.transpose(2, 0, 1)).unsqueeze(0)
# GFPGAN 始终 FP32,避免 FP16 精度导致色偏;归一化到 [-1, 1]
_face_t = ((_face_t - 0.5) / 0.5).to(device)
with torch.no_grad():
_out = gfpgan_enhancer(_face_t, return_rgb=True, weight=0.35)[0]
# 输出 tensor: RGB, [-1, 1] 范围 → clamp → 映射到 [0, 255] uint8
_out = _out.squeeze(0).float().cpu().clamp_(-1.0, 1.0)
_out = ((_out + 1.0) / 2.0 * 255.0).numpy().transpose(1, 2, 0)
_out_rgb = _out.astype(np.uint8)
# RGB → BGR,与 ori_frame 保持一致,确保 face_parsing 融合时通道正确
_out_bgr = cv2.cvtColor(_out_rgb, cv2.COLOR_RGB2BGR)
res_frame_resized = cv2.resize(_out_bgr, (_fw, _fh),
interpolation=cv2.INTER_LANCZOS4)
del _face_t, _out, _out_rgb, _out_bgr
except Exception as _gfpgan_err:
logger.warning("GFPGAN 增强失败(帧 %d),使用原图: %s", i, _gfpgan_err)
# face parsing 融合
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:
combined = ori_frame.copy()
try: combined[y1:y2, x1:x2] = res_frame_resized
except Exception: combined = ori_frame
_ff_proc.stdin.write(combined.tobytes())
_ff_proc.stdin.close()
_ff_ret = _ff_proc.wait(timeout=120)
# 从全局列表中移除已完成的 ffmpeg 进程
with _active_ffmpeg_lock:
try:
_active_ffmpeg_procs.remove(_ff_proc)
except ValueError:
pass
if _ff_ret != 0:
_ff_err = _ff_proc.stderr.read().decode(errors="ignore") if _ff_proc.stderr else ""
raise RuntimeError(f"ffmpeg编码失败(exit={_ff_ret}): {_ff_err[-300:]}")
except Exception:
try: _ff_proc.kill()
except Exception: pass
with _active_ffmpeg_lock:
try:
_active_ffmpeg_procs.remove(_ff_proc)
except ValueError:
pass
raise
# ── Step 8: 高帧率视频处理 ──
# 对口型数字人视频,25fps 完全够用。
# MCI 运动补偿插帧极其耗时(343帧>6分钟),已移除。
# 直接输出 25fps 推理结果,后续封装音频后播放器自动适配帧率。
final_video = silent_video_path
if video_downsampled:
logger.info("输入视频 %.1f fps,降帧至 %.1f fps 推理后直接输出(不做插帧还原)", original_fps, inference_fps)
_mux_video_with_audio(final_video, audio_path, output_path)
try:
torch.cuda.empty_cache()
if audio_wav_path.exists():
audio_wav_path.unlink()
# 清理中间文件
for _tmp in [silent_video_path] + ([video_path] if video_downsampled else []):
if _tmp.exists() and str(_tmp) != str(output_path):
_tmp.unlink()
except Exception as e:
logger.warning("清理中间文件失败: %s", e)
logger.info("MuseTalk 推理完成: output=%s, duration=%.2fs, input_fps=%.1f, inference_fps=%.1f",
output_path.name, _get_media_duration(output_path), original_fps, inference_fps)
# ── 路由 ──────────────────────────────────────────────────────────────
@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,
"gfpgan_loaded": _gfpgan_loaded,
"gfpgan_load_error": _gfpgan_load_error,
"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}
abort_event = threading.Event()
global current_abort_event, current_thread
current_abort_event = abort_event
def inference_thread():
try:
_run_inference(video_path, audio_path, output_path, bbox_shift=bbox_shift, abort_event=abort_event)
except Exception as exc:
logger.exception("推理异常: %s", exc)
result_container["error"] = str(exc)
thread = threading.Thread(target=inference_thread, daemon=True)
current_thread = thread
thread.start()
# 动态超时:基础 120s + 预估帧数×0.15s,上限 1800s
video_fps = _get_video_fps(video_path)
video_duration = _get_media_duration(video_path)
estimated_frames = int(video_duration * video_fps) if video_duration > 0 else 500
dynamic_timeout = min(max(Config.inference_timeout, 120 + estimated_frames * 0.15), 1800)
thread.join(timeout=dynamic_timeout)
if thread.is_alive():
# 超时处理:发终止信号 + 杀 ffmpeg 子进程 + 等待线程退出
logger.error("推理超时 (>%ds),终止任务 %s", dynamic_timeout, task_id)
abort_event.set()
# 终止所有正在运行的 ffmpeg 子进程
with _active_ffmpeg_lock:
for proc in _active_ffmpeg_procs[:]:
try:
proc.kill()
proc.wait(timeout=5)
except Exception:
pass
_active_ffmpeg_procs.clear()
# 等待线程退出(daemon 线程会在主进程退出时自动终止,但这里给一定时间让它清理)
thread.join(timeout=10)
# 标记超时错误,由 finally 统一清理
result_container["error"] = f"推理超时(>{dynamic_timeout:.0f}s)"
result_container["timeout"] = True
if result_container.get("timeout"):
return jsonify({"error": result_container["error"], "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
current_abort_event = None
current_thread = None
# 清理临时文件
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():
"""终止当前正在进行的推理任务."""
global current_abort_event, current_thread
if current_task["task_id"] is None:
return jsonify({"message": "当前无正在运行的任务"})
task_id = current_task["task_id"]
logger.info("收到取消请求,终止任务 %s", task_id)
# 发送终止信号给推理线程
if current_abort_event is not None:
current_abort_event.set()
logger.info("已发送终止信号给推理线程")
# 终止所有正在运行的 ffmpeg 子进程
with _active_ffmpeg_lock:
for proc in _active_ffmpeg_procs[:]:
try:
proc.kill()
proc.wait(timeout=5)
logger.info("已终止 ffmpeg 子进程")
except Exception as exc:
logger.warning("终止 ffmpeg 子进程失败: %s", exc)
_active_ffmpeg_procs.clear()
# 等待推理线程退出(daemon 线程会在主进程退出时自动终止)
if current_thread is not None and current_thread.is_alive():
current_thread.join(timeout=10)
if current_thread.is_alive():
logger.warning("推理线程未能在 10s 内退出,将作为 daemon 线程自动终止")
# 清理临时文件
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
current_abort_event = None
current_thread = None
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()