15 Commits
Author SHA1 Message Date
Tung Nguyen (Blockchain) 50abaace75 Merge pull request #58 from sipherxyz/develop
Release v1.0.6
2024-11-04 21:05:04 +07:00
Tung Nguyen d0cf2abdaf Bump version 1.0.6 2024-11-04 21:01:54 +07:00
Tung Nguyen 36ed264742 add new LLM models 2024-11-04 21:01:02 +07:00
Tung Nguyen fb8bc917ba add TextSwitchCase node 2024-11-04 21:01:02 +07:00
Tung Nguyen (Blockchain) 8d538c9678 Merge pull request #56 from sipherxyz/develop
Fix invalid path aux.py in windows
2024-10-31 20:34:42 +07:00
Tung Nguyen 8034af6478 Bump version to 1.0.5 2024-10-31 13:33:19 +00:00
Tung Nguyen 713f6de761 rename aux.py to preprocessor.py to fix windows issue 2024-10-31 13:32:42 +00:00
Tung Nguyen (Blockchain) 83c3732201 Update pyproject.toml
Bump version to 1.0.4
2024-10-31 15:00:55 +07:00
Tung Nguyen (Blockchain) de986044b3 Merge pull request #54 from sipherxyz/develop
fix controlnet preprocessor node
2024-10-31 15:00:05 +07:00
Tung Nguyen bb6994f677 fix controlnet preprocessor node 2024-10-31 07:57:50 +00:00
Tung Nguyen (Blockchain) 4405434b9c Merge pull request #53 from sipherxyz/develop
Update model download with sha256 validate
2024-10-30 17:25:21 +07:00
Tung Nguyen e7e9e58c66 bump version to 1.0.3 2024-10-30 17:22:14 +07:00
Tung Nguyen 51dd8fcb7c update model download with sha validate 2024-10-30 17:06:33 +07:00
Tung Nguyen (Blockchain) 4be544aa9e Merge pull request #52 from sipherxyz/develop
Release 1.0.2
2024-10-30 12:30:07 +07:00
Tung Nguyen (Blockchain) 08fa873f2c Merge pull request #49 from sipherxyz/develop
use torch.jit to load Lama model
2024-10-25 14:35:20 +07:00
19 changed files with 618 additions and 376 deletions
+16 -1
View File
@@ -67,6 +67,21 @@ Randomizes the order of lines in a multiline string.
![text random multiline](https://github.com/user-attachments/assets/86f811e3-579e-4ccc-81a3-e216cd851d3c)
#### TextSwitchCase
Switch between multiple cases based on a condition.
**Inputs:**
- `switch_cases`: Switch cases, separated by new lines
- `condition`: Condition to switch on
- `default_value`: Default value when no condition matches
- `delimiter`: Delimiter between case and value, default is `:`
The `switch_cases` format is `case<delimiter>value`, where `case` is the condition to match and `value` is the value to return when the condition matches. You can have new lines in the value to return multiple lines.
![text switch case](https://github.com/user-attachments/assets/4c5450a8-6a3a-4d3c-8c2a-c6e3a33cb95f)
### Inpainting Nodes
#### PrepareImageAndMaskForInpaint
@@ -199,4 +214,4 @@ Handles completion requests to LLMs.
- `config`: Model configuration
- `seed`: Random seed
![Screenshot 2024-10-30 at 11 20 12](https://github.com/user-attachments/assets/45b8d4fd-57cd-4bd9-8274-d3e6ac4ef938)
![Screenshot 2024-10-30 at 11 20 12](https://github.com/user-attachments/assets/45b8d4fd-57cd-4bd9-8274-d3e6ac4ef938)
+8 -43
View File
@@ -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"}),
+106
View File
@@ -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)
-142
View File
@@ -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,)
+4 -8
View File
@@ -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
+3 -1
View File
@@ -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",
}
+41 -21
View File
@@ -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)
+8 -11
View File
@@ -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":
+3 -1
View File
@@ -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
View File
@@ -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)
+12 -2
View File
@@ -11,6 +11,8 @@ from ..utils import ensure_package, tensor2pil, pil2base64
gpt_models = [
"gpt-3.5-turbo",
"gpt-3.5-turbo-16k",
"gpt-4o",
"gpt-4o-mini",
"gpt-4-turbo",
"gpt-4-vision-preview",
"gpt-4-turbo-preview",
@@ -20,9 +22,16 @@ gpt_models = [
"gpt-4",
]
gpt_vision_models = ["gpt-4-turbo", "gpt-4-turbo-preview", "gpt-4-vision-preview"]
gpt_vision_models = ["gpt-4o", "gpt-4o-mini", "gpt-4-turbo", "gpt-4-turbo-preview", "gpt-4-vision-preview"]
claude3_models = ["claude-3-opus-20240229", "claude-3-sonnet-20240229", "claude-3-haiku-20240307"]
claude3_models = [
"claude-3-5-sonnet-latest",
"claude-3-5-sonnet-20241022",
"claude-3-opus-latest",
"claude-3-opus-20240229",
"claude-3-sonnet-20240229",
"claude-3-haiku-20240307",
]
claude2_models = ["claude-2.1"]
aws_regions = [
@@ -40,6 +49,7 @@ aws_regions = [
bedrock_anthropic_versions = ["bedrock-2023-05-31"]
bedrock_claude3_models = [
"anthropic.claude-3-5-sonnet-20241022-v2:0",
"anthropic.claude-3-haiku-20240307-v1:0",
"anthropic.claude-3-sonnet-20240229-v1:0",
"anthropic.claude-3-opus-20240229-v1:0",
+81 -31
View File
@@ -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):
-2
View File
@@ -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
+88 -1
View File
@@ -20,6 +20,42 @@ from .utils import pil2tensor, tensor2pil, ensure_package, get_dict_attribute
MAX_RESOLUTION = 8192
class AnyType(str):
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
def __ne__(self, __value: object) -> bool:
return False
class FlexibleOptionalInputType(dict):
"""A special class to make flexible nodes that pass data to our python handlers.
Enables both flexible/dynamic input types (like for Any Switch) or a dynamic number of inputs
(like for Any Switch, Context Switch, Context Merge, Power Lora Loader, etc).
Note, for ComfyUI, all that's needed is the `__contains__` override below, which tells ComfyUI
that our node will handle the input, regardless of what it is.
However, with https://github.com/comfyanonymous/ComfyUI/pull/2666 a large change would occur
requiring more details on the input itself. There, we need to return a list/tuple where the first
item is the type. This can be a real type, or use the AnyType for additional flexibility.
This should be forwards compatible unless more changes occur in the PR.
"""
def __init__(self, type):
self.type = type
def __getitem__(self, key):
return (self.type,)
def __contains__(self, key):
return True
any_type = AnyType("*")
def prepare_image_for_preview(image: Image.Image, output_dir: str, prefix=None):
if prefix is None:
prefix = "preview_" + "".join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5))
@@ -242,7 +278,7 @@ class UtilLoadImageAsMaskFromUrl(UtilLoadImageFromUrl):
mask = np.array(mask, dtype=np.float32) / 255.0
mask = torch.from_numpy(mask)
if channel == "alpha":
mask = 1. - mask
mask = 1.0 - mask
else:
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
@@ -531,6 +567,55 @@ class UtilBooleanPrimitive:
return (value, str(value))
class UtilTextSwitchCase:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"switch_cases": (
"STRING",
{
"default": "",
"multiline": True,
"dynamicPrompts": False,
"placeholder": "case_1:output_1\ncase_2:output_2\nthat span multiple lines\ncase_3:output_3",
},
),
"condition": ("STRING", {"default": ""}),
"default_value": ("STRING", {"default": ""}),
"delimiter": ("STRING", {"default": ":"}),
}
}
RETURN_TYPES = ("STRING",)
CATEGORY = "Art Venture/Utils"
FUNCTION = "text_switch_case"
def text_switch_case(self, switch_cases: str, condition: str, default_value: str, delimiter: str = ":"):
# Split into cases first
cases = switch_cases.split("\n")
current_case = None
current_output = []
for line in cases:
if delimiter in line:
# Process previous case if exists
if current_case is not None and condition == current_case:
return ("\n".join(current_output),)
# Start new case
current_case, output = line.split(delimiter, 1)
current_output = [output]
elif current_case is not None:
current_output.append(line)
# Check last case
if current_case is not None and condition == current_case:
return ("\n".join(current_output),)
return (default_value,)
class UtilImageMuxer:
@classmethod
def INPUT_TYPES(s):
@@ -1194,6 +1279,7 @@ NODE_CLASS_MAPPINGS = {
"NumberScaler": UtilNumberScaler,
"MergeModels": UtilModelMerge,
"TextRandomMultiline": UtilTextRandomMultiline,
"TextSwitchCase": UtilTextSwitchCase,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LoadImageFromUrl": "Load Image From URL",
@@ -1229,4 +1315,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"NumberScaler": "Number Scaler",
"MergeModels": "Merge Models",
"TextRandomMultiline": "Text Random Multiline",
"TextSwitchCase": "Text Switch Case",
}
-17
View File
@@ -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
View File
@@ -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.6"
license = "LICENSE"
dependencies = ["timm==0.6.13", "transformers", "fairscale", "pycocoevalcap", "opencv-python", "qrcode[pil]", "pytorch_lightning", "kornia", "pydantic", "segment_anything", "boto3>=1.34.101"]
+53
View File
@@ -0,0 +1,53 @@
import { app } from '../../../scripts/app.js';
import { ComfyWidgets } from '../../../scripts/widgets.js';
import {
addKVState,
chainCallback,
hideWidgetForGood,
addWidgetChangeCallback,
} from './utils.js';
function addTextSwitchCaseWidget(nodeType) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
const dataWidget = this.widgets.find((w) => w.name === 'switch_cases');
const delimiterWidget = this.widgets.find((w) => w.name === 'delimiter');
this.widgets = this.widgets.filter((w) => w.name !== 'condition');
let conditionCombo = null;
const updateConditionCombo = () => {
if (!delimiterWidget.value) return;
const cases = (dataWidget.value ?? '')
.split('\n')
.filter((line) => line.includes(delimiterWidget.value))
.map((line) => line.split(delimiterWidget.value)[0]);
if (!conditionCombo) {
conditionCombo = ComfyWidgets['COMBO'](this, 'condition', [
['__default__', ...(cases ?? [])],
]).widget;
} else {
conditionCombo.options.values = ['__default__', ...cases];
}
};
updateConditionCombo();
dataWidget.inputEl.addEventListener('input', updateConditionCombo);
addWidgetChangeCallback(delimiterWidget, updateConditionCombo);
});
}
app.registerExtension({
name: 'ArtVenture.TextSwitchCase',
async beforeRegisterNodeDef(nodeType, nodeData) {
if (!nodeData) return;
if (nodeData.name !== 'TextSwitchCase') {
return;
}
addKVState(nodeType);
addTextSwitchCaseWidget(nodeType);
},
});
+2 -71
View File
@@ -3,6 +3,8 @@ import { api } from '../../../scripts/api.js';
import { $el } from '../../../scripts/ui.js';
import { createImageHost } from '../../../scripts/ui/imagePreview.js';
import { chainCallback, addKVState } from './utils.js';
const style = `
.comfy-img-preview video {
object-fit: contain;
@@ -15,24 +17,6 @@ const URL_REGEX = /^((blob:)?https?:\/\/|\/view\?|\/api\/view\?|data:image\/)/
const supportedNodes = ['LoadImageFromUrl', 'LoadImageAsMaskFromUrl', 'LoadVideoFromUrl'];
function chainCallback(object, property, callback) {
if (object == undefined) {
//This should not happen.
console.error('Tried to add callback to non-existant object');
return;
}
if (property in object) {
const callback_orig = object[property];
object[property] = function () {
const r = callback_orig.apply(this, arguments);
callback.apply(this, arguments);
return r;
};
} else {
object[property] = callback;
}
}
function injectHidden(widget) {
widget.computeSize = (target_width) => {
if (widget.hidden) {
@@ -54,59 +38,6 @@ function injectHidden(widget) {
});
}
function addKVState(nodeType) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
chainCallback(this, 'onConfigure', function (info) {
if (!this.widgets) {
//Node has no widgets, there is nothing to restore
return;
}
if (typeof info.widgets_values != 'object') {
//widgets_values is in some unknown inactionable format
return;
}
let widgetDict = info.widgets_values;
if (widgetDict.length == undefined) {
for (let w of this.widgets) {
if (w.name in widgetDict) {
w.value = widgetDict[w.name];
} else {
//attempt to restore default value
let inputs = LiteGraph.getNodeType(this.type).nodeData.input;
let initialValue = null;
if (inputs?.required?.hasOwnProperty(w.name)) {
if (inputs.required[w.name][1]?.hasOwnProperty('default')) {
initialValue = inputs.required[w.name][1].default;
} else if (inputs.required[w.name][0].length) {
initialValue = inputs.required[w.name][0][0];
}
} else if (inputs?.optional?.hasOwnProperty(w.name)) {
if (inputs.optional[w.name][1]?.hasOwnProperty('default')) {
initialValue = inputs.optional[w.name][1].default;
} else if (inputs.optional[w.name][0].length) {
initialValue = inputs.optional[w.name][0][0];
}
}
if (initialValue) {
w.value = initialValue;
}
}
}
}
});
chainCallback(this, 'onSerialize', function (info) {
info.widgets_values = {};
if (!this.widgets) {
//object has no widgets, there is nothing to store
return;
}
for (let w of this.widgets) {
info.widgets_values[w.name] = w.value;
}
});
});
}
function migrateWidget(nodeType, oldWidgetName, newWidgetName) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
if (!this.widgets) return;
+146
View File
@@ -0,0 +1,146 @@
export const CONVERTED_TYPE = "converted-widget";
export function hideWidgetForGood(node, widget, suffix = "") {
widget.origType = widget.type;
widget.origComputeSize = widget.computeSize;
widget.computeSize = () => [0, -4]; // -4 is due to the gap litegraph adds between widgets automatically
widget.type = CONVERTED_TYPE + suffix;
// Hide any linked widgets, e.g. seed+seedControl
if (widget.linkedWidgets) {
for (const w of widget.linkedWidgets) {
hideWidgetForGood(node, w, ":" + widget.name);
}
}
}
const doesInputWithNameExist = (node, name) => {
return node.inputs ? node.inputs.some((input) => input.name === name) : false;
};
const HIDDEN_TAG = "tschide";
const origProps = {};
// Toggle Widget + change size
export function toggleWidget(node, widget, show = false, suffix = "", updateSize = true) {
if (!widget || doesInputWithNameExist(node, widget.name)) return;
// Store the original properties of the widget if not already stored
if (!origProps[widget.name]) {
origProps[widget.name] = {
origType: widget.type,
origComputeSize: widget.computeSize,
};
}
const origSize = node.size;
// Set the widget type and computeSize based on the show flag
widget.type = show ? origProps[widget.name].origType : HIDDEN_TAG + suffix;
widget.computeSize = show
? origProps[widget.name].origComputeSize
: () => [0, -4];
// Recursively handle linked widgets if they exist
widget.linkedWidgets?.forEach((w) =>
toggleWidget(node, w, ":" + widget.name, show)
);
// Calculate the new height for the node based on its computeSize method
if (updateSize) {
const newHeight = node.computeSize()[1];
node.setSize([node.size[0], newHeight]);
}
}
export function addWidgetChangeCallback(widget, callback) {
let widgetValue = widget.value;
let originalDescriptor = Object.getOwnPropertyDescriptor(widget, "value");
Object.defineProperty(widget, "value", {
get() {
return originalDescriptor && originalDescriptor.get
? originalDescriptor.get.call(widget)
: widgetValue;
},
set(newVal) {
if (originalDescriptor && originalDescriptor.set) {
originalDescriptor.set.call(widget, newVal);
} else {
widgetValue = newVal;
}
callback(newVal);
},
});
}
export function chainCallback(object, property, callback) {
if (object == undefined) {
//This should not happen.
console.error("Tried to add callback to non-existant object");
return;
}
if (property in object) {
const callback_orig = object[property];
object[property] = function () {
const r = callback_orig.apply(this, arguments);
callback.apply(this, arguments);
return r;
};
} else {
object[property] = callback;
}
}
export function addKVState(nodeType) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
chainCallback(this, 'onConfigure', function (info) {
if (!this.widgets) {
//Node has no widgets, there is nothing to restore
return;
}
if (typeof info.widgets_values != 'object') {
//widgets_values is in some unknown inactionable format
return;
}
let widgetDict = info.widgets_values;
if (widgetDict.length == undefined) {
for (let w of this.widgets) {
if (w.name in widgetDict) {
w.value = widgetDict[w.name];
} else {
//attempt to restore default value
let inputs = LiteGraph.getNodeType(this.type).nodeData.input;
let initialValue = null;
if (inputs?.required?.hasOwnProperty(w.name)) {
if (inputs.required[w.name][1]?.hasOwnProperty('default')) {
initialValue = inputs.required[w.name][1].default;
} else if (inputs.required[w.name][0].length) {
initialValue = inputs.required[w.name][0][0];
}
} else if (inputs?.optional?.hasOwnProperty(w.name)) {
if (inputs.optional[w.name][1]?.hasOwnProperty('default')) {
initialValue = inputs.optional[w.name][1].default;
} else if (inputs.optional[w.name][0].length) {
initialValue = inputs.optional[w.name][0][0];
}
}
if (initialValue) {
w.value = initialValue;
}
}
}
}
});
chainCallback(this, 'onSerialize', function (info) {
info.widgets_values = {};
if (!this.widgets) {
//object has no widgets, there is nothing to store
return;
}
for (let w of this.widgets) {
info.widgets_values[w.name] = w.value;
}
});
});
}