2 Commits
Author SHA1 Message Date
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 11 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)(),
)
)
+9 -15
View File
@@ -51,7 +51,6 @@ class SAMLoader:
}
}
RELOAD_INST = True
RETURN_TYPES = ("SAM_MODEL",)
FUNCTION = "load_model"
@@ -70,7 +69,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)
@@ -88,7 +86,6 @@ class InitSAMPredictor:
}
}
RELOAD_INST = True
RETURN_TYPES = ("SAM_PREDICTOR",)
FUNCTION = "init_predictor"
@@ -110,7 +107,6 @@ class LoadDINOModel:
}
}
RELOAD_INST = True
RETURN_TYPES = ("DINO_MODEL",)
FUNCTION = "load_model"
@@ -178,7 +174,6 @@ class DinoBoxes:
}
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE", "DINO_BOXES")
FUNCTION = "get_boxes"
@@ -218,7 +213,6 @@ class SAMMaskFromBoxes:
}
}
RELOAD_INST = True
RETURN_TYPES = ("MASK",)
FUNCTION = "get_mask"
@@ -256,7 +250,6 @@ class MaskToTrimap:
}
}
RELOAD_INST = True
RETURN_TYPES = ("TRIMAP","MASK")
FUNCTION = "get_trimap"
@@ -279,16 +272,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)
device = comfy.model_management.get_torch_device()
vitmatte.to(device)
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)
@@ -305,7 +301,6 @@ class LoadVITMatteModel:
}
}
RELOAD_INST = True
RETURN_TYPES = ("VIT_MATTE_MODEL",)
FUNCTION = "load_model"
@@ -313,7 +308,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:
@@ -327,15 +322,14 @@ 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),