diff --git a/README.MD b/README.MD index 998e56f..1b97fc4 100644 --- a/README.MD +++ b/README.MD @@ -80,6 +80,7 @@ When this error has occurred, please check the network environment. ## Update **If the dependency package error after updating, please reinstall the relevant dependency packages.
+* Optimize performance of the ```vitmate``` method for Ultra nodes when processing large-size image. * [CropByMaskV2](#CropByMaskV2) add option to round the cutting size by multiples. * Commit [CheckMask](#CheckMask) node, it detect whether the mask contains sufficient effective areas. Commit [HSVValue](#HSVValue) node, it convert color values to HSV values. * [BooleanOperatorV2](#BooleanOperatorV2), [NumberCalculatorV2](#NumberCalculatorV2), [Integer](#Integer), [Float](#Float), [Boolean](#Boolean) nodes add string output to output the value as a string for use with [SwitchCase](#SwitchCase). diff --git a/README_CN.MD b/README_CN.MD index f114dc8..0884583 100644 --- a/README_CN.MD +++ b/README_CN.MD @@ -80,6 +80,7 @@ git clone https://github.com/chflame163/ComfyUI_LayerStyle.git ## 更新说明 **如果本插件更新后出现依赖包错误,请重新安装相关依赖包。 +* 优化Ultra节点的```vitmatte```方法在处理大尺寸图片时的性能。 * [CropByMaskV2](#CropByMaskV2) 增加裁切尺寸按倍数取整选项。 * 添加 [CheckMask](#CheckMask) 节点, 用于检测遮罩是否包含足够的有效区域。 * 添加 [HSVValue](#HSVValue) 节点, 用于转换色值为HSV值。 diff --git a/py/imagefunc.py b/py/imagefunc.py index dc24e87..47b35a9 100644 --- a/py/imagefunc.py +++ b/py/imagefunc.py @@ -1402,6 +1402,11 @@ def generate_VITMatte(image:Image, trimap:Image, local_files_only:bool=False) -> image = image.convert('RGB') if trimap.mode != 'L': trimap = trimap.convert('L') + max_size = 2048 + width, height = image.size + if width * height > max_size * max_size: + image = image.resize((max_size, max_size), Image.BILINEAR) + trimap = trimap.resize((max_size, max_size), Image.BILINEAR) model_name = "hustvl/vitmatte-small-composition-1k" vit_matte_model = load_VITMatte_model(model_name=model_name, local_files_only=local_files_only) inputs = vit_matte_model.processor(images=image, trimaps=trimap, return_tensors="pt") @@ -1410,25 +1415,28 @@ def generate_VITMatte(image:Image, trimap:Image, local_files_only:bool=False) -> mask = tensor2pil(predictions).convert('L') mask = mask.crop( (0, 0, image.width, image.height)) # remove padding that the prediction appends (works in 32px tiles) + if width * height > max_size * max_size: + mask = mask.resize((width, height), Image.BILINEAR) return mask def generate_VITMatte_trimap(mask:torch.Tensor, erode_kernel_size:int, dilate_kernel_size:int) -> Image: + def g_trimap(mask, erode_kernel_size=10, dilate_kernel_size=10): + erode_kernel = np.ones((erode_kernel_size, erode_kernel_size), np.uint8) + dilate_kernel = np.ones((dilate_kernel_size, dilate_kernel_size), np.uint8) + eroded = cv2.erode(mask, erode_kernel, iterations=5) + dilated = cv2.dilate(mask, dilate_kernel, iterations=5) + trimap = np.zeros_like(mask) + trimap[dilated == 255] = 128 + trimap[eroded == 255] = 255 + return trimap + mask = mask.squeeze(0).cpu().detach().numpy().astype(np.uint8) * 255 - trimap = __generate_trimap(mask, erode_kernel_size, dilate_kernel_size).astype(np.float32) + trimap = g_trimap(mask, erode_kernel_size, dilate_kernel_size).astype(np.float32) trimap[trimap == 128] = 0.5 trimap[trimap == 255] = 1 trimap = torch.from_numpy(trimap).unsqueeze(0) - return tensor2pil(trimap).convert('L') -def __generate_trimap(mask, erode_kernel_size=10, dilate_kernel_size=10): - erode_kernel = np.ones((erode_kernel_size, erode_kernel_size), np.uint8) - dilate_kernel = np.ones((dilate_kernel_size, dilate_kernel_size), np.uint8) - eroded = cv2.erode(mask, erode_kernel, iterations=5) - dilated = cv2.dilate(mask, dilate_kernel, iterations=5) - trimap = np.zeros_like(mask) - trimap[dilated == 255] = 128 - trimap[eroded == 255] = 255 - return trimap + return tensor2pil(trimap).convert('L') def get_a_person_mask_generator_model_path() -> str: diff --git a/py/purge_vram.py b/py/purge_vram.py index 7797d40..4c1fb01 100644 --- a/py/purge_vram.py +++ b/py/purge_vram.py @@ -1,9 +1,8 @@ +import comfy.model_management as mm from .imagefunc import AnyType any = AnyType("*") - NODE_NAME = 'PurgeVRAM' - class PurgeVRAM: def __init__(self): @@ -30,14 +29,14 @@ class PurgeVRAM: import torch.cuda import gc import comfy.model_management + gc.collect() if purge_cache: if torch.cuda.is_available(): - gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() if purge_models: comfy.model_management.unload_all_models() - + comfy.model_management.soft_empty_cache() return (None,)