4 Commits
Author SHA1 Message Date
Michael Poutre 0c097988df feat(Nodes): Return mask from GenerateVITMatte 2023-09-26 18:16:42 -07:00
Michael Poutre 90eabb8306 refactor: Set RELOAD_INST on all nodes 2023-09-26 18:12:36 -07:00
Michael Poutre d8ce8c1f00 Merge pull request #1 from seanlynch/unhardcode-device
Unhardcode the device for VitMatte models
2023-09-26 18:11:41 -07:00
Sean Lynch 01e1f673e5 Unhardcode the device for VitMatte models
Also fix some pathname issues.
2023-09-26 17:48:47 -04:00
2 changed files with 26 additions and 17 deletions
+2 -2
View File
@@ -1,7 +1,7 @@
import torch.nn as nn
from functools import partial
from detectron2.config import LazyCall as L
from custom_nodes.Comfy_MatteAnything.Matte_Anything.modeling import ViTMatte, MattingCriterion, Detail_Capture, ViT
from custom_nodes.Comfy_KepMatteAnything.Matte_Anything.modeling import ViTMatte, MattingCriterion, Detail_Capture, ViT
# Base
embed_dim, num_heads = 384, 6
@@ -42,4 +42,4 @@ model = L(ViTMatte)(
input_format = "RGB",
size_divisibility=32,
decoder=L(Detail_Capture)(),
)
)
+24 -15
View File
@@ -51,6 +51,7 @@ class SAMLoader:
}
}
RELOAD_INST = True
RETURN_TYPES = ("SAM_MODEL",)
FUNCTION = "load_model"
@@ -69,7 +70,6 @@ class SAMLoader:
sam = sam_model_registry[model_kind](checkpoint=modelname)
# Unless user explicitly wants to use CPU, we use GPU
device = comfy.model_management.get_torch_device() if device_mode == "Prefer GPU" else "CPU"
if device_mode == "Prefer GPU":
sam.to(device=device)
@@ -87,6 +87,7 @@ class InitSAMPredictor:
}
}
RELOAD_INST = True
RETURN_TYPES = ("SAM_PREDICTOR",)
FUNCTION = "init_predictor"
@@ -108,6 +109,7 @@ class LoadDINOModel:
}
}
RELOAD_INST = True
RETURN_TYPES = ("DINO_MODEL",)
FUNCTION = "load_model"
@@ -175,6 +177,7 @@ class DinoBoxes:
}
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE", "DINO_BOXES")
FUNCTION = "get_boxes"
@@ -214,6 +217,7 @@ class SAMMaskFromBoxes:
}
}
RELOAD_INST = True
RETURN_TYPES = ("MASK",)
FUNCTION = "get_mask"
@@ -251,6 +255,7 @@ class MaskToTrimap:
}
}
RELOAD_INST = True
RETURN_TYPES = ("TRIMAP","MASK")
FUNCTION = "get_trimap"
@@ -273,16 +278,19 @@ class MaskToTrimap:
vitmatte_config = {
'vit_b': './custom_nodes/Comfy_MatteAnything/Matte_Anything/configs/matte_anything.py',
'vit_b': './custom_nodes/Comfy_KepMatteAnything/Matte_Anything/configs/matte_anything.py',
}
def init_vitmatte(model_type, model_path):
def init_vitmatte(model_type, model_path, device_mode):
"""
Initialize the vitmatte with model_type in ['vit_s', 'vit_b']
"""
cfg = LazyConfig.load(os.path.abspath(vitmatte_config[model_type]))
vitmatte = instantiate(cfg.model)
# TODO: add device mode
vitmatte.to("mps")
device = comfy.model_management.get_torch_device() if device_mode in ("Prefer GPU", "AUTO") else "CPU"
if device != "CPU":
vitmatte.to(device)
vitmatte.eval()
DetectionCheckpointer(vitmatte).load(model_path)
@@ -299,6 +307,7 @@ class LoadVITMatteModel:
}
}
RELOAD_INST = True
RETURN_TYPES = ("VIT_MATTE_MODEL",)
FUNCTION = "load_model"
@@ -306,7 +315,7 @@ class LoadVITMatteModel:
def load_model(self, model_name, device_mode="auto"):
model_path = folder_paths.get_full_path("matte", model_name)
vitmatte = init_vitmatte('vit_b', model_path)
vitmatte = init_vitmatte('vit_b', model_path, device_mode)
return (vitmatte,)
class GenerateVITMatte:
@@ -320,24 +329,24 @@ class GenerateVITMatte:
}
}
RETURN_TYPES = ("IMAGE",)
RELOAD_INST = True
RETURN_TYPES = ("IMAGE", "MASK")
FUNCTION = "generate_matte"
CATEGORY = "Matte Anything"
def generate_matte(self, image, trimap, vit_matte_model):
device = vit_matte_model.device
image_in = pil2cv(tensor2pil(image))
input = {
"image": torch.from_numpy(image_in).permute(2, 0, 1).unsqueeze(0).to("mps")/255,
"trimap": trimap.unsqueeze(0).to("mps"),
"image": torch.from_numpy(image_in).permute(2, 0, 1).unsqueeze(0).to(device)/255,
"trimap": trimap.unsqueeze(0).to(device),
# "trimap": torch.from_numpy(trimap).unsqueeze(0).unsqueeze(0),
}
alpha = vit_matte_model(input)['phas'].flatten(0, 2)
alpha = alpha.detach().cpu().numpy()
alpha_np = alpha.detach().cpu().numpy()
# Converts alpha matte to RGBA image using the image tensor
image = cv2.cvtColor(image_in, cv2.COLOR_BGR2RGBA)
alpha = np.clip(255. * alpha, 0, 255).astype(np.uint8)
image[:, :, 3] = alpha
image = Image.fromarray(image)
image = pil2tensor(image)
return (image,)
alpha_np = np.clip(255. * alpha_np, 0, 255).astype(np.uint8)
image[:, :, 3] = alpha_np
return pil2tensor(Image.fromarray(image)), alpha