Files
rui40000-RUI-Nodes/sdmatte/nn_utils.py
T
rui40000andClaude Opus 4.8 e1a575fcda feat: 新增 SDMatte 精细抠图节点,实测复现官方效果
基于 SDMatte(vivo 相机研究院,ICCV 2025)的交互式抠图节点,
擅长发丝、绒毛、玻璃、烟雾等常规抠图模型处理不好的边缘。

实现要点:
- 严格照搬官方 configs/SDMatte.py 的推理配置(bbox 视觉提示、fp32、1024 分辨率),
  不做任何启发式后处理,输出即模型原始 alpha
- 内置官方 LongfeiHuang/SDMatte 的配置文件,无需下载 SD 2.1 权重,也无需联网。
  官方 load_weight=False 只用 config 搭骨架,全部权重由 checkpoint 覆盖;
  原版 SD 2.1 的 config 缺 bbox_time_embed_dim 等三个专有字段,缺字段时直接报错而非猜测
- 同时支持官方 .pth(12.1GB)与社区 .safetensors(5.19GB)。两者模型权重实测
  逐像素完全相同,pth 多出的 6.5GB 是 detectron2 的优化器状态;读 pth 时以受限
  Unpickler 只解析 model 段,内存占用与 safetensors 相当
- 自动适配 transformers 5.x 移除 text_model 包装层导致的键名漂移,
  避免 text_encoder 的 372 个权重被静默丢弃
- 加载后校验 1316 个张量全部对齐,有任何未覆盖/未使用的权重即中止报错
- 默认开启注意力分片,1024 下显存峰值由约 15.5GB 降至 9.1GB,速度反而略快

实测(官方效果图中的羊驼,对比官方公布 alpha):MAD=0.0113。
同图同权重下 ComfyUI-SDMatte 为 MAD=0.0884,相差 7.8 倍,主因是其
aux_input="trimap" —— 官方 aux_input_list 只含 point_mask/bbox_mask/mask,
trimap 从未作为视觉提示参与训练,且该分支坐标恒为 [0,0,1,1]、定位信息丢失。

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-16 11:16:36 +08:00

63 lines
2.8 KiB
Python

# -*- coding: utf-8 -*-
"""
SDMatte 的网络结构改造工具。
摘自 vivoCameraResearch/SDMatte 的 utils/utils.py,仅保留推理必需的三个函数,
去掉了训练期才用到的 get_unknown_tensor_from_pred(其内部硬编码了 .cuda())。
函数体与官方保持逐行一致,改动会直接影响权重能否对上号。
"""
import torch.nn as nn
from torch.nn import Conv2d
from torch.nn.parameter import Parameter
from diffusers.models.attention_processor import Attention, AttnProcessor
from .replace import custom_prepare_attention_mask, custom_get_attention_scores
def replace_unet_conv_in(unet, num):
"""把 conv_in 从 4 通道扩成 4*num 通道,权重按 num 复制并等比缩小。"""
_weight = unet.conv_in.weight.clone() # [320, 4, 3, 3]
_bias = unet.conv_in.bias.clone() # [320]
_weight = _weight.repeat((1, num, 1, 1))
# half the activation magnitude
_weight = _weight / num
_n_convin_out_channel = unet.conv_in.out_channels
_new_conv_in = Conv2d(4 * num, _n_convin_out_channel, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
_new_conv_in.weight = Parameter(_weight)
_new_conv_in.bias = Parameter(_bias)
unet.conv_in = _new_conv_in
# 官方此处会改写 unet.config["in_channels"];新版 diffusers 的 config 是
# FrozenDict,赋值会抛异常,而该字段在推理期并不被读取,故安全跳过。
try:
unet.config["in_channels"] = 4 * num
except Exception:
pass
return unet
def add_aux_conv_in(unet):
"""新增 aux_conv_in:把视觉提示的 latent 编码成 1024 维,充当 cross-attention 的条件。"""
aux_conv_in = nn.Conv2d(in_channels=4, out_channels=1024, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
aux_conv_in.weight.data[:320, :, :, :] = unet.conv_in.weight.data.clone()
aux_conv_in.weight.data[320:, :, :, :] = 0.0
aux_conv_in.bias.data[:320] = unet.conv_in.bias.data.clone()
aux_conv_in.bias.data[320:] = 0.0
unet.aux_conv_in = aux_conv_in
return unet
def replace_attention_mask_method(module, residual_connection):
"""递归替换注意力的 mask 处理逻辑,使其支持 SDMatte 的空间掩码自注意力。"""
if isinstance(module, Attention):
module.processor = AttnProcessor()
if hasattr(module, "prepare_attention_mask"):
module.prepare_attention_mask = custom_prepare_attention_mask.__get__(module)
if hasattr(module, "cross_attention_dim") and module.cross_attention_dim == 320:
module.residual_connection = residual_connection
if hasattr(module, "get_attention_scores"):
module.get_attention_scores = custom_get_attention_scores.__get__(module)
for child_name, child_module in module.named_children():
replace_attention_mask_method(child_module, residual_connection)