255 lines
9.9 KiB
Python
255 lines
9.9 KiB
Python
# -*- 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
|