Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
336f092d6f | ||
|
|
6a0d4123d1 |
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user