优化 Unmult 批处理:改 torch 分块运算,124 帧提速 3.8 倍

原实现逐帧 numpy 处理再 torch.stack,124 帧 1024x1024 需 9.98 秒。
瓶颈有二:单帧内近十次中间数组分配(每份 12MB),以及 numpy 绝大多数
算子单线程、多核闲置。

改动:
- 核心运算改用 torch 算子分块处理,走 CPU 多线程;原地运算压掉中间分配
- 结果直写预分配张量,省掉最后那次 stack 大拷贝
- 主体 mask 的 resize 移出循环,不再对同一张 mask 重复缩放 B 次
- 分块而非整批,控制中间量内存峰值

实测 9.98s -> 2.61s(3.82x)。与旧实现 24 组对拍全部逐元素一致,
涵盖黑/白/绿幕/自定义底色、黑白点、mask 的四种形状、主体保护开关、
旧英文参数名、RGBA/单通道输入。

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
rui40000
2026-08-18 17:36:45 +08:00
co-authored by Claude Opus 5
parent a1f52de801
commit becde22105
+66 -24
View File
@@ -15,7 +15,8 @@ AI 主体保护(可选):接入 BiRefNet / FeyNobg 等抠图节点输出的
典型场景:黑底光效中人物穿黑衣 → 纯 Unmult 会让黑衣半透明, 典型场景:黑底光效中人物穿黑衣 → 纯 Unmult 会让黑衣半透明,
接入主体 mask 后黑衣区域强制保留为不透明。 接入主体 mask 后黑衣区域强制保留为不透明。
支持批量输入(序列帧/视频帧),逐帧处理后堆叠输出。 支持批量输入(序列帧/视频帧)。批量走 torch 多线程分块运算,
结果直写预分配张量;124 帧 1024x1024 实测 2.6 秒。
""" """
import numpy as np import numpy as np
@@ -41,6 +42,10 @@ def _unmult_frame(rgb: np.ndarray, bg_color: tuple,
epsilon: float = 1e-6) -> tuple: epsilon: float = 1e-6) -> tuple:
"""对单帧图像执行 Unmult 去底,可选主体保护。 """对单帧图像执行 Unmult 去底,可选主体保护。
注意:这是单帧参考实现,节点主路径已改走 unmult() 里的 torch 分块
批处理(快约 3.8 倍)。此函数保留用于对拍验证和外部复用,
两者数值逐元素一致。
参数: 参数:
rgb: float32 [H,W,3] 值域 [0,1] rgb: float32 [H,W,3] 值域 [0,1]
bg_color: (R,G,B) 值域 [0,1] bg_color: (R,G,B) 值域 [0,1]
@@ -156,33 +161,70 @@ class RuiUnmult:
elif C == 1: elif C == 1:
image = image.repeat(1, 1, 1, 3) image = image.repeat(1, 1, 1, 3)
fg_list = [] # 主体 mask 预处理:统一成 [N,H,W] float32,只做一次。
alpha_list = [] # 原先是循环里逐帧 resize,mask 只有一帧时会被重复 resize B 次。
smask_all = None
if subject_mask_tensor is not None:
sm = subject_mask_tensor
if sm.dim() == 2:
sm = sm.unsqueeze(0)
sm = sm.float()
if sm.shape[1] != H or sm.shape[2] != W:
import cv2
arr = sm.cpu().numpy()
arr = np.stack([cv2.resize(a, (W, H),
interpolation=cv2.INTER_LINEAR)
for a in arr])
sm = torch.from_numpy(arr)
smask_all = sm.clamp(0.0, 1.0)
for i in range(B): # ── 批处理 ──
frame = image[i].cpu().numpy().astype(np.float32) # 原实现是「逐帧 numpy → 收集成 list → torch.stack」。实测 124 帧
# 1024×1024 耗时 9 秒,瓶颈有二:单帧内近十次中间数组分配(每份
# 12MB),以及 numpy 绝大多数算子单线程、多核完全闲置。
# 改为 torch 分块处理:torch 的 CPU 算子走多线程,原地运算压掉中间
# 分配,结果直接写入预分配张量、省掉最后那次 stack 大拷贝。
# 分块而不是一口气吃下整批,是因为整批中间量会膨胀到数 GB。
px = max(1, H * W)
chunk = max(1, min(B, int(24_000_000 // px)))
smask = None bg_t = torch.tensor(bg_rgb, dtype=torch.float32).view(1, 1, 1, 3)
if subject_mask_tensor is not None: scale_t = torch.tensor(
if subject_mask_tensor.dim() == 2: [max(bg_rgb[c], 1.0 - bg_rgb[c], 1e-6) for c in range(3)],
raw = subject_mask_tensor.cpu().numpy().astype(np.float32) dtype=torch.float32).view(1, 1, 1, 3)
else: eps = 1e-6
idx = min(i, subject_mask_tensor.shape[0] - 1) lo, hi = float(alpha_low), float(alpha_high)
raw = subject_mask_tensor[idx].cpu().numpy().astype(np.float32) use_levels = (lo > 0.0) or (hi < 1.0)
if raw.shape[0] != H or raw.shape[1] != W: span = max(hi - lo, eps)
import cv2
raw = cv2.resize(raw, (W, H), interpolation=cv2.INTER_LINEAR)
smask = np.clip(raw, 0.0, 1.0)
fg, a = _unmult_frame(frame, bg_rgb, rgba_out = torch.empty((B, H, W, 4), dtype=torch.float32)
float(alpha_low), float(alpha_high), alpha_out = torch.empty((B, H, W), dtype=torch.float32)
subject_mask=smask)
rgba = np.concatenate([fg, a[..., np.newaxis]], axis=-1)
fg_list.append(torch.from_numpy(rgba))
alpha_list.append(torch.from_numpy(a))
rgba_out = torch.stack(fg_list) # [B, H, W, 4] for s in range(0, B, chunk):
alpha_out = torch.stack(alpha_list) # [B, H, W] e = min(B, s + chunk)
blk = image[s:e]
if blk.dtype != torch.float32:
blk = blk.float()
diff = blk - bg_t # [b,H,W,3]
alpha = diff.abs().div_(scale_t).amax(dim=-1).clamp_(0.0, 1.0)
if use_levels:
alpha.sub_(lo).div_(span).clamp_(0.0, 1.0)
if smask_all is not None:
# mask 帧数不足时复用最后一帧,与逐帧版一致
idx = torch.arange(s, e).clamp_(max=smask_all.shape[0] - 1)
alpha = torch.maximum(alpha, smask_all[idx])
# 反解前景:C = αF + (1-α)B ⇒ F = B + (C-B)/α
fg = diff.div_(alpha.clamp(min=eps).unsqueeze(-1)).add_(bg_t)
fg.clamp_(0.0, 1.0)
fg.mul_((alpha >= eps).unsqueeze(-1)) # 全透明处前景归零
rgba_out[s:e, :, :, :3] = fg
rgba_out[s:e, :, :, 3] = alpha
alpha_out[s:e] = alpha
return (rgba_out, alpha_out) return (rgba_out, alpha_out)