feat(gpu): add GFPGAN face enhancement + ffmpeg pipe optimization
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m43s
AI Code Review / AI Code Review (pull_request) Successful in 7m33s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 6m26s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 15s
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 10s
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 / 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
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m40s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Validate - Style (pull_request) Successful in 22m19s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 18m23s
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 / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Unit Tests (pull_request) Successful in 26m31s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 29m43s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 10s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 7m54s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 1h28m58s
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 / CI Gate (pull_request) Successful in 1s
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped

- Add GFPGANv1.4 face super-resolution (FP16, +170MB VRAM, ~64ms/frame)
  - Controllable via MUSE_USE_GFPGAN env var (default enabled)
  - Enhances generated face region before blending with original frame
- Replace PNG disk I/O with ffmpeg stdin pipe encoding (~2s faster)
  - Frames streamed directly to ffmpeg as raw BGR24, no temp files
  - Eliminates output_frames_dir/ PNG write/read cycle
- Total warm inference: 54s for 5s video (was 56s with PNG, 48s without GFPGAN)
- VRAM usage: ~3.3GB with MuseTalk+GFPGAN loaded (RTX 2060 6GB OK)
This commit is contained in:
ying
2026-09-20 22:51:02 +08:00
parent e7bf85ca86
commit 93cb3e12a0
+102 -44
View File
@@ -87,6 +87,8 @@ class Config:
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"
# ── 全局状态 ──────────────────────────────────────────────────────────
@@ -477,6 +479,36 @@ def _load_musetalk_models():
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()
if Config.use_float16:
gfpgan_model = gfpgan_model.half()
gfpgan_model = gfpgan_model.to(device)
del gfpgan_ckpt
logger.info("GFPGAN 加载完成 (FP16=%s)", Config.use_float16)
else:
logger.warning("GFPGAN 模型不存在: %s,跳过人脸增强", gfpgan_path)
except Exception as e:
logger.warning("GFPGAN 加载失败,跳过人脸增强: %s", e)
gfpgan_model = None
else:
logger.info("GFPGAN 已禁用 (MUSE_USE_GFPGAN=0)")
_muse_models = {
"vae": vae,
"unet": unet,
@@ -484,6 +516,7 @@ def _load_musetalk_models():
"timesteps": timesteps,
"audio_processor": audio_processor,
"face_parsing": face_parsing,
"gfpgan": gfpgan_model,
"device": device,
"model_version": model_version,
}
@@ -776,66 +809,91 @@ def _run_inference(
logger.info("推理完成,生成 %d 帧", len(res_frame_list))
# ── Step 7: 合成最终帧并写PNG ──
output_frames_dir = video_path.parent / "output_frames"
output_frames_dir.mkdir(parents=True, exist_ok=True)
# ── 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(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)
n_out = min(len(res_frame_list), num_output_frames)
try:
for i in tqdm(range(n_out), desc="合成帧"):
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]
for i, res_frame in enumerate(tqdm(res_frame_list[:num_output_frames], desc="合成帧")):
cyc_i = i % len(coord_cycle)
bbox = coord_cycle[cyc_i]
ori_frame = copy.deepcopy(frame_list_cycle[cyc_i])
if bbox is None:
combined = ori_frame
else:
x1, y1, x2, y2 = bbox
try:
res_frame_resized = cv2.resize(
res_frame.astype(np.uint8), (x2-x1, y2-y1)
res_frame.astype(np.uint8), (x2-x1, y2-y1),
interpolation=cv2.INTER_LANCZOS4
)
except Exception:
combined = ori_frame
cv2.imwrite(str(output_frames_dir / f"{i:08d}.png"), combined)
_ff_proc.stdin.write(ori_frame.tobytes())
continue
# face parsing 融合(fp 实例存在时传入)
# GFPGAN 人脸超分增强
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)
_face_t = ((_face_t - 0.5) / 0.5).to(device)
if Config.use_float16:
_face_t = _face_t.half()
with torch.no_grad():
_out = gfpgan_enhancer(_face_t, return_rgb=False, weight=0.5)[0]
_out = _out.squeeze(0).float().cpu().clamp_(-1,1)
_out = ((_out + 1)/2*255).numpy().transpose(1,2,0)
_out_bgr = cv2.cvtColor(_out.astype(np.uint8), cv2.COLOR_RGB2BGR)
res_frame_resized = cv2.resize(_out_bgr, (_fw, _fh),
interpolation=cv2.INTER_LANCZOS4)
del _face_t, _out, _out_bgr
except Exception:
pass
# face parsing 融合
try:
if face_parsing is not None:
combined = get_image(ori_frame, res_frame_resized,
[x1, y1, x2, y2], fp=face_parsing)
[x1,y1,x2,y2], fp=face_parsing)
else:
combined = get_image(ori_frame, res_frame_resized,
[x1, y1, x2, y2])
except Exception as e:
logger.warning("face parsing 融合失败,使用简单粘贴: %s", e)
# 简单粘贴 fallback
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
cv2.imwrite(str(output_frames_dir / f"{i:08d}.png"), combined)
try: combined[y1:y2, x1:x2] = res_frame_resized
except Exception: combined = ori_frame
# 帧序列 → 无声视频
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,
)
_ff_proc.stdin.write(combined.tobytes())
_ff_proc.stdin.close()
_ff_ret = _ff_proc.wait(timeout=120)
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
raise
# 将无声视频复制到输出路径
shutil.copy2(str(silent_video_path), str(output_path))
# 清理中间文件
try:
shutil.rmtree(str(output_frames_dir))
torch.cuda.empty_cache()
if audio_wav_path.exists():
audio_wav_path.unlink()
if silent_video_path.exists() and str(silent_video_path) != str(output_path):