feat: initial project setup
This commit is contained in:
@@ -0,0 +1,254 @@
|
||||
# -*- 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
|
||||
Reference in New Issue
Block a user