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,)