diff --git a/README.md b/README.md index f57b0b9..811e679 100644 --- a/README.md +++ b/README.md @@ -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 用户提供实用、高效的节点工具集。🐶 是我们的项目标志,代表着忠诚、友好和可靠。 diff --git a/__init__.py b/__init__.py index 5ad098c..02b2111 100644 --- a/__init__.py +++ b/__init__.py @@ -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'] diff --git a/feynobg/__init__.py b/feynobg/__init__.py new file mode 100644 index 0000000..4467c7c --- /dev/null +++ b/feynobg/__init__.py @@ -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"] diff --git a/feynobg/birefnet/__init__.py b/feynobg/birefnet/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/feynobg/birefnet/image_processing_birefnet.py b/feynobg/birefnet/image_processing_birefnet.py new file mode 100644 index 0000000..d61acf6 --- /dev/null +++ b/feynobg/birefnet/image_processing_birefnet.py @@ -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 diff --git a/feynobg/birefnet/modeling_birefnet.py b/feynobg/birefnet/modeling_birefnet.py new file mode 100644 index 0000000..ded038d --- /dev/null +++ b/feynobg/birefnet/modeling_birefnet.py @@ -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 diff --git a/feynobg/loss.py b/feynobg/loss.py new file mode 100644 index 0000000..2aeb6c8 --- /dev/null +++ b/feynobg/loss.py @@ -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 diff --git a/feynobg/mixin.py b/feynobg/mixin.py new file mode 100644 index 0000000..1a73552 --- /dev/null +++ b/feynobg/mixin.py @@ -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, + ) diff --git a/feynobg/utils.py b/feynobg/utils.py new file mode 100644 index 0000000..9ba1d54 --- /dev/null +++ b/feynobg/utils.py @@ -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 }}}} +--- + +

+ +

+ +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 + +""" diff --git a/feynobg_node.py b/feynobg_node.py new file mode 100644 index 0000000..11abed4 --- /dev/null +++ b/feynobg_node.py @@ -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", +} diff --git a/watermark_node.py b/watermark_node.py new file mode 100644 index 0000000..79da98e --- /dev/null +++ b/watermark_node.py @@ -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", +}