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
+13 -4
View File
@@ -51,6 +51,7 @@ class SAMLoader:
}
}
RELOAD_INST = True
RETURN_TYPES = ("SAM_MODEL",)
FUNCTION = "load_model"
@@ -87,6 +88,7 @@ class InitSAMPredictor:
}
}
RELOAD_INST = True
RETURN_TYPES = ("SAM_PREDICTOR",)
FUNCTION = "init_predictor"
@@ -108,6 +110,7 @@ class LoadDINOModel:
}
}
RELOAD_INST = True
RETURN_TYPES = ("DINO_MODEL",)
FUNCTION = "load_model"
@@ -175,6 +178,7 @@ class DinoBoxes:
}
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE", "DINO_BOXES")
FUNCTION = "get_boxes"
@@ -214,6 +218,7 @@ class SAMMaskFromBoxes:
}
}
RELOAD_INST = True
RETURN_TYPES = ("MASK",)
FUNCTION = "get_mask"
@@ -251,6 +256,7 @@ class MaskToTrimap:
}
}
RELOAD_INST = True
RETURN_TYPES = ("TRIMAP","MASK")
FUNCTION = "get_trimap"
@@ -281,8 +287,8 @@ def init_vitmatte(model_type, model_path):
"""
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()
vitmatte.to(device)
vitmatte.eval()
DetectionCheckpointer(vitmatte).load(model_path)
@@ -299,6 +305,7 @@ class LoadVITMatteModel:
}
}
RELOAD_INST = True
RETURN_TYPES = ("VIT_MATTE_MODEL",)
FUNCTION = "load_model"
@@ -320,6 +327,7 @@ class GenerateVITMatte:
}
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE",)
FUNCTION = "generate_matte"
@@ -327,9 +335,10 @@ class GenerateVITMatte:
def generate_matte(self, image, trimap, vit_matte_model):
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("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)