optimize performance of vitmate method when processing large-size image

This commit is contained in:
chflame163
2024-06-26 12:01:17 +08:00
parent 268a75bf60
commit fc239893aa
4 changed files with 24 additions and 15 deletions
+1
View File
@@ -80,6 +80,7 @@ When this error has occurred, please check the network environment.
## Update
<font size="4">**If the dependency package error after updating, please reinstall the relevant dependency packages. </font><br />
* 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).
+1
View File
@@ -80,6 +80,7 @@ git clone https://github.com/chflame163/ComfyUI_LayerStyle.git
## 更新说明
<font size="4">**如果本插件更新后出现依赖包错误,请重新安装相关依赖包。
* 优化Ultra节点的```vitmatte```方法在处理大尺寸图片时的性能。
* [CropByMaskV2](#CropByMaskV2) 增加裁切尺寸按倍数取整选项。
* 添加 [CheckMask](#CheckMask) 节点, 用于检测遮罩是否包含足够的有效区域。
* 添加 [HSVValue](#HSVValue) 节点, 用于转换色值为HSV值。
+19 -11
View File
@@ -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:
+3 -4
View File
@@ -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,)