# -*- coding: utf-8 -*- """图像预处理 / 后处理。 所有数值行为对齐 ComfyUI 节点 ``comfyui_birefnet_ll``: * 预处理 —— 与 birefnetNode.ImagePreprocessor 相同的 Resize + ImageNet Normalize * 遮罩还原 —— 等价 comfy.utils.common_upscale(F.interpolate,无抗锯齿) * 前景精修 —— fast-foreground-estimation(Photoroom)的 box-blur 版本 """ from __future__ import annotations import os from functools import lru_cache from typing import Dict, Iterable, Optional, Sequence, Tuple import numpy as np import torch from PIL import Image, ImageOps # 与推理无关的通用工具(后缀集合、文件名清洗、尺寸换算)放在 fsutil 里, # 这里重新导出,保持 `imageops.safe_stem` 这类既有调用点不变。 from .fsutil import ( # noqa: F401 ALPHA_SUFFIXES, IMAGE_SUFFIXES, fit_longest_side, safe_stem, ) try: # cv2 的 box blur 速度更快、背景更纯(与节点一致) import cv2 _HAS_CV2 = True except Exception: # pragma: no cover cv2 = None # type: ignore[assignment] _HAS_CV2 = False IMAGENET_MEAN = [0.485, 0.456, 0.406] IMAGENET_STD = [0.229, 0.224, 0.225] #: 与节点 interpolation_modes_mapping 对齐(torchvision InterpolationMode 的数值) INTERPOLATIONS: Dict[str, int] = { "nearest": 0, "bilinear": 2, "bicubic": 3, "nearest-exact": 0, } #: 输出遮罩还原时允许的 F.interpolate 模式 UPSCALE_METHODS = ("bilinear", "nearest", "nearest-exact", "bicubic") #: 预处理尺寸必须是 32 的倍数(Swin 骨干下采样 32 倍) SIZE_ALIGN = 32 # --------------------------------------------------------------------------- # # 基础转换 # --------------------------------------------------------------------------- # def pil_to_rgb_array(image: Image.Image) -> np.ndarray: """PIL -> uint8 RGB 数组,同时按 EXIF 方向自动旋转。""" if image.mode == "RGBA": # 有透明通道时先合成到白底,避免 alpha 变黑 bg = Image.new("RGBA", image.size, (255, 255, 255, 255)) image = Image.alpha_composite(bg, image) elif image.mode not in ("RGB", "L"): image = image.convert("RGB") image = ImageOps.exif_transpose(image) if image.mode != "RGB": image = image.convert("RGB") return np.asarray(image, dtype=np.uint8) def to_pil_rgba(rgb: np.ndarray, alpha: Optional[np.ndarray] = None) -> Image.Image: """uint8 RGB (+ float alpha[0,1]) -> PIL 图片。""" if alpha is None: return Image.fromarray(np.ascontiguousarray(rgb), "RGB") a8 = np.clip(np.rint(alpha * 255.0), 0, 255).astype(np.uint8) rgba = np.dstack([rgb, a8]) return Image.fromarray(np.ascontiguousarray(rgba), "RGBA") def to_pil_mask(alpha: np.ndarray) -> Image.Image: """float mask[0,1] -> 8bit 灰度图。""" m8 = np.clip(np.rint(alpha * 255.0), 0, 255).astype(np.uint8) return Image.fromarray(np.ascontiguousarray(m8), "L") def parse_color(value: object, default: Tuple[int, int, int] = (255, 255, 255)) -> Tuple[int, int, int]: """解析 ``#rgb`` / ``#rrggbb`` / ``(r,g,b)`` / int 为 RGB 三元组。""" if value is None: return default if isinstance(value, (list, tuple)) and len(value) >= 3: return tuple(int(np.clip(v, 0, 255)) for v in value[:3]) # type: ignore[return-value] if isinstance(value, int): return ((value >> 16) & 0xFF, (value >> 8) & 0xFF, value & 0xFF) if isinstance(value, str): s = value.strip().lstrip("#") if len(s) == 3: s = "".join(ch * 2 for ch in s) if len(s) == 6: try: v = int(s, 16) return ((v >> 16) & 0xFF, (v >> 8) & 0xFF, v & 0xFF) except ValueError: pass return default # --------------------------------------------------------------------------- # # 输入尺寸 # --------------------------------------------------------------------------- # def _align(v: float, align: int = SIZE_ALIGN, minimum: int = SIZE_ALIGN) -> int: v = int(round(v / align) * align) return max(minimum, v) def build_input_size( src_h: int, src_w: int, mode: str = "square", width: int = 1024, height: int = 1024, longest_side: int = 1024, ) -> Tuple[int, int]: """计算网络输入尺寸 (h, w)。 mode: square 固定 1024x1024(与节点默认一致,也是官方推荐分辨率) longest 保持宽高比,长边 = longest_side,边长对齐到 32 的倍数 custom 使用自定义 width / height """ if src_h <= 0 or src_w <= 0: return max(SIZE_ALIGN, int(height)), max(SIZE_ALIGN, int(width)) if mode == "custom": h, w = _align(int(height)), _align(int(width)) elif mode == "longest": target = max(SIZE_ALIGN, int(longest_side)) scale = target / float(max(src_h, src_w)) h, w = _align(src_h * scale), _align(src_w * scale) else: # square h, w = _align(int(height)), _align(int(width)) return h, w # --------------------------------------------------------------------------- # # 预处理 # --------------------------------------------------------------------------- # @lru_cache(maxsize=32) def _transform(size_hw: Tuple[int, int], method: str): """缓存 torchvision 变换(与节点 ImagePreprocessor 完全一致)。""" from torchvision import transforms interp = INTERPOLATIONS.get(method, 2) return transforms.Compose( [ transforms.Resize(size_hw, interpolation=interp), transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD), ] ) def preprocess(rgb: np.ndarray, size_hw: Tuple[int, int], method: str = "bilinear") -> torch.Tensor: """uint8 RGB (h,w,3) -> 归一化 float32 张量 (1,3,H,W)。""" arr = np.asarray(rgb, dtype=np.float32) / 255.0 # 新建可写数组,避免 from_numpy 告警 tensor = torch.from_numpy(np.ascontiguousarray(arr)).permute(2, 0, 1).unsqueeze(0) return _transform(tuple(size_hw), method)(tensor) def upscale_mask(mask_bchw: torch.Tensor, height: int, width: int, method: str = "bilinear") -> torch.Tensor: """把网络输出的低分辨率遮罩还原到原图尺寸(等价 comfy.utils.common_upscale)。""" mode = method if method in UPSCALE_METHODS else "bilinear" if tuple(mask_bchw.shape[-2:]) == (height, width): return mask_bchw return torch.nn.functional.interpolate(mask_bchw, size=(height, width), mode=mode) def filter_mask(mask: torch.Tensor, threshold: float) -> torch.Tensor: """低于阈值的概率直接置零(与节点 util.filter_mask 一致)。""" if threshold <= 0: return mask return mask * (mask > threshold).to(mask.dtype) # --------------------------------------------------------------------------- # # 前景精修(fast-foreground-estimation) # --------------------------------------------------------------------------- # def _box_blur(arr: np.ndarray, r: int) -> np.ndarray: """box blur;r 为核尺寸。cv2 可用时优先使用,保证与节点结果一致。 注意:cv2 对 (h, w, 1) 的输入会返回 (h, w),这里统一补齐通道维, 否则后续与 (h, w, 3) 广播会报错(原节点也是靠 `[:, :, None]` 兜住这一点)。 """ r = int(max(1, r)) limit = 2 * min(arr.shape[0], arr.shape[1]) + 1 r = int(min(r, limit)) if _HAS_CV2: out = cv2.blur(np.ascontiguousarray(arr, dtype=np.float32), (r, r)) return out[..., None] if (out.ndim == 2 and arr.ndim == 3) else out # 兜底:torchvision 高斯模糊 from torchvision.transforms import functional as TF if r % 2 == 0: r += 1 t = torch.from_numpy(np.ascontiguousarray(arr)).permute(2, 0, 1).unsqueeze(0).float() out = TF.gaussian_blur(t, r) if r > 1 else t return out.squeeze(0).permute(1, 2, 0).numpy() def _fb_step(image: np.ndarray, F: np.ndarray, B: np.ndarray, alpha: np.ndarray, r: int): a = alpha[:, :, None] if alpha.ndim == 2 else alpha blurred_alpha = _box_blur(a, r) blurred_F = _box_blur(F * a, r) / (blurred_alpha + 1e-5) blurred_B = _box_blur(B * (1.0 - a), r) / ((1.0 - blurred_alpha) + 1e-5) out = blurred_F + a * (image - a * blurred_F - (1.0 - a) * blurred_B) return np.clip(out, 0.0, 1.0), blurred_B def refine_foreground(rgb01: np.ndarray, alpha: np.ndarray, r1: int = 90, r2: int = 6) -> np.ndarray: """估计前景色,消除半透明边缘残留的原背景色。 参考 https://github.com/Photoroom/fast-foreground-estimation """ image = np.ascontiguousarray(rgb01, dtype=np.float32) a = alpha.astype(np.float32) F, blur_B = _fb_step(image, image, image, a, r1) F2, _ = _fb_step(image, F, blur_B, a, r2) return np.clip(F2, 0.0, 1.0) # --------------------------------------------------------------------------- # # 合成输出 # --------------------------------------------------------------------------- # def compose_on_color(rgb: np.ndarray, alpha: np.ndarray, color: Sequence[int]) -> Image.Image: """按遮罩把前景合成到纯色背景上,返回 RGB 图。""" a = alpha[:, :, None].astype(np.float32) bg = np.array(color, dtype=np.float32).reshape(1, 1, 3) out = rgb.astype(np.float32) * a + bg * (1.0 - a) return Image.fromarray(np.clip(np.rint(out), 0, 255).astype(np.uint8), "RGB") def save_image(image: Image.Image, path: str, quality: int = 95) -> str: """保存图片(自动建目录)。PNG/WebP 支持透明,JPEG 不支持。""" os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) ext = os.path.splitext(path)[1].lower() if ext in (".jpg", ".jpeg"): if image.mode == "RGBA": bg = Image.new("RGBA", image.size, (255, 255, 255, 255)) image = Image.alpha_composite(bg, image).convert("RGB") image.save(path, quality=quality, subsampling=0, optimize=True) elif ext == ".webp": image.save(path, quality=quality, method=4) else: image.save(path, optimize=False) return path