优化 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:
co-authored by
Claude Opus 5
parent
a1f52de801
commit
becde22105
+66
-24
@@ -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)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user