基于 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>
205 lines
7.3 KiB
Python
205 lines
7.3 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
SDMatte 权重读取。
|
||
|
||
官方发布的 SDMatte_plus.pth 是 detectron2 的训练检查点,顶层结构为
|
||
{"model": {...1316 个张量...}, "trainer": {...优化器状态...}, "iteration": int}
|
||
其中 trainer 约占 6GB,推理完全用不到,且内部引用了 omegaconf 的类型,
|
||
直接 torch.load 会连带反序列化它、白白吃掉一倍内存,还平添一个依赖。
|
||
|
||
因此这里自己解析 zip 容器:先用受限的 Unpickler 还原出对象骨架(张量一律替换成
|
||
惰性占位),再只把 model 段的张量真正读进内存。附带收益是全程不执行 pickle 里的
|
||
任意代码,比 torch.load(weights_only=False) 更安全。
|
||
"""
|
||
|
||
import io
|
||
import os
|
||
import pickle
|
||
import zipfile
|
||
from collections import OrderedDict
|
||
|
||
import torch
|
||
|
||
# torch storage 类名 -> dtype
|
||
_STORAGE_DTYPES = {
|
||
"FloatStorage": torch.float32,
|
||
"HalfStorage": torch.float16,
|
||
"DoubleStorage": torch.float64,
|
||
"BFloat16Storage": torch.bfloat16,
|
||
"LongStorage": torch.int64,
|
||
"IntStorage": torch.int32,
|
||
"ShortStorage": torch.int16,
|
||
"CharStorage": torch.int8,
|
||
"ByteStorage": torch.uint8,
|
||
"BoolStorage": torch.bool,
|
||
}
|
||
|
||
|
||
class _Dummy:
|
||
"""吞掉 trainer 段里那些我们不关心的对象(omegaconf 容器等)。"""
|
||
|
||
def __init__(self, *args, **kwargs):
|
||
pass
|
||
|
||
def __setstate__(self, state):
|
||
pass
|
||
|
||
def __setitem__(self, key, value):
|
||
pass
|
||
|
||
def append(self, *args):
|
||
pass
|
||
|
||
def __reduce__(self):
|
||
return (_Dummy, ())
|
||
|
||
|
||
class _LazyStorage:
|
||
__slots__ = ("key", "dtype", "numel")
|
||
|
||
def __init__(self, key, dtype, numel):
|
||
self.key = key
|
||
self.dtype = dtype
|
||
self.numel = numel
|
||
|
||
|
||
class _LazyTensor:
|
||
"""记下重建一个张量所需的全部信息,但不读取数据。"""
|
||
|
||
__slots__ = ("storage", "offset", "size", "stride")
|
||
|
||
def __init__(self, storage, offset, size, stride):
|
||
self.storage = storage
|
||
self.offset = offset
|
||
self.size = size
|
||
self.stride = stride
|
||
|
||
|
||
def _lazy_rebuild(storage, storage_offset, size, stride, *args, **kwargs):
|
||
return _LazyTensor(storage, storage_offset, tuple(size), tuple(stride))
|
||
|
||
|
||
class _StorageType:
|
||
"""代表 torch.FloatStorage 之类的类对象,只用于查 dtype。"""
|
||
|
||
def __init__(self, name):
|
||
self.dtype = _STORAGE_DTYPES.get(name, torch.float32)
|
||
|
||
|
||
class _RestrictedUnpickler(pickle.Unpickler):
|
||
def find_class(self, module, name):
|
||
if module == "collections" and name == "OrderedDict":
|
||
return OrderedDict
|
||
if module == "torch._utils" and name in ("_rebuild_tensor_v2", "_rebuild_tensor"):
|
||
return _lazy_rebuild
|
||
if module == "torch" and name in _STORAGE_DTYPES:
|
||
return _StorageType(name)
|
||
# 其余一律替换成惰性替身,绝不 import 外部模块、绝不执行其代码
|
||
return _Dummy
|
||
|
||
def persistent_load(self, saved_id):
|
||
# 形如 ('storage', <storage_type>, key, location, numel)
|
||
if not (isinstance(saved_id, tuple) and len(saved_id) >= 5 and saved_id[0] == "storage"):
|
||
return _Dummy()
|
||
storage_type, key, location, numel = saved_id[1], saved_id[2], saved_id[3], saved_id[4]
|
||
dtype = getattr(storage_type, "dtype", torch.float32)
|
||
return _LazyStorage(str(key), dtype, numel)
|
||
|
||
|
||
def _materialize(zf, prefix, lazy, name):
|
||
"""把一个 _LazyTensor 真正读成 torch.Tensor。"""
|
||
st = lazy.storage
|
||
raw = zf.read(f"{prefix}data/{st.key}")
|
||
# frombuffer 需要可写缓冲区,bytearray 在此产生唯一一次拷贝
|
||
flat = torch.frombuffer(bytearray(raw), dtype=st.dtype)
|
||
if not lazy.size: # 0 维标量
|
||
return flat[lazy.offset].clone()
|
||
return torch.as_strided(flat, lazy.size, lazy.stride, lazy.offset).clone()
|
||
|
||
|
||
def _load_pth_model_only(path):
|
||
with zipfile.ZipFile(path) as zf:
|
||
names = zf.namelist()
|
||
pkl_name = next((n for n in names if n.endswith("data.pkl")), None)
|
||
if pkl_name is None:
|
||
raise ValueError(f"不是 torch 的 zip 检查点:{path}")
|
||
prefix = pkl_name[: -len("data.pkl")]
|
||
|
||
# 字节序校验:权重按原始字节读入,端序不符会得到垃圾数据
|
||
if f"{prefix}byteorder" in names:
|
||
bo = zf.read(f"{prefix}byteorder").decode().strip()
|
||
if bo != "little":
|
||
raise ValueError(f"检查点字节序为 {bo},本加载器仅支持小端")
|
||
|
||
skeleton = _RestrictedUnpickler(io.BytesIO(zf.read(pkl_name))).load()
|
||
|
||
if isinstance(skeleton, dict) and isinstance(skeleton.get("model"), dict):
|
||
raw_sd = skeleton["model"]
|
||
elif isinstance(skeleton, dict) and isinstance(skeleton.get("state_dict"), dict):
|
||
raw_sd = skeleton["state_dict"]
|
||
elif isinstance(skeleton, dict):
|
||
raw_sd = skeleton
|
||
else:
|
||
raise ValueError(f"无法识别的检查点结构:{type(skeleton)}")
|
||
|
||
state_dict = {}
|
||
for k, v in raw_sd.items():
|
||
if isinstance(v, _LazyTensor):
|
||
state_dict[k] = _materialize(zf, prefix, v, k)
|
||
return state_dict
|
||
|
||
|
||
def load_sdmatte_state_dict(path):
|
||
"""读入 SDMatte 权重,返回纯张量的 state_dict。支持 .pth 与 .safetensors。"""
|
||
ext = os.path.splitext(path)[1].lower()
|
||
if ext == ".safetensors":
|
||
from safetensors.torch import load_file
|
||
|
||
return load_file(path, device="cpu")
|
||
if ext in (".pth", ".pt", ".ckpt", ".bin"):
|
||
return _load_pth_model_only(path)
|
||
raise ValueError(f"不支持的权重格式:{ext}")
|
||
|
||
|
||
_TM_PREFIX = "text_encoder.text_model."
|
||
_TE_PREFIX = "text_encoder."
|
||
|
||
|
||
def adapt_state_dict_to_model(state_dict, model):
|
||
"""
|
||
抹平 transformers 版本差异导致的 text_encoder 键名错位。
|
||
|
||
官方权重用 transformers 4.x 保存,CLIPTextModel 内部裹了一层 text_model,
|
||
键形如 text_encoder.text_model.encoder.layers.0...;
|
||
transformers 5.x 起该包装层被移除,键变成 text_encoder.encoder.layers.0...。
|
||
|
||
两者不匹配时 load_state_dict(strict=False) 会把 text_encoder 的 372 个权重
|
||
悄悄全部丢掉、让它停留在随机初始化状态 —— 不报错,但输出质量明显劣化。
|
||
这里按当前环境的实际结构做一次前缀增删。
|
||
"""
|
||
model_keys = set(model.state_dict().keys())
|
||
model_wraps = any(k.startswith(_TM_PREFIX) for k in model_keys)
|
||
ckpt_wraps = any(k.startswith(_TM_PREFIX) for k in state_dict)
|
||
|
||
if model_wraps == ckpt_wraps:
|
||
return state_dict
|
||
|
||
adapted = {}
|
||
if ckpt_wraps and not model_wraps:
|
||
# 权重带 text_model,当前 transformers 不带 -> 剥离
|
||
for k, v in state_dict.items():
|
||
if k.startswith(_TM_PREFIX):
|
||
adapted[_TE_PREFIX + k[len(_TM_PREFIX):]] = v
|
||
else:
|
||
adapted[k] = v
|
||
print(f"[Ruinode-SDMatte] 已适配 text_encoder 键名(剥离 text_model 层,transformers>=5)")
|
||
else:
|
||
# 权重不带 text_model,当前 transformers 需要 -> 补上
|
||
for k, v in state_dict.items():
|
||
if k.startswith(_TE_PREFIX) and not k.startswith(_TM_PREFIX):
|
||
adapted[_TM_PREFIX + k[len(_TE_PREFIX):]] = v
|
||
else:
|
||
adapted[k] = v
|
||
print(f"[Ruinode-SDMatte] 已适配 text_encoder 键名(补回 text_model 层,transformers<5)")
|
||
return adapted
|