Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8034af6478 | ||
|
|
713f6de761 | ||
|
|
83c3732201 | ||
|
|
de986044b3 | ||
|
|
bb6994f677 | ||
|
|
4405434b9c | ||
|
|
e7e9e58c66 | ||
|
|
51dd8fcb7c | ||
|
|
4be544aa9e | ||
|
|
08fa873f2c |
@@ -3,13 +3,10 @@ from typing import List
|
||||
import folder_paths
|
||||
from nodes import ControlNetLoader, ControlNetApply, ControlNetApplyAdvanced
|
||||
|
||||
from .preprocessors import control_net_preprocessors, DummyPreprocessor
|
||||
from .preprocessor import preprocessors, apply_preprocessor
|
||||
from .advanced import comfy_load_controlnet
|
||||
|
||||
|
||||
control_net_preprocessors["tile"] = (DummyPreprocessor, [])
|
||||
|
||||
|
||||
def load_controlnet(control_net_name, control_net_override="None", timestep_keyframe=None):
|
||||
if control_net_override != "None":
|
||||
if control_net_override not in folder_paths.get_filename_list("controlnet"):
|
||||
@@ -23,33 +20,6 @@ def load_controlnet(control_net_name, control_net_override="None", timestep_keyf
|
||||
return comfy_load_controlnet(control_net_name, timestep_keyframe=timestep_keyframe)
|
||||
|
||||
|
||||
def apply_preprocessor(image, preprocessor, resolution=512):
|
||||
if preprocessor == "None":
|
||||
return image
|
||||
|
||||
if preprocessor not in control_net_preprocessors:
|
||||
raise Exception(f"Preprocessor {preprocessor} is not implemented")
|
||||
|
||||
preprocessor_class, default_args = control_net_preprocessors[preprocessor]
|
||||
default_args: List = default_args.copy()
|
||||
|
||||
required_args = preprocessor_class.INPUT_TYPES()["required"].keys()
|
||||
optional_args = preprocessor_class.INPUT_TYPES().get("optional", {}).keys()
|
||||
resolution_idx = list(optional_args).index("resolution")
|
||||
default_args.insert(resolution_idx, resolution)
|
||||
default_args.insert(0, image)
|
||||
|
||||
preprocessor_args = {key: default_args[i] for i, key in enumerate(required_args)}
|
||||
preprocessor_args.update({key: default_args[i + len(required_args)] for i, key in enumerate(optional_args)})
|
||||
|
||||
function_name = preprocessor_class.FUNCTION
|
||||
res = getattr(preprocessor_class(), function_name)(**preprocessor_args)
|
||||
if isinstance(res, dict):
|
||||
res = res["result"]
|
||||
|
||||
return res[0]
|
||||
|
||||
|
||||
def detect_controlnet(preprocessor: str, sd_version: str):
|
||||
controlnets = folder_paths.get_filename_list("controlnet")
|
||||
controlnets = filter(lambda x: sd_version in x, controlnets)
|
||||
@@ -103,15 +73,13 @@ class AVControlNetLoader(ControlNetLoader):
|
||||
|
||||
|
||||
class AV_ControlNetPreprocessor:
|
||||
preprocessors = list(control_net_preprocessors.keys())
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"preprocessor": (["None", "tile"] + s.preprocessors,),
|
||||
"sd_version": (["sd15", "sdxl", "sdxl_t2i"],),
|
||||
"preprocessor": (["None"] + preprocessors,),
|
||||
"sd_version": (["sd15", "sdxl"],),
|
||||
},
|
||||
"optional": {
|
||||
"resolution": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
|
||||
@@ -126,7 +94,7 @@ class AV_ControlNetPreprocessor:
|
||||
|
||||
def detect_controlnet(self, image, preprocessor, sd_version, resolution=512, preprocessor_override="None"):
|
||||
if preprocessor_override != "None":
|
||||
if preprocessor_override not in control_net_preprocessors:
|
||||
if preprocessor_override not in preprocessors:
|
||||
print(
|
||||
f"Warning: Not found ControlNet preprocessor {preprocessor_override}. Use {preprocessor} instead."
|
||||
)
|
||||
@@ -141,7 +109,6 @@ class AV_ControlNetPreprocessor:
|
||||
|
||||
class AVControlNetEfficientStacker:
|
||||
controlnets = folder_paths.get_filename_list("controlnet")
|
||||
preprocessors = list(control_net_preprocessors.keys())
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -155,7 +122,7 @@ class AVControlNetEfficientStacker:
|
||||
),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"preprocessor": (["None"] + s.preprocessors,),
|
||||
"preprocessor": (["None"] + preprocessors,),
|
||||
},
|
||||
"optional": {
|
||||
"cnet_stack": ("CONTROL_NET_STACK",),
|
||||
@@ -219,7 +186,7 @@ class AVControlNetEfficientStackerSimple(AVControlNetEfficientStacker):
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01},
|
||||
),
|
||||
"preprocessor": (["None"] + s.preprocessors,),
|
||||
"preprocessor": (["None"] + preprocessors,),
|
||||
},
|
||||
"optional": {
|
||||
"cnet_stack": ("CONTROL_NET_STACK",),
|
||||
@@ -242,7 +209,6 @@ class AVControlNetEfficientStackerSimple(AVControlNetEfficientStacker):
|
||||
|
||||
class AVControlNetEfficientLoader(ControlNetApply):
|
||||
controlnets = folder_paths.get_filename_list("controlnet")
|
||||
preprocessors = list(control_net_preprocessors.keys())
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -255,7 +221,7 @@ class AVControlNetEfficientLoader(ControlNetApply):
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01},
|
||||
),
|
||||
"preprocessor": (["None"] + s.preprocessors,),
|
||||
"preprocessor": (["None"] + preprocessors,),
|
||||
},
|
||||
"optional": {
|
||||
"control_net_override": ("STRING", {"default": "None"}),
|
||||
@@ -295,7 +261,6 @@ class AVControlNetEfficientLoader(ControlNetApply):
|
||||
|
||||
class AVControlNetEfficientLoaderAdvanced(ControlNetApplyAdvanced):
|
||||
controlnets = folder_paths.get_filename_list("controlnet")
|
||||
preprocessors = list(control_net_preprocessors.keys())
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -311,7 +276,7 @@ class AVControlNetEfficientLoaderAdvanced(ControlNetApplyAdvanced):
|
||||
),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"preprocessor": (["None"] + s.preprocessors,),
|
||||
"preprocessor": (["None"] + preprocessors,),
|
||||
},
|
||||
"optional": {
|
||||
"control_net_override": ("STRING", {"default": "None"}),
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
import os
|
||||
from typing import Dict
|
||||
|
||||
import folder_paths
|
||||
|
||||
from ..utils import load_module
|
||||
|
||||
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
|
||||
preprocessors_dir_names = ["ControlNetPreprocessors", "comfyui_controlnet_aux"]
|
||||
|
||||
preprocessors: list[str] = []
|
||||
_preprocessors_map = {
|
||||
"canny": "CannyEdgePreprocessor",
|
||||
"canny_pyra": "PyraCannyPreprocessor",
|
||||
"lineart": "LineArtPreprocessor",
|
||||
"lineart_anime": "AnimeLineArtPreprocessor",
|
||||
"lineart_manga": "Manga2Anime_LineArt_Preprocessor",
|
||||
"lineart_any": "AnyLineArtPreprocessor_aux",
|
||||
"scribble": "ScribblePreprocessor",
|
||||
"scribble_xdog": "Scribble_XDoG_Preprocessor",
|
||||
"scribble_pidi": "Scribble_PiDiNet_Preprocessor",
|
||||
"scribble_hed": "FakeScribblePreprocessor",
|
||||
"hed": "HEDPreprocessor",
|
||||
"pidi": "PiDiNetPreprocessor",
|
||||
"mlsd": "M-LSDPreprocessor",
|
||||
"pose": "DWPreprocessor",
|
||||
"openpose": "OpenposePreprocessor",
|
||||
"dwpose": "DWPreprocessor",
|
||||
"pose_dense": "DensePosePreprocessor",
|
||||
"pose_animal": "AnimalPosePreprocessor",
|
||||
"normalmap_bae": "BAE-NormalMapPreprocessor",
|
||||
"normalmap_dsine": "DSINE-NormalMapPreprocessor",
|
||||
"normalmap_midas": "MiDaS-NormalMapPreprocessor",
|
||||
"depth": "DepthAnythingV2Preprocessor",
|
||||
"depth_anything": "DepthAnythingPreprocessor",
|
||||
"depth_anything_v2": "DepthAnythingV2Preprocessor",
|
||||
"depth_anything_zoe": "Zoe_DepthAnythingPreprocessor",
|
||||
"depth_zoe": "Zoe-DepthMapPreprocessor",
|
||||
"depth_midas": "MiDaS-DepthMapPreprocessor",
|
||||
"depth_leres": "LeReS-DepthMapPreprocessor",
|
||||
"depth_metric3d": "Metric3D-DepthMapPreprocessor",
|
||||
"depth_meshgraphormer": "MeshGraphormer-DepthMapPreprocessor",
|
||||
"seg_ofcoco": "OneFormer-COCO-SemSegPreprocessor",
|
||||
"seg_ofade20k": "OneFormer-ADE20K-SemSegPreprocessor",
|
||||
"seg_ufade20k": "UniFormer-SemSegPreprocessor",
|
||||
"seg_animeface": "AnimeFace_SemSegPreprocessor",
|
||||
"shuffle": "ShufflePreprocessor",
|
||||
"teed": "TEEDPreprocessor",
|
||||
"color": "ColorPreprocessor",
|
||||
"sam": "SAMPreprocessor",
|
||||
"tile": "TilePreprocessor"
|
||||
}
|
||||
|
||||
|
||||
def apply_preprocessor(image, preprocessor, resolution=512):
|
||||
raise NotImplementedError("apply_preprocessor is not implemented")
|
||||
|
||||
|
||||
try:
|
||||
module_path = None
|
||||
|
||||
for custom_node in custom_nodes:
|
||||
custom_node = custom_node if not os.path.islink(custom_node) else os.readlink(custom_node)
|
||||
for module_dir in preprocessors_dir_names:
|
||||
if module_dir in os.listdir(custom_node):
|
||||
module_path = os.path.abspath(os.path.join(custom_node, module_dir))
|
||||
break
|
||||
|
||||
if module_path is None:
|
||||
raise Exception("Could not find ControlNetPreprocessors nodes")
|
||||
|
||||
module = load_module(module_path)
|
||||
print("Loaded ControlNetPreprocessors nodes from", module_path)
|
||||
|
||||
nodes: Dict = getattr(module, "NODE_CLASS_MAPPINGS")
|
||||
available_preprocessors: list[str] = getattr(module, "PREPROCESSOR_OPTIONS")
|
||||
|
||||
AIO_Preprocessor = nodes.get("AIO_Preprocessor", None)
|
||||
if AIO_Preprocessor is None:
|
||||
raise Exception("Could not find AIO_Preprocessor node")
|
||||
|
||||
for name, preprocessor in _preprocessors_map.items():
|
||||
if preprocessor in available_preprocessors:
|
||||
preprocessors.append(name)
|
||||
|
||||
aio_preprocessor = AIO_Preprocessor()
|
||||
|
||||
def apply_preprocessor(image, preprocessor, resolution=512):
|
||||
if preprocessor == "None":
|
||||
return image
|
||||
|
||||
if preprocessor not in preprocessors:
|
||||
raise Exception(f"Preprocessor {preprocessor} is not implemented")
|
||||
|
||||
preprocessor_cls = _preprocessors_map[preprocessor]
|
||||
args = {"preprocessor": preprocessor_cls, "image": image, "resolution": resolution}
|
||||
|
||||
function_name = AIO_Preprocessor.FUNCTION
|
||||
res = getattr(aio_preprocessor, function_name)(**args)
|
||||
if isinstance(res, dict):
|
||||
res = res["result"]
|
||||
|
||||
return res[0]
|
||||
|
||||
except Exception as e:
|
||||
print(e)
|
||||
@@ -1,142 +0,0 @@
|
||||
import os
|
||||
import math
|
||||
from typing import Dict
|
||||
|
||||
import folder_paths
|
||||
|
||||
from ..utils import load_module
|
||||
|
||||
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
|
||||
preprocessors_dir_names = ["ControlNetPreprocessors", "comfyui_controlnet_aux"]
|
||||
|
||||
control_net_preprocessors = {}
|
||||
|
||||
try:
|
||||
module_path = None
|
||||
|
||||
for custom_node in custom_nodes:
|
||||
custom_node = (
|
||||
custom_node if not os.path.islink(custom_node) else os.readlink(custom_node)
|
||||
)
|
||||
for module_dir in preprocessors_dir_names:
|
||||
if module_dir in os.listdir(custom_node):
|
||||
module_path = os.path.abspath(os.path.join(custom_node, module_dir))
|
||||
break
|
||||
|
||||
if module_path is None:
|
||||
raise Exception("Could not find ControlNetPreprocessors nodes")
|
||||
|
||||
module = load_module(module_path)
|
||||
print("Loaded ControlNetPreprocessors nodes from", module_path)
|
||||
|
||||
nodes: Dict = getattr(module, "NODE_CLASS_MAPPINGS")
|
||||
|
||||
if "CannyEdgePreprocessor" in nodes:
|
||||
control_net_preprocessors["canny"] = (
|
||||
nodes["CannyEdgePreprocessor"],
|
||||
[100, 200],
|
||||
)
|
||||
if "LineArtPreprocessor" in nodes:
|
||||
control_net_preprocessors["lineart"] = (
|
||||
nodes["LineArtPreprocessor"],
|
||||
["disable"],
|
||||
)
|
||||
control_net_preprocessors["lineart_coarse"] = (
|
||||
nodes["LineArtPreprocessor"],
|
||||
["enable"],
|
||||
)
|
||||
if "AnimeLineArtPreprocessor" in nodes:
|
||||
control_net_preprocessors["lineart_anime"] = (
|
||||
nodes["AnimeLineArtPreprocessor"],
|
||||
[],
|
||||
)
|
||||
if "Manga2Anime_LineArt_Preprocessor" in nodes:
|
||||
control_net_preprocessors["lineart_manga"] = (
|
||||
nodes["Manga2Anime_LineArt_Preprocessor"],
|
||||
[],
|
||||
)
|
||||
if "ScribblePreprocessor" in nodes:
|
||||
control_net_preprocessors["scribble"] = (nodes["ScribblePreprocessor"], [])
|
||||
if "FakeScribblePreprocessor" in nodes:
|
||||
control_net_preprocessors["scribble_hed"] = (
|
||||
nodes["FakeScribblePreprocessor"],
|
||||
["enable"],
|
||||
)
|
||||
if "HEDPreprocessor" in nodes:
|
||||
control_net_preprocessors["hed"] = (nodes["HEDPreprocessor"], ["disable"])
|
||||
control_net_preprocessors["hed_safe"] = (nodes["HEDPreprocessor"], ["enable"])
|
||||
if "PiDiNetPreprocessor" in nodes:
|
||||
control_net_preprocessors["pidi"] = (
|
||||
nodes["PiDiNetPreprocessor"],
|
||||
["disable"],
|
||||
)
|
||||
control_net_preprocessors["pidi_safe"] = (
|
||||
nodes["PiDiNetPreprocessor"],
|
||||
["enable"],
|
||||
)
|
||||
if "M-LSDPreprocessor" in nodes:
|
||||
control_net_preprocessors["mlsd"] = (nodes["M-LSDPreprocessor"], [0.1, 0.1])
|
||||
if "OpenposePreprocessor" in nodes:
|
||||
control_net_preprocessors["openpose"] = (
|
||||
nodes["OpenposePreprocessor"],
|
||||
["enable", "enable", "enable"],
|
||||
)
|
||||
control_net_preprocessors["pose"] = control_net_preprocessors["openpose"]
|
||||
if "DWPreprocessor" in nodes:
|
||||
control_net_preprocessors["dwpose"] = (
|
||||
nodes["DWPreprocessor"],
|
||||
["enable", "enable", "enable", "yolox_l.onnx", "dw-ll_ucoco_384.onnx"],
|
||||
)
|
||||
# use DWPreprocessor for pose by default if available
|
||||
control_net_preprocessors["pose"] = control_net_preprocessors["dwpose"]
|
||||
if "BAE-NormalMapPreprocessor" in nodes:
|
||||
control_net_preprocessors["normalmap_bae"] = (
|
||||
nodes["BAE-NormalMapPreprocessor"],
|
||||
[],
|
||||
)
|
||||
if "MiDaS-NormalMapPreprocessor" in nodes:
|
||||
control_net_preprocessors["normalmap_midas"] = (
|
||||
nodes["MiDaS-NormalMapPreprocessor"],
|
||||
[math.pi * 2.0, 0.1],
|
||||
)
|
||||
if "MiDaS-DepthMapPreprocessor" in nodes:
|
||||
control_net_preprocessors["depth_midas"] = (
|
||||
nodes["MiDaS-DepthMapPreprocessor"],
|
||||
[math.pi * 2.0, 0.4],
|
||||
)
|
||||
if "Zoe-DepthMapPreprocessor" in nodes:
|
||||
control_net_preprocessors["depth"] = (nodes["Zoe-DepthMapPreprocessor"], [])
|
||||
control_net_preprocessors["depth_zoe"] = (nodes["Zoe-DepthMapPreprocessor"], [])
|
||||
if "OneFormer-COCO-SemSegPreprocessor" in nodes:
|
||||
control_net_preprocessors["seg_ofcoco"] = (
|
||||
nodes["OneFormer-COCO-SemSegPreprocessor"],
|
||||
[],
|
||||
)
|
||||
if "OneFormer-ADE20K-SemSegPreprocessor" in nodes:
|
||||
control_net_preprocessors["seg_ofade20k"] = (
|
||||
nodes["OneFormer-ADE20K-SemSegPreprocessor"],
|
||||
[],
|
||||
)
|
||||
if "UniFormer-SemSegPreprocessor" in nodes:
|
||||
control_net_preprocessors["seg_ufade20k"] = (
|
||||
nodes["UniFormer-SemSegPreprocessor"],
|
||||
[],
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
|
||||
class DummyPreprocessor:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",)
|
||||
}
|
||||
}
|
||||
|
||||
FUNCTION = "process"
|
||||
|
||||
def process(self, image):
|
||||
return (image,)
|
||||
@@ -6,7 +6,7 @@ import torch.nn.functional as F
|
||||
import folder_paths
|
||||
import comfy.model_management as model_management
|
||||
|
||||
from ...model_utils import download_model
|
||||
from ...model_utils import download_file
|
||||
|
||||
|
||||
lama = None
|
||||
@@ -33,14 +33,10 @@ def pad_tensor_to_modulo(img, mod):
|
||||
def load_model():
|
||||
global lama
|
||||
if lama is None:
|
||||
files = download_model(
|
||||
model_path=model_dir,
|
||||
model_url=model_url,
|
||||
ext_filter=[".pt"],
|
||||
download_name="big-lama.pt",
|
||||
)
|
||||
model_path = os.path.join(model_dir, "big-lama.pt")
|
||||
download_file(model_url, model_path, model_sha)
|
||||
|
||||
lama = torch.jit.load(files[0], map_location="cpu")
|
||||
lama = torch.jit.load(model_path, map_location="cpu")
|
||||
lama.eval()
|
||||
|
||||
return lama
|
||||
|
||||
@@ -1,14 +1,16 @@
|
||||
from .blip_node import BlipLoader, BlipCaption
|
||||
from .blip_node import BlipLoader, BlipCaption, DownloadAndLoadBlip
|
||||
from .danbooru import DeepDanbooruCaption
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"BLIPLoader": BlipLoader,
|
||||
"BLIPCaption": BlipCaption,
|
||||
"DownloadAndLoadBlip": DownloadAndLoadBlip,
|
||||
"DeepDanbooruCaption": DeepDanbooruCaption,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"BLIPLoader": "BLIP Loader",
|
||||
"BLIPCaption": "BLIP Caption",
|
||||
"DownloadAndLoadBlip": "Download and Load BLIP Model",
|
||||
"DeepDanbooruCaption": "Deep Danbooru Caption",
|
||||
}
|
||||
|
||||
|
||||
@@ -7,17 +7,25 @@ from torchvision.transforms.functional import InterpolationMode
|
||||
import folder_paths
|
||||
from comfy.model_management import text_encoder_device, text_encoder_offload_device, soft_empty_cache
|
||||
|
||||
from ..model_utils import download_model
|
||||
from ..model_utils import download_file
|
||||
from ..utils import tensor2pil
|
||||
|
||||
blips = {}
|
||||
blip_size = 384
|
||||
gpu = text_encoder_device()
|
||||
cpu = text_encoder_offload_device()
|
||||
model_url = (
|
||||
"https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_caption_capfilt_large.pth"
|
||||
)
|
||||
model_dir = os.path.join(folder_paths.models_dir, "blip")
|
||||
models = {
|
||||
"model_base_caption_capfilt_large.pth": {
|
||||
"url": "https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_caption_capfilt_large.pth",
|
||||
"sha": "96ac8749bd0a568c274ebe302b3a3748ab9be614c737f3d8c529697139174086",
|
||||
},
|
||||
"model_base_capfilt_large.pth": {
|
||||
"url": "https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_capfilt_large.pth",
|
||||
"sha": "8f5187458d4d47bb87876faf3038d5947eff17475edf52cf47b62e84da0b235f",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
folder_paths.folder_names_and_paths["blip"] = (
|
||||
[model_dir],
|
||||
@@ -90,12 +98,7 @@ def join_caption(caption, prefix, suffix):
|
||||
|
||||
def blip_caption(model, image, min_length, max_length):
|
||||
image = tensor2pil(image)
|
||||
|
||||
if "transformers==4.26.1" in packages(True):
|
||||
print("Using Legacy `transformImaage()`")
|
||||
tensor = transformImage_legacy(image)
|
||||
else:
|
||||
tensor = transformImage(image)
|
||||
tensor = transformImage(image)
|
||||
|
||||
with torch.no_grad():
|
||||
caption = model.generate(
|
||||
@@ -122,8 +125,32 @@ class BlipLoader:
|
||||
CATEGORY = "Art Venture/Captioning"
|
||||
|
||||
def load_blip(self, model_name):
|
||||
model = load_blip(model_name)
|
||||
return (model,)
|
||||
return (load_blip(model_name),)
|
||||
|
||||
|
||||
class DownloadAndLoadBlip:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model_name": (list(models.keys()),),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("BLIP_MODEL",)
|
||||
FUNCTION = "download_and_load_blip"
|
||||
CATEGORY = "Art Venture/Captioning"
|
||||
|
||||
def download_and_load_blip(self, model_name):
|
||||
if model_name not in folder_paths.get_filename_list("blip"):
|
||||
model_info = models[model_name]
|
||||
download_file(
|
||||
model_info["url"],
|
||||
os.path.join(model_dir, model_name),
|
||||
model_info["sha"],
|
||||
)
|
||||
|
||||
return (load_blip(model_name),)
|
||||
|
||||
|
||||
class BlipCaption:
|
||||
@@ -173,15 +200,8 @@ class BlipCaption:
|
||||
return ([join_caption("", prefix, suffix)],)
|
||||
|
||||
if blip_model is None:
|
||||
ckpts = folder_paths.get_filename_list("blip")
|
||||
if len(ckpts) == 0:
|
||||
ckpts = download_model(
|
||||
model_path=model_dir,
|
||||
model_url=model_url,
|
||||
ext_filter=[".pth"],
|
||||
download_name="model_base_caption_capfilt_large.pth",
|
||||
)
|
||||
blip_model = load_blip(ckpts[0])
|
||||
downloader = DownloadAndLoadBlip()
|
||||
blip_model = downloader.download_and_load_blip("model_base_caption_capfilt_large.pth")[0]
|
||||
|
||||
device = gpu if device_mode != "CPU" else cpu
|
||||
blip_model = blip_model.to(device)
|
||||
|
||||
@@ -7,7 +7,7 @@ import folder_paths
|
||||
from comfy.model_management import text_encoder_device, text_encoder_offload_device, soft_empty_cache
|
||||
|
||||
from ..image_utils import resize_image
|
||||
from ..model_utils import download_model
|
||||
from ..model_utils import download_file
|
||||
from ..utils import is_junction, tensor2pil
|
||||
from .blip_node import join_caption
|
||||
|
||||
@@ -15,28 +15,25 @@ danbooru = None
|
||||
blip_size = 384
|
||||
gpu = text_encoder_device()
|
||||
cpu = text_encoder_offload_device()
|
||||
model_dir = os.path.join(folder_paths.models_dir, "blip")
|
||||
model_url = "https://github.com/AUTOMATIC1111/TorchDeepDanbooru/releases/download/v1/model-resnet_custom_v3.pt"
|
||||
model_sha = "3841542cda4dd037da12a565e854b3347bb2eec8fbcd95ea3941b2c68990a355"
|
||||
re_special = re.compile(r"([\\()])")
|
||||
|
||||
|
||||
def load_danbooru(device_mode):
|
||||
global danbooru
|
||||
if danbooru is None:
|
||||
blip_dir = os.path.join(folder_paths.models_dir, "blip")
|
||||
if not os.path.exists(blip_dir) and not is_junction(blip_dir):
|
||||
os.makedirs(blip_dir, exist_ok=True)
|
||||
if not os.path.exists(model_dir) and not is_junction(model_dir):
|
||||
os.makedirs(model_dir, exist_ok=True)
|
||||
|
||||
files = download_model(
|
||||
model_path=blip_dir,
|
||||
model_url=model_url,
|
||||
ext_filter=[".pt"],
|
||||
download_name="model-resnet_custom_v3.pt",
|
||||
)
|
||||
model_path = os.path.join(model_dir, "model-resnet_custom_v3.pt")
|
||||
download_file(model_url, model_path, model_sha)
|
||||
|
||||
from .models.deepbooru_model import DeepDanbooruModel
|
||||
|
||||
danbooru = DeepDanbooruModel()
|
||||
danbooru.load_state_dict(torch.load(files[0], map_location="cpu"))
|
||||
danbooru.load_state_dict(torch.load(model_path, map_location="cpu"))
|
||||
danbooru.eval()
|
||||
|
||||
if device_mode != "CPU":
|
||||
|
||||
@@ -1,12 +1,14 @@
|
||||
from .segmenter import ISNetLoader, ISNetSegment
|
||||
from .segmenter import ISNetLoader, ISNetSegment, DownloadISNetModel
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ISNetLoader": ISNetLoader,
|
||||
"ISNetSegment": ISNetSegment,
|
||||
"DownloadISNetModel": DownloadISNetModel,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ISNetLoader": "ISNet Loader",
|
||||
"ISNetSegment": "ISNet Segment",
|
||||
"DownloadISNetModel": "Download and Load ISNet Model",
|
||||
}
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
|
||||
+46
-23
@@ -12,17 +12,30 @@ import folder_paths
|
||||
import comfy.model_management as model_management
|
||||
import comfy.utils
|
||||
|
||||
from ..model_utils import download_model
|
||||
from ..utils import pil2tensor, tensor2pil, numpy2pil
|
||||
from ..model_utils import download_file
|
||||
from ..utils import pil2tensor, tensor2pil
|
||||
from ..logger import logger
|
||||
|
||||
|
||||
isnets = {}
|
||||
cache_size = [1024, 1024]
|
||||
gpu = model_management.get_torch_device()
|
||||
cpu = torch.device("cpu")
|
||||
model_dir = os.path.join(folder_paths.models_dir, "isnet")
|
||||
model_url = "https://huggingface.co/NimaBoscarino/IS-Net_DIS-general-use/resolve/main/isnet-general-use.pth"
|
||||
cache_size = [1024, 1024]
|
||||
models = {
|
||||
"isnet-general-use.pth": {
|
||||
"url": "https://huggingface.co/NimaBoscarino/IS-Net_DIS-general-use/resolve/main/isnet-general-use.pth",
|
||||
"sha": "9e1aafea58f0b55d0c35077e0ceade6ba1ba2bce372fd4f8f77215391f3fac13",
|
||||
},
|
||||
"isnetis.pth": {
|
||||
"url": "https://github.com/Sanster/models/releases/download/isnetis/isnetis.pth",
|
||||
"sha": "90a970badbd99ca7839b4e0beb09a36565d24edba7e4a876de23c761981e79e0",
|
||||
},
|
||||
"RMBG-1.4.bin": {
|
||||
"url": "https://huggingface.co/briaai/RMBG-1.4/resolve/main/pytorch_model.bin",
|
||||
"sha": "59569acdb281ac9fc9f78f9d33b6f9f17f68e25086b74f9025c35bb5f2848967",
|
||||
},
|
||||
}
|
||||
|
||||
folder_paths.folder_names_and_paths["isnet"] = (
|
||||
[model_dir],
|
||||
@@ -134,7 +147,6 @@ class ISNetLoader:
|
||||
return {
|
||||
"required": {
|
||||
"model_name": (folder_paths.get_filename_list("isnet"),),
|
||||
"model_override": ("STRING", {"default": "None"}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -142,15 +154,33 @@ class ISNetLoader:
|
||||
FUNCTION = "load_isnet"
|
||||
CATEGORY = "Art Venture/Segmentation"
|
||||
|
||||
def load_isnet(self, model_name, model_override="None"):
|
||||
if model_override != "None":
|
||||
if model_override not in folder_paths.get_filename_list("isnet"):
|
||||
logger.warning(f"Model override {model_override} not found. Use {model_name} instead.")
|
||||
else:
|
||||
model_name = model_override
|
||||
def load_isnet(self, model_name):
|
||||
return (load_isnet_model(model_name),)
|
||||
|
||||
model = load_isnet_model(model_name)
|
||||
return (model,)
|
||||
|
||||
class DownloadISNetModel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model_name": (list(models.keys()),),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("ISNET_MODEL",)
|
||||
FUNCTION = "download_isnet"
|
||||
CATEGORY = "Art Venture/Segmentation"
|
||||
|
||||
def download_isnet(self, model_name):
|
||||
if model_name not in folder_paths.get_filename_list("isnet"):
|
||||
model_info = models[model_name]
|
||||
download_file(
|
||||
model_info["url"],
|
||||
os.path.join(model_dir, model_name),
|
||||
model_info["sha"],
|
||||
)
|
||||
|
||||
return (load_isnet_model(model_name),)
|
||||
|
||||
|
||||
class ISNetSegment:
|
||||
@@ -179,15 +209,8 @@ class ISNetSegment:
|
||||
return (images, masks)
|
||||
|
||||
if isnet_model is None:
|
||||
ckpts = folder_paths.get_filename_list("isnet")
|
||||
if len(ckpts) == 0:
|
||||
ckpts = download_model(
|
||||
model_path=model_dir,
|
||||
model_url=model_url,
|
||||
ext_filter=[".pth"],
|
||||
download_name="isnet-general-use.pth",
|
||||
)
|
||||
isnet_model = load_isnet_model(ckpts[0])
|
||||
downloader = DownloadISNetModel()
|
||||
isnet_model = downloader.download_isnet("isnet-general-use.pth")[0]
|
||||
|
||||
device = gpu if device_mode != "CPU" else cpu
|
||||
isnet_model = isnet_model.to(device)
|
||||
@@ -198,7 +221,7 @@ class ISNetSegment:
|
||||
for image in images:
|
||||
mask = predict(isnet_model, image, device)
|
||||
mask_im = tensor2pil(mask.permute(1, 2, 0))
|
||||
cropped = Image.new("RGBA", mask_im.size, (0,0,0,0))
|
||||
cropped = Image.new("RGBA", mask_im.size, (0, 0, 0, 0))
|
||||
cropped.paste(tensor2pil(image), mask=mask_im)
|
||||
|
||||
masks.append(mask)
|
||||
|
||||
+81
-31
@@ -1,7 +1,12 @@
|
||||
import os
|
||||
import re
|
||||
import torch
|
||||
import hashlib
|
||||
import urllib.request
|
||||
import urllib.error
|
||||
from tqdm import tqdm
|
||||
from urllib.parse import urlparse
|
||||
from typing import Dict, Optional
|
||||
|
||||
|
||||
def natural_sort_key(s, regex=re.compile("([0-9]+)")):
|
||||
@@ -56,44 +61,89 @@ def load_file_from_url(
|
||||
return cached_file
|
||||
|
||||
|
||||
def download_model(
|
||||
model_path: str,
|
||||
model_url: str = None,
|
||||
ext_filter=None,
|
||||
download_name=None,
|
||||
ext_blacklist=None,
|
||||
) -> list:
|
||||
"""
|
||||
A one-and done loader to try finding the desired models in specified directories.
|
||||
def calculate_sha(file: str, force=False) -> Optional[str]:
|
||||
sha_file = f"{file}.sha"
|
||||
|
||||
@param download_name: Specify to download from model_url immediately.
|
||||
@param model_url: If no other models are found, this will be downloaded on upscale.
|
||||
@param model_path: The location to store/find models in.
|
||||
@param ext_filter: An optional list of filename extensions to filter by
|
||||
@return: A list of paths containing the desired model(s)
|
||||
# Check if the .sha file exists
|
||||
if not force and os.path.exists(sha_file):
|
||||
try:
|
||||
with open(sha_file, "r") as f:
|
||||
stored_hash = f.read().strip()
|
||||
if stored_hash:
|
||||
return stored_hash
|
||||
except IOError as e:
|
||||
print(f"Failed to read hash: {e}")
|
||||
|
||||
# Calculate the hash if the .sha file doesn't exist or is empty
|
||||
try:
|
||||
with open(file, "rb") as fp:
|
||||
file_hash = hashlib.sha256()
|
||||
while chunk := fp.read(8192):
|
||||
file_hash.update(chunk)
|
||||
calculated_hash = file_hash.hexdigest()
|
||||
|
||||
# Write the calculated hash to the .sha file
|
||||
try:
|
||||
with open(sha_file, "w") as f:
|
||||
f.write(calculated_hash)
|
||||
except IOError as e:
|
||||
print(f"Failed to write hash to {sha_file}: {e}")
|
||||
|
||||
return calculated_hash
|
||||
except IOError as e:
|
||||
print(f"Failed to read file {file}: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def download_file(url: str, dst: str, sha256sum: Optional[str] = None) -> Dict[str, Optional[str]]:
|
||||
"""
|
||||
output = []
|
||||
Downloads a file from a URL to a destination path, optionally verifying its SHA-256 checksum.
|
||||
|
||||
:param url: URL of the file to download
|
||||
:param dst: Destination path to save the downloaded file
|
||||
:param sha256sum: Optional SHA-256 checksum to verify the downloaded file
|
||||
:return: Dictionary with file path, download status, calculated checksum, and checksum match status
|
||||
"""
|
||||
# Ensure the directory exists
|
||||
os.makedirs(os.path.dirname(dst), exist_ok=True)
|
||||
|
||||
file_exists = os.path.isfile(dst)
|
||||
file_checksum = None
|
||||
checksum_match = None
|
||||
downloaded = False
|
||||
|
||||
try:
|
||||
for full_path in walk_files(model_path, allowed_extensions=ext_filter):
|
||||
if os.path.islink(full_path) and not os.path.exists(full_path):
|
||||
print(f"Skipping broken symlink: {full_path}")
|
||||
continue
|
||||
if ext_blacklist is not None and any(full_path.endswith(x) for x in ext_blacklist):
|
||||
continue
|
||||
if full_path not in output:
|
||||
output.append(full_path)
|
||||
if file_exists:
|
||||
file_checksum = calculate_sha(dst)
|
||||
if sha256sum:
|
||||
checksum_match = file_checksum == sha256sum
|
||||
if not checksum_match:
|
||||
os.remove(dst)
|
||||
|
||||
if model_url is not None and len(output) == 0:
|
||||
if download_name is not None:
|
||||
output.append(load_file_from_url(model_url, model_dir=model_path, file_name=download_name))
|
||||
else:
|
||||
output.append(model_url)
|
||||
if not file_exists or checksum_match == False:
|
||||
with tqdm(unit="B", unit_scale=True, unit_divisor=1024, miniters=1, desc=dst.split("/")[-1]) as t:
|
||||
|
||||
except Exception:
|
||||
pass
|
||||
def reporthook(blocknum, blocksize, totalsize):
|
||||
if t.total is None and totalsize > 0:
|
||||
t.total = totalsize
|
||||
read_so_far = blocknum * blocksize
|
||||
t.update(max(0, read_so_far - t.n))
|
||||
|
||||
return output
|
||||
urllib.request.urlretrieve(url, dst, reporthook=reporthook)
|
||||
downloaded = True
|
||||
|
||||
file_checksum = calculate_sha(dst, force=True)
|
||||
if sha256sum:
|
||||
checksum_match = file_checksum == sha256sum
|
||||
|
||||
except urllib.error.URLError as ex:
|
||||
print("Download failed:", ex)
|
||||
if os.path.isfile(dst):
|
||||
os.remove(dst)
|
||||
except Exception as ex:
|
||||
print("An error occurred:", ex)
|
||||
finally:
|
||||
return {"file": dst, "downloaded": downloaded, "sha": file_checksum, "match": checksum_match}
|
||||
|
||||
|
||||
def load_jit_torch_file(model_path: str):
|
||||
|
||||
@@ -62,8 +62,6 @@ from .llm import (
|
||||
NODE_DISPLAY_NAME_MAPPINGS as LLM_NODE_DISPLAY_NAME_MAPPINGS,
|
||||
)
|
||||
|
||||
from .model_utils import load_file_from_url
|
||||
|
||||
|
||||
class AVVAELoader(VAELoader):
|
||||
@classmethod
|
||||
|
||||
@@ -54,23 +54,6 @@ def _construct_pip_command(package_name, version=None):
|
||||
return pip_install + [package_name]
|
||||
|
||||
|
||||
# modified from https://stackoverflow.com/questions/22058048/hashing-a-file-in-python
|
||||
def calculate_file_hash(filename: str, hash_every_n: int = 1):
|
||||
import hashlib
|
||||
|
||||
h = hashlib.sha256()
|
||||
b = bytearray(10 * 1024 * 1024) # read 10 megabytes at a time
|
||||
mv = memoryview(b)
|
||||
with open(filename, "rb", buffering=0) as f:
|
||||
i = 0
|
||||
# don't hash entire file, only portions of it if requested
|
||||
while n := f.readinto(mv):
|
||||
if i % hash_every_n == 0:
|
||||
h.update(mv[:n])
|
||||
i += 1
|
||||
return h.hexdigest()
|
||||
|
||||
|
||||
def get_dict_attribute(dict_inst: dict, name_string: str, default=None):
|
||||
nested_keys = name_string.split(".")
|
||||
value = dict_inst
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-art-venture"
|
||||
description = "A comprehensive set of custom nodes for ComfyUI, focusing on utilities for image processing, JSON manipulation, model operations and working with object via URLs"
|
||||
version = "1.0.2"
|
||||
version = "1.0.5"
|
||||
license = "LICENSE"
|
||||
dependencies = ["timm==0.6.13", "transformers", "fairscale", "pycocoevalcap", "opencv-python", "qrcode[pil]", "pytorch_lightning", "kornia", "pydantic", "segment_anything", "boto3>=1.34.101"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user