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>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
1e42521fec
commit
e1a575fcda
@@ -0,0 +1,4 @@
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.egg-info/
|
||||
.DS_Store
|
||||
@@ -23,6 +23,7 @@ Rui-Node🐶 是一个功能丰富的 ComfyUI 节点集合,提供图像处理
|
||||
|
||||
### 🤖 AI模型类
|
||||
- [千问编辑图像生成 / Qwen Edit Image Generation](#4-千问编辑图像生成--qwen-edit-image-generation)
|
||||
- [SDMatte 精细抠图 / SDMatte Interactive Matting](#17-sdmatte-精细抠图--sdmatte-interactive-matting)
|
||||
|
||||
### 📝 文本处理类
|
||||
- [镜头分词器 / Shot Splitter](#5-镜头分词器--shot-splitter)
|
||||
@@ -537,6 +538,151 @@ Rui-Node🐶 是一个功能丰富的 ComfyUI 节点集合,提供图像处理
|
||||
6. 所有节点名称采用中英双语显示
|
||||
7. 遮罩处理节点自动处理尺寸不匹配问题
|
||||
|
||||
### 17. SDMatte 精细抠图 / SDMatte Interactive Matting
|
||||
|
||||
基于 [SDMatte](https://github.com/vivoCameraResearch/SDMatte)(vivo 相机研究院,ICCV 2025)的交互式抠图节点。
|
||||
擅长发丝、绒毛、玻璃、烟雾等常规抠图模型处理不好的边缘。
|
||||
|
||||
包含两个节点:
|
||||
|
||||
| 节点 | 作用 |
|
||||
|---|---|
|
||||
| **SDMatte 加载器** | 载入权重,构建网络并常驻显存 |
|
||||
| **SDMatte 精细抠图** | 用视觉提示(框/掩码/点)驱动模型输出 alpha |
|
||||
|
||||
#### 模型准备
|
||||
|
||||
把权重放到 `ComfyUI/models/SDMatte/` 下即可,两种格式任选其一:
|
||||
|
||||
- `SDMatte_plus.pth` — 官方发布,12.1GB,[LongfeiHuang/SDMatte](https://huggingface.co/LongfeiHuang/SDMatte)
|
||||
- `SDMatte_plus.safetensors` — 社区转换,5.19GB,[1038lab/SDMatte](https://huggingface.co/1038lab/SDMatte)
|
||||
|
||||
> **这两个文件的模型权重逐比特完全相同**,不必纠结选哪个。
|
||||
> 已实测比对全部 1316 个张量:键名、形状、精度(均为 F32)、数值全部一致,无一例外。
|
||||
> 官方 pth 是 detectron2 的训练检查点,顶层为 `{"model", "trainer", "iteration"}`,
|
||||
> 多出的约 6.9GB 是 `trainer` 里的优化器状态与梯度缩放器,推理不参与。
|
||||
> 换用 pth **不会带来任何质量提升**。本节点两种格式都支持,读 pth 时只解析 `model` 段,
|
||||
> 内存占用与 safetensors 相当。
|
||||
|
||||
**不需要下载 Stable Diffusion 2.1 的权重。** SDMatte 虽以 SD 2.1 为骨架,但官方推理配置
|
||||
(`configs/SDMatte.py` 中 `load_weight=False`)只用配置文件搭出网络结构,全部权重随后由
|
||||
SDMatte 检查点覆盖。官方 HuggingFace 仓库本身也只发布 `.pth` 加若干 `config.json`,
|
||||
不含任何 SD 权重。所需配置已随本节点一起分发,开箱即用、无需联网。
|
||||
|
||||
#### 参数说明
|
||||
|
||||
**SDMatte 加载器**
|
||||
|
||||
| 参数 | 说明 |
|
||||
|---|---|
|
||||
| `ckpt_name` | `models/SDMatte/` 下的权重文件 |
|
||||
| `precision` | `fp32`(默认,与官方测试配置一致)/ `fp16`(省显存,但 SD 2.1 的 VAE 半精度下易溢出) |
|
||||
| `device` | `auto` / `cpu` |
|
||||
| `attention_slicing` | 默认开启。1024 下显存峰值从约 **15.5GB 降到 9.1GB**,实测速度反而略快,输出差异仅 1e-6 量级 |
|
||||
|
||||
> 显存参考(fp32 @ 1024,实测于 RTX 5090):开分片约 **9.1GB**,关分片约 **15.5GB**。
|
||||
> 12GB 显存的卡请保持分片开启。
|
||||
|
||||
**SDMatte 精细抠图**
|
||||
|
||||
| 参数 | 说明 |
|
||||
|---|---|
|
||||
| `mask` | 指示抠哪个目标的提示掩码,**不必精确**,粗略覆盖主体即可 |
|
||||
| `prompt_type` | 视觉提示类型,见下表 |
|
||||
| `inference_size` | 默认 `1024`,与官方测试一致 |
|
||||
| `is_transparent` | 玻璃、纱、烟雾等透明物体**务必打开** |
|
||||
| `caption` | 可选文本描述,留空即官方默认行为 |
|
||||
| `point_radius` | 仅 `point_mask` 生效,默认 35 |
|
||||
|
||||
`prompt_type` 选择:
|
||||
|
||||
| 取值 | 含义 | 适用 |
|
||||
|---|---|---|
|
||||
| `bbox_mask` | 取掩码外接框作为提示 | **默认,官方测试脚本的主路径,通常最稳** |
|
||||
| `mask` | 直接用掩码本身 | 已有较准的粗分割时 |
|
||||
| `point_mask` | 在掩码内随机取 10 个点 | 复现论文的点提示实验 |
|
||||
| `auto_mask` | 不给定位信息 | 画面只有单一主体 |
|
||||
|
||||
#### 典型接法
|
||||
|
||||
```
|
||||
加载图像 ──────────────┬──> SDMatte 精细抠图 ──> alpha (MASK)
|
||||
│ ▲ └──> cutout (IMAGE)
|
||||
任意分割节点 ──> mask ──┘ │
|
||||
SDMatte 加载器 ───────────────────┘
|
||||
```
|
||||
|
||||
`mask` 可以来自任何粗分割来源(SAM、rembg、手绘遮罩皆可)——SDMatte 的职责正是把粗糙边缘细化。
|
||||
|
||||
#### 实测数据
|
||||
|
||||
用官方效果图中的羊驼原图(绒毛边缘)跑本节点,与官方给出的 GT alpha 对比:
|
||||
|
||||
| 指标 | 数值 |
|
||||
|---|---|
|
||||
| MAD(平均绝对误差) | 0.0113 |
|
||||
| MSE | 0.0026 |
|
||||
| SAD | 0.807 千像素 |
|
||||
|
||||
(GT 取自官方效果图截图,含有损压缩与水印,故存在固有误差下限。)
|
||||
|
||||
各配置对输出的实际影响(透明玻璃杯,差异像素指偏差 > 0.05 的占比):
|
||||
|
||||
| 对照项 | 平均差 | 差异像素占比 |
|
||||
|---|---|---|
|
||||
| `inference_size` 1024 vs 512 | 0.082 | 32.4% |
|
||||
| `is_transparent` 关 vs 开 | 0.059 | 25.0% |
|
||||
| 官方 `[F,T,F]` vs 误用 `[T,T,T]` 条件分配 | 0.026 | 18.8% |
|
||||
|
||||
结论:**分辨率影响最大,建议保持 1024**;抠透明物体时 `is_transparent` 必须打开。
|
||||
|
||||
#### 与 ComfyUI-SDMatte 的横向实测
|
||||
|
||||
同一张图、同一份权重、同一台机器,对跑 [ComfyUI-SDMatte](https://github.com/flybirdxx/ComfyUI-SDMatte)
|
||||
与本节点,以官方公布的 alpha 为参照:
|
||||
|
||||
| 实现 | 配置 | MAD ↓ |
|
||||
|---|---|---|
|
||||
| **本节点** | 官方 `configs/SDMatte.py`,bbox 提示,fp32 | **0.0113** |
|
||||
| ComfyUI-SDMatte | 默认(trimap 提示 + `mask_refine`) | 0.0884 |
|
||||
| ComfyUI-SDMatte | 关闭 `mask_refine` | 0.0885 |
|
||||
|
||||
**相差 7.8 倍**,且其输出肉眼可见地发灰、边缘晕开。
|
||||
|
||||
主因是**视觉提示类型**:官方 `configs/SDMatte.py` 固定 `aux_input="bbox_mask"`,
|
||||
而其 `aux_input_list` 只含 `point_mask` / `bbox_mask` / `mask` —— **trimap 从未作为视觉提示参与训练**。
|
||||
ComfyUI-SDMatte 传 `aux_input="trimap"`,把模型推到了没训练过的输入模式上,
|
||||
且该分支的 `trimap_coords` 恒为 `[0,0,1,1]`,定位信息全部丢失。
|
||||
开不开它的 `mask_refine` 几乎不影响这一结论(0.0884 vs 0.0885),说明问题不在后处理。
|
||||
|
||||
#### 实现要点
|
||||
|
||||
若与其它 SDMatte 实现效果对不上,按影响从大到小排查:
|
||||
|
||||
1. **视觉提示类型**(影响最大)。必须用官方训练过的 `bbox_mask` / `mask` / `point_mask`,
|
||||
并传入真实的归一化坐标。用 trimap 当视觉提示是模型没见过的用法。
|
||||
|
||||
2. **UNet 配置来源**。SDMatte 在标准 SD 2.1 的 UNet 配置上额外定义了
|
||||
`bbox_time_embed_dim` / `point_embeddings_input_dim` / `bbox_embeddings_input_dim` 三个字段。
|
||||
误用原版 SD 2.1 的 `config.json` 会缺这些字段,只能猜默认值,猜错则相应权重被
|
||||
`strict=False` 静默丢弃。本节点直接分发官方配置,并在缺字段时**直接报错而非猜测**。
|
||||
|
||||
3. **transformers 版本**。官方权重用 transformers 4.x 保存,`CLIPTextModel` 内部裹了一层
|
||||
`text_model`;transformers 5.x 起该层被移除,导致 text_encoder 的 372 个权重键名对不上、
|
||||
被整体静默丢弃、停留在随机初始化。本节点会按当前环境自动增删该前缀。
|
||||
|
||||
4. **条件分配**。官方 `use_encoder_hidden_states_list=[False, True, False]` 决定 UNet
|
||||
下采样/中间/上采样三段各接收哪种条件,漏传会退化成 `[True, True, True]`。
|
||||
实测单独影响不大(羊驼 MAD 0.01135 → 0.01148),透明物体上更明显。
|
||||
|
||||
5. **权重对齐校验**。本节点在加载后校验键的完整性,一旦有权重未被覆盖或未被使用就**中止并报错**。
|
||||
这类问题不会让模型崩溃,只会让输出质量悄悄下降,是最难排查的一类,因此宁可停下也不放行。
|
||||
|
||||
6. 全程 fp32、1024 分辨率,且**不做任何启发式后处理**(不做阈值裁剪、对比度拉伸之类的"优化"),
|
||||
输出即模型原始 alpha。
|
||||
|
||||
---
|
||||
|
||||
## 🐕 关于 Rui-Node🐶
|
||||
|
||||
Rui-Node🐶 致力于为 ComfyUI 用户提供实用、高效的节点工具集。🐶 是我们的项目标志,代表着忠诚、友好和可靠。
|
||||
|
||||
+11
@@ -32,6 +32,15 @@ from .color_matcher_node import NODE_DISPLAY_NAME_MAPPINGS as COLORMATCHER_NODE_
|
||||
# 新增:素材拆分节点
|
||||
from .image_splitter_node import NODE_CLASS_MAPPINGS as IMAGESPLITTER_NODE_CLASS_MAPPINGS
|
||||
from .image_splitter_node import NODE_DISPLAY_NAME_MAPPINGS as IMAGESPLITTER_NODE_DISPLAY_NAME_MAPPINGS
|
||||
# 新增:SDMatte 精细抠图节点(依赖 diffusers/transformers,缺失时不影响其余节点加载)
|
||||
try:
|
||||
from .sdmatte_node import NODE_CLASS_MAPPINGS as SDMATTE_NODE_CLASS_MAPPINGS
|
||||
from .sdmatte_node import NODE_DISPLAY_NAME_MAPPINGS as SDMATTE_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as _e:
|
||||
print(f"[Ruinode] SDMatte 节点未加载:{_e}")
|
||||
print("[Ruinode] 如需使用,请安装:pip install diffusers transformers safetensors scipy opencv-python")
|
||||
SDMATTE_NODE_CLASS_MAPPINGS = {}
|
||||
SDMATTE_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
# 合并节点映射字典
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
@@ -51,6 +60,7 @@ NODE_CLASS_MAPPINGS.update(UTF8_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(OPENAI_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(COLORMATCHER_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(IMAGESPLITTER_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(SDMATTE_NODE_CLASS_MAPPINGS)
|
||||
|
||||
# 合并节点显示名称映射
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
@@ -70,5 +80,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(UTF8_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(OPENAI_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(COLORMATCHER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(IMAGESPLITTER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(SDMATTE_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
|
||||
@@ -2,3 +2,9 @@ torch
|
||||
numpy
|
||||
Pillow
|
||||
requests
|
||||
# 以下为 SDMatte 抠图节点所需(ComfyUI 环境通常已自带 opencv/scipy/safetensors)
|
||||
diffusers>=0.30.0
|
||||
transformers>=4.40.0
|
||||
safetensors
|
||||
scipy
|
||||
opencv-python
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
SDMatte 推理代码(移植自 vivoCameraResearch/SDMatte,MIT License)。
|
||||
|
||||
此处刻意不在包级别导入 meta_arch —— 它依赖 diffusers/transformers,
|
||||
若环境缺失会让整个 Ruinode 节点包加载失败。改由 sdmatte_node 在真正用到时再导入。
|
||||
"""
|
||||
@@ -0,0 +1,204 @@
|
||||
# -*- 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
|
||||
@@ -0,0 +1,14 @@
|
||||
{
|
||||
"_class_name": "DDIMScheduler",
|
||||
"_diffusers_version": "0.8.0",
|
||||
"beta_end": 0.012,
|
||||
"beta_schedule": "scaled_linear",
|
||||
"beta_start": 0.00085,
|
||||
"clip_sample": false,
|
||||
"num_train_timesteps": 1000,
|
||||
"prediction_type": "v_prediction",
|
||||
"set_alpha_to_one": false,
|
||||
"skip_prk_steps": true,
|
||||
"steps_offset": 1,
|
||||
"trained_betas": null
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
{
|
||||
"_name_or_path": "hf-models/stable-diffusion-v2-768x768/text_encoder",
|
||||
"architectures": [
|
||||
"CLIPTextModel"
|
||||
],
|
||||
"attention_dropout": 0.0,
|
||||
"bos_token_id": 0,
|
||||
"dropout": 0.0,
|
||||
"eos_token_id": 2,
|
||||
"hidden_act": "gelu",
|
||||
"hidden_size": 1024,
|
||||
"initializer_factor": 1.0,
|
||||
"initializer_range": 0.02,
|
||||
"intermediate_size": 4096,
|
||||
"layer_norm_eps": 1e-05,
|
||||
"max_position_embeddings": 77,
|
||||
"model_type": "clip_text_model",
|
||||
"num_attention_heads": 16,
|
||||
"num_hidden_layers": 23,
|
||||
"pad_token_id": 1,
|
||||
"projection_dim": 512,
|
||||
"torch_dtype": "float32",
|
||||
"transformers_version": "4.25.0.dev0",
|
||||
"vocab_size": 49408
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,24 @@
|
||||
{
|
||||
"bos_token": {
|
||||
"content": "<|startoftext|>",
|
||||
"lstrip": false,
|
||||
"normalized": true,
|
||||
"rstrip": false,
|
||||
"single_word": false
|
||||
},
|
||||
"eos_token": {
|
||||
"content": "<|endoftext|>",
|
||||
"lstrip": false,
|
||||
"normalized": true,
|
||||
"rstrip": false,
|
||||
"single_word": false
|
||||
},
|
||||
"pad_token": "!",
|
||||
"unk_token": {
|
||||
"content": "<|endoftext|>",
|
||||
"lstrip": false,
|
||||
"normalized": true,
|
||||
"rstrip": false,
|
||||
"single_word": false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
{
|
||||
"add_prefix_space": false,
|
||||
"bos_token": {
|
||||
"__type": "AddedToken",
|
||||
"content": "<|startoftext|>",
|
||||
"lstrip": false,
|
||||
"normalized": true,
|
||||
"rstrip": false,
|
||||
"single_word": false
|
||||
},
|
||||
"do_lower_case": true,
|
||||
"eos_token": {
|
||||
"__type": "AddedToken",
|
||||
"content": "<|endoftext|>",
|
||||
"lstrip": false,
|
||||
"normalized": true,
|
||||
"rstrip": false,
|
||||
"single_word": false
|
||||
},
|
||||
"errors": "replace",
|
||||
"model_max_length": 77,
|
||||
"name_or_path": "hf-models/stable-diffusion-v2-768x768/tokenizer",
|
||||
"pad_token": "<|endoftext|>",
|
||||
"special_tokens_map_file": "./special_tokens_map.json",
|
||||
"tokenizer_class": "CLIPTokenizer",
|
||||
"unk_token": {
|
||||
"__type": "AddedToken",
|
||||
"content": "<|endoftext|>",
|
||||
"lstrip": false,
|
||||
"normalized": true,
|
||||
"rstrip": false,
|
||||
"single_word": false
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,47 @@
|
||||
{
|
||||
"_class_name": "UNet2DConditionModel",
|
||||
"_diffusers_version": "0.8.0",
|
||||
"_name_or_path": "hf-models/stable-diffusion-v2-768x768/unet",
|
||||
"act_fn": "silu",
|
||||
"bbox_time_embed_dim": 320,
|
||||
"point_embeddings_input_dim": 1680,
|
||||
"bbox_embeddings_input_dim": 1280,
|
||||
"attention_head_dim": [
|
||||
5,
|
||||
10,
|
||||
20,
|
||||
20
|
||||
],
|
||||
"block_out_channels": [
|
||||
320,
|
||||
640,
|
||||
1280,
|
||||
1280
|
||||
],
|
||||
"center_input_sample": false,
|
||||
"cross_attention_dim": 1024,
|
||||
"down_block_types": [
|
||||
"CrossAttnDownBlock2D",
|
||||
"CrossAttnDownBlock2D",
|
||||
"CrossAttnDownBlock2D",
|
||||
"DownBlock2D"
|
||||
],
|
||||
"downsample_padding": 1,
|
||||
"dual_cross_attention": false,
|
||||
"flip_sin_to_cos": true,
|
||||
"freq_shift": 0,
|
||||
"in_channels": 4,
|
||||
"layers_per_block": 2,
|
||||
"mid_block_scale_factor": 1,
|
||||
"norm_eps": 1e-05,
|
||||
"norm_num_groups": 32,
|
||||
"out_channels": 4,
|
||||
"sample_size": 96,
|
||||
"up_block_types": [
|
||||
"UpBlock2D",
|
||||
"CrossAttnUpBlock2D",
|
||||
"CrossAttnUpBlock2D",
|
||||
"CrossAttnUpBlock2D"
|
||||
],
|
||||
"use_linear_projection": true
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
{
|
||||
"_class_name": "AutoencoderKL",
|
||||
"_diffusers_version": "0.8.0",
|
||||
"_name_or_path": "hf-models/stable-diffusion-v2-768x768/vae",
|
||||
"act_fn": "silu",
|
||||
"block_out_channels": [
|
||||
128,
|
||||
256,
|
||||
512,
|
||||
512
|
||||
],
|
||||
"down_block_types": [
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D"
|
||||
],
|
||||
"in_channels": 3,
|
||||
"latent_channels": 4,
|
||||
"layers_per_block": 2,
|
||||
"norm_num_groups": 32,
|
||||
"out_channels": 3,
|
||||
"sample_size": 768,
|
||||
"up_block_types": [
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D"
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,250 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
SDMatte 主干网络。
|
||||
|
||||
移植自 vivoCameraResearch/SDMatte 的 modeling/SDMatte/meta_arch.py。
|
||||
相对官方版本只做了两类改动,计算图本身逐行保持一致:
|
||||
|
||||
1. 官方把 .cuda() 硬编码在 forward 各处,这里改为跟随模型自身所在设备,
|
||||
以便在 ComfyUI 里支持 CPU / 多卡 / 显存卸载。
|
||||
2. 官方 init_submodule 在 load_weight=False 时用 os.path.join 拼子目录,
|
||||
这里同样只读配置、不读权重,但对路径做了存在性校验并给出可读的报错。
|
||||
|
||||
注意 load_weight 恒为 False 是官方推理配置(configs/SDMatte.py)的原意:
|
||||
网络结构只从 config.json 构建,全部权重随后由 checkpoint 覆盖,
|
||||
因此 Stable Diffusion 2.1 的权重文件在整个流程中不参与。
|
||||
"""
|
||||
|
||||
import os
|
||||
import random
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from diffusers import DDIMScheduler, AutoencoderKL
|
||||
from diffusers.models.embeddings import get_timestep_embedding
|
||||
from transformers import CLIPTextModel, CLIPTokenizer, CLIPTextConfig
|
||||
|
||||
from .nn_utils import replace_unet_conv_in, replace_attention_mask_method, add_aux_conv_in
|
||||
from .replace import CustomUNet
|
||||
|
||||
AUX_INPUT_DIT = {
|
||||
"auto_mask": "auto_coords",
|
||||
"point_mask": "point_coords",
|
||||
"bbox_mask": "bbox_coords",
|
||||
"mask": "mask_coords",
|
||||
"trimap": "trimap_coords",
|
||||
}
|
||||
|
||||
|
||||
class SDMatte(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
pretrained_model_name_or_path,
|
||||
conv_scale=3,
|
||||
num_inference_steps=1,
|
||||
aux_input="bbox_mask",
|
||||
use_aux_input=False,
|
||||
use_coor_input=True,
|
||||
use_dis_loss=True,
|
||||
use_attention_mask=True,
|
||||
use_encoder_attention_mask=False,
|
||||
add_noise=False,
|
||||
attn_mask_aux_input=["point_mask", "bbox_mask", "mask"],
|
||||
aux_input_list=["point_mask", "bbox_mask", "mask"],
|
||||
use_encoder_hidden_states=True,
|
||||
residual_connection=False,
|
||||
use_attention_mask_list=[True, True, True],
|
||||
use_encoder_hidden_states_list=[True, True, True],
|
||||
load_weight=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.init_submodule(pretrained_model_name_or_path, load_weight)
|
||||
self.num_inference_steps = num_inference_steps
|
||||
self.aux_input = aux_input
|
||||
self.use_aux_input = use_aux_input
|
||||
self.use_coor_input = use_coor_input
|
||||
self.use_dis_loss = use_dis_loss
|
||||
self.use_attention_mask = use_attention_mask
|
||||
self.use_encoder_attention_mask = use_encoder_attention_mask
|
||||
self.add_noise = add_noise
|
||||
self.attn_mask_aux_input = attn_mask_aux_input
|
||||
self.aux_input_list = aux_input_list
|
||||
self.use_encoder_hidden_states = use_encoder_hidden_states
|
||||
if use_encoder_hidden_states:
|
||||
self.unet = add_aux_conv_in(self.unet)
|
||||
if not add_noise:
|
||||
conv_scale -= 1
|
||||
if not use_aux_input:
|
||||
conv_scale -= 1
|
||||
if conv_scale > 1:
|
||||
self.unet = replace_unet_conv_in(self.unet, conv_scale)
|
||||
replace_attention_mask_method(self.unet, residual_connection)
|
||||
self.text_encoder.requires_grad_(False)
|
||||
self.vae.requires_grad_(False)
|
||||
self.unet.use_attention_mask_list = use_attention_mask_list
|
||||
self.unet.use_encoder_hidden_states_list = use_encoder_hidden_states_list
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.unet.parameters()).device
|
||||
|
||||
def init_submodule(self, pretrained_model_name_or_path, load_weight):
|
||||
if load_weight:
|
||||
self.text_encoder = CLIPTextModel.from_pretrained(pretrained_model_name_or_path, subfolder="text_encoder")
|
||||
self.vae = AutoencoderKL.from_pretrained(pretrained_model_name_or_path, subfolder="vae")
|
||||
self.unet = CustomUNet.from_pretrained(
|
||||
pretrained_model_name_or_path, subfolder="unet", low_cpu_mem_usage=True, ignore_mismatched_sizes=False
|
||||
)
|
||||
self.noise_scheduler = DDIMScheduler.from_pretrained(pretrained_model_name_or_path, subfolder="scheduler")
|
||||
self.tokenizer = CLIPTokenizer.from_pretrained(pretrained_model_name_or_path, subfolder="tokenizer")
|
||||
return
|
||||
|
||||
# 仅从 config 构建骨架,不加载任何 SD 权重
|
||||
unet_path = os.path.join(pretrained_model_name_or_path, "unet")
|
||||
unet_cfg_file = os.path.join(unet_path, "config.json")
|
||||
if not os.path.isfile(unet_cfg_file):
|
||||
raise FileNotFoundError(f"缺少 SDMatte 的 unet 配置:{unet_cfg_file}")
|
||||
|
||||
unet_config = CustomUNet.load_config(unet_path)
|
||||
# SDMatte 在标准 SD 2.1 的 UNet 配置上追加了三个字段,用于坐标嵌入与不透明度嵌入。
|
||||
# 缺任何一个都说明配置来源不对(例如误用了原版 SD 2.1 的 config.json),
|
||||
# 此时若放任默认值继续,权重会因形状不符被静默丢弃,必须直接报错。
|
||||
for field in ("bbox_time_embed_dim", "point_embeddings_input_dim", "bbox_embeddings_input_dim"):
|
||||
if field not in unet_config:
|
||||
raise ValueError(
|
||||
f"unet 配置缺少 SDMatte 专有字段 '{field}':{unet_cfg_file}\n"
|
||||
"这通常是误用了原版 Stable Diffusion 2.1 的 config.json。"
|
||||
"请使用 LongfeiHuang/SDMatte 提供的配置文件。"
|
||||
)
|
||||
|
||||
text_config = CLIPTextConfig.from_pretrained(pretrained_model_name_or_path, subfolder="text_encoder")
|
||||
self.text_encoder = CLIPTextModel(text_config)
|
||||
|
||||
vae_path = os.path.join(pretrained_model_name_or_path, "vae")
|
||||
self.vae = AutoencoderKL.from_config(AutoencoderKL.load_config(vae_path))
|
||||
|
||||
self.unet = CustomUNet.from_config(unet_config, low_cpu_mem_usage=True, ignore_mismatched_sizes=False)
|
||||
|
||||
scheduler_path = os.path.join(pretrained_model_name_or_path, "scheduler", "scheduler_config.json")
|
||||
self.noise_scheduler = DDIMScheduler.from_config(DDIMScheduler.load_config(scheduler_path))
|
||||
|
||||
self.tokenizer = CLIPTokenizer.from_pretrained(pretrained_model_name_or_path, subfolder="tokenizer")
|
||||
|
||||
def forward(self, data):
|
||||
device = self.device
|
||||
rgb = data["image"].to(device)
|
||||
B = rgb.shape[0]
|
||||
|
||||
if self.aux_input is None and self.training:
|
||||
aux_input_type = random.choice(self.aux_input_list)
|
||||
elif self.aux_input is None:
|
||||
aux_input_type = "point_mask"
|
||||
else:
|
||||
aux_input_type = self.aux_input
|
||||
|
||||
# 视觉提示 -> latent
|
||||
if self.use_aux_input:
|
||||
aux_input = data[aux_input_type].to(device)
|
||||
aux_input = aux_input.repeat(1, 3, 1, 1)
|
||||
aux_input_h = self.vae.encoder(aux_input.to(rgb.dtype))
|
||||
aux_input_moments = self.vae.quant_conv(aux_input_h)
|
||||
aux_input_mean, _ = torch.chunk(aux_input_moments, 2, dim=1)
|
||||
aux_input_latent = aux_input_mean * self.vae.config.scaling_factor
|
||||
else:
|
||||
aux_input_latent = None
|
||||
|
||||
# 视觉提示的坐标嵌入
|
||||
coor_name = AUX_INPUT_DIT[aux_input_type]
|
||||
coor = data[coor_name].to(device)
|
||||
if coor_name == "point_coords":
|
||||
N = coor.shape[1]
|
||||
for i in range(N, 1680):
|
||||
if 1680 % i == 0:
|
||||
num_channels = 1680 // i
|
||||
pad_size = i - N
|
||||
padding = torch.zeros((B, pad_size), dtype=coor.dtype, device=coor.device)
|
||||
coor = torch.cat([coor, padding], dim=1)
|
||||
zero_coor = torch.zeros((B, pad_size + N), dtype=coor.dtype, device=coor.device)
|
||||
break
|
||||
if self.use_coor_input:
|
||||
coor = get_timestep_embedding(
|
||||
coor.flatten(),
|
||||
num_channels,
|
||||
flip_sin_to_cos=True,
|
||||
downscale_freq_shift=0,
|
||||
)
|
||||
else:
|
||||
coor = get_timestep_embedding(
|
||||
zero_coor.flatten(),
|
||||
num_channels,
|
||||
flip_sin_to_cos=True,
|
||||
downscale_freq_shift=0,
|
||||
)
|
||||
added_cond_kwargs = {"point_coords": coor}
|
||||
else:
|
||||
if self.use_coor_input:
|
||||
added_cond_kwargs = {"bbox_mask_coords": coor}
|
||||
else:
|
||||
coor = torch.tensor([[0, 0, 1, 1]] * B, device=device)
|
||||
added_cond_kwargs = {"bbox_mask_coords": coor}
|
||||
|
||||
# 掩码自注意力
|
||||
if self.use_attention_mask and aux_input_type in self.attn_mask_aux_input:
|
||||
attention_mask = data[aux_input_type].to(device)
|
||||
attention_mask = (attention_mask + 1) / 2
|
||||
attention_mask = F.interpolate(attention_mask, scale_factor=1 / 8, mode="nearest")
|
||||
attention_mask = attention_mask.flatten(start_dim=1)
|
||||
else:
|
||||
attention_mask = None
|
||||
|
||||
# 原图 -> latent
|
||||
rgb_h = self.vae.encoder(rgb)
|
||||
rgb_moments = self.vae.quant_conv(rgb_h)
|
||||
rgb_mean, _ = torch.chunk(rgb_moments, 2, dim=1)
|
||||
rgb_latent = rgb_mean * self.vae.config.scaling_factor
|
||||
|
||||
# 视觉提示驱动的交互条件
|
||||
encoder_hidden_states = None
|
||||
if self.use_encoder_hidden_states and aux_input_latent is not None:
|
||||
encoder_hidden_states = self.unet.aux_conv_in(aux_input_latent)
|
||||
encoder_hidden_states = encoder_hidden_states.view(B, 1024, -1)
|
||||
encoder_hidden_states = encoder_hidden_states.permute(0, 2, 1)
|
||||
|
||||
if "caption" in data:
|
||||
prompt = data["caption"]
|
||||
else:
|
||||
prompt = [""] * B
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=self.tokenizer.model_max_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids.to(device)
|
||||
text_embed = self.text_encoder(text_input_ids)[0]
|
||||
encoder_hidden_states_2 = text_embed
|
||||
|
||||
# 不透明度嵌入:透明物体走另一条分支
|
||||
is_trans = data["is_trans"].to(device)
|
||||
trans = 1 - is_trans
|
||||
|
||||
# 官方在此处构造了 timestep 张量却传入 None,即单步、不加噪。此处保持一致。
|
||||
label_latent = self.unet(
|
||||
sample=torch.cat([rgb_latent, aux_input_latent], dim=1),
|
||||
trans=trans,
|
||||
timestep=None,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
encoder_hidden_states_2=encoder_hidden_states_2,
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
attention_mask=attention_mask,
|
||||
).sample
|
||||
label_latent = label_latent / self.vae.config.scaling_factor
|
||||
z = self.vae.post_quant_conv(label_latent)
|
||||
stacked = self.vae.decoder(z)
|
||||
label_mean = stacked.mean(dim=1, keepdim=True)
|
||||
output = torch.clip(label_mean, -1.0, 1.0)
|
||||
output = (output + 1.0) / 2.0
|
||||
return output
|
||||
@@ -0,0 +1,62 @@
|
||||
# -*- 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)
|
||||
@@ -0,0 +1,157 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
视觉提示构造。
|
||||
|
||||
SDMatte 是交互式抠图:除了原图,还要一个"指哪抠哪"的视觉提示(框/掩码/点),
|
||||
以及该提示的归一化坐标。官方在 data/dataset.py 里用 GenBBox / GenMask / GenPoint
|
||||
从 GT alpha 现场生成这些提示;在 ComfyUI 里没有 GT,改由用户给的 mask 充当提示来源,
|
||||
这正是官方设计的交互用法。
|
||||
|
||||
除此之外,各函数的行为与官方测试期逐行对齐(含 1024 分辨率、连通域筛选、
|
||||
sigma=radius 的高斯点扩散、以及 x2-1 归一化),偏离任何一处都会让结果偏离论文指标。
|
||||
"""
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import scipy.ndimage
|
||||
from scipy.ndimage import label
|
||||
|
||||
# 官方 configs/SDMatte.py: psm="gauss", radius=25;测试期用 radius + 10
|
||||
DEFAULT_POINT_RADIUS = 35
|
||||
DEFAULT_POINT_THRES = 0.8
|
||||
NUM_POINTS = 10
|
||||
|
||||
|
||||
def resize_image(img, size):
|
||||
"""等价于官方 Resize:双线性缩放到 size x size。"""
|
||||
return cv2.resize(img, (size, size), interpolation=cv2.INTER_LINEAR)
|
||||
|
||||
|
||||
def gen_bbox(ref, coe_scale=0.0, rng=None):
|
||||
"""
|
||||
复刻官方 GenBBox(测试期 coe_scale=0,即不做随机扰动)。
|
||||
|
||||
返回矩形提示掩码与归一化坐标 [x_min, y_min, x_max, y_max]。
|
||||
当存在一个显著大于其余的连通域时,官方只取该主体的外接框,
|
||||
以免零星噪点把框撑大。
|
||||
"""
|
||||
height, width = ref.shape
|
||||
coords = np.nonzero(ref)
|
||||
if coords[0].size == 0 or coords[1].size == 0:
|
||||
return np.zeros_like(ref, dtype=np.float32), np.array([0, 0, 1, 1], dtype=np.float32)
|
||||
|
||||
binary_mask = ref > 0
|
||||
labeled_array, num_features = label(binary_mask)
|
||||
y_min, x_min = np.argwhere(binary_mask).min(axis=0)
|
||||
y_max, x_max = np.argwhere(binary_mask).max(axis=0)
|
||||
|
||||
if num_features > 0:
|
||||
component_coords = [np.argwhere(labeled_array == i) for i in range(1, num_features + 1)]
|
||||
areas = [c.shape[0] for c in component_coords]
|
||||
sorted_areas_idx = np.argsort(areas)[::-1]
|
||||
max_area_idx = sorted_areas_idx[0]
|
||||
second_max_area_idx = sorted_areas_idx[1] if len(sorted_areas_idx) > 1 else None
|
||||
max_area = areas[max_area_idx]
|
||||
second_max_area = areas[second_max_area_idx] if second_max_area_idx is not None else 0
|
||||
if max_area >= 10 * second_max_area:
|
||||
max_coords = component_coords[max_area_idx]
|
||||
y_min, x_min = max_coords.min(axis=0)
|
||||
y_max, x_max = max_coords.max(axis=0)
|
||||
|
||||
if coe_scale:
|
||||
rng = rng or np.random
|
||||
coe = rng.uniform(0, coe_scale)
|
||||
padding_y = int(coe * (y_max - y_min))
|
||||
padding_x = int(coe * (x_max - x_min))
|
||||
y_min_p = padding_y if rng.random() < 0.5 else -padding_y
|
||||
y_max_p = padding_y if rng.random() < 0.5 else -padding_y
|
||||
x_min_p = padding_x if rng.random() < 0.5 else -padding_x
|
||||
x_max_p = padding_x if rng.random() < 0.5 else -padding_x
|
||||
y_min, y_max = max(0, y_min + y_min_p), min(height, y_max + y_max_p)
|
||||
x_min, x_max = max(0, x_min + x_min_p), min(width, x_max + x_max_p)
|
||||
|
||||
bbox_mask = np.zeros_like(ref, dtype=np.float32)
|
||||
bbox_mask[y_min:y_max, x_min:x_max] = 1
|
||||
|
||||
coords_norm = np.array(
|
||||
[x_min / width, y_min / height, x_max / width, y_max / height], dtype=np.float32
|
||||
)
|
||||
return bbox_mask, coords_norm
|
||||
|
||||
|
||||
def gen_mask(ref):
|
||||
"""复刻官方 GenMask 在"掩码已给定"时的分支:掩码原样,坐标取其外接框。"""
|
||||
h, w = ref.shape
|
||||
coords = np.nonzero(ref)
|
||||
if coords[0].size == 0 or coords[1].size == 0:
|
||||
mask_coords = np.array([0, 0, 1, 1], dtype=np.float32)
|
||||
else:
|
||||
y_min, x_min = np.argwhere(ref).min(axis=0)
|
||||
y_max, x_max = np.argwhere(ref).max(axis=0)
|
||||
mask_coords = np.array([x_min / w, y_min / h, x_max / w, y_max / h], dtype=np.float32)
|
||||
return ref.astype(np.float32), mask_coords
|
||||
|
||||
|
||||
def gen_point(ref, thres=DEFAULT_POINT_THRES, radius=DEFAULT_POINT_RADIUS, seed=0):
|
||||
"""
|
||||
复刻官方 GenPoint 的 gauss 模式:在提示区域内随机取 10 个点,
|
||||
每点摊开成一个 sigma=radius 的高斯斑,取逐像素最大值合并。
|
||||
|
||||
官方用的是全局 np.random,这里换成带种子的独立发生器,让结果可复现。
|
||||
"""
|
||||
height, width = ref.shape
|
||||
alpha_mask = (ref > thres).astype(np.float32)
|
||||
y_coords, x_coords = np.where(alpha_mask == 1)
|
||||
|
||||
if len(y_coords) < NUM_POINTS:
|
||||
return np.zeros_like(ref, dtype=np.float32), np.zeros(20, dtype=np.float32)
|
||||
|
||||
rng = np.random.default_rng(seed)
|
||||
selected_indices = rng.choice(len(y_coords), size=NUM_POINTS, replace=False)
|
||||
|
||||
point_mask = np.zeros_like(ref, dtype=np.float32)
|
||||
point_coords = []
|
||||
for idx in selected_indices:
|
||||
y_center = y_coords[idx]
|
||||
x_center = x_coords[idx]
|
||||
tmp_mask = np.zeros_like(ref, dtype=np.float32)
|
||||
tmp_mask[y_center, x_center] = 1
|
||||
tmp_mask = scipy.ndimage.gaussian_filter(tmp_mask, sigma=radius)
|
||||
peak = np.max(tmp_mask)
|
||||
if peak > 0:
|
||||
tmp_mask = tmp_mask / peak
|
||||
point_mask = np.maximum(point_mask, tmp_mask)
|
||||
point_coords.append(x_center / width)
|
||||
point_coords.append(y_center / height)
|
||||
|
||||
if len(point_coords) < 20:
|
||||
point_coords = np.concatenate([point_coords, np.zeros(20 - len(point_coords))])
|
||||
|
||||
return point_mask, np.array(point_coords[:20], dtype=np.float32)
|
||||
|
||||
|
||||
def gen_auto(ref):
|
||||
"""auto 模式:整幅图都是提示区域,不提供任何定位信息。"""
|
||||
return np.ones_like(ref, dtype=np.float32), np.array([0, 0, 1, 1], dtype=np.float32)
|
||||
|
||||
|
||||
def build_prompt(ref, prompt_type, point_radius=DEFAULT_POINT_RADIUS, seed=0):
|
||||
"""
|
||||
按提示类型产出 (提示掩码, 坐标)。
|
||||
|
||||
ref 为 [H, W] 的 float32,取值 [0, 1],且已缩放到推理分辨率。
|
||||
"""
|
||||
if prompt_type == "bbox_mask":
|
||||
return gen_bbox(ref)
|
||||
if prompt_type == "mask":
|
||||
return gen_mask(ref)
|
||||
if prompt_type == "point_mask":
|
||||
return gen_point(ref, radius=point_radius, seed=seed)
|
||||
if prompt_type == "auto_mask":
|
||||
return gen_auto(ref)
|
||||
raise ValueError(f"未知的提示类型:{prompt_type}")
|
||||
|
||||
|
||||
def normalize(x):
|
||||
"""官方 Normalize:把 [0,1] 映射到 [-1,1]。"""
|
||||
return x.astype(np.float32) * 2 - 1
|
||||
@@ -0,0 +1,550 @@
|
||||
import math
|
||||
from typing import Any, Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
import math
|
||||
from diffusers import UNet2DConditionModel
|
||||
from diffusers.models.embeddings import Timesteps, TimestepEmbedding
|
||||
from diffusers.models.unets.unet_2d_blocks import (
|
||||
get_down_block,
|
||||
get_up_block,
|
||||
get_mid_block,
|
||||
)
|
||||
from diffusers.models.activations import get_activation
|
||||
from diffusers.models.unets.unet_2d_condition import UNet2DConditionOutput
|
||||
from diffusers.utils import USE_PEFT_BACKEND, scale_lora_layers, unscale_lora_layers
|
||||
|
||||
|
||||
def custom_prepare_attention_mask(
|
||||
self, attention_mask: torch.Tensor, target_length: int, batch_size: int, out_dim: int = 3
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
Prepare the attention mask for the attention computation.
|
||||
|
||||
Args:
|
||||
attention_mask (`torch.Tensor`):
|
||||
The attention mask to prepare.
|
||||
target_length (`int`):
|
||||
The target length of the attention mask. This is the length of the attention mask after padding.
|
||||
batch_size (`int`):
|
||||
The batch size, which is used to repeat the attention mask.
|
||||
out_dim (`int`, *optional*, defaults to `3`):
|
||||
The output dimension of the attention mask. Can be either `3` or `4`.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`: The prepared attention mask.
|
||||
"""
|
||||
head_size = self.heads
|
||||
if attention_mask is None:
|
||||
return attention_mask
|
||||
|
||||
current_length: int = attention_mask.shape[-1]
|
||||
if current_length != target_length:
|
||||
if attention_mask.device.type == "mps":
|
||||
# HACK: MPS: Does not support padding by greater than dimension of input tensor.
|
||||
# Instead, we can manually construct the padding tensor.
|
||||
padding_shape = (attention_mask.shape[0], attention_mask.shape[1], target_length)
|
||||
padding = torch.zeros(padding_shape, dtype=attention_mask.dtype, device=attention_mask.device)
|
||||
attention_mask = torch.cat([attention_mask, padding], dim=2)
|
||||
else:
|
||||
# TODO: for pipelines such as stable-diffusion, padding cross-attn mask:
|
||||
# we want to instead pad by (0, remaining_length), where remaining_length is:
|
||||
# remaining_length: int = target_length - current_length
|
||||
# TODO: re-enable tests/models/test_models_unet_2d_condition.py#test_model_xattn_padding
|
||||
B = attention_mask.shape[0]
|
||||
current_size = int(math.sqrt(current_length))
|
||||
target_size = int(math.sqrt(target_length))
|
||||
assert current_size**2 == current_length, f"current_length ({current_length}) cannot be squared to an integer size"
|
||||
assert target_size**2 == target_length, f"target_length ({target_length}) cannot be squared to an integer size"
|
||||
attention_mask = attention_mask.view(B, -1, current_size, current_size)
|
||||
attention_mask = F.interpolate(attention_mask, size=(target_size, target_size), mode="nearest")
|
||||
attention_mask = attention_mask.view(B, 1, target_length)
|
||||
|
||||
if out_dim == 3:
|
||||
if attention_mask.shape[0] < batch_size * head_size:
|
||||
attention_mask = attention_mask.repeat_interleave(head_size, dim=0)
|
||||
elif out_dim == 4:
|
||||
attention_mask = attention_mask.unsqueeze(1)
|
||||
attention_mask = attention_mask.repeat_interleave(head_size, dim=1)
|
||||
|
||||
return attention_mask
|
||||
|
||||
|
||||
def custom_get_attention_scores(self, query: torch.Tensor, key: torch.Tensor, attention_mask: torch.Tensor = None) -> torch.Tensor:
|
||||
r"""
|
||||
Compute the attention scores.
|
||||
|
||||
Args:
|
||||
query (`torch.Tensor`): The query tensor.
|
||||
key (`torch.Tensor`): The key tensor.
|
||||
attention_mask (`torch.Tensor`, *optional*): The attention mask to use. If `None`, no mask is applied.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`: The attention probabilities/scores.
|
||||
"""
|
||||
dtype = query.dtype
|
||||
if self.upcast_attention:
|
||||
query = query.float()
|
||||
key = key.float()
|
||||
|
||||
# if attention_mask is not None and len(torch.unique(attention_mask)) <= 2:
|
||||
if attention_mask is not None:
|
||||
baddbmm_input = attention_mask
|
||||
beta = 1
|
||||
else:
|
||||
baddbmm_input = torch.empty(query.shape[0], query.shape[1], key.shape[1], dtype=query.dtype, device=query.device)
|
||||
beta = 0
|
||||
|
||||
attention_scores = torch.baddbmm(
|
||||
baddbmm_input,
|
||||
query,
|
||||
key.transpose(-1, -2),
|
||||
beta=beta,
|
||||
alpha=self.scale,
|
||||
)
|
||||
|
||||
# if attention_mask is not None and len(torch.unique(attention_mask)) > 2:
|
||||
# m = 1 - (attention_mask / -10000.0)
|
||||
# attention_scores = m * attention_scores
|
||||
|
||||
del baddbmm_input
|
||||
|
||||
if self.upcast_softmax:
|
||||
attention_scores = attention_scores.float()
|
||||
|
||||
attention_probs = attention_scores.softmax(dim=-1)
|
||||
del attention_scores
|
||||
|
||||
attention_probs = attention_probs.to(dtype)
|
||||
|
||||
return attention_probs
|
||||
|
||||
|
||||
class CustomUNet(UNet2DConditionModel):
|
||||
def __init__(
|
||||
self,
|
||||
sample_size: Optional[int] = None,
|
||||
in_channels: int = 4,
|
||||
out_channels: int = 4,
|
||||
flip_sin_to_cos: bool = True,
|
||||
freq_shift: int = 0,
|
||||
down_block_types: Tuple[str] = (
|
||||
"CrossAttnDownBlock2D",
|
||||
"CrossAttnDownBlock2D",
|
||||
"CrossAttnDownBlock2D",
|
||||
"DownBlock2D",
|
||||
),
|
||||
mid_block_type: Optional[str] = "UNetMidBlock2DCrossAttn",
|
||||
up_block_types: Tuple[str] = ("UpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D"),
|
||||
only_cross_attention: Union[bool, Tuple[bool]] = False,
|
||||
block_out_channels: Tuple[int] = (320, 640, 1280, 1280),
|
||||
layers_per_block: Union[int, Tuple[int]] = 2,
|
||||
downsample_padding: int = 1,
|
||||
mid_block_scale_factor: float = 1,
|
||||
dropout: float = 0.0,
|
||||
act_fn: str = "silu",
|
||||
norm_num_groups: Optional[int] = 32,
|
||||
norm_eps: float = 1e-5,
|
||||
cross_attention_dim: Union[int, Tuple[int]] = 1280,
|
||||
transformer_layers_per_block: Union[int, Tuple[int], Tuple[Tuple]] = 1,
|
||||
reverse_transformer_layers_per_block: Optional[Tuple[Tuple[int]]] = None,
|
||||
attention_head_dim: Union[int, Tuple[int]] = 8,
|
||||
num_attention_heads: Optional[Union[int, Tuple[int]]] = None,
|
||||
dual_cross_attention: bool = False,
|
||||
use_linear_projection: bool = False,
|
||||
upcast_attention: bool = False,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
resnet_skip_time_act: bool = False,
|
||||
resnet_out_scale_factor: int = 1.0,
|
||||
time_embedding_dim: Optional[int] = None,
|
||||
timestep_post_act: Optional[str] = None,
|
||||
time_cond_proj_dim: Optional[int] = None,
|
||||
conv_in_kernel: int = 3,
|
||||
conv_out_kernel: int = 3,
|
||||
bbox_time_embed_dim: Optional[int] = None,
|
||||
point_embeddings_input_dim: Optional[int] = None,
|
||||
bbox_embeddings_input_dim: Optional[int] = None,
|
||||
attention_type: str = "default",
|
||||
class_embeddings_concat: bool = False,
|
||||
mid_block_only_cross_attention: Optional[bool] = None,
|
||||
cross_attention_norm: Optional[str] = None,
|
||||
use_attention_mask_list=[True, True, True],
|
||||
use_encoder_hidden_states_list=[True, True, True],
|
||||
):
|
||||
super().__init__()
|
||||
self.use_attention_mask_list = use_attention_mask_list
|
||||
self.use_encoder_hidden_states_list = use_encoder_hidden_states_list
|
||||
self.sample_size = sample_size
|
||||
num_attention_heads = num_attention_heads or attention_head_dim
|
||||
|
||||
# input
|
||||
conv_in_padding = (conv_in_kernel - 1) // 2
|
||||
self.conv_in = nn.Conv2d(in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding)
|
||||
|
||||
# time
|
||||
time_embed_dim = time_embedding_dim or block_out_channels[0] * 4
|
||||
self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift)
|
||||
timestep_input_dim = block_out_channels[0]
|
||||
self.time_embedding = TimestepEmbedding(
|
||||
timestep_input_dim,
|
||||
time_embed_dim,
|
||||
act_fn=act_fn,
|
||||
post_act_fn=timestep_post_act,
|
||||
cond_proj_dim=time_cond_proj_dim,
|
||||
)
|
||||
|
||||
self.point_embedding = TimestepEmbedding(point_embeddings_input_dim, time_embed_dim)
|
||||
self.bbox_time_proj = Timesteps(bbox_time_embed_dim, flip_sin_to_cos, freq_shift)
|
||||
self.bbox_embedding = TimestepEmbedding(bbox_embeddings_input_dim, time_embed_dim)
|
||||
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
if isinstance(only_cross_attention, bool):
|
||||
if mid_block_only_cross_attention is None:
|
||||
mid_block_only_cross_attention = only_cross_attention
|
||||
only_cross_attention = [only_cross_attention] * len(down_block_types)
|
||||
|
||||
if mid_block_only_cross_attention is None:
|
||||
mid_block_only_cross_attention = False
|
||||
|
||||
if isinstance(num_attention_heads, int):
|
||||
num_attention_heads = (num_attention_heads,) * len(down_block_types)
|
||||
|
||||
if isinstance(attention_head_dim, int):
|
||||
attention_head_dim = (attention_head_dim,) * len(down_block_types)
|
||||
|
||||
if isinstance(cross_attention_dim, int):
|
||||
cross_attention_dim = (cross_attention_dim,) * len(down_block_types)
|
||||
|
||||
if isinstance(layers_per_block, int):
|
||||
layers_per_block = [layers_per_block] * len(down_block_types)
|
||||
|
||||
if isinstance(transformer_layers_per_block, int):
|
||||
transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types)
|
||||
|
||||
if class_embeddings_concat:
|
||||
blocks_time_embed_dim = time_embed_dim * 2
|
||||
else:
|
||||
blocks_time_embed_dim = time_embed_dim
|
||||
|
||||
# down
|
||||
output_channel = block_out_channels[0]
|
||||
for i, down_block_type in enumerate(down_block_types):
|
||||
input_channel = output_channel
|
||||
output_channel = block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
|
||||
down_block = get_down_block(
|
||||
down_block_type,
|
||||
num_layers=layers_per_block[i],
|
||||
transformer_layers_per_block=transformer_layers_per_block[i],
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
temb_channels=blocks_time_embed_dim,
|
||||
add_downsample=not is_final_block,
|
||||
resnet_eps=norm_eps,
|
||||
resnet_act_fn=act_fn,
|
||||
resnet_groups=norm_num_groups,
|
||||
cross_attention_dim=cross_attention_dim[i],
|
||||
num_attention_heads=num_attention_heads[i],
|
||||
downsample_padding=downsample_padding,
|
||||
dual_cross_attention=dual_cross_attention,
|
||||
use_linear_projection=use_linear_projection,
|
||||
only_cross_attention=only_cross_attention[i],
|
||||
upcast_attention=upcast_attention,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
attention_type=attention_type,
|
||||
resnet_skip_time_act=resnet_skip_time_act,
|
||||
resnet_out_scale_factor=resnet_out_scale_factor,
|
||||
cross_attention_norm=cross_attention_norm,
|
||||
attention_head_dim=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel,
|
||||
dropout=dropout,
|
||||
)
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
# mid
|
||||
self.mid_block = get_mid_block(
|
||||
mid_block_type,
|
||||
temb_channels=blocks_time_embed_dim,
|
||||
in_channels=block_out_channels[-1],
|
||||
resnet_eps=norm_eps,
|
||||
resnet_act_fn=act_fn,
|
||||
resnet_groups=norm_num_groups,
|
||||
output_scale_factor=mid_block_scale_factor,
|
||||
transformer_layers_per_block=transformer_layers_per_block[-1],
|
||||
num_attention_heads=num_attention_heads[-1],
|
||||
cross_attention_dim=cross_attention_dim[-1],
|
||||
dual_cross_attention=dual_cross_attention,
|
||||
use_linear_projection=use_linear_projection,
|
||||
mid_block_only_cross_attention=mid_block_only_cross_attention,
|
||||
upcast_attention=upcast_attention,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
attention_type=attention_type,
|
||||
resnet_skip_time_act=resnet_skip_time_act,
|
||||
cross_attention_norm=cross_attention_norm,
|
||||
attention_head_dim=attention_head_dim[-1],
|
||||
dropout=dropout,
|
||||
)
|
||||
|
||||
# count how many layers upsample the images
|
||||
self.num_upsamplers = 0
|
||||
|
||||
# up
|
||||
reversed_block_out_channels = list(reversed(block_out_channels))
|
||||
reversed_num_attention_heads = list(reversed(num_attention_heads))
|
||||
reversed_layers_per_block = list(reversed(layers_per_block))
|
||||
reversed_cross_attention_dim = list(reversed(cross_attention_dim))
|
||||
reversed_transformer_layers_per_block = (
|
||||
list(reversed(transformer_layers_per_block))
|
||||
if reverse_transformer_layers_per_block is None
|
||||
else reverse_transformer_layers_per_block
|
||||
)
|
||||
only_cross_attention = list(reversed(only_cross_attention))
|
||||
|
||||
output_channel = reversed_block_out_channels[0]
|
||||
for i, up_block_type in enumerate(up_block_types):
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
|
||||
prev_output_channel = output_channel
|
||||
output_channel = reversed_block_out_channels[i]
|
||||
input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)]
|
||||
|
||||
# add upsample block for all BUT final layer
|
||||
if not is_final_block:
|
||||
add_upsample = True
|
||||
self.num_upsamplers += 1
|
||||
else:
|
||||
add_upsample = False
|
||||
|
||||
up_block = get_up_block(
|
||||
up_block_type,
|
||||
num_layers=reversed_layers_per_block[i] + 1,
|
||||
transformer_layers_per_block=reversed_transformer_layers_per_block[i],
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
prev_output_channel=prev_output_channel,
|
||||
temb_channels=blocks_time_embed_dim,
|
||||
add_upsample=add_upsample,
|
||||
resnet_eps=norm_eps,
|
||||
resnet_act_fn=act_fn,
|
||||
resolution_idx=i,
|
||||
resnet_groups=norm_num_groups,
|
||||
cross_attention_dim=reversed_cross_attention_dim[i],
|
||||
num_attention_heads=reversed_num_attention_heads[i],
|
||||
dual_cross_attention=dual_cross_attention,
|
||||
use_linear_projection=use_linear_projection,
|
||||
only_cross_attention=only_cross_attention[i],
|
||||
upcast_attention=upcast_attention,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
attention_type=attention_type,
|
||||
resnet_skip_time_act=resnet_skip_time_act,
|
||||
resnet_out_scale_factor=resnet_out_scale_factor,
|
||||
cross_attention_norm=cross_attention_norm,
|
||||
attention_head_dim=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel,
|
||||
dropout=dropout,
|
||||
)
|
||||
self.up_blocks.append(up_block)
|
||||
prev_output_channel = output_channel
|
||||
|
||||
# out
|
||||
if norm_num_groups is not None:
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=norm_eps)
|
||||
|
||||
self.conv_act = get_activation(act_fn)
|
||||
|
||||
else:
|
||||
self.conv_norm_out = None
|
||||
self.conv_act = None
|
||||
|
||||
conv_out_padding = (conv_out_kernel - 1) // 2
|
||||
self.conv_out = nn.Conv2d(block_out_channels[0], out_channels, kernel_size=conv_out_kernel, padding=conv_out_padding)
|
||||
|
||||
# distillation
|
||||
self.feature_map = []
|
||||
|
||||
def _get_value(self, use_list, true_value, false_value):
|
||||
down_value = mid_value = up_value = false_value
|
||||
|
||||
if use_list[0]:
|
||||
down_value = true_value
|
||||
if use_list[1]:
|
||||
mid_value = true_value
|
||||
if use_list[2]:
|
||||
up_value = true_value
|
||||
|
||||
return down_value, mid_value, up_value
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
timestep: Union[torch.Tensor, float, int],
|
||||
trans: Union[torch.Tensor, float, int],
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_2: Optional[torch.Tensor] = None,
|
||||
timestep_cond: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
|
||||
encoder_attention_mask: Optional[torch.Tensor] = None,
|
||||
) -> Union[UNet2DConditionOutput, Tuple]:
|
||||
default_overall_up_factor = 2**self.num_upsamplers
|
||||
forward_upsample_size = False
|
||||
upsample_size = None
|
||||
|
||||
for dim in sample.shape[-2:]:
|
||||
if dim % default_overall_up_factor != 0:
|
||||
forward_upsample_size = True
|
||||
break
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0
|
||||
attention_mask = attention_mask.unsqueeze(1)
|
||||
|
||||
if encoder_attention_mask is not None:
|
||||
encoder_attention_mask = (1 - encoder_attention_mask.to(sample.dtype)) * -10000.0
|
||||
encoder_attention_mask = encoder_attention_mask.unsqueeze(1)
|
||||
|
||||
# 0. center input if necessary
|
||||
if self.config.center_input_sample:
|
||||
sample = 2 * sample - 1.0
|
||||
|
||||
down_attn_mask, mid_attn_mask, up_attn_mask = self._get_value(self.use_attention_mask_list, attention_mask, None)
|
||||
down_encoder_hidden_states, mid_encoder_hidden_states, up_encoder_hidden_states = self._get_value(
|
||||
self.use_encoder_hidden_states_list, encoder_hidden_states, encoder_hidden_states_2
|
||||
)
|
||||
|
||||
# 1. time
|
||||
t_emb, op_emb, aug_emb = None, None, None
|
||||
|
||||
if timestep is not None:
|
||||
timesteps = timestep
|
||||
timesteps = timesteps.expand(sample.shape[0])
|
||||
t_emb = self.time_proj(timesteps)
|
||||
t_emb = t_emb.to(dtype=sample.dtype)
|
||||
|
||||
t_emb = self.time_embedding(t_emb, timestep_cond)
|
||||
|
||||
# opacity
|
||||
if trans is not None:
|
||||
trans = trans.expand(sample.shape[0])
|
||||
op_emb = self.time_proj(trans)
|
||||
op_emb = op_emb.to(dtype=sample.dtype)
|
||||
|
||||
op_emb = self.time_embedding(op_emb, timestep_cond)
|
||||
|
||||
if t_emb is not None and op_emb is not None:
|
||||
emb = t_emb + op_emb
|
||||
elif op_emb is not None:
|
||||
emb = op_emb
|
||||
elif t_emb is not None:
|
||||
emb = t_emb
|
||||
else:
|
||||
raise ValueError("Missing required field: 'timestep' and 'trans'. Please ensure it is included in your input.")
|
||||
|
||||
if "point_coords" in added_cond_kwargs:
|
||||
coords_embeds = added_cond_kwargs.get("point_coords")
|
||||
coords_embeds = coords_embeds.reshape((sample.shape[0], -1))
|
||||
coords_embeds = coords_embeds.to(emb.dtype)
|
||||
aug_emb = self.point_embedding(coords_embeds)
|
||||
elif "bbox_mask_coords" in added_cond_kwargs:
|
||||
coords = added_cond_kwargs.get("bbox_mask_coords")
|
||||
coords_embeds = self.bbox_time_proj(coords.flatten())
|
||||
coords_embeds = coords_embeds.reshape((sample.shape[0], -1))
|
||||
coords_embeds = coords_embeds.to(emb.dtype)
|
||||
aug_emb = self.bbox_embedding(coords_embeds)
|
||||
else:
|
||||
raise ValueError(f"{self.__class__} cannot find point_coords or bbox_coords in added_cond_kwargs.")
|
||||
|
||||
emb = emb + aug_emb if aug_emb is not None else emb
|
||||
|
||||
# 2. pre-process
|
||||
sample = self.conv_in(sample)
|
||||
|
||||
# distillation
|
||||
self.feature_map = []
|
||||
|
||||
# 3. down
|
||||
lora_scale = cross_attention_kwargs.get("scale", 1.0) if cross_attention_kwargs is not None else 1.0
|
||||
if USE_PEFT_BACKEND:
|
||||
scale_lora_layers(self, lora_scale)
|
||||
|
||||
down_block_res_samples = (sample,)
|
||||
for downsample_block in self.down_blocks:
|
||||
if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention:
|
||||
additional_residuals = {}
|
||||
sample, res_samples = downsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
encoder_hidden_states=down_encoder_hidden_states,
|
||||
attention_mask=down_attn_mask,
|
||||
cross_attention_kwargs=cross_attention_kwargs,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
**additional_residuals,
|
||||
)
|
||||
else:
|
||||
sample, res_samples = downsample_block(hidden_states=sample, temb=emb, scale=lora_scale)
|
||||
|
||||
down_block_res_samples += res_samples
|
||||
|
||||
self.feature_map.append(sample)
|
||||
|
||||
# 4. mid
|
||||
if self.mid_block is not None:
|
||||
if hasattr(self.mid_block, "has_cross_attention") and self.mid_block.has_cross_attention:
|
||||
sample = self.mid_block(
|
||||
sample,
|
||||
emb,
|
||||
encoder_hidden_states=mid_encoder_hidden_states,
|
||||
attention_mask=mid_attn_mask,
|
||||
cross_attention_kwargs=cross_attention_kwargs,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
)
|
||||
else:
|
||||
sample = self.mid_block(sample, emb)
|
||||
|
||||
self.feature_map.append(sample)
|
||||
|
||||
# 5. up
|
||||
for i, upsample_block in enumerate(self.up_blocks):
|
||||
is_final_block = i == len(self.up_blocks) - 1
|
||||
|
||||
res_samples = down_block_res_samples[-len(upsample_block.resnets) :]
|
||||
down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)]
|
||||
|
||||
if not is_final_block and forward_upsample_size:
|
||||
upsample_size = down_block_res_samples[-1].shape[2:]
|
||||
|
||||
if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention:
|
||||
sample = upsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
res_hidden_states_tuple=res_samples,
|
||||
encoder_hidden_states=up_encoder_hidden_states,
|
||||
cross_attention_kwargs=cross_attention_kwargs,
|
||||
upsample_size=upsample_size,
|
||||
attention_mask=up_attn_mask,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
)
|
||||
else:
|
||||
sample = upsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
res_hidden_states_tuple=res_samples,
|
||||
upsample_size=upsample_size,
|
||||
scale=lora_scale,
|
||||
)
|
||||
|
||||
self.feature_map.append(sample)
|
||||
|
||||
# 6. post-process
|
||||
if self.conv_norm_out:
|
||||
sample = self.conv_norm_out(sample)
|
||||
sample = self.conv_act(sample)
|
||||
sample = self.conv_out(sample)
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
return UNet2DConditionOutput(sample=sample)
|
||||
+331
@@ -0,0 +1,331 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
SDMatte 精细化抠图节点(Ruinode)
|
||||
|
||||
对应论文:SDMatte: Grafting Diffusion Models for Interactive Matting (ICCV 2025)
|
||||
官方实现:https://github.com/vivoCameraResearch/SDMatte
|
||||
|
||||
与社区已有实现的关键差别(按实测影响从大到小排列):
|
||||
|
||||
1. 视觉提示走官方的 bbox_mask 路径,并真正传入归一化坐标。
|
||||
官方 configs/SDMatte.py 固定 aux_input="bbox_mask",其 aux_input_list 只有
|
||||
point_mask / bbox_mask / mask —— trimap 从未作为视觉提示参与训练。
|
||||
ComfyUI-SDMatte 传 aux_input="trimap",等于把模型推到没训练过的输入模式上,
|
||||
且该分支的 trimap_coords 恒为 [0,0,1,1],定位信息全部丢失。
|
||||
实测(官方效果图里的羊驼,与官方公布 alpha 比):本节点 MAD=0.0113,
|
||||
ComfyUI-SDMatte MAD=0.0884,相差 7.8 倍,且其输出明显发灰、边缘晕开。
|
||||
|
||||
2. UNet 结构用官方 LongfeiHuang/SDMatte 的 config.json 构建,而非原版 SD 2.1 的。
|
||||
官方配置额外定义了 bbox_time_embed_dim=320 等三个字段;用 SD 2.1 的配置会缺字段,
|
||||
只能靠猜默认值,一旦猜错,对应权重会被 strict=False 静默丢弃。
|
||||
|
||||
3. 照搬官方 configs/SDMatte.py 的 model_kwargs,含
|
||||
use_encoder_hidden_states_list=[False, True, False](漏传会退化成 [True,True,True])。
|
||||
实测该项单独影响不大(羊驼 MAD 0.01135 -> 0.01148),透明物体上更明显;
|
||||
影响虽小,但没有任何理由偏离官方配置。
|
||||
|
||||
4. 全程 fp32、1024 分辨率推理,与官方测试配置一致,且不做任何启发式后处理。
|
||||
ComfyUI-SDMatte 的 mask_refine 会做阈值截断与 *1.2 提亮,实测反而把边缘打成硬边。
|
||||
|
||||
Stable Diffusion 2.1 的权重在本流程中不需要:官方 load_weight=False,
|
||||
网络只从 config 建骨架,全部权重来自 SDMatte 检查点。
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
import folder_paths
|
||||
|
||||
# ---------------------------------------------------------------- 模型目录注册
|
||||
|
||||
SDMATTE_DIR = os.path.join(folder_paths.models_dir, "SDMatte")
|
||||
os.makedirs(SDMATTE_DIR, exist_ok=True)
|
||||
|
||||
# 与 ComfyUI-SDMatte 共用同一目录,已下载过的权重可直接复用
|
||||
if "SDMatte" in folder_paths.folder_names_and_paths:
|
||||
_paths, _exts = folder_paths.folder_names_and_paths["SDMatte"]
|
||||
if SDMATTE_DIR not in _paths:
|
||||
_paths.append(SDMATTE_DIR)
|
||||
_exts.update({".pth", ".safetensors", ".pt", ".ckpt"})
|
||||
else:
|
||||
folder_paths.folder_names_and_paths["SDMatte"] = (
|
||||
[SDMATTE_DIR],
|
||||
{".pth", ".safetensors", ".pt", ".ckpt"},
|
||||
)
|
||||
|
||||
# 官方配置(unet/vae/text_encoder/scheduler/tokenizer)已随节点一起分发,无需联网
|
||||
CONFIG_DIR = os.path.join(os.path.dirname(os.path.realpath(__file__)), "sdmatte", "configs")
|
||||
|
||||
# 官方 configs/SDMatte.py -> hy_dict.model_kwargs,逐字对齐
|
||||
OFFICIAL_MODEL_KWARGS = dict(
|
||||
load_weight=False,
|
||||
conv_scale=3,
|
||||
num_inference_steps=1,
|
||||
aux_input="bbox_mask",
|
||||
add_noise=False,
|
||||
use_dis_loss=True,
|
||||
use_aux_input=True,
|
||||
use_coor_input=True,
|
||||
use_attention_mask=True,
|
||||
residual_connection=False,
|
||||
use_encoder_hidden_states=True,
|
||||
use_attention_mask_list=[True, True, True],
|
||||
use_encoder_hidden_states_list=[False, True, False],
|
||||
)
|
||||
|
||||
_MODEL_CACHE = {}
|
||||
|
||||
|
||||
def _list_checkpoints():
|
||||
try:
|
||||
files = folder_paths.get_filename_list("SDMatte")
|
||||
except Exception:
|
||||
files = []
|
||||
if not files:
|
||||
files = [
|
||||
f for f in os.listdir(SDMATTE_DIR)
|
||||
if f.lower().endswith((".pth", ".safetensors", ".pt", ".ckpt"))
|
||||
] if os.path.isdir(SDMATTE_DIR) else []
|
||||
return sorted(files) if files else ["未找到权重,请放入 models/SDMatte"]
|
||||
|
||||
|
||||
def _resolve_ckpt(name):
|
||||
path = folder_paths.get_full_path("SDMatte", name)
|
||||
if path and os.path.isfile(path):
|
||||
return path
|
||||
direct = os.path.join(SDMATTE_DIR, name)
|
||||
if os.path.isfile(direct):
|
||||
return direct
|
||||
raise FileNotFoundError(
|
||||
f"找不到权重 '{name}'。请把 SDMatte_plus.pth 放到:{SDMATTE_DIR}\n"
|
||||
"官方下载地址:https://huggingface.co/LongfeiHuang/SDMatte"
|
||||
)
|
||||
|
||||
|
||||
def _build_model(ckpt_path, dtype, device, attention_slicing=True):
|
||||
from .sdmatte.ckpt_io import load_sdmatte_state_dict, adapt_state_dict_to_model
|
||||
from .sdmatte.meta_arch import SDMatte
|
||||
|
||||
print(f"[Ruinode-SDMatte] 按官方配置构建网络:{CONFIG_DIR}")
|
||||
model = SDMatte(pretrained_model_name_or_path=CONFIG_DIR, **OFFICIAL_MODEL_KWARGS)
|
||||
|
||||
print(f"[Ruinode-SDMatte] 读取权重:{os.path.basename(ckpt_path)}")
|
||||
state_dict = load_sdmatte_state_dict(ckpt_path)
|
||||
print(f"[Ruinode-SDMatte] 权重张量数:{len(state_dict)}")
|
||||
|
||||
state_dict = adapt_state_dict_to_model(state_dict, model)
|
||||
missing, unexpected = model.load_state_dict(state_dict, strict=False)
|
||||
|
||||
# load_state_dict(strict=False) 会把对不上的权重悄悄丢掉,模型照样能跑,
|
||||
# 只是输出质量下降 —— 这是最难排查的一类问题,所以这里必须叫停而不是继续。
|
||||
if missing or unexpected:
|
||||
print(f"[Ruinode-SDMatte] 权重未对齐:缺失 {len(missing)} 个,多余 {len(unexpected)} 个")
|
||||
for k in list(missing)[:10]:
|
||||
print(f" 未被覆盖: {k}")
|
||||
for k in list(unexpected)[:10]:
|
||||
print(f" 未被使用: {k}")
|
||||
raise RuntimeError(
|
||||
f"权重与网络结构不匹配(缺失 {len(missing)},多余 {len(unexpected)})。"
|
||||
"继续推理会得到质量劣化的结果,故中止。请确认权重文件是否为官方 SDMatte / SDMatte_plus。"
|
||||
)
|
||||
|
||||
print(f"[Ruinode-SDMatte] 权重与网络完全匹配({len(state_dict)} 个张量)")
|
||||
|
||||
# 1024 分辨率下最浅一层的自注意力是 16384x16384,一次性算完峰值约 15.5GB。
|
||||
# 分片逐块计算同一批注意力,实测显存降到 9.1GB 且更快(省下的搬运多于分片开销),
|
||||
# 数值仅因浮点累加次序不同产生 ~1e-6 的偏差,肉眼不可见。
|
||||
if attention_slicing:
|
||||
try:
|
||||
from diffusers.models.attention_processor import SlicedAttnProcessor
|
||||
|
||||
model.unet.set_attn_processor(SlicedAttnProcessor(slice_size=1))
|
||||
print("[Ruinode-SDMatte] 已启用注意力分片(显存约降 40%)")
|
||||
except Exception as e:
|
||||
print(f"[Ruinode-SDMatte] 注意力分片启用失败,按不分片继续:{e}")
|
||||
|
||||
model.eval()
|
||||
model.to(device=device, dtype=dtype)
|
||||
return model
|
||||
|
||||
|
||||
class RuiSDMatteLoader:
|
||||
"""加载 SDMatte 权重。支持官方 .pth(12.1GB)与社区转换的 .safetensors(5.19GB)。"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"ckpt_name": (_list_checkpoints(), {
|
||||
"tooltip": "放在 models/SDMatte 下的权重。\n"
|
||||
"官方 SDMatte_plus.pth 与社区 SDMatte_plus.safetensors 的模型权重完全等价,\n"
|
||||
"pth 多出的约 6GB 是训练用的优化器状态,推理不参与。"
|
||||
}),
|
||||
"precision": (["fp32", "fp16"], {
|
||||
"default": "fp32",
|
||||
"tooltip": "官方测试配置为 fp32(amp.enabled=False)。\n"
|
||||
"fp16 省显存但 SD 2.1 的 VAE 在半精度下容易溢出,可能出现黑图或噪点。"
|
||||
}),
|
||||
"device": (["auto", "cpu"], {"default": "auto"}),
|
||||
"attention_slicing": ("BOOLEAN", {
|
||||
"default": True,
|
||||
"tooltip": "分片计算注意力。1024 分辨率下显存峰值从约 15.5GB 降到 9.1GB,\n"
|
||||
"实测速度反而略快,输出差异在 1e-6 量级、肉眼不可见。\n"
|
||||
"显存充裕且想严格对齐官方数值时可关闭。"
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SDMATTE_MODEL",)
|
||||
RETURN_NAMES = ("sdmatte_model",)
|
||||
FUNCTION = "load"
|
||||
CATEGORY = "Ruinode/SDMatte"
|
||||
|
||||
def load(self, ckpt_name, precision, device, attention_slicing=True):
|
||||
import comfy.model_management
|
||||
|
||||
ckpt_path = _resolve_ckpt(ckpt_name)
|
||||
dev = torch.device("cpu") if device == "cpu" else comfy.model_management.get_torch_device()
|
||||
dtype = torch.float32 if precision == "fp32" else torch.float16
|
||||
|
||||
key = (ckpt_path, str(dtype), str(dev), bool(attention_slicing))
|
||||
cached = _MODEL_CACHE.get(key)
|
||||
if cached is not None:
|
||||
return (cached,)
|
||||
|
||||
_MODEL_CACHE.clear() # 单份 5GB 起步,不做多份缓存
|
||||
model = _build_model(ckpt_path, dtype, dev, attention_slicing)
|
||||
_MODEL_CACHE[key] = model
|
||||
return (model,)
|
||||
|
||||
|
||||
class RuiSDMatte:
|
||||
"""用视觉提示(框/掩码/点)驱动 SDMatte,输出精细 alpha。"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"sdmatte_model": ("SDMATTE_MODEL",),
|
||||
"image": ("IMAGE",),
|
||||
"mask": ("MASK", {
|
||||
"tooltip": "指示要抠哪个目标的提示掩码,不必精确,粗略覆盖主体即可。"
|
||||
}),
|
||||
"prompt_type": (["bbox_mask", "mask", "point_mask", "auto_mask"], {
|
||||
"default": "bbox_mask",
|
||||
"tooltip": "视觉提示类型。\n"
|
||||
"bbox_mask:取掩码外接框作为提示,官方测试脚本的默认路径,通常最稳;\n"
|
||||
"mask:直接用掩码本身,适合已有较准的粗分割;\n"
|
||||
"point_mask:在掩码内随机取 10 个点;\n"
|
||||
"auto_mask:不给定位信息,全图自动,画面只有单一主体时可用。"
|
||||
}),
|
||||
"inference_size": ([512, 640, 768, 896, 1024, 1152, 1280], {
|
||||
"default": 1024,
|
||||
"tooltip": "官方测试固定用 1024,降低会明显损失边缘细节。"
|
||||
}),
|
||||
"is_transparent": ("BOOLEAN", {
|
||||
"default": False,
|
||||
"tooltip": "目标是否为玻璃、纱、烟雾等透明/半透明物体。\n"
|
||||
"该开关会切换模型的不透明度嵌入分支,抠透明物时务必打开。"
|
||||
}),
|
||||
},
|
||||
"optional": {
|
||||
"caption": ("STRING", {
|
||||
"default": "", "multiline": False,
|
||||
"tooltip": "可选的文本描述(对应 RefMatte 的表达式)。留空即为官方测试时的默认行为。"
|
||||
}),
|
||||
"point_radius": ("INT", {
|
||||
"default": 35, "min": 5, "max": 100,
|
||||
"tooltip": "仅 point_mask 生效。官方测试期取 35(训练 radius 25 + 10)。"
|
||||
}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFF}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK", "IMAGE")
|
||||
RETURN_NAMES = ("alpha", "cutout")
|
||||
FUNCTION = "apply"
|
||||
CATEGORY = "Ruinode/SDMatte"
|
||||
|
||||
def apply(self, sdmatte_model, image, mask, prompt_type, inference_size,
|
||||
is_transparent, caption="", point_radius=35, seed=0):
|
||||
from .sdmatte import prompts as P
|
||||
|
||||
model = sdmatte_model
|
||||
device = model.device
|
||||
dtype = next(model.unet.parameters()).dtype
|
||||
size = int(inference_size)
|
||||
|
||||
B, H, W, _ = image.shape
|
||||
|
||||
# 掩码可能与图像批次数不一致,按 ComfyUI 惯例广播
|
||||
if mask.dim() == 2:
|
||||
mask = mask.unsqueeze(0)
|
||||
if mask.shape[0] != B:
|
||||
mask = mask[:1].repeat(B, 1, 1)
|
||||
|
||||
# 提示类型只是 forward 里的一个分支选择,切换它无需重建模型
|
||||
model.aux_input = prompt_type
|
||||
|
||||
images_t, aux_t, coords_t = [], [], []
|
||||
coor_name = None
|
||||
|
||||
for b in range(B):
|
||||
img_np = image[b].detach().cpu().float().numpy() # [H,W,3] in [0,1]
|
||||
msk_np = mask[b].detach().cpu().float().numpy() # [H,W] in [0,1]
|
||||
|
||||
img_r = P.resize_image(img_np, size)
|
||||
# 官方 Resize 对 alpha 用双线性;GenMask 的既有掩码分支用最近邻
|
||||
interp = cv2.INTER_NEAREST if prompt_type == "mask" else cv2.INTER_LINEAR
|
||||
msk_r = cv2.resize(msk_np, (size, size), interpolation=interp)
|
||||
msk_r = np.clip(msk_r, 0.0, 1.0)
|
||||
|
||||
aux_np, coords_np = P.build_prompt(
|
||||
msk_r, prompt_type, point_radius=point_radius, seed=seed + b
|
||||
)
|
||||
|
||||
images_t.append(torch.from_numpy(P.normalize(img_r)).permute(2, 0, 1))
|
||||
aux_t.append(torch.from_numpy(P.normalize(aux_np)).unsqueeze(0))
|
||||
coords_t.append(torch.from_numpy(coords_np))
|
||||
|
||||
from .sdmatte.meta_arch import AUX_INPUT_DIT
|
||||
coor_name = AUX_INPUT_DIT[prompt_type]
|
||||
|
||||
data = {
|
||||
"image": torch.stack(images_t).to(device=device, dtype=dtype),
|
||||
prompt_type: torch.stack(aux_t).to(device=device, dtype=dtype),
|
||||
coor_name: torch.stack(coords_t).to(device=device, dtype=dtype),
|
||||
"is_trans": torch.tensor([1 if is_transparent else 0] * B, dtype=torch.long),
|
||||
"caption": [caption] * B,
|
||||
}
|
||||
|
||||
with torch.no_grad():
|
||||
pred = model(data) # [B,1,size,size] in [0,1]
|
||||
|
||||
pred = pred.detach().float().cpu()
|
||||
|
||||
# 缩放回原图尺寸。官方 inference.py 用 cv2 双线性,并会量化到 uint8;
|
||||
# 这里保留浮点,避免白白丢掉 8bit 之外的过渡信息。
|
||||
alphas = []
|
||||
for b in range(B):
|
||||
a = pred[b, 0].numpy()
|
||||
a = cv2.resize(a, (W, H), interpolation=cv2.INTER_LINEAR)
|
||||
alphas.append(torch.from_numpy(np.clip(a, 0.0, 1.0)))
|
||||
alpha = torch.stack(alphas) # [B,H,W]
|
||||
|
||||
cutout = image.detach().cpu().float() * alpha.unsqueeze(-1)
|
||||
|
||||
return (alpha, cutout)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"RuiSDMatteLoader": RuiSDMatteLoader,
|
||||
"RuiSDMatte": RuiSDMatte,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"RuiSDMatteLoader": "SDMatte 加载器",
|
||||
"RuiSDMatte": "SDMatte 精细抠图",
|
||||
}
|
||||
Reference in New Issue
Block a user