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:
co-authored by
Claude Opus 4.8
parent
9dd4fbd890
commit
1bd90c71b6
@@ -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
@@ -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']
|
||||
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
@@ -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",
|
||||
}
|
||||
@@ -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",
|
||||
}
|
||||
Reference in New Issue
Block a user