Files
xiaoxia-saas/deploy/gpu_worker/musetalk_server.py
T
xiaoxia 8ebccd61d9 fix: 集成真实MuseTalk推理 + TTS style/speed/volume参数全链路修复
MuseTalk推理修复(P0):
- 替换_run_inference() stub代码为真实MuseTalk推理逻辑
- 新增音频预处理:任意格式→16kHz mono 16bit WAV
- 模型懒加载(VAE+UNet+PE+Whisper),首次推理后复用
- 使用mirror indexing循环帧,消除视频循环边界跳变
- bbox_shift参数可通过请求配置(范围-5~5)
- FP16推理支持,节省RTX2060显存
- /health接口增加模型加载状态信息

TTS参数透传修复(P0):
- schemas新增style字段(TTSSynthesizeRequest/TTSPreviewRequest/CreateLipsyncJobRequest/AiAvatarTtsPreviewRequest)
- schemas新增volume字段(TTSSynthesizeRequest/CreateLipsyncJobRequest)
- cosyvoice_service新增STYLE_INSTRUCTION_MAP和build_style_instruction()
- submit_synthesize_task新增style/pitch参数支持
- workflow start_synthesis/resynthesize/submit_segments全链路传递style/volume/pitch
- tts.py路由/lipsync.py路由/lipsync_service.py全链路透传新参数
- lipsync_tts.py Celery task支持style/volume参数
2026-09-20 15:52:53 +08:00

929 lines
34 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 生成口型同步视频。
#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, read_imgs
from musetalk.utils.utils import datagen
from musetalk.utils.blending import get_image
# 加载模型(首次调用时加载,后续复用)
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)
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,
)
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)
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 + 10, ori_frame.shape[0])
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],
)
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():
"""启动 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()