feat: 新增满屏文字水印与 FeyNobg 抠图两个节点

满屏文字水印(watermark_node.py):
- 文案/字体/字号/旋转角度/密度/字间距/透明度/颜色 八项可调
- 交错网格平铺后整体旋转再从中心裁切,任意角度下四角均无空白

FeyNobg 抠图(feynobg_node.py + feynobg/):
- 内嵌 nobg 推理子集(Apache-2.0),全自动去背景,输出 alpha 与去背景图
- 重写预处理去掉对 transformers>=5.4 的依赖,4.x 环境可直接使用
- 修复权重键名与 transformers 4.x 的 SwinBackbone 命名不兼容:
  不处理时 958 个参数仅 405 个对得上,backbone 形同随机初始化,
  模型不报错但 alpha 几乎全黑;现按环境自动重映射并严格校验,
  除确定性 buffer 外任何失配都直接中止

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
rui40000
2026-07-29 09:26:05 +08:00
co-authored by Claude Opus 4.8
parent 9dd4fbd890
commit 1bd90c71b6
11 changed files with 1540 additions and 0 deletions
+73
View File
@@ -16,6 +16,7 @@ Rui-Node🐶 是一个功能丰富的 ComfyUI 节点集合,提供图像处理
- [颜色匹配器 / Color Matcher](#13-颜色匹配器--color-matcher)
- [素材拆分 / Sprite Splitter](#14-素材拆分--sprite-splitter)
- [素材拆分(带透明通道) / Sprite Splitter RGBA](#15-素材拆分带透明通道--sprite-splitter-rgba)
- [满屏文字水印 / Full-Screen Text Watermark](#21-满屏文字水印--full-screen-text-watermark)
### 📁 文件存储与加载类
- [按路径加载图像 / Load Image By Path](#3-按路径加载图像--load-image-by-path)
@@ -25,6 +26,7 @@ Rui-Node🐶 是一个功能丰富的 ComfyUI 节点集合,提供图像处理
- [千问编辑图像生成 / Qwen Edit Image Generation](#4-千问编辑图像生成--qwen-edit-image-generation)
- [SDMatte 精细抠图 / SDMatte Interactive Matting](#17-sdmatte-精细抠图--sdmatte-interactive-matting)
- [ZenMux API 连接 / ZenMux API Connector](#18-zenmux-api-连接--zenmux-api-connector)
- [FeyNobg 抠图 / FeyNobg Matting](#22-feynobg-抠图--feynobg-matting)
### 📝 文本处理类
- [镜头分词器 / Shot Splitter](#5-镜头分词器--shot-splitter)
@@ -816,6 +818,77 @@ WAS Node Suite 的「Text Multiline」会把 `#` 开头的行**当注释删除**
---
### 21. 满屏文字水印 / Full-Screen Text Watermark
**分类**: `Rui-Node🐶/图像调节🎨`
**功能描述**:
给输入图像铺满一层平铺的文字水印,常用于版权标注、样图防盗、批量打标。文字按**交错网格**平铺,整体旋转后从中心裁切与原图等大的区域,因此任意旋转角度下四角也都被水印覆盖,不留空白。支持中英文混排与多行文案(用换行分隔),逐字符绘制以支持字间距。批量图像逐张处理。
**输入参数**:
- `image` (IMAGE): 输入图像
- `text` (STRING, 多行): 水印文案,支持换行分隔的多行文本
- `font` (选择): 字体,列表来自 `Ruinode/font` 目录(ttf/otf/ttc,放入后刷新页面即出现,与 Markdown 节点共用同一套字体扫描)
- `font_size` (INT): 文字大小(像素),默认 48,范围 8~500
- `angle` (FLOAT): 水印整体旋转角度(度),默认 30,范围 -180~180
- `density` (FLOAT): 水印密度,综合控制**行间距与同行水印之间的间距**,值越大越密,默认 1.0,范围 0.1~5.0
- `letter_spacing` (INT): 字间距,单条文案内相邻字符的额外间距(像素,可为负),默认 0,范围 -20~200
- `opacity` (FLOAT): 水印透明程度,0=完全透明(原样返回),100=完全不透明,默认 35,范围 0~100
- `color` (STRING): 水印文字颜色,支持 `#RRGGBB` / `#RGB` / `"r,g,b"` / 常见英文色名(white、red、yellow…),默认 `#FFFFFF`
**输出**:
- `image` (IMAGE): 叠加水印后的图像
**使用场景**:
- 给出图 / 样片加满屏防盗水印
- 批量素材统一打上版权或"仅供参考"标注
---
### 22. FeyNobg 抠图 / FeyNobg Matting
**分类**: `Rui-Node🐶/抠图✂️`
**功能描述**:
全自动去背景抠图,**不需要任何提示**,输入图像直接输出 alpha。模型为 feyn 开源的 [FeyNobg](https://huggingface.co/feyninc/FeyNobg)(Apache-2.0),在 BiRefNet(CAAI AIR 2024)基础上扩展:Swin-Large 主干 + 梯度注意力 / 图像块注入 / 多尺度输入三项增强,原生 1024×1024 推理,权重约 1.05GB。
**与 [SDMatte 精细抠图](#17-sdmatte-精细抠图--sdmatte-interactive-matting) 的分工**:
- **FeyNobg**:全自动、一步出图、速度快,适合批量去背景(画面主体明确时首选)
- **SDMatte**:需要框/掩码提示指定目标,适合画面里有多个主体、要精确抠其中一个
**模型准备**:
首次运行会自动从 HuggingFace 下载到 `ComfyUI/models/nobg/FeyNobg`(约 1.05GB)。也可手动下载 `config.json`、`preprocessor_config.json`、`model.safetensors` 放入该目录。
**输入参数**:
- `image` (IMAGE): 输入图像
- `model_name` (选择): `models/nobg` 下的模型目录,未找到时自动下载
- `resolution` (选择): 推理分辨率,默认 **1024**(模型原生训练分辨率)。调低省显存但边缘变粗;调高不一定更好,可能出现结构断裂
- `precision` (选择): `fp32`(默认)/ `fp16`。**实测两者输出一致**(同图 alpha 均值均为 0.657),fp16 显存减半且明显更快,推荐优先用 fp16
- `device` (选择): `auto` / `cpu`
- `invert_mask` (BOOLEAN, 可选): 反转 alpha,默认前景为白
**输出**:
- `alpha` (MASK): 抠图 alpha,值域 [0,1]
- `cutout` (IMAGE): 去背景图(黑底)。需要透明 PNG 时,把 `alpha` 接到 `JoinImageWithAlpha` 一类节点
**实测数据**(1139×1280 人物插画,RTX 显卡):
| 配置 | 耗时 | 前景占比 |
|:-----|-----:|--------:|
| fp32 @1024 | 12.4s(含首次加载) | 0.659 |
| fp16 @1024 | 2.4s | 0.659 |
| fp32 @768 | 2.2s | 0.656 |
发丝、飘带、细链条等高频细节均能完整分离,边缘为自然的半透明过渡而非硬边。
**实现说明**(两个坑,都已在节点内处理):
1. **预处理依赖**:上游 `nobg` 的预处理模块继承 `transformers>=5.4` 的 `TorchvisionBackend`,而 ComfyUI 常见环境仍是 transformers 4.x,直接引入会报 `No module named 'transformers.image_processing_backends'`。本节点内嵌了 nobg 推理子集(`feynobg/`)并**重写了预处理**,数值规格与官方逐项对齐(1024 双线性抗锯齿缩放 + ImageNet 标准化;后处理先 sigmoid 再缩放),**无需升级 transformers**。同时绕开了上游 `AutoModel` 里会联网查 tags 的 `model_info()`,保证离线可用。
2. **权重键名不兼容(更隐蔽)**:FeyNobg 的权重用 transformers 5.x 导出,其 `SwinBackbone` 的模块命名与 4.x 不同(`bb.swin.*` 多一层、attention 从 `self.query/key/value` 重构为 `q/k/v_proj`、前馈层 `mlp.fc1/fc2` 对应 `intermediate.dense`/`output.dense`)。若不处理,958 个参数只有 405 个能对上,**整个 backbone 形同随机初始化——模型照样跑完不报错,但输出的 alpha 几乎全黑**(实测 max 0.02、mean 0.000)。节点内做了键名重映射(按环境自动判断是否需要),并**严格校验**:除确定性 buffer `relative_position_index` 与 backbone 末端未使用的 `bb.layernorm` 外,任何缺失/多余都直接报错中止,绝不接受静默劣化的结果。
---
## 🐕 关于 Rui-Node🐶
Rui-Node🐶 致力于为 ComfyUI 用户提供实用、高效的节点工具集。🐶 是我们的项目标志,代表着忠诚、友好和可靠。
+20
View File
@@ -61,6 +61,22 @@ except Exception as _e:
# 新增:多行文本框(原样输出,规避 WAS Text Multiline 吞 "#" 行的问题)
from .text_box_node import NODE_CLASS_MAPPINGS as TEXTBOX_NODE_CLASS_MAPPINGS
from .text_box_node import NODE_DISPLAY_NAME_MAPPINGS as TEXTBOX_NODE_DISPLAY_NAME_MAPPINGS
# 新增:FeyNobg 全自动抠图节点(BiRefNet 系,内嵌 nobg 推理子集)
try:
from .feynobg_node import NODE_CLASS_MAPPINGS as FEYNOBG_NODE_CLASS_MAPPINGS
from .feynobg_node import NODE_DISPLAY_NAME_MAPPINGS as FEYNOBG_NODE_DISPLAY_NAME_MAPPINGS
except Exception as _e:
print(f"[Ruinode] FeyNobg 抠图 节点未加载:{_e}")
FEYNOBG_NODE_CLASS_MAPPINGS = {}
FEYNOBG_NODE_DISPLAY_NAME_MAPPINGS = {}
# 新增:满屏文字水印节点(平铺文字 + 旋转/密度/透明度/颜色可调)
try:
from .watermark_node import NODE_CLASS_MAPPINGS as WATERMARK_NODE_CLASS_MAPPINGS
from .watermark_node import NODE_DISPLAY_NAME_MAPPINGS as WATERMARK_NODE_DISPLAY_NAME_MAPPINGS
except Exception as _e:
print(f"[Ruinode] 满屏文字水印 节点未加载:{_e}")
WATERMARK_NODE_CLASS_MAPPINGS = {}
WATERMARK_NODE_DISPLAY_NAME_MAPPINGS = {}
# 合并节点映射字典
NODE_CLASS_MAPPINGS = {}
@@ -84,6 +100,8 @@ NODE_CLASS_MAPPINGS.update(SDMATTE_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(ZENMUX_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MDIMG_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(TEXTBOX_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(WATERMARK_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FEYNOBG_NODE_CLASS_MAPPINGS)
# 合并节点显示名称映射
NODE_DISPLAY_NAME_MAPPINGS = {}
@@ -107,5 +125,7 @@ NODE_DISPLAY_NAME_MAPPINGS.update(SDMATTE_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(ZENMUX_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MDIMG_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(TEXTBOX_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(WATERMARK_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FEYNOBG_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+22
View File
@@ -0,0 +1,22 @@
# -*- coding: utf-8 -*-
"""
FeyNobg 推理子集(内嵌自 https://github.com/feyninc/nobg ,Apache-2.0)
=====================================================================
只保留推理必需的部分,供 Ruinode 的 FeyNobg 抠图节点使用:
- modeling_birefnet.py / loss.py / mixin.py / utils.py —— 与上游逐字节一致
- image_processing_birefnet.py —— 已重写,去掉对 transformers>=5.4 的依赖
(上游继承 TorchvisionBackend,在 transformers 4.x 上直接 ImportError)
上游的 auto.py 未内嵌:其 AutoModel.from_pretrained 会调 model_info() 联网查
仓库 tags,而 ComfyUI 场景下模型是本地目录,联网检查既没必要又会在离线时失败。
节点改为直接调用 BiRefNet.from_pretrained(本地目录)。
"""
# ruff: noqa: F401
from .birefnet.image_processing_birefnet import BiRefNetImageProcessor
from .birefnet.modeling_birefnet import BiRefNet, BiRefNetConfig
# 对应上游 nobg 版本
__version__ = "0.2.4"
__all__ = ["BiRefNet", "BiRefNetConfig", "BiRefNetImageProcessor"]
View File
@@ -0,0 +1,140 @@
# -*- coding: utf-8 -*-
"""
BiRefNet 预处理 / 后处理(Ruinode 内嵌版)
==========================================
官方 nobg 的同名文件继承 transformers>=5.4 的 TorchvisionBackend,而 ComfyUI
常见环境仍是 transformers 4.x(本机 T8 为 4.49.0),直接引入会报:
ModuleNotFoundError: No module named 'transformers.image_processing_backends'
升级 transformers 到 5.x 极可能连累其他自定义节点,因此这里用纯 torch/PIL
复刻官方的同一套数值规格,去掉对 transformers 的依赖:
- 预处理:转 RGB → 双线性缩放到 size(抗锯齿)→ 归一到 [0,1] → ImageNet 标准化
- 后处理:logits.sigmoid() → 双线性缩放回原尺寸
(**先 sigmoid 再缩放**,与官方 post_process_alpha_matting 及其 eval 脚本一致,
顺序颠倒会让边缘过渡发生偏移)
- cutout:把 alpha 写入原图的 alpha 通道,输出 RGBA
除去掉 transformers 依赖外,逐项对齐官方实现;训练用的 segmentation_maps /
labels 分支推理不需要,未移植。
"""
import json
import os
import torch
import torch.nn.functional as F
# 与 transformers.image_utils 的 IMAGENET_DEFAULT_MEAN / STD 相同,
# 也与 FeyNobg 的 preprocessor_config.json 一致
IMAGENET_DEFAULT_MEAN = (0.485, 0.456, 0.406)
IMAGENET_DEFAULT_STD = (0.229, 0.224, 0.225)
class BiRefNetImageProcessor:
"""BiRefNet 图像处理器(预处理 + 后处理)。"""
def __init__(self, size=None, image_mean=None, image_std=None,
rescale_factor=1 / 255, **kwargs):
size = size or {"height": 1024, "width": 1024}
self.size = {
"height": int(size.get("height", 1024)),
"width": int(size.get("width", 1024)),
}
self.image_mean = tuple(image_mean or IMAGENET_DEFAULT_MEAN)
self.image_std = tuple(image_std or IMAGENET_DEFAULT_STD)
self.rescale_factor = rescale_factor
# ------------------------------------------------------------------ 构造
@classmethod
def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
"""从本地目录的 preprocessor_config.json 读取配置(不联网)。"""
cfg = {}
if os.path.isdir(pretrained_model_name_or_path):
p = os.path.join(pretrained_model_name_or_path,
"preprocessor_config.json")
if os.path.isfile(p):
with open(p, "r", encoding="utf-8") as f:
cfg = json.load(f)
cfg.update(kwargs)
return cls(
size=cfg.get("size"),
image_mean=cfg.get("image_mean"),
image_std=cfg.get("image_std"),
rescale_factor=cfg.get("rescale_factor", 1 / 255),
)
# ---------------------------------------------------------------- 预处理
def preprocess_tensor(self, images, device=None, dtype=torch.float32):
"""
ComfyUI 原生张量入口,避免 PIL 往返。
images: (B, H, W, 3) float,值域 [0,1](ComfyUI 的 IMAGE 约定)
返回: (B, 3, size_h, size_w),已 ImageNet 标准化
"""
x = images.permute(0, 3, 1, 2).contiguous().float() # BHWC -> BCHW
# 官方走 torchvision 的 resize,对张量默认开抗锯齿;缩小到 1024 时
# 是否抗锯齿对细边缘影响可见,这里保持一致
x = F.interpolate(
x, size=(self.size["height"], self.size["width"]),
mode="bilinear", align_corners=False, antialias=True,
)
mean = torch.tensor(self.image_mean, dtype=torch.float32).view(1, 3, 1, 1)
std = torch.tensor(self.image_std, dtype=torch.float32).view(1, 3, 1, 1)
x = (x - mean) / std
return x.to(device=device, dtype=dtype) if device is not None \
else x.to(dtype=dtype)
def __call__(self, images, return_tensors="pt", **kwargs):
"""PIL 入口,兼容官方 README 的用法:processor(image, return_tensors='pt')。"""
import numpy as np
from PIL import Image as PILImage
if isinstance(images, PILImage.Image):
images = [images]
arrs = []
for im in images:
im = im.convert("RGB")
arrs.append(np.asarray(im, dtype=np.float32) * self.rescale_factor)
batch = torch.from_numpy(np.stack(arrs)) # (B,H,W,3) in [0,1]
return {"pixel_values": self.preprocess_tensor(batch)}
# ---------------------------------------------------------------- 后处理
def post_process_alpha_matting(self, outputs, target_sizes=None):
"""
原始 logits -> 每图 alpha([0,1],形状 (H, W))。
outputs: 含 "logits" 的 dict 或带 .logits 的对象,logits 形状 (B,1,H,W)
target_sizes: [(h, w), ...],逐图缩放回原尺寸
"""
logits = outputs["logits"] if isinstance(outputs, dict) else outputs.logits
if target_sizes is not None and len(logits) != len(target_sizes):
raise ValueError(
f"给了 {len(target_sizes)} 个目标尺寸,但批次里有 {len(logits)} 张图"
)
# 官方注释明确:先 sigmoid 再缩放,与 eval/benchmark 脚本一致
probs = logits.sigmoid()
mattes = []
for idx in range(len(probs)):
alpha = probs[idx].unsqueeze(0) # (1,1,H,W)
if target_sizes is not None:
alpha = F.interpolate(
alpha, size=target_sizes[idx],
mode="bilinear", align_corners=False,
)
mattes.append(alpha[0, 0])
return mattes
@staticmethod
def cutout(image, alpha):
"""把 alpha 合成到原图,返回 RGBA。alpha 可为 (H,W) 张量或 PIL 图。"""
from PIL import Image as PILImage
if isinstance(alpha, torch.Tensor):
arr = (alpha.detach().clamp(0, 1) * 255).to(torch.uint8).cpu().numpy()
alpha = PILImage.fromarray(arr, mode="L")
if alpha.size != image.size:
alpha = alpha.resize(image.size, PILImage.Resampling.BILINEAR)
cutout = image.convert("RGBA")
cutout.putalpha(alpha)
return cutout
+617
View File
@@ -0,0 +1,617 @@
import logging
import os
import re
from dataclasses import dataclass, field, fields
from importlib.metadata import PackageNotFoundError, version
import torch
import torch.nn.functional as F
from torch import nn
from torchvision.ops import deform_conv2d
from transformers import SwinConfig
from transformers.models.swin.modeling_swin import SwinBackbone
from ..loss import birefnet_loss
from ..mixin import Revised_Mixin
from ..utils import model_card_template
logger = logging.getLogger(__name__)
try:
NOBG_VERSION = version("nobg")
except PackageNotFoundError: # running from source without an installed dist
NOBG_VERSION = "0.0.0+unknown"
BIREFNET_CITATION = """@article{zheng2024birefnet,
title={Bilateral Reference for High-Resolution Dichotomous Image Segmentation},
author={Zheng, Peng and Gao, Dehong and Fan, Deng-Ping and Liu, Li and Laaksonen, Jorma and Ouyang, Wanli and Sebe, Nicu},
journal={CAAI Artificial Intelligence Research},
volume={3},
pages={9150038},
year={2024},
url={https://arxiv.org/abs/2401.03407},
}"""
@dataclass
class BiRefNetConfig:
"""Configuration for BiRefNet (Bilateral Reference Network) with Swin backbone."""
image_size: int = 1024
patch_size: int = 4
embed_dim: int = 192
num_layers: int = 4
depths: list = field(default_factory=lambda: [2, 2, 18, 2])
num_heads: list = field(default_factory=lambda: [6, 12, 24, 48])
window_size: int = 12
mlp_ratio: float = 4.0
drop_path_rate: float = 0.2
dec_channels_inter: int = 64
use_multi_scale_input: bool = True
use_gradient_attention: bool = True
use_image_patch_injection: bool = True
nobg_version: str = NOBG_VERSION
def __post_init__(self):
if len(self.depths) != self.num_layers:
raise ValueError(
f"depths has {len(self.depths)} entries but num_layers={self.num_layers}"
)
if len(self.num_heads) != self.num_layers:
raise ValueError(
f"num_heads has {len(self.num_heads)} entries but num_layers={self.num_layers}"
)
class DeformableConv2d(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: int = 3,
stride: int = 1,
padding: int = 1,
bias: bool = False,
):
super().__init__()
ks = (kernel_size, kernel_size)
self.stride = (stride, stride)
self.padding = (padding, padding)
self.offset_conv = nn.Conv2d(
in_channels,
2 * ks[0] * ks[1],
kernel_size=ks,
stride=stride,
padding=padding,
bias=True,
)
nn.init.constant_(self.offset_conv.weight, 0.0)
assert self.offset_conv.bias is not None
nn.init.constant_(self.offset_conv.bias, 0.0)
self.modulator_conv = nn.Conv2d(
in_channels,
1 * ks[0] * ks[1],
kernel_size=ks,
stride=stride,
padding=padding,
bias=True,
)
nn.init.constant_(self.modulator_conv.weight, 0.0)
assert self.modulator_conv.bias is not None
nn.init.constant_(self.modulator_conv.bias, 0.0)
self.regular_conv = nn.Conv2d(
in_channels,
out_channels,
kernel_size=ks,
stride=stride,
padding=padding,
bias=bias,
)
def forward(self, x):
offset = self.offset_conv(x)
modulator = 2.0 * torch.sigmoid(self.modulator_conv(x))
x = deform_conv2d(
input=x,
offset=offset,
weight=self.regular_conv.weight,
bias=self.regular_conv.bias,
padding=self.padding,
mask=modulator,
stride=self.stride,
)
return x
class _ASPPModuleDeformable(nn.Module):
def __init__(self, in_channels, planes, kernel_size, padding):
super().__init__()
self.atrous_conv = DeformableConv2d(
in_channels,
planes,
kernel_size=kernel_size,
stride=1,
padding=padding,
bias=False,
)
self.bn = nn.BatchNorm2d(planes)
self.relu = nn.ReLU(inplace=True)
def forward(self, x):
x = self.atrous_conv(x)
x = self.bn(x)
return self.relu(x)
class ASPPDeformable(nn.Module):
def __init__(self, in_channels, out_channels=None):
super().__init__()
if out_channels is None:
out_channels = in_channels
inter_channels = 256
parallel_block_sizes = [1, 3, 7]
self.aspp1 = _ASPPModuleDeformable(in_channels, inter_channels, 1, padding=0)
self.aspp_deforms = nn.ModuleList(
[
_ASPPModuleDeformable(
in_channels, inter_channels, conv_size, padding=conv_size // 2
)
for conv_size in parallel_block_sizes
]
)
self.global_avg_pool = nn.Sequential(
nn.AdaptiveAvgPool2d((1, 1)),
nn.Conv2d(in_channels, inter_channels, 1, stride=1, bias=False),
nn.BatchNorm2d(inter_channels),
nn.ReLU(inplace=True),
)
self.conv1 = nn.Conv2d(
inter_channels * (2 + len(self.aspp_deforms)), out_channels, 1, bias=False
)
self.bn1 = nn.BatchNorm2d(out_channels)
self.relu = nn.ReLU(inplace=True)
self.dropout = nn.Dropout(0.5)
def forward(self, x):
x1 = self.aspp1(x)
x_aspp_deforms = [aspp_deform(x) for aspp_deform in self.aspp_deforms]
x5 = self.global_avg_pool(x)
x5 = F.interpolate(x5, size=x1.shape[2:], mode="bilinear", align_corners=True)
x = torch.cat((x1, *x_aspp_deforms, x5), dim=1)
x = self.conv1(x)
x = self.bn1(x)
x = self.relu(x)
return self.dropout(x)
class BasicDecBlk(nn.Module):
def __init__(self, in_channels=64, out_channels=64, inter_channels=64):
super().__init__()
self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, padding=1)
self.relu_in = nn.ReLU(inplace=True)
self.dec_att = ASPPDeformable(in_channels=inter_channels)
self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, padding=1)
self.bn_in = nn.BatchNorm2d(inter_channels)
self.bn_out = nn.BatchNorm2d(out_channels)
def forward(self, x):
x = self.conv_in(x)
x = self.bn_in(x)
x = self.relu_in(x)
x = self.dec_att(x)
x = self.conv_out(x)
x = self.bn_out(x)
return x
class BasicLatBlk(nn.Module):
def __init__(self, in_channels, out_channels):
super().__init__()
self.conv = nn.Conv2d(in_channels, out_channels, 1)
def forward(self, x):
return self.conv(x)
class SimpleConvs(nn.Module):
def __init__(self, in_channels: int, out_channels: int, inter_channels=64):
super().__init__()
self.conv1 = nn.Conv2d(in_channels, inter_channels, 3, 1, 1)
self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, 1)
def forward(self, x):
return self.conv_out(self.conv1(x))
def image2patches(image, patch_ref):
grid_h = image.shape[-2] // patch_ref.shape[-2]
grid_w = image.shape[-1] // patch_ref.shape[-1]
B, C, H, W = image.shape
h = H // grid_h
w = W // grid_w
x = image.view(B, C, grid_h, h, grid_w, w)
x = x.permute(0, 1, 2, 4, 3, 5).contiguous()
x = x.view(B, C * grid_h * grid_w, h, w)
return x
class Decoder(nn.Module):
"""Generic BiRefNet decoder over `num_layers` stages.
`channels` is deep-to-shallow: channels[0] is the (squeezed) deepest
feature width, channels[-1] the shallowest. Stage i of the loop consumes
channels[i] (+ optional image-patch injection) and produces channels[i+1],
except the last stage which produces channels[-1] // 2.
"""
def __init__(self, channels: list[int], config: BiRefNetConfig):
super().__init__()
inter = config.dec_channels_inter
n = config.num_layers
self.num_layers = n
self.use_gradient_attention = config.use_gradient_attention
self.use_image_patch_injection = config.use_image_patch_injection
# Injection output widths: stages 0 and 1 use channels[0] // 8; stage
# j >= 2 uses channels[j - 1] // 8; the final full-resolution injection
# uses channels[-1] // 8.
ipt_out = [channels[0] // 8] + [
channels[max(j - 1, 0)] // 8 for j in range(1, n)
]
ipt_out.append(channels[-1] // 8)
if config.use_image_patch_injection:
# image2patches gives 3 * grid^2 channels; the grid at stage j is
# patch_size * 2**(n - 1 - j), and 1 at full resolution.
ipt_in = [
3 * (config.patch_size * 2 ** (n - 1 - j)) ** 2 for j in range(n)
] + [3]
self.ipt_blks = nn.ModuleList(
[
SimpleConvs(ipt_in[j], ipt_out[j], inter_channels=inter)
for j in range(n + 1)
]
)
dec_in = [channels[i] + ipt_out[i] for i in range(n)]
dec_out = [channels[i + 1] for i in range(n - 1)] + [channels[-1] // 2]
self.decoder_blocks = nn.ModuleList(
[BasicDecBlk(dec_in[i], dec_out[i], inter) for i in range(n)]
)
self.lateral_blocks = nn.ModuleList(
[BasicLatBlk(channels[i + 1], channels[i + 1]) for i in range(n - 1)]
)
self.conv_ms_spvn = nn.ModuleList(
[nn.Conv2d(channels[i + 1], 1, 1, 1, 0) for i in range(n - 1)]
)
if config.use_gradient_attention:
_N = 16
self.gdt_convs = nn.ModuleList(
[
nn.Sequential(
nn.Conv2d(channels[i + 1], _N, 3, 1, 1),
nn.BatchNorm2d(_N),
nn.ReLU(inplace=True),
)
for i in range(n - 1)
]
)
self.gdt_convs_pred = nn.ModuleList(
[nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0)) for _ in range(n - 1)]
)
self.gdt_convs_attn = nn.ModuleList(
[nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0)) for _ in range(n - 1)]
)
self.conv_out1 = nn.Sequential(
nn.Conv2d(channels[-1] // 2 + channels[-1] // 8, 1, 1, 1, 0)
)
def _inject(self, image, p, blk, size):
patches = image2patches(image, patch_ref=p)
patches = F.interpolate(patches, size=size, mode="bilinear", align_corners=True)
return torch.cat((p, blk(patches)), 1)
def forward(self, features: list[torch.Tensor]) -> list[torch.Tensor]:
image = features[0]
feats = features[1:] # shallow -> deep, length num_layers
n = self.num_layers
p = feats[-1] # deepest (already squeezed)
outs = []
for i in range(n - 1):
if self.use_image_patch_injection:
p = self._inject(image, p, self.ipt_blks[i], p.shape[2:])
p = self.decoder_blocks[i](p)
if self.use_gradient_attention:
p_gdt = self.gdt_convs[i](p)
gdt_attn = self.gdt_convs_attn[i](p_gdt).sigmoid()
p = p * gdt_attn
outs.append(self.conv_ms_spvn[i](p))
skip = feats[n - 2 - i]
p = F.interpolate(
p, size=skip.shape[2:], mode="bilinear", align_corners=True
)
p = p + self.lateral_blocks[i](skip)
if self.use_image_patch_injection:
p = self._inject(image, p, self.ipt_blks[n - 1], p.shape[2:])
p = self.decoder_blocks[n - 1](p)
p = F.interpolate(p, size=image.shape[2:], mode="bilinear", align_corners=True)
if self.use_image_patch_injection:
p = self._inject(image, p, self.ipt_blks[n], image.shape[2:])
outs.append(self.conv_out1(p))
return outs
def _remap_legacy_state_dict(
state_dict: dict[str, torch.Tensor],
) -> dict[str, torch.Tensor]:
"""Remap the pre-0.2.0 custom-Swin key layout to the current layout.
Backbone: custom Swin keys -> transformers SwinBackbone keys (the fused
attn.qkv is split into thirds for q/k/v). Decoder: numbered attributes
(decoder_block4..1, ipt_blk5..1, ...) -> ModuleList indices.
"""
out: dict[str, torch.Tensor] = {}
for k, v in state_dict.items():
if "relative_position_index" in k:
continue # non-persistent buffer in the new layout
nk = k
if k.startswith("bb."):
rest = k[len("bb.") :]
if rest.startswith("patch_embed.proj."):
nk = (
"bb.swin.embeddings.patch_embeddings.projection."
+ rest.rsplit(".", 1)[-1]
)
elif rest.startswith("patch_embed.norm."):
nk = "bb.swin.embeddings.norm." + rest.rsplit(".", 1)[-1]
elif m := re.match(r"norm(\d+)\.(weight|bias)$", rest):
nk = f"bb.hidden_states_norms.stage{int(m.group(1)) + 1}.{m.group(2)}"
elif m := re.match(r"layers\.(\d+)\.blocks\.(\d+)\.(.*)$", rest):
pre = f"bb.swin.encoder.layers.{m.group(1)}.blocks.{m.group(2)}."
sub = m.group(3)
if m2 := re.match(r"attn\.qkv\.(weight|bias)$", sub):
q, kk, vv = v.chunk(3, dim=0)
out[pre + f"attention.q_proj.{m2.group(1)}"] = q
out[pre + f"attention.k_proj.{m2.group(1)}"] = kk
out[pre + f"attention.v_proj.{m2.group(1)}"] = vv
continue
elif sub.startswith("attn.proj."):
nk = pre + "attention.o_proj." + sub.rsplit(".", 1)[-1]
elif sub == "attn.relative_position_bias_table":
nk = (
pre
+ "attention.relative_position_bias.relative_position_bias_table"
)
elif sub.startswith("norm1."):
nk = pre + "layernorm_before." + sub.rsplit(".", 1)[-1]
elif sub.startswith("norm2."):
nk = pre + "layernorm_after." + sub.rsplit(".", 1)[-1]
else: # mlp.fc1 / mlp.fc2
nk = pre + sub
elif m := re.match(r"layers\.(\d+)\.downsample\.(.*)$", rest):
nk = f"bb.swin.encoder.layers.{m.group(1)}.downsample.{m.group(2)}"
elif k.startswith("decoder."):
rest = k[len("decoder.") :]
if m := re.match(r"decoder_block(\d)\.(.*)$", rest):
nk = f"decoder.decoder_blocks.{4 - int(m.group(1))}.{m.group(2)}"
elif m := re.match(r"lateral_block(\d)\.(.*)$", rest):
nk = f"decoder.lateral_blocks.{4 - int(m.group(1))}.{m.group(2)}"
elif m := re.match(r"ipt_blk(\d)\.(.*)$", rest):
nk = f"decoder.ipt_blks.{5 - int(m.group(1))}.{m.group(2)}"
elif m := re.match(r"conv_ms_spvn_(\d)\.(.*)$", rest):
nk = f"decoder.conv_ms_spvn.{4 - int(m.group(1))}.{m.group(2)}"
elif m := re.match(r"gdt_convs_pred_(\d)\.(.*)$", rest):
nk = f"decoder.gdt_convs_pred.{4 - int(m.group(1))}.{m.group(2)}"
elif m := re.match(r"gdt_convs_attn_(\d)\.(.*)$", rest):
nk = f"decoder.gdt_convs_attn.{4 - int(m.group(1))}.{m.group(2)}"
elif m := re.match(r"gdt_convs_(\d)\.(.*)$", rest):
nk = f"decoder.gdt_convs.{4 - int(m.group(1))}.{m.group(2)}"
out[nk] = v
return out
class BiRefNet(
nn.Module,
Revised_Mixin,
library_name="nobg",
repo_url="https://github.com/feyninc/nobg",
paper_url="https://arxiv.org/abs/2401.03407",
license="apache-2.0",
tags=["nobg", "birefnet"],
model_card_template=model_card_template(
class_name="BiRefNet",
default_repo="nobg/birefnet",
citation=BIREFNET_CITATION,
),
):
"""Bilateral Reference Network for high-resolution dichotomous image segmentation."""
def __init__(self, config: BiRefNetConfig | None = None):
super().__init__()
self.config = config or BiRefNetConfig()
swin_config = SwinConfig(
image_size=self.config.image_size,
patch_size=self.config.patch_size,
embed_dim=self.config.embed_dim,
depths=self.config.depths,
num_heads=self.config.num_heads,
window_size=self.config.window_size,
mlp_ratio=self.config.mlp_ratio,
drop_path_rate=self.config.drop_path_rate,
out_features=[f"stage{i + 1}" for i in range(self.config.num_layers)], # ty: ignore[unknown-argument]
)
self.bb = SwinBackbone(swin_config)
base_channels = [
self.config.embed_dim * (2**i) for i in range(self.config.num_layers)
]
if self.config.use_multi_scale_input:
channels = [c * 2 for c in base_channels]
else:
channels = base_channels
# channels = [C1 .. Cn] shallow->deep; the decoder expects deep->shallow
dec_channels = list(reversed(channels))
# Squeeze module: deepest feature with all shallower features
# downsampled and concatenated onto it.
squeeze_in = channels[-1] + sum(channels[:-1])
self.squeeze_module = nn.Sequential(
BasicDecBlk(squeeze_in, dec_channels[0], self.config.dec_channels_inter)
)
self.decoder = Decoder(dec_channels, self.config)
# Loss used when `labels` is passed to forward(). Plain function
# attribute (not a submodule), so it never enters the state dict and
# can be hot-swapped: `model.criterion = my_loss`.
self.criterion = birefnet_loss
def forward(
self, pixel_values: torch.Tensor, labels: torch.Tensor | None = None
) -> dict[str, torch.Tensor | list[torch.Tensor]]:
x = pixel_values
feats = list(self.bb(x).feature_maps)
if self.config.use_multi_scale_input:
_, _, H, W = x.shape
x_half = F.interpolate(
x, size=(H // 2, W // 2), mode="bilinear", align_corners=True
)
feats_half = self.bb(x_half).feature_maps
feats = [
torch.cat(
[
f,
F.interpolate(
fh, size=f.shape[2:], mode="bilinear", align_corners=True
),
],
dim=1,
)
for f, fh in zip(feats, feats_half)
]
# Context aggregation: all shallower features downsampled onto the deepest
deepest = feats[-1]
context = torch.cat(
[
F.interpolate(
f, size=deepest.shape[2:], mode="bilinear", align_corners=True
)
for f in feats[:-1]
]
+ [deepest],
dim=1,
)
feats[-1] = self.squeeze_module(context)
scaled_preds = self.decoder([pixel_values, *feats])
logits = scaled_preds[-1]
if labels is not None:
loss = self.criterion(scaled_preds, labels)
return {
"loss": loss,
"logits": logits,
"intermediate_logits": scaled_preds[:-1],
}
return {"logits": logits, "intermediate_logits": scaled_preds[:-1]}
@classmethod
def from_origin(
cls,
origin: str | os.PathLike | nn.Module,
config: BiRefNetConfig | None = None,
*,
token: str | None = None,
**overrides,
) -> "BiRefNet":
"""Build a (possibly re-parameterized) BiRefNet from a previous model.
`origin` is a Hub repo id, a local directory containing
`config.json` + `model.safetensors`, or an `nn.Module` instance.
The origin's config fields are merged with `overrides` (or replaced
entirely by `config`), the new model is constructed, and every origin
weight whose (remapped) key and shape match is injected; anything new
or reshaped keeps its fresh initialization. Understands both the
current key layout and the pre-0.2.0 custom-Swin layout.
"""
if isinstance(origin, nn.Module):
state_dict = origin.state_dict()
origin_config = getattr(origin, "config", None)
origin_fields = (
{f.name: getattr(origin_config, f.name) for f in fields(origin_config)}
if origin_config is not None
and hasattr(origin_config, "__dataclass_fields__")
else {}
)
else:
import json
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
origin = str(origin)
if os.path.isdir(origin):
config_path = os.path.join(origin, "config.json")
weights_path = os.path.join(origin, "model.safetensors")
else:
config_path = hf_hub_download(origin, "config.json", token=token)
weights_path = hf_hub_download(origin, "model.safetensors", token=token)
with open(config_path) as f:
origin_fields = json.load(f)
state_dict = load_file(weights_path)
if config is None:
known = {f.name for f in fields(BiRefNetConfig)}
merged = {k: v for k, v in origin_fields.items() if k in known}
merged.update(overrides)
merged["nobg_version"] = overrides.get("nobg_version", NOBG_VERSION)
if "depths" in merged and "num_layers" not in overrides:
merged["num_layers"] = len(merged["depths"])
config = BiRefNetConfig(**merged)
elif overrides:
raise ValueError(
"pass either an explicit `config` or field overrides, not both"
)
model = cls(config)
is_legacy = any(k.startswith("bb.patch_embed.") for k in state_dict)
if is_legacy:
state_dict = _remap_legacy_state_dict(state_dict)
target = model.state_dict()
compatible = {}
skipped_shape = []
for k, v in state_dict.items():
if k in target and target[k].shape == v.shape:
compatible[k] = v
elif k in target:
skipped_shape.append(k)
missing = [k for k in target if k not in compatible]
model.load_state_dict(compatible, strict=False)
logger.info(
"from_origin: injected %d/%d tensors (%d shape-mismatched, "
"%d fresh-initialized)%s",
len(compatible),
len(target),
len(skipped_shape),
len(missing),
" [legacy layout remapped]" if is_legacy else "",
)
return model
+37
View File
@@ -0,0 +1,37 @@
import torch
import torch.nn.functional as F
def iou_loss(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
inter = (pred * target).sum(dim=(2, 3))
union = pred.sum(dim=(2, 3)) + target.sum(dim=(2, 3)) - inter
return (1 - inter / (union + 1e-8)).mean()
def ssim_loss(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
C1 = 0.01**2
C2 = 0.03**2
mu_x = F.avg_pool2d(pred, 3, 1, 1)
mu_y = F.avg_pool2d(target, 3, 1, 1)
sigma_x = F.avg_pool2d(pred * pred, 3, 1, 1) - mu_x * mu_x
sigma_y = F.avg_pool2d(target * target, 3, 1, 1) - mu_y * mu_y
sigma_xy = F.avg_pool2d(pred * target, 3, 1, 1) - mu_x * mu_y
ssim_map = ((2 * mu_x * mu_y + C1) * (2 * sigma_xy + C2)) / (
(mu_x * mu_x + mu_y * mu_y + C1) * (sigma_x + sigma_y + C2)
)
return torch.clamp((1 - ssim_map) / 2, 0, 1).mean()
def birefnet_loss(scaled_preds: list[torch.Tensor], gt: torch.Tensor) -> torch.Tensor:
"""Multi-scale pixel loss matching BiRefNet training: weighted BCE + IoU + SSIM."""
loss = torch.tensor(0.0, device=gt.device)
for pred in scaled_preds:
if pred.shape[2:] != gt.shape[2:]:
pred = F.interpolate(
pred, size=gt.shape[2:], mode="bilinear", align_corners=True
)
pred_sig = pred.sigmoid()
loss = loss + 30 * F.binary_cross_entropy_with_logits(pred, gt)
loss = loss + 0.5 * iou_loss(pred_sig, gt)
loss = loss + 10 * ssim_loss(pred_sig, gt)
return loss
+48
View File
@@ -0,0 +1,48 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any
from huggingface_hub import PyTorchModelHubMixin, whoami
from .utils import set_doc
if TYPE_CHECKING:
from _typeshed import DataclassInstance
class Revised_Mixin(PyTorchModelHubMixin):
@set_doc(PyTorchModelHubMixin.push_to_hub.__doc__)
def push_to_hub(
self,
repo_id: str,
*,
config: dict | DataclassInstance | None = None,
commit_message: str = "Push model using huggingface_hub.",
private: bool | None = None,
token: str | None = None,
branch: str | None = None,
create_pr: bool | None = None,
allow_patterns: list[str] | str | None = None,
ignore_patterns: list[str] | str | None = None,
delete_patterns: list[str] | str | None = None,
model_card_kwargs: dict[str, Any] | None = None,
) -> str:
if model_card_kwargs is None:
model_card_kwargs = {}
if "/" not in repo_id:
username = whoami()["name"]
repo_id = f"{username}/{repo_id}"
model_card_kwargs["repo_id"] = repo_id
return super().push_to_hub(
repo_id,
config=config,
commit_message=commit_message,
private=private,
token=token,
branch=branch,
create_pr=create_pr,
allow_patterns=allow_patterns,
ignore_patterns=ignore_patterns,
delete_patterns=delete_patterns,
model_card_kwargs=model_card_kwargs,
)
+64
View File
@@ -0,0 +1,64 @@
from functools import wraps
# update docs for specific functions
def set_doc(doc):
def decorator(func):
@wraps(func)
def wrapper(*args, **kwargs):
return func(*args, **kwargs)
wrapper.__doc__ = doc
return wrapper
return decorator
# general template for models
def model_card_template(
*, class_name: str, default_repo: str, citation: str | None = None
) -> str:
citation_section = (
f"""
## Citation
If you use this model, please cite:
```bibtex
{citation}
```
"""
if citation
else ""
)
return f"""---
{{{{ card_data }}}}
---
<p align="center">
<img src="https://usefeyn.com/feyn/feyn_mark.svg"/>
</p>
This model has been pushed to the Hub using the [PytorchModelHubMixin](https://huggingface.co/docs/huggingface_hub/package_reference/mixins#huggingface_hub.PyTorchModelHubMixin) integration.
Library: [nobg]({{{{repo_url}}}})
## how to load
```
pip install nobg
```
use the AutoModel class
```python
from nobg import AutoModel
model = AutoModel.from_pretrained("{{{{ repo_id | default("{default_repo}", true) }}}}")
```
or you can use the model class directly
```python
from nobg import {class_name}
model = {class_name}.from_pretrained("{{{{ repo_id | default("{default_repo}", true) }}}}")
```
{citation_section}
## Contributions
Any contributions are welcome at https://github.com/feyninc/nobg
"""
+301
View File
@@ -0,0 +1,301 @@
# -*- coding: utf-8 -*-
"""
FeyNobg 抠图节点(Ruinode)
===========================
对应模型:https://huggingface.co/feyninc/FeyNobg (Apache-2.0)
上游库: https://github.com/feyninc/nobg
FeyNobg 是 feyn 在 BiRefNet(CAAI AIR 2024)基础上扩展的通用抠图模型:
Swin-Large 主干 + 三项自定义增强(梯度注意力 use_gradient_attention、
图像块注入 use_image_patch_injection、多尺度输入 use_multi_scale_input),
原生 1024×1024 推理,权重约 1.05GB。
与 SDMatte 的定位差异:
- FeyNobg 全自动,不需要任何提示,一步出 alpha,适合批量去背景;
- SDMatte 需要框/掩码提示来指定抠哪个目标,适合画面里有多个主体时精确取一个。
实现说明:上游 nobg 的预处理模块继承 transformers>=5.4 的 TorchvisionBackend,
在 ComfyUI 常见的 transformers 4.x 上会 ImportError,故 Ruinode 内嵌了推理子集
并重写了预处理(见 feynobg/birefnet/image_processing_birefnet.py),
数值规格与官方对齐,且不需要升级 transformers。
"""
import os
import numpy as np
import torch
import folder_paths
# ---------------------------------------------------------------- 模型目录注册
NOBG_DIR = os.path.join(folder_paths.models_dir, "nobg")
os.makedirs(NOBG_DIR, exist_ok=True)
HF_REPO = "feyninc/FeyNobg"
DEFAULT_MODEL_DIRNAME = "FeyNobg"
_MODEL_CACHE = {}
_NO_MODEL_HINT = "未找到模型,将自动下载"
def _is_model_dir(p):
"""一个模型快照目录需同时具备 config.json 与权重文件。"""
if not os.path.isdir(p):
return False
if not os.path.isfile(os.path.join(p, "config.json")):
return False
return any(f.endswith((".safetensors", ".bin"))
for f in os.listdir(p))
def _list_models():
"""扫描 models/nobg 下的模型快照目录。"""
found = []
if os.path.isdir(NOBG_DIR):
for name in sorted(os.listdir(NOBG_DIR)):
if _is_model_dir(os.path.join(NOBG_DIR, name)):
found.append(name)
return found or [_NO_MODEL_HINT]
def _ensure_model(name):
"""返回模型目录;不存在时从 HuggingFace 拉取到 models/nobg/FeyNobg。"""
if name and name != _NO_MODEL_HINT:
path = os.path.join(NOBG_DIR, name)
if _is_model_dir(path):
return path
target = os.path.join(NOBG_DIR, DEFAULT_MODEL_DIRNAME)
if _is_model_dir(target):
return target
print(f"[Ruinode-FeyNobg] 本地未找到模型,开始从 HuggingFace 下载 {HF_REPO}"
f"(约 1.05GB)→ {target}")
try:
from huggingface_hub import snapshot_download
snapshot_download(
HF_REPO, local_dir=target,
allow_patterns=["*.json", "*.safetensors"],
)
except Exception as e:
raise RuntimeError(
f"自动下载失败:{e}\n"
f"请手动下载 https://huggingface.co/{HF_REPO} 的 "
f"config.json / preprocessor_config.json / model.safetensors,"
f"放到:{target}"
)
if not _is_model_dir(target):
raise RuntimeError(f"下载完成但目录不完整,请检查:{target}")
print(f"[Ruinode-FeyNobg] 下载完成:{target}")
return target
def _remap_swin_keys(state_dict):
"""
把 transformers 5.x 保存的 Swin 权重键名,改写成 4.x 的命名。
FeyNobg 的权重是用 transformers 5.x 导出的,而 ComfyUI 环境普遍还是 4.x,
两版的 SwinBackbone 模块命名不同(数学结构一致,纯改名):
bb.swin.X -> bb.X (5.x 多包一层 swin)
attention.{q,k,v}_proj -> attention.self.{query,key,value}
attention.o_proj -> attention.output.dense
attention.relative_position_bias
.relative_position_bias_table -> attention.self.relative_position_bias_table
mlp.fc1 / mlp.fc2 -> intermediate.dense / output.dense
不改名直接 load_state_dict(strict=False) 的后果:958 个参数里只有 405 个能对上,
整个 backbone 形同随机初始化 —— 模型照样能跑完,但输出的 alpha 几乎全黑
(实测 max 0.02、mean 0.000)。这类"不报错的错"最难排查,故此处显式处理。
"""
out = {}
for k, v in state_dict.items():
nk = k
if nk.startswith("bb.swin."):
nk = "bb." + nk[len("bb.swin."):]
nk = nk.replace(".attention.q_proj.", ".attention.self.query.")
nk = nk.replace(".attention.k_proj.", ".attention.self.key.")
nk = nk.replace(".attention.v_proj.", ".attention.self.value.")
nk = nk.replace(".attention.o_proj.", ".attention.output.dense.")
nk = nk.replace(".attention.relative_position_bias.relative_position_bias_table",
".attention.self.relative_position_bias_table")
nk = nk.replace(".mlp.fc1.", ".intermediate.dense.")
nk = nk.replace(".mlp.fc2.", ".output.dense.")
out[nk] = v
return out
def _load_weights_strict(model, ckpt_path):
"""加载权重并严格校验,宁可报错也不接受"能跑但结果劣化"的静默失配。"""
from safetensors.torch import load_file
sd = load_file(ckpt_path)
own = set(model.state_dict().keys())
direct = len(own & set(sd.keys()))
remapped = _remap_swin_keys(sd)
after = len(own & set(remapped.keys()))
# 权重与本机 transformers 命名不一致时才改写;若环境本就是 5.x,直接匹配更优
if after > direct:
print(f"[Ruinode-FeyNobg] 权重键名按 transformers 4.x 改写:"
f"{direct} -> {after} / {len(own)} 匹配")
sd = remapped
result = model.load_state_dict(sd, strict=False)
# relative_position_index 是由窗口大小推出的确定性 buffer,构造时已算好,不需要权重;
# bb.layernorm 是 backbone 末端的 norm,BiRefNet 只取中间各级特征图,用不到。
# 这两类之外的任何缺失/多余,都意味着结构对不上,必须叫停。
bad_missing = [k for k in result.missing_keys
if not k.endswith("relative_position_index")]
bad_unexpected = [k for k in result.unexpected_keys
if not k.startswith("bb.layernorm.")]
if bad_missing or bad_unexpected:
for k in bad_missing[:8]:
print(f" 未被覆盖: {k}")
for k in bad_unexpected[:8]:
print(f" 未被使用: {k}")
raise RuntimeError(
f"权重与网络结构不匹配(异常缺失 {len(bad_missing)},异常多余 "
f"{len(bad_unexpected)})。继续推理只会得到劣化结果,故中止。\n"
f"请确认权重为 feyninc/FeyNobg,且 transformers 版本受支持"
f"(已在 4.57 与 5.x 上验证)。"
)
print(f"[Ruinode-FeyNobg] 权重加载完成({len(own) - len(result.missing_keys)}"
f"/{len(own)} 个张量,另 {len(result.missing_keys)} 个为确定性 buffer)")
def _load_model(model_dir, dtype, device):
key = (model_dir, str(dtype), str(device))
cached = _MODEL_CACHE.get(key)
if cached is not None:
return cached
import dataclasses
import json
from .feynobg import BiRefNet, BiRefNetConfig
print(f"[Ruinode-FeyNobg] 加载模型:{model_dir}({dtype},{device})")
with open(os.path.join(model_dir, "config.json"), "r", encoding="utf-8") as f:
raw_cfg = json.load(f)
# config.json 里含 nobg_version 等非模型字段,按 dataclass 的字段名过滤
valid = {f.name for f in dataclasses.fields(BiRefNetConfig)}
cfg = BiRefNetConfig(**{k: v for k, v in raw_cfg.items() if k in valid})
ckpt = os.path.join(model_dir, "model.safetensors")
if not os.path.isfile(ckpt):
raise FileNotFoundError(f"缺少权重文件:{ckpt}")
_MODEL_CACHE.clear() # 单份约 1GB,不做多份缓存
# 不走 BiRefNet.from_pretrained:它内部按 huggingface_hub 的宽松策略加载,
# 键名对不上时会静默丢弃权重(正是上面 _remap_swin_keys 说明的那个坑)
model = BiRefNet(cfg)
_load_weights_strict(model, ckpt)
model.eval()
model.to(device=device, dtype=dtype)
_MODEL_CACHE[key] = model
return model
class RuiFeyNobg:
"""FeyNobg 全自动抠图:输入图像,输出 alpha 与去背景图。"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"model_name": (_list_models(), {
"tooltip": "放在 models/nobg 下的模型目录(需含 config.json 与 "
"model.safetensors)。\n"
"留空或未找到时,首次运行会自动从 HuggingFace 下载 "
"feyninc/FeyNobg(约 1.05GB)。"
}),
"resolution": ([512, 768, 1024, 1280, 1536], {
"default": 1024,
"tooltip": "推理分辨率。模型原生训练分辨率为 1024,改动会影响细节与显存:\n"
"调低更省显存但边缘变粗;调高不一定更好,可能出现结构断裂。"
}),
"precision": (["fp32", "fp16"], {
"default": "fp32",
"tooltip": "实测两者输出一致(同图 alpha 均值都是 0.657),\n"
"但 fp16 显存减半且明显更快,推荐优先用 fp16。\n"
"保留 fp32 作默认,是为老显卡上半精度异常时有个退路。"
}),
"device": (["auto", "cpu"], {"default": "auto"}),
},
"optional": {
"invert_mask": ("BOOLEAN", {
"default": False,
"tooltip": "反转 alpha:默认前景为白(1),开启后前景为黑。"
}),
},
}
RETURN_TYPES = ("MASK", "IMAGE")
RETURN_NAMES = ("alpha", "cutout")
FUNCTION = "matting"
CATEGORY = "Rui-Node🐶/抠图✂️"
@classmethod
def VALIDATE_INPUTS(cls, model_name, **kwargs):
"""模型列表随目录变化,宽松放行,运行时兜底(含自动下载)。"""
return True
def matting(self, image, model_name, resolution, precision, device,
invert_mask=False):
import comfy.model_management
from .feynobg import BiRefNetImageProcessor
model_dir = _ensure_model(model_name)
dev = torch.device("cpu") if device == "cpu" \
else comfy.model_management.get_torch_device()
dtype = torch.float32 if precision == "fp32" else torch.float16
model = _load_model(model_dir, dtype, dev)
proc = BiRefNetImageProcessor.from_pretrained(model_dir)
res = int(resolution)
proc.size = {"height": res, "width": res}
# ComfyUI 的 IMAGE 约定是 RGB,但上游抠图类节点常输出 RGBA
B, H, W, C = image.shape
if C == 4:
image = image[..., :3]
elif C == 1:
image = image.repeat(1, 1, 1, 3)
elif C != 3:
raise ValueError(f"image 需要 1 / 3 / 4 通道,实际收到 {C} 通道")
image = image.contiguous()
alphas = []
for i in range(B):
# 逐张推理:1024 分辨率下 Swin-Large 峰值显存不低,整批一次容易 OOM
pixel_values = proc.preprocess_tensor(
image[i:i + 1], device=dev, dtype=dtype)
with torch.no_grad():
outputs = model(pixel_values=pixel_values)
# fp16 推理出的 logits 先转回 fp32 再 sigmoid/缩放,避免精度损失
if isinstance(outputs, dict):
outputs = {"logits": outputs["logits"].float()}
alpha = proc.post_process_alpha_matting(
outputs, target_sizes=[(H, W)])[0]
alphas.append(alpha.clamp(0, 1).cpu())
alpha = torch.stack(alphas) # [B,H,W]
if invert_mask:
alpha = 1.0 - alpha
cutout = image.detach().cpu().float() * alpha.unsqueeze(-1)
return (alpha, cutout)
NODE_CLASS_MAPPINGS = {
"RuiFeyNobg": RuiFeyNobg,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"RuiFeyNobg": "FeyNobg 抠图 / FeyNobg Matting",
}
+218
View File
@@ -0,0 +1,218 @@
# -*- coding: utf-8 -*-
"""
满屏文字水印节点
================
给输入图像铺满一层平铺的文字水印,常用于版权标注 / 防盗图。
可调参数:
1. text —— 水印文案(支持多行,用换行分隔)
2. font —— 字体样式(下拉,来自 Ruinode/font 目录,同 Markdown 节点)
3. font_size —— 文字大小(像素)
4. angle —— 水印整体旋转角度(度,-180~180)
5. density —— 水印密度(综合控制行间距与同行水印之间的间距,值越大越密)
6. letter_spacing —— 字间距(单条文案内相邻字符的额外间距,像素,可为负)
7. opacity —— 水印透明程度(0=完全透明,100=完全不透明)
8. color —— 水印文字颜色(#RRGGBB / #RGB / "r,g,b" / 常见英文色名)
实现要点:在一张比原图对角线还大的透明画布上按交错网格平铺文字,
整体旋转后从中心裁出与原图等大的区域,再 alpha 合成回原图,
这样任意旋转角度下都能保证四角也被水印覆盖,不留空白。
"""
import math
import numpy as np
import torch
from PIL import Image, ImageDraw
from .mdimg.fonts import scan_fonts, load as load_font, DEFAULT_FONT_LABEL
# 常见英文色名兜底
_COLOR_NAMES = {
"white": (255, 255, 255), "black": (0, 0, 0), "red": (255, 0, 0),
"green": (0, 128, 0), "blue": (0, 0, 255), "gray": (128, 128, 128),
"grey": (128, 128, 128), "yellow": (255, 255, 0), "orange": (255, 165, 0),
"cyan": (0, 255, 255), "magenta": (255, 0, 255), "silver": (192, 192, 192),
}
def _parse_color(s):
"""把颜色字符串解析为 (r, g, b);无法识别时回退灰色。"""
s = str(s or "").strip()
if not s:
return (128, 128, 128)
if s.startswith("#"):
h = s[1:].strip()
if len(h) == 3: # #RGB -> #RRGGBB
h = "".join(c * 2 for c in h)
if len(h) == 6:
try:
return tuple(int(h[i:i + 2], 16) for i in (0, 2, 4))
except ValueError:
pass
if "," in s: # "r,g,b"
parts = [p.strip() for p in s.split(",")]
try:
vals = [max(0, min(255, int(float(p)))) for p in parts[:3]]
if len(vals) == 3:
return tuple(vals)
except ValueError:
pass
return _COLOR_NAMES.get(s.lower(), (128, 128, 128))
def _line_width(font, text, spacing):
"""含字间距时单行文本的像素宽度。"""
if not text:
return 0.0
w = 0.0
for ch in text:
w += font.getlength(ch) + spacing
return max(0.0, w - spacing) # 末字符不加尾随间距
def _draw_spaced_line(draw, xy, text, font, fill, spacing):
"""逐字符绘制一行文本,字与字之间插入 spacing 像素间距。"""
x, y = xy
for ch in text:
draw.text((x, y), ch, font=font, fill=fill)
x += font.getlength(ch) + spacing
class ImageWatermarkNode:
"""满屏文字水印。"""
@classmethod
def INPUT_TYPES(cls):
font_labels = list(scan_fonts().keys())
default_font = DEFAULT_FONT_LABEL if DEFAULT_FONT_LABEL in font_labels \
else (font_labels[0] if font_labels else "")
return {
"required": {
"image": ("IMAGE",),
"text": ("STRING", {
"default": "仅供参考 / SAMPLE",
"multiline": True,
}),
"font": (font_labels, {"default": default_font}),
"font_size": ("INT", {
"default": 48, "min": 8, "max": 500, "step": 1,
}),
"angle": ("FLOAT", {
"default": 30.0, "min": -180.0, "max": 180.0, "step": 1.0,
}),
"density": ("FLOAT", {
"default": 1.0, "min": 0.1, "max": 5.0, "step": 0.05,
}),
"letter_spacing": ("INT", {
"default": 0, "min": -20, "max": 200, "step": 1,
}),
"opacity": ("FLOAT", {
"default": 35.0, "min": 0.0, "max": 100.0, "step": 1.0,
}),
"color": ("STRING", {"default": "#FFFFFF", "multiline": False}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "apply_watermark"
CATEGORY = "Rui-Node🐶/图像调节🎨"
@classmethod
def VALIDATE_INPUTS(cls, font, **kwargs):
"""字体列表会随 font 目录变化,宽松放行,运行时兜底回退。"""
return True
# ---- 生成单张图的水印层并合成 ----
def _watermark_one(self, pil_img, text, font, font_size, angle,
density, letter_spacing, opacity, color):
w, h = pil_img.size
rgb = _parse_color(color)
alpha = int(round(max(0.0, min(100.0, opacity)) / 100.0 * 255))
if alpha <= 0 or not text.strip():
return pil_img.convert("RGB") # 全透明或空文案:原样返回
fill = (rgb[0], rgb[1], rgb[2], alpha)
lines = text.split("\n")
ascent, descent = font.getmetrics()
line_h = ascent + descent
inner_gap = max(1, int(line_h * 0.25)) # 文案内部多行的行距
tile_w = max((_line_width(font, ln, letter_spacing) for ln in lines),
default=0.0)
tile_h = line_h * len(lines) + inner_gap * (len(lines) - 1)
tile_w = max(tile_w, 1.0)
tile_h = max(tile_h, 1.0)
# density 越大间距越小:以字号为基准的空隙除以密度
gap = max(2.0, font_size * 2.0 / max(0.1, density))
step_x = tile_w + gap
step_y = tile_h + gap
# 画布要盖住旋转后的对角线,四周再留一圈平铺余量
diag = math.ceil(math.sqrt(w * w + h * h))
margin = int(max(step_x, step_y) + max(tile_w, tile_h)) + 10
cw = diag + 2 * margin
ch = diag + 2 * margin
layer = Image.new("RGBA", (cw, ch), (0, 0, 0, 0))
draw = ImageDraw.Draw(layer)
# 交错平铺(隔行水平偏移半步),视觉更自然
row = 0
y = -step_y
while y < ch + step_y:
x_off = (step_x / 2.0) if (row % 2) else 0.0
x = -step_x + x_off
while x < cw + step_x:
ty = y
for ln in lines:
if ln:
_draw_spaced_line(draw, (x, ty), ln, font, fill,
letter_spacing)
ty += line_h + inner_gap
x += step_x
y += step_y
row += 1
# 整体旋转后从中心裁出与原图等大的区域
rotated = layer.rotate(angle, resample=Image.BICUBIC, expand=False)
left = (cw - w) // 2
top = (ch - h) // 2
crop = rotated.crop((left, top, left + w, top + h))
base = pil_img.convert("RGBA")
out = Image.alpha_composite(base, crop)
return out.convert("RGB")
def apply_watermark(self, image, text, font, font_size, angle,
density, letter_spacing, opacity, color):
# 加载字体(沿用 Markdown 节点的扫描/回退逻辑)
fmap = scan_fonts()
font_path = fmap.get(font)
if font_path is None:
font_path = next(iter(fmap.values()), "")
print(f"[Rui-Node] 水印字体 '{font}' 不在列表中,已回退 "
f"{font_path or '内置默认'}")
pil_font = load_font(font_path, font_size)
results = []
batch = image.shape[0]
for i in range(batch):
arr = np.clip(image[i].cpu().numpy(), 0.0, 1.0)
pil_img = Image.fromarray((arr * 255).astype(np.uint8), "RGB")
out = self._watermark_one(
pil_img, text, pil_font, int(font_size), float(angle),
float(density), int(letter_spacing), float(opacity), color)
out_np = np.asarray(out, dtype=np.float32) / 255.0
results.append(torch.from_numpy(out_np))
return (torch.stack(results),)
NODE_CLASS_MAPPINGS = {
"ImageWatermark": ImageWatermarkNode,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ImageWatermark": "满屏文字水印 / Full-Screen Text Watermark",
}