diff --git a/README.md b/README.md index ca49404..2ae8bac 100644 --- a/README.md +++ b/README.md @@ -58,13 +58,18 @@ The following models are supported: - **Input**: Mask - **Output**: Processed Mask - **Parameters**: - - Blur Radius: Gaussian blur radius for mask smoothing - - Feather Radius: Edge feathering radius - - Contrast: Mask contrast adjustment - - Brightness: Mask brightness adjustment + - Detail Method: Choose from VITMatte, PyMatting, or GuidedFilter + - Erode/Dilate: Control the trimap generation + - Black/White Point: Adjust mask levels + - Max Megapixels: Control processing resolution ## Changelog +### v1.1.1 +- Fixed VITMatte processing quality issues +- Optimized image size handling for different processing methods +- Improved mask processing workflow + ### v1.1.0 - Optimized dependency management - Removed version constraints for better compatibility diff --git a/__init__.py b/__init__.py index c24b776..504a6f6 100644 --- a/__init__.py +++ b/__init__.py @@ -2,4 +2,4 @@ from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] -__version__ = 'v1.1.0' \ No newline at end of file +__version__ = 'v1.1.1' \ No newline at end of file diff --git a/imagefunc.py b/imagefunc.py index 450596d..6857303 100644 --- a/imagefunc.py +++ b/imagefunc.py @@ -61,6 +61,18 @@ def mask_edge_detail(image, mask, radius, black_point, white_point): mask = tensor2pil(mask) image = np.array(image) mask = np.array(mask) + # 修正 image 类型和通道数 + if image.dtype not in [np.uint8, np.float32]: + image = image.astype(np.uint8) + if image.ndim == 2: + pass # 灰度 + elif image.shape[2] > 3: + image = image[:, :, :3] # 只保留前三通道 + # 修正 mask 类型和通道数 + if mask.dtype not in [np.uint8, np.float32]: + mask = mask.astype(np.uint8) + if mask.ndim == 3 and mask.shape[2] > 1: + mask = mask[:, :, 0] # 只保留单通道 mask = cv2.ximgproc.guidedFilter(image, mask, radius, 1e-6) mask = adjust_levels(Image.fromarray(mask), black_point, white_point) return torch.from_numpy(np.array(mask)).unsqueeze(0) diff --git a/nodes.py b/nodes.py index da23e50..4c6c8ba 100644 --- a/nodes.py +++ b/nodes.py @@ -267,21 +267,51 @@ class ProcessDetails_RBS: local_files_only = detail_method == "VITMatte(local)" trimap = generate_VITMatte_trimap(orig_mask, detail_erode, detail_dilate) processed_mask = generate_VITMatte(orig_image, trimap, local_files_only, device, max_megapixels) - elif detail_method == "PyMatting": - trimap = generate_VITMatte_trimap(orig_mask, detail_erode, detail_dilate) - try: - from pymatting import estimate_alpha_lkm - except ImportError: - raise RuntimeError("请先安装 pymatting 库: pip install pymatting scikit-image") - import numpy as np - image_np = np.array(orig_image.convert('RGB')) - trimap_np = np.array(trimap.convert('L')) / 255.0 - # PyMatting只用LKM算法,参数风格与LayerStyle_Advance一致 - alpha = estimate_alpha_lkm(image_np, trimap_np) - processed_mask = Image.fromarray((alpha * 255).astype(np.uint8)) - else: # GuidedFilter - processed_mask = mask_edge_detail(i, m, detail_erode, black_point, white_point) - processed_mask = tensor2pil(processed_mask) + else: + # 计算目标尺寸 + width, height = orig_image.size + max_pixels = int(max_megapixels * 1_048_576) + orig_pixels = width * height + patch_size = 32 # 保证分辨率为32的倍数 + + if orig_pixels > max_pixels: + scale = (max_pixels / orig_pixels) ** 0.5 + new_width = max(1, int(width * scale)) + new_height = max(1, int(height * scale)) + else: + new_width, new_height = width, height + + # 向下取整为patch_size的倍数 + new_width = (new_width // patch_size) * patch_size + new_height = (new_height // patch_size) * patch_size + new_width = max(patch_size, new_width) + new_height = max(patch_size, new_height) + inference_image_size = (new_width, new_height) + + log(f"[ProcessDetails_RBS] 原始尺寸: {width}x{height}, 实际推理尺寸: {inference_image_size[0]}x{inference_image_size[1]}", message_type='info') + + # 调整图像大小 + resized_image = orig_image.resize(inference_image_size, Image.BILINEAR) + resized_mask = orig_mask.resize(inference_image_size, Image.BILINEAR) + + if detail_method == "PyMatting": + trimap = generate_VITMatte_trimap(resized_mask, detail_erode, detail_dilate) + try: + from pymatting import estimate_alpha_lkm + except ImportError: + raise RuntimeError("请先安装 pymatting 库: pip install pymatting scikit-image") + import numpy as np + image_np = np.array(resized_image.convert('RGB')) + trimap_np = np.array(trimap.convert('L')) / 255.0 + # PyMatting只用LKM算法,参数风格与LayerStyle_Advance一致 + alpha = estimate_alpha_lkm(image_np, trimap_np) + processed_mask = Image.fromarray((alpha * 255).astype(np.uint8)) + else: # GuidedFilter + processed_mask = mask_edge_detail(pil2tensor(resized_image), image2mask(resized_mask), detail_erode, black_point, white_point) + processed_mask = tensor2pil(processed_mask) + + # 还原到原始尺寸 + processed_mask = processed_mask.resize(orig_image.size, Image.BILINEAR) # 应用处理后的蒙版到图像 processed_image = RGB2RGBA(orig_image, processed_mask) diff --git a/pyproject.toml b/pyproject.toml index bac9347..de97c60 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "removebackgroundsuite" description = "A matting toolkit based on ComfyUI, supporting multiple matting models and detail processing methods." -version = "1.1.0" +version = "1.1.1" license = {file = "LICENSE"} dependencies = ["torch", "numpy", "Pillow", "torchvision", "opencv-python", "scipy", "transformers", "tqdm"]