dd7a5ad658
P0: 集成GFPGAN面部增强后处理 - 懒加载GFPGAN模型,首次推理时自动加载并复用 - 对嘴部/面部区域做增强,提升锐度和纹理细节 - 支持mouth/face两种增强区域模式 - 高斯模糊遮罩平滑融合,消除接缝 P1: 60fps视频帧率适配 - 检测输入视频帧率,>30fps时自动降至25fps推理 - 推理完成后使用minterpolate运动补偿插帧恢复原始帧率 - 消除MuseTalk训练帧率不匹配导致的重复帧问题 P2: 高质量音频重采样 - 使用ffmpeg soxr重采样器(精度28)替代默认线性插值 - 不可用时自动回退默认重采样器 P1: 融合参数优化 - v1.5模型扩大bbox融合区域(下巴+15px,两侧各+2px) - 减少AI生成区域与原画面的接缝感 /inference接口新增face_enhance参数,支持请求级开关 /health接口新增GFPGAN状态和face_enhance_enabled信息
1255 lines
47 KiB
Python
1255 lines
47 KiB
Python
"""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
|
||
MUSE_FACE_ENHANCE 启用 GFPGAN 面部增强后处理,默认 1(开启)
|
||
MUSE_ENHANCE_REGION 增强区域:mouth(仅嘴部,快)或 face(全脸,更自然),默认 mouth
|
||
MUSE_ENHANCE_UPSCALE 面部增强上采样倍数,默认 1(不放大)
|
||
|
||
接口:
|
||
GET /health 健康检查 + GPU 显存信息 + MuseTalk/GFPGAN 模型状态
|
||
POST /inference 推理请求(multipart: video + audio, form: bbox_shift, face_enhance)
|
||
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):启用开关、区域、上采样倍数
|
||
face_enhance: bool = _env("MUSE_FACE_ENHANCE", "1") == "1"
|
||
enhance_region: str = _env("MUSE_ENHANCE_REGION", "mouth") or "mouth"
|
||
enhance_upscale: int = max(1, min(4, int(_env("MUSE_ENHANCE_UPSCALE", "1"))))
|
||
# GFPGAN 模型权重路径(相对 MUSE_DIR 或绝对路径)
|
||
gfpgan_model_path: str = _env(
|
||
"MUSE_GFPGAN_MODEL", "models/GFPGAN/GFPGANv1.4.pth"
|
||
)
|
||
|
||
|
||
# ── 全局状态 ──────────────────────────────────────────────────────────
|
||
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
|
||
|
||
# ── GFPGAN 面部增强模型懒加载 ───────────────────────────────────────
|
||
_gfpgan = None
|
||
_gfpgan_lock = threading.Lock()
|
||
_gfpgan_loaded = False
|
||
_gfpgan_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 频谱错位、
|
||
音素特征提取错误,口型只跟能量不跟音素。
|
||
|
||
使用 soxr 高质量重采样器(如果可用),避免线性插值带来的频谱失真。
|
||
"""
|
||
# 尝试使用 soxr 高质量重采样;不可用时回退默认 resampler
|
||
cmd = [
|
||
"ffmpeg", "-y", "-v", "warning",
|
||
"-i", str(input_path),
|
||
"-af", "aresample=resampler=soxr:precision=28", # 高质量重采样
|
||
"-ar", str(target_sr), # 重采样到 16kHz
|
||
"-ac", "1", # 单声道
|
||
"-sample_fmt", "s16", # 16bit PCM
|
||
str(output_path),
|
||
]
|
||
try:
|
||
_run_ffmpeg(cmd, timeout=60)
|
||
except RuntimeError:
|
||
# soxr 不可用时回退到默认重采样器
|
||
logger.warning("soxr 重采样器不可用,回退默认重采样器")
|
||
cmd_fallback = [
|
||
"ffmpeg", "-y", "-v", "warning",
|
||
"-i", str(input_path),
|
||
"-ar", str(target_sr),
|
||
"-ac", "1",
|
||
"-sample_fmt", "s16",
|
||
str(output_path),
|
||
]
|
||
_run_ffmpeg(cmd_fallback, 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
|
||
|
||
|
||
# ── GFPGAN 面部增强 ───────────────────────────────────────────────────
|
||
|
||
|
||
def _load_gfpgan():
|
||
"""懒加载 GFPGAN 面部增强模型(全局单例,首次调用时加载).
|
||
|
||
GFPGAN 用于提升 MuseTalk 输出的面部清晰度和纹理细节,
|
||
特别改善嘴部区域模糊问题。
|
||
优先从 MuseTalk 自带的 models/GFPGAN 目录加载权重,
|
||
也支持通过 MUSE_GFPGAN_MODEL 环境变量指定绝对路径。
|
||
"""
|
||
global _gfpgan, _gfpgan_loaded, _gfpgan_load_error
|
||
|
||
if _gfpgan_loaded:
|
||
return _gfpgan
|
||
if _gfpgan_load_error is not None:
|
||
raise _gfpgan_load_error
|
||
|
||
with _gfpgan_lock:
|
||
if _gfpgan_loaded:
|
||
return _gfpgan
|
||
|
||
try:
|
||
import torch
|
||
|
||
# 解析模型权重路径
|
||
model_path = Config.gfpgan_model_path
|
||
if not os.path.isabs(model_path):
|
||
model_path = os.path.join(Config.muse_dir, model_path)
|
||
|
||
if not os.path.isfile(model_path):
|
||
raise FileNotFoundError(
|
||
f"GFPGAN 模型权重不存在: {model_path}\n"
|
||
f"请下载到 MuseTalk 仓库的 models/GFPGAN/ 目录下,\n"
|
||
f"或通过 MUSE_GFPGAN_MODEL 环境变量指定路径。\n"
|
||
f"下载地址: https://github.com/TencentARC/GFPGAN/releases"
|
||
)
|
||
|
||
# 尝试多种导入方式
|
||
gfpgan_model = None
|
||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||
|
||
# 方式 1: 从 gfpgan 官方包加载
|
||
try:
|
||
from gfpgan import GFPGANer
|
||
gfpgan_model = GFPGANer(
|
||
model_path=model_path,
|
||
upscale=Config.enhance_upscale,
|
||
arch="clean",
|
||
channel_multiplier=2,
|
||
device=device,
|
||
)
|
||
logger.info("GFPGAN 加载成功(gfpgan 官方包)")
|
||
except ImportError:
|
||
pass
|
||
|
||
# 方式 2: 从 MuseTalk 自带的 facelib/gfpgan 加载
|
||
if gfpgan_model is None:
|
||
muse_str = str(Path(Config.muse_dir))
|
||
if muse_str not in sys.path:
|
||
sys.path.insert(0, muse_str)
|
||
try:
|
||
from gfpgan import GFPGANer
|
||
gfpgan_model = GFPGANer(
|
||
model_path=model_path,
|
||
upscale=Config.enhance_upscale,
|
||
arch="clean",
|
||
channel_multiplier=2,
|
||
device=device,
|
||
)
|
||
logger.info("GFPGAN 加载成功(MuseTalk 环境)")
|
||
except ImportError:
|
||
pass
|
||
|
||
if gfpgan_model is None:
|
||
raise ImportError(
|
||
"无法导入 GFPGAN。请先安装: pip install gfpgan\n"
|
||
"并确保模型权重文件存在于指定路径。"
|
||
)
|
||
|
||
_gfpgan = {
|
||
"model": gfpgan_model,
|
||
"device": device,
|
||
}
|
||
_gfpgan_loaded = True
|
||
logger.info("GFPGAN 面部增强模型加载完成 (device=%s)", device)
|
||
|
||
return _gfpgan
|
||
|
||
except Exception as exc:
|
||
_gfpgan_load_error = exc
|
||
logger.error("GFPGAN 模型加载失败: %s", exc)
|
||
raise
|
||
|
||
|
||
def _enhance_face(frame, bbox, enhance_region="mouth", padding_ratio=0.3):
|
||
"""对帧中的人脸/嘴部区域应用 GFPGAN 增强.
|
||
|
||
Args:
|
||
frame: BGR numpy 数组
|
||
bbox: [x1, y1, x2, y2] 嘴部/面部区域
|
||
enhance_region: 'mouth' 仅增强嘴部区域(更快), 'face' 增强全脸(更自然)
|
||
padding_ratio: 区域扩展比例,用于提供足够的上下文给 GFPGAN
|
||
|
||
Returns:
|
||
增强后的帧 (BGR numpy 数组)
|
||
"""
|
||
import cv2
|
||
import numpy as np
|
||
|
||
models = _load_gfpgan()
|
||
gfpgan_model = models["model"]
|
||
|
||
x1, y1, x2, y2 = bbox
|
||
h, w = frame.shape[:2]
|
||
|
||
if enhance_region == "mouth":
|
||
# 嘴部区域:扩展 bbox 提供足够上下文
|
||
pad_x = int((x2 - x1) * padding_ratio)
|
||
pad_y = int((y2 - y1) * padding_ratio * 1.5) # 嘴部上下多留一些空间
|
||
crop_x1 = max(0, x1 - pad_x)
|
||
crop_y1 = max(0, y1 - pad_y)
|
||
crop_x2 = min(w, x2 + pad_x)
|
||
crop_y2 = min(h, y2 + pad_y)
|
||
else:
|
||
# 全脸区域:使用更大的 padding
|
||
face_w = x2 - x1
|
||
face_h = y2 - y1
|
||
pad_x = int(face_w * 0.5)
|
||
pad_y = int(face_h * 0.6)
|
||
crop_x1 = max(0, x1 - pad_x)
|
||
crop_y1 = max(0, y1 - pad_y)
|
||
crop_x2 = min(w, x2 + pad_x)
|
||
crop_y2 = min(h, y2 + pad_y)
|
||
|
||
# 裁剪区域
|
||
crop = frame[crop_y1:crop_y2, crop_x1:crop_x2].copy()
|
||
if crop.size == 0:
|
||
return frame
|
||
|
||
# 确保最小尺寸(GFPGAN 需要至少 64x64)
|
||
crop_h, crop_w = crop.shape[:2]
|
||
if crop_h < 64 or crop_w < 64:
|
||
return frame
|
||
|
||
try:
|
||
# GFPGAN 增强
|
||
_, _, enhanced_crop = gfpgan_model.enhance(
|
||
crop,
|
||
has_aligned=False,
|
||
only_center_face=False,
|
||
paste_back=True,
|
||
)
|
||
|
||
# 确保尺寸一致(upscale 可能导致尺寸变化)
|
||
if enhanced_crop.shape[:2] != crop.shape[:2]:
|
||
enhanced_crop = cv2.resize(enhanced_crop, (crop_w, crop_h))
|
||
|
||
# 创建高斯模糊遮罩用于平滑融合
|
||
mask_h, mask_w = enhanced_crop.shape[:2]
|
||
# 使用较大的模糊核以获得更柔和的过渡
|
||
ksize = max(mask_h, mask_w) // 3
|
||
if ksize % 2 == 0:
|
||
ksize += 1
|
||
ksize = max(ksize, 11) # 最小核大小
|
||
mask = np.zeros((mask_h, mask_w), dtype=np.float32)
|
||
# 中心区域为 1
|
||
margin_y = int(mask_h * 0.15)
|
||
margin_x = int(mask_w * 0.15)
|
||
mask[margin_y:mask_h-margin_y, margin_x:mask_w-margin_x] = 1.0
|
||
mask = cv2.GaussianBlur(mask, (ksize, ksize), 0)
|
||
|
||
# 融合:增强区域与原图平滑过渡
|
||
mask_3ch = mask[:, :, np.newaxis]
|
||
blended_crop = (enhanced_crop.astype(np.float32) * mask_3ch +
|
||
crop.astype(np.float32) * (1 - mask_3ch))
|
||
blended_crop = np.clip(blended_crop, 0, 255).astype(np.uint8)
|
||
|
||
# 放回原帧
|
||
frame[crop_y1:crop_y2, crop_x1:crop_x2] = blended_crop
|
||
|
||
except Exception as e:
|
||
logger.warning("GFPGAN 面部增强失败,使用原帧: %s", e)
|
||
|
||
return frame
|
||
|
||
|
||
# ── 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,
|
||
face_enhance: bool = None,
|
||
) -> None:
|
||
"""执行 MuseTalk 真实推理.
|
||
|
||
流程:
|
||
1. 音频预处理:任意格式 → 16kHz mono 16bit WAV(高质量 soxr 重采样)
|
||
2. 帧率适配:高帧率视频(>30fps)降至 25fps 推理,输出时恢复原始帧率
|
||
3. 加载/复用 MuseTalk 模型(VAE + UNet + PE + Whisper)
|
||
4. 视频预处理:提取帧 → 人脸检测 → 获取 bbox → VAE 编码 latent
|
||
5. 音频特征提取:whisper 提取 audio features (50×384 per chunk)
|
||
6. 批量推理:UNet 去噪 → VAE 解码 → 得到口型同步的人脸帧
|
||
7. 帧合成:将生成的人脸贴回原帧(使用 face parsing + 高斯模糊边缘融合)
|
||
8. GFPGAN 面部增强后处理(可选):提升嘴部/面部区域的清晰度和纹理
|
||
9. 输出无声视频(后续由 _mux_video_with_audio 封装 TTS 音频)
|
||
|
||
Args:
|
||
video_path: 输入视频路径
|
||
audio_path: 输入音频路径(任意格式,会被预处理为 16kHz WAV)
|
||
output_path: 输出无声视频路径
|
||
bbox_shift: 口型区域垂直偏移量,默认 0,范围 [-5, 5]
|
||
face_enhance: 是否启用 GFPGAN 面部增强;None 时使用 Config.face_enhance 配置
|
||
"""
|
||
import cv2
|
||
import numpy as np
|
||
import torch
|
||
from tqdm import tqdm
|
||
|
||
from musetalk.utils.preprocessing import get_landmark_and_bbox, read_imgs
|
||
from musetalk.utils.utils import datagen
|
||
from musetalk.utils.blending import get_image
|
||
|
||
# 是否启用面部增强
|
||
if face_enhance is None:
|
||
face_enhance = Config.face_enhance
|
||
|
||
# 加载模型(首次调用时加载,后续复用)
|
||
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"]
|
||
|
||
fps = _get_video_fps(video_path)
|
||
audio_duration = _get_media_duration(audio_path)
|
||
video_duration = _get_media_duration(video_path)
|
||
|
||
# ── 帧率适配:MuseTalk 训练在 25fps,高帧率输入需降采样 ──
|
||
original_fps = fps
|
||
inference_fps = 25.0 # MuseTalk 标准帧率
|
||
need_fps_conversion = fps > 30 # 超过 30fps 的输入需要处理
|
||
|
||
if need_fps_conversion:
|
||
logger.info(
|
||
"检测到高帧率输入 (%.1f fps),将先降至 %.1f fps 推理,推理后恢复至原始帧率",
|
||
fps, inference_fps,
|
||
)
|
||
# 使用 ffmpeg 将高帧率视频降至 25fps
|
||
downsampled_video = video_path.parent / "video_25fps.mp4"
|
||
_run_ffmpeg(
|
||
[
|
||
"ffmpeg", "-y", "-v", "warning",
|
||
"-i", str(video_path),
|
||
"-r", str(int(inference_fps)),
|
||
"-vsync", "cfr",
|
||
"-c:v", "libx264",
|
||
"-preset", "fast",
|
||
"-crf", "18",
|
||
"-an", # 去除音频,后续单独处理
|
||
str(downsampled_video),
|
||
],
|
||
timeout=300,
|
||
)
|
||
work_video_path = downsampled_video
|
||
fps = inference_fps
|
||
logger.info("视频降采样至 25fps 完成")
|
||
else:
|
||
work_video_path = video_path
|
||
|
||
logger.info(
|
||
"MuseTalk 推理开始: video=%.2fs, audio=%.2fs, original_fps=%.1f, "
|
||
"inference_fps=%.1f, bbox_shift=%d, face_enhance=%s",
|
||
video_duration, audio_duration, original_fps, fps, bbox_shift, face_enhance,
|
||
)
|
||
|
||
# ── 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(work_video_path))
|
||
total_frames = len(input_frames)
|
||
if total_frames == 0:
|
||
raise RuntimeError("未能从视频中提取到任何帧")
|
||
logger.info("提取到 %d 帧视频画面 (工作帧率 %.1f fps)", total_frames, fps)
|
||
|
||
# ── Step 3: 人脸检测 & bbox 计算 ──
|
||
coord_list, coord_placeholder = get_landmark_and_bbox(
|
||
input_frames, vid_pts=0, bbox_shift=bbox_shift
|
||
)
|
||
valid_bboxes = sum(1 for c in coord_list if c is not coord_placeholder)
|
||
logger.info("人脸检测完成,有效 bbox: %d/%d", valid_bboxes, total_frames)
|
||
|
||
# 使用 mirror indexing 循环帧和坐标(避免硬切跳变)
|
||
num_output_frames = int(audio_duration * fps)
|
||
if num_output_frames <= 0:
|
||
num_output_frames = total_frames
|
||
|
||
# ── Step 4: 音频特征提取 ──
|
||
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,
|
||
)
|
||
logger.info("音频特征提取完成: %d 个 chunk", len(whisper_features))
|
||
|
||
# ── Step 5: 视频帧 VAE 编码为 latent ──
|
||
input_latent_list = []
|
||
with torch.no_grad():
|
||
for frame in input_frames:
|
||
frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||
frame_resized = cv2.resize(frame_rgb, (256, 256))
|
||
frame_tensor = torch.from_numpy(frame_resized).float() / 127.5 - 1.0
|
||
frame_tensor = frame_tensor.permute(2, 0, 1).unsqueeze(0).to(device)
|
||
if Config.use_float16:
|
||
frame_tensor = frame_tensor.half()
|
||
latent = vae.encode_latents(frame_tensor)
|
||
input_latent_list.append(latent)
|
||
logger.info("视频帧 VAE 编码完成: %d 个 latent", len(input_latent_list))
|
||
|
||
# ── Step 6: 批量推理 ──
|
||
result_frames = []
|
||
total_batches = (num_output_frames + Config.batch_size - 1) // Config.batch_size
|
||
|
||
for batch_idx in tqdm(range(total_batches), desc="MuseTalk 推理"):
|
||
# 构建 whisper batch
|
||
whisper_batch = whisper_features[
|
||
batch_idx * Config.batch_size : (batch_idx + 1) * Config.batch_size
|
||
]
|
||
if len(whisper_batch) == 0:
|
||
break
|
||
|
||
# 构建 latent batch(使用 mirror indexing 循环)
|
||
latent_indices = []
|
||
for i in range(len(whisper_batch)):
|
||
global_idx = batch_idx * Config.batch_size + i
|
||
frame_idx = _mirror_index(total_frames, global_idx)
|
||
latent_indices.append(frame_idx)
|
||
|
||
latent_batch = torch.cat(
|
||
[input_latent_list[idx] for idx in latent_indices], dim=0
|
||
)
|
||
latent_batch = latent_batch.to(device)
|
||
if Config.use_float16:
|
||
latent_batch = latent_batch.half()
|
||
|
||
# 位置编码
|
||
audio_feature_batch = pe(whisper_batch)
|
||
|
||
# UNet 推理
|
||
with torch.no_grad():
|
||
latent_batch = latent_batch.to(dtype=unet.model.dtype)
|
||
pred_latents = unet.model(
|
||
latent_batch,
|
||
timesteps,
|
||
encoder_hidden_states=audio_feature_batch,
|
||
).sample
|
||
|
||
# VAE 解码
|
||
recon_frames = vae.decode_latents(pred_latents)
|
||
result_frames.extend(recon_frames)
|
||
|
||
logger.info("推理完成,生成 %d 帧", len(result_frames))
|
||
|
||
# ── Step 7: 合成最终帧并写视频 ──
|
||
output_frames_dir = video_path.parent / "output_frames"
|
||
output_frames_dir.mkdir(parents=True, exist_ok=True)
|
||
|
||
# 预加载 GFPGAN(如果启用面部增强)
|
||
if face_enhance:
|
||
try:
|
||
_load_gfpgan()
|
||
logger.info("GFPGAN 面部增强已启用,区域: %s", Config.enhance_region)
|
||
except Exception as e:
|
||
logger.warning("GFPGAN 加载失败,跳过面部增强: %s", e)
|
||
face_enhance = False
|
||
|
||
for i, res_frame in enumerate(tqdm(result_frames[:num_output_frames], desc="合成帧")):
|
||
orig_idx = _mirror_index(total_frames, i)
|
||
ori_frame = copy.deepcopy(input_frames[orig_idx])
|
||
bbox = coord_list[orig_idx] if orig_idx < len(coord_list) else coord_placeholder
|
||
|
||
if bbox is coord_placeholder:
|
||
# 没有检测到人脸的帧,保持原样
|
||
combined = ori_frame
|
||
else:
|
||
x1, y1, x2, y2 = bbox
|
||
if model_version == "v15":
|
||
# v1.5 额外扩展下边界(下巴区域),扩大融合范围
|
||
y2 = min(y2 + 15, ori_frame.shape[0])
|
||
x1 = max(0, x1 - 2)
|
||
x2 = min(ori_frame.shape[1], x2 + 2)
|
||
|
||
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 做边缘融合
|
||
combined = get_image(
|
||
ori_frame,
|
||
res_frame_resized,
|
||
[x1, y1, x2, y2],
|
||
)
|
||
|
||
# ── P0: GFPGAN 面部增强后处理 ──
|
||
if face_enhance:
|
||
try:
|
||
combined = _enhance_face(
|
||
combined,
|
||
[x1, y1, x2, y2],
|
||
enhance_region=Config.enhance_region,
|
||
)
|
||
except Exception as e:
|
||
logger.warning("第 %d 帧面部增强失败,使用未增强结果: %s", i, e)
|
||
|
||
cv2.imwrite(str(output_frames_dir / f"{i:08d}.png"), combined)
|
||
|
||
# ── Step 8: 帧序列 → 无声视频 ──
|
||
silent_video_path = video_path.parent / "silent_output.mp4"
|
||
|
||
if need_fps_conversion:
|
||
# 需要恢复原始帧率:使用 minterpolate 做运动补偿插帧
|
||
logger.info("将推理结果从 25fps 恢复至 %.1f fps(运动补偿插帧)", original_fps)
|
||
_run_ffmpeg(
|
||
[
|
||
"ffmpeg", "-y", "-v", "warning",
|
||
"-r", str(int(inference_fps)),
|
||
"-f", "image2",
|
||
"-i", str(output_frames_dir / "%08d.png"),
|
||
"-vf", f"minterpolate=fps={int(original_fps)}:mi_mode=mci:mc_mode=aobmc:me_mode=bidir:vsbmc=1,format=yuv420p",
|
||
"-c:v", "libx264",
|
||
"-preset", "medium",
|
||
"-crf", "18",
|
||
str(silent_video_path),
|
||
],
|
||
timeout=600,
|
||
)
|
||
else:
|
||
_run_ffmpeg(
|
||
[
|
||
"ffmpeg", "-y", "-v", "warning",
|
||
"-r", str(int(original_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()
|
||
if need_fps_conversion:
|
||
downsampled_video = video_path.parent / "video_25fps.mp4"
|
||
if downsampled_video.exists():
|
||
downsampled_video.unlink()
|
||
except Exception as e:
|
||
logger.warning("清理中间文件失败: %s", e)
|
||
|
||
logger.info(
|
||
"MuseTalk 推理完成: output=%s, duration=%.2fs, fps=%.1f, face_enhance=%s",
|
||
output_path.name, _get_media_duration(output_path), original_fps, face_enhance,
|
||
)
|
||
|
||
|
||
# ── 路由 ──────────────────────────────────────────────────────────────
|
||
|
||
|
||
@app.route("/health", methods=["GET"])
|
||
def health():
|
||
"""健康检查 + GPU 显存信息 + MuseTalk/GFPGAN 模型状态."""
|
||
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": str(_gfpgan_load_error) if _gfpgan_load_error else None,
|
||
"face_enhance_enabled": Config.face_enhance,
|
||
"timestamp": time.time(),
|
||
}
|
||
)
|
||
|
||
|
||
@app.route("/inference", methods=["POST"])
|
||
def inference():
|
||
"""推理请求:multipart form 包含 video 和 audio 文件.
|
||
|
||
可选 form 参数:
|
||
bbox_shift: 口型区域垂直偏移量,默认 0,范围 [-5, 5]
|
||
face_enhance: 是否启用 GFPGAN 面部增强后处理,"0" 或 "1",默认跟随配置
|
||
"""
|
||
# 并发控制:检查锁
|
||
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)) # 限制范围
|
||
|
||
# 解析 face_enhance 参数:支持请求级覆盖
|
||
face_enhance_param = request.form.get("face_enhance", None)
|
||
if face_enhance_param is not None:
|
||
face_enhance = face_enhance_param.strip() in ("1", "true", "yes", "on")
|
||
else:
|
||
face_enhance = None # 跟随 Config 默认配置
|
||
|
||
# 文件大小检查
|
||
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, face_enhance=%s",
|
||
task_id, video_path.name, audio_path.name, bbox_shift, face_enhance)
|
||
|
||
# 更新当前任务信息
|
||
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,
|
||
face_enhance=face_enhance,
|
||
)
|
||
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():
|
||
"""启动 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()
|