Files
2026-10-08 09:46:47 +08:00

255 lines
9.9 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- 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