feat: initial project setup

This commit is contained in:
2026-10-08 09:46:47 +08:00
commit 5833a303bc
65 changed files with 11628 additions and 0 deletions
+254
View File
@@ -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