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
Co-authored-by: xiaoxia <dev@xiaoxiajianji.com> Co-committed-by: xiaoxia <dev@xiaoxiajianji.com>
1251 lines
51 KiB
Python
1251 lines
51 KiB
Python
"""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()
|