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 }}}} +--- + +
+
+