基于 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>
63 lines
2.8 KiB
Python
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)
|