2 Commits
Author SHA1 Message Date
Michael Poutre 336f092d6f refactor: Set RELOAD_INST on all nodes 2023-09-25 22:14:21 -07:00
Michael Poutre 6a0d4123d1 feat: Get device automatically 2023-09-25 21:54:40 -07:00
2 changed files with 17 additions and 11 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_KepMatteAnything.Matte_Anything.modeling import ViTMatte, MattingCriterion, Detail_Capture, ViT
from custom_nodes.Comfy_MatteAnything.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)(),
)
)
+15 -9
View File
@@ -51,6 +51,7 @@ class SAMLoader:
}
}
RELOAD_INST = True
RETURN_TYPES = ("SAM_MODEL",)
FUNCTION = "load_model"
@@ -69,6 +70,7 @@ 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)
@@ -86,6 +88,7 @@ class InitSAMPredictor:
}
}
RELOAD_INST = True
RETURN_TYPES = ("SAM_PREDICTOR",)
FUNCTION = "init_predictor"
@@ -107,6 +110,7 @@ class LoadDINOModel:
}
}
RELOAD_INST = True
RETURN_TYPES = ("DINO_MODEL",)
FUNCTION = "load_model"
@@ -174,6 +178,7 @@ class DinoBoxes:
}
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE", "DINO_BOXES")
FUNCTION = "get_boxes"
@@ -213,6 +218,7 @@ class SAMMaskFromBoxes:
}
}
RELOAD_INST = True
RETURN_TYPES = ("MASK",)
FUNCTION = "get_mask"
@@ -250,6 +256,7 @@ class MaskToTrimap:
}
}
RELOAD_INST = True
RETURN_TYPES = ("TRIMAP","MASK")
FUNCTION = "get_trimap"
@@ -272,19 +279,16 @@ class MaskToTrimap:
vitmatte_config = {
'vit_b': './custom_nodes/Comfy_KepMatteAnything/Matte_Anything/configs/matte_anything.py',
'vit_b': './custom_nodes/Comfy_MatteAnything/Matte_Anything/configs/matte_anything.py',
}
def init_vitmatte(model_type, model_path, device_mode):
def init_vitmatte(model_type, model_path):
"""
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)
device = comfy.model_management.get_torch_device() if device_mode in ("Prefer GPU", "AUTO") else "CPU"
if device != "CPU":
vitmatte.to(device)
device = comfy.model_management.get_torch_device()
vitmatte.to(device)
vitmatte.eval()
DetectionCheckpointer(vitmatte).load(model_path)
@@ -301,6 +305,7 @@ class LoadVITMatteModel:
}
}
RELOAD_INST = True
RETURN_TYPES = ("VIT_MATTE_MODEL",)
FUNCTION = "load_model"
@@ -308,7 +313,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, device_mode)
vitmatte = init_vitmatte('vit_b', model_path)
return (vitmatte,)
class GenerateVITMatte:
@@ -322,14 +327,15 @@ class GenerateVITMatte:
}
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE",)
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))
device = comfy.model_management.get_torch_device()
input = {
"image": torch.from_numpy(image_in).permute(2, 0, 1).unsqueeze(0).to(device)/255,
"trimap": trimap.unsqueeze(0).to(device),