13 Commits
Author SHA1 Message Date
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 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) ![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 ### Inpainting Nodes
#### PrepareImageAndMaskForInpaint #### PrepareImageAndMaskForInpaint
@@ -199,4 +214,4 @@ Handles completion requests to LLMs.
- `config`: Model configuration - `config`: Model configuration
- `seed`: Random seed - `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 import folder_paths
from nodes import ControlNetLoader, ControlNetApply, ControlNetApplyAdvanced 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 from .advanced import comfy_load_controlnet
control_net_preprocessors["tile"] = (DummyPreprocessor, [])
def load_controlnet(control_net_name, control_net_override="None", timestep_keyframe=None): def load_controlnet(control_net_name, control_net_override="None", timestep_keyframe=None):
if control_net_override != "None": if control_net_override != "None":
if control_net_override not in folder_paths.get_filename_list("controlnet"): 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) 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): def detect_controlnet(preprocessor: str, sd_version: str):
controlnets = folder_paths.get_filename_list("controlnet") controlnets = folder_paths.get_filename_list("controlnet")
controlnets = filter(lambda x: sd_version in x, controlnets) controlnets = filter(lambda x: sd_version in x, controlnets)
@@ -103,15 +73,13 @@ class AVControlNetLoader(ControlNetLoader):
class AV_ControlNetPreprocessor: class AV_ControlNetPreprocessor:
preprocessors = list(control_net_preprocessors.keys())
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
"required": { "required": {
"image": ("IMAGE",), "image": ("IMAGE",),
"preprocessor": (["None", "tile"] + s.preprocessors,), "preprocessor": (["None"] + preprocessors,),
"sd_version": (["sd15", "sdxl", "sdxl_t2i"],), "sd_version": (["sd15", "sdxl"],),
}, },
"optional": { "optional": {
"resolution": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}), "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"): def detect_controlnet(self, image, preprocessor, sd_version, resolution=512, preprocessor_override="None"):
if preprocessor_override != "None": if preprocessor_override != "None":
if preprocessor_override not in control_net_preprocessors: if preprocessor_override not in preprocessors:
print( print(
f"Warning: Not found ControlNet preprocessor {preprocessor_override}. Use {preprocessor} instead." f"Warning: Not found ControlNet preprocessor {preprocessor_override}. Use {preprocessor} instead."
) )
@@ -141,7 +109,6 @@ class AV_ControlNetPreprocessor:
class AVControlNetEfficientStacker: class AVControlNetEfficientStacker:
controlnets = folder_paths.get_filename_list("controlnet") controlnets = folder_paths.get_filename_list("controlnet")
preprocessors = list(control_net_preprocessors.keys())
@classmethod @classmethod
def INPUT_TYPES(s): 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}), "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}), "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"preprocessor": (["None"] + s.preprocessors,), "preprocessor": (["None"] + preprocessors,),
}, },
"optional": { "optional": {
"cnet_stack": ("CONTROL_NET_STACK",), "cnet_stack": ("CONTROL_NET_STACK",),
@@ -219,7 +186,7 @@ class AVControlNetEfficientStackerSimple(AVControlNetEfficientStacker):
"FLOAT", "FLOAT",
{"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}, {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01},
), ),
"preprocessor": (["None"] + s.preprocessors,), "preprocessor": (["None"] + preprocessors,),
}, },
"optional": { "optional": {
"cnet_stack": ("CONTROL_NET_STACK",), "cnet_stack": ("CONTROL_NET_STACK",),
@@ -242,7 +209,6 @@ class AVControlNetEfficientStackerSimple(AVControlNetEfficientStacker):
class AVControlNetEfficientLoader(ControlNetApply): class AVControlNetEfficientLoader(ControlNetApply):
controlnets = folder_paths.get_filename_list("controlnet") controlnets = folder_paths.get_filename_list("controlnet")
preprocessors = list(control_net_preprocessors.keys())
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -255,7 +221,7 @@ class AVControlNetEfficientLoader(ControlNetApply):
"FLOAT", "FLOAT",
{"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}, {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01},
), ),
"preprocessor": (["None"] + s.preprocessors,), "preprocessor": (["None"] + preprocessors,),
}, },
"optional": { "optional": {
"control_net_override": ("STRING", {"default": "None"}), "control_net_override": ("STRING", {"default": "None"}),
@@ -295,7 +261,6 @@ class AVControlNetEfficientLoader(ControlNetApply):
class AVControlNetEfficientLoaderAdvanced(ControlNetApplyAdvanced): class AVControlNetEfficientLoaderAdvanced(ControlNetApplyAdvanced):
controlnets = folder_paths.get_filename_list("controlnet") controlnets = folder_paths.get_filename_list("controlnet")
preprocessors = list(control_net_preprocessors.keys())
@classmethod @classmethod
def INPUT_TYPES(s): 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}), "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}), "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"preprocessor": (["None"] + s.preprocessors,), "preprocessor": (["None"] + preprocessors,),
}, },
"optional": { "optional": {
"control_net_override": ("STRING", {"default": "None"}), "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 folder_paths
import comfy.model_management as model_management import comfy.model_management as model_management
from ...model_utils import download_model from ...model_utils import download_file
lama = None lama = None
@@ -33,14 +33,10 @@ def pad_tensor_to_modulo(img, mod):
def load_model(): def load_model():
global lama global lama
if lama is None: if lama is None:
files = download_model( model_path = os.path.join(model_dir, "big-lama.pt")
model_path=model_dir, download_file(model_url, model_path, model_sha)
model_url=model_url,
ext_filter=[".pt"],
download_name="big-lama.pt",
)
lama = torch.jit.load(files[0], map_location="cpu") lama = torch.jit.load(model_path, map_location="cpu")
lama.eval() lama.eval()
return lama 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 from .danbooru import DeepDanbooruCaption
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"BLIPLoader": BlipLoader, "BLIPLoader": BlipLoader,
"BLIPCaption": BlipCaption, "BLIPCaption": BlipCaption,
"DownloadAndLoadBlip": DownloadAndLoadBlip,
"DeepDanbooruCaption": DeepDanbooruCaption, "DeepDanbooruCaption": DeepDanbooruCaption,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"BLIPLoader": "BLIP Loader", "BLIPLoader": "BLIP Loader",
"BLIPCaption": "BLIP Caption", "BLIPCaption": "BLIP Caption",
"DownloadAndLoadBlip": "Download and Load BLIP Model",
"DeepDanbooruCaption": "Deep Danbooru Caption", "DeepDanbooruCaption": "Deep Danbooru Caption",
} }
+41 -21
View File
@@ -7,17 +7,25 @@ from torchvision.transforms.functional import InterpolationMode
import folder_paths import folder_paths
from comfy.model_management import text_encoder_device, text_encoder_offload_device, soft_empty_cache 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 from ..utils import tensor2pil
blips = {} blips = {}
blip_size = 384 blip_size = 384
gpu = text_encoder_device() gpu = text_encoder_device()
cpu = text_encoder_offload_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") 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"] = ( folder_paths.folder_names_and_paths["blip"] = (
[model_dir], [model_dir],
@@ -90,12 +98,7 @@ def join_caption(caption, prefix, suffix):
def blip_caption(model, image, min_length, max_length): def blip_caption(model, image, min_length, max_length):
image = tensor2pil(image) image = tensor2pil(image)
tensor = transformImage(image)
if "transformers==4.26.1" in packages(True):
print("Using Legacy `transformImaage()`")
tensor = transformImage_legacy(image)
else:
tensor = transformImage(image)
with torch.no_grad(): with torch.no_grad():
caption = model.generate( caption = model.generate(
@@ -122,8 +125,32 @@ class BlipLoader:
CATEGORY = "Art Venture/Captioning" CATEGORY = "Art Venture/Captioning"
def load_blip(self, model_name): def load_blip(self, model_name):
model = load_blip(model_name) return (load_blip(model_name),)
return (model,)
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: class BlipCaption:
@@ -173,15 +200,8 @@ class BlipCaption:
return ([join_caption("", prefix, suffix)],) return ([join_caption("", prefix, suffix)],)
if blip_model is None: if blip_model is None:
ckpts = folder_paths.get_filename_list("blip") downloader = DownloadAndLoadBlip()
if len(ckpts) == 0: blip_model = downloader.download_and_load_blip("model_base_caption_capfilt_large.pth")[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])
device = gpu if device_mode != "CPU" else cpu device = gpu if device_mode != "CPU" else cpu
blip_model = blip_model.to(device) 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 comfy.model_management import text_encoder_device, text_encoder_offload_device, soft_empty_cache
from ..image_utils import resize_image 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 ..utils import is_junction, tensor2pil
from .blip_node import join_caption from .blip_node import join_caption
@@ -15,28 +15,25 @@ danbooru = None
blip_size = 384 blip_size = 384
gpu = text_encoder_device() gpu = text_encoder_device()
cpu = text_encoder_offload_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_url = "https://github.com/AUTOMATIC1111/TorchDeepDanbooru/releases/download/v1/model-resnet_custom_v3.pt"
model_sha = "3841542cda4dd037da12a565e854b3347bb2eec8fbcd95ea3941b2c68990a355"
re_special = re.compile(r"([\\()])") re_special = re.compile(r"([\\()])")
def load_danbooru(device_mode): def load_danbooru(device_mode):
global danbooru global danbooru
if danbooru is None: if danbooru is None:
blip_dir = os.path.join(folder_paths.models_dir, "blip") if not os.path.exists(model_dir) and not is_junction(model_dir):
if not os.path.exists(blip_dir) and not is_junction(blip_dir): os.makedirs(model_dir, exist_ok=True)
os.makedirs(blip_dir, exist_ok=True)
files = download_model( model_path = os.path.join(model_dir, "model-resnet_custom_v3.pt")
model_path=blip_dir, download_file(model_url, model_path, model_sha)
model_url=model_url,
ext_filter=[".pt"],
download_name="model-resnet_custom_v3.pt",
)
from .models.deepbooru_model import DeepDanbooruModel from .models.deepbooru_model import DeepDanbooruModel
danbooru = 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() danbooru.eval()
if device_mode != "CPU": 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 = { NODE_CLASS_MAPPINGS = {
"ISNetLoader": ISNetLoader, "ISNetLoader": ISNetLoader,
"ISNetSegment": ISNetSegment, "ISNetSegment": ISNetSegment,
"DownloadISNetModel": DownloadISNetModel,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"ISNetLoader": "ISNet Loader", "ISNetLoader": "ISNet Loader",
"ISNetSegment": "ISNet Segment", "ISNetSegment": "ISNet Segment",
"DownloadISNetModel": "Download and Load ISNet Model",
} }
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] __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.model_management as model_management
import comfy.utils import comfy.utils
from ..model_utils import download_model from ..model_utils import download_file
from ..utils import pil2tensor, tensor2pil, numpy2pil from ..utils import pil2tensor, tensor2pil
from ..logger import logger from ..logger import logger
isnets = {} isnets = {}
cache_size = [1024, 1024]
gpu = model_management.get_torch_device() gpu = model_management.get_torch_device()
cpu = torch.device("cpu") cpu = torch.device("cpu")
model_dir = os.path.join(folder_paths.models_dir, "isnet") 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" models = {
cache_size = [1024, 1024] "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"] = ( folder_paths.folder_names_and_paths["isnet"] = (
[model_dir], [model_dir],
@@ -134,7 +147,6 @@ class ISNetLoader:
return { return {
"required": { "required": {
"model_name": (folder_paths.get_filename_list("isnet"),), "model_name": (folder_paths.get_filename_list("isnet"),),
"model_override": ("STRING", {"default": "None"}),
}, },
} }
@@ -142,15 +154,33 @@ class ISNetLoader:
FUNCTION = "load_isnet" FUNCTION = "load_isnet"
CATEGORY = "Art Venture/Segmentation" CATEGORY = "Art Venture/Segmentation"
def load_isnet(self, model_name, model_override="None"): def load_isnet(self, model_name):
if model_override != "None": return (load_isnet_model(model_name),)
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
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: class ISNetSegment:
@@ -179,15 +209,8 @@ class ISNetSegment:
return (images, masks) return (images, masks)
if isnet_model is None: if isnet_model is None:
ckpts = folder_paths.get_filename_list("isnet") downloader = DownloadISNetModel()
if len(ckpts) == 0: isnet_model = downloader.download_isnet("isnet-general-use.pth")[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])
device = gpu if device_mode != "CPU" else cpu device = gpu if device_mode != "CPU" else cpu
isnet_model = isnet_model.to(device) isnet_model = isnet_model.to(device)
@@ -198,7 +221,7 @@ class ISNetSegment:
for image in images: for image in images:
mask = predict(isnet_model, image, device) mask = predict(isnet_model, image, device)
mask_im = tensor2pil(mask.permute(1, 2, 0)) 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) cropped.paste(tensor2pil(image), mask=mask_im)
masks.append(mask) masks.append(mask)
+12 -2
View File
@@ -11,6 +11,8 @@ from ..utils import ensure_package, tensor2pil, pil2base64
gpt_models = [ gpt_models = [
"gpt-3.5-turbo", "gpt-3.5-turbo",
"gpt-3.5-turbo-16k", "gpt-3.5-turbo-16k",
"gpt-4o",
"gpt-4o-mini",
"gpt-4-turbo", "gpt-4-turbo",
"gpt-4-vision-preview", "gpt-4-vision-preview",
"gpt-4-turbo-preview", "gpt-4-turbo-preview",
@@ -20,9 +22,16 @@ gpt_models = [
"gpt-4", "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"] claude2_models = ["claude-2.1"]
aws_regions = [ aws_regions = [
@@ -40,6 +49,7 @@ aws_regions = [
bedrock_anthropic_versions = ["bedrock-2023-05-31"] bedrock_anthropic_versions = ["bedrock-2023-05-31"]
bedrock_claude3_models = [ bedrock_claude3_models = [
"anthropic.claude-3-5-sonnet-20241022-v2:0",
"anthropic.claude-3-haiku-20240307-v1:0", "anthropic.claude-3-haiku-20240307-v1:0",
"anthropic.claude-3-sonnet-20240229-v1:0", "anthropic.claude-3-sonnet-20240229-v1:0",
"anthropic.claude-3-opus-20240229-v1:0", "anthropic.claude-3-opus-20240229-v1:0",
+81 -31
View File
@@ -1,7 +1,12 @@
import os import os
import re import re
import torch import torch
import hashlib
import urllib.request
import urllib.error
from tqdm import tqdm
from urllib.parse import urlparse from urllib.parse import urlparse
from typing import Dict, Optional
def natural_sort_key(s, regex=re.compile("([0-9]+)")): def natural_sort_key(s, regex=re.compile("([0-9]+)")):
@@ -56,44 +61,89 @@ def load_file_from_url(
return cached_file return cached_file
def download_model( def calculate_sha(file: str, force=False) -> Optional[str]:
model_path: str, sha_file = f"{file}.sha"
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.
@param download_name: Specify to download from model_url immediately. # Check if the .sha file exists
@param model_url: If no other models are found, this will be downloaded on upscale. if not force and os.path.exists(sha_file):
@param model_path: The location to store/find models in. try:
@param ext_filter: An optional list of filename extensions to filter by with open(sha_file, "r") as f:
@return: A list of paths containing the desired model(s) 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: try:
for full_path in walk_files(model_path, allowed_extensions=ext_filter): if file_exists:
if os.path.islink(full_path) and not os.path.exists(full_path): file_checksum = calculate_sha(dst)
print(f"Skipping broken symlink: {full_path}") if sha256sum:
continue checksum_match = file_checksum == sha256sum
if ext_blacklist is not None and any(full_path.endswith(x) for x in ext_blacklist): if not checksum_match:
continue os.remove(dst)
if full_path not in output:
output.append(full_path)
if model_url is not None and len(output) == 0: if not file_exists or checksum_match == False:
if download_name is not None: with tqdm(unit="B", unit_scale=True, unit_divisor=1024, miniters=1, desc=dst.split("/")[-1]) as t:
output.append(load_file_from_url(model_url, model_dir=model_path, file_name=download_name))
else:
output.append(model_url)
except Exception: def reporthook(blocknum, blocksize, totalsize):
pass 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): 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, NODE_DISPLAY_NAME_MAPPINGS as LLM_NODE_DISPLAY_NAME_MAPPINGS,
) )
from .model_utils import load_file_from_url
class AVVAELoader(VAELoader): class AVVAELoader(VAELoader):
@classmethod @classmethod
+88 -1
View File
@@ -20,6 +20,42 @@ from .utils import pil2tensor, tensor2pil, ensure_package, get_dict_attribute
MAX_RESOLUTION = 8192 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): def prepare_image_for_preview(image: Image.Image, output_dir: str, prefix=None):
if prefix is None: if prefix is None:
prefix = "preview_" + "".join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5)) 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 = np.array(mask, dtype=np.float32) / 255.0
mask = torch.from_numpy(mask) mask = torch.from_numpy(mask)
if channel == "alpha": if channel == "alpha":
mask = 1. - mask mask = 1.0 - mask
else: else:
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu") mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
@@ -531,6 +567,55 @@ class UtilBooleanPrimitive:
return (value, str(value)) 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: class UtilImageMuxer:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -1194,6 +1279,7 @@ NODE_CLASS_MAPPINGS = {
"NumberScaler": UtilNumberScaler, "NumberScaler": UtilNumberScaler,
"MergeModels": UtilModelMerge, "MergeModels": UtilModelMerge,
"TextRandomMultiline": UtilTextRandomMultiline, "TextRandomMultiline": UtilTextRandomMultiline,
"TextSwitchCase": UtilTextSwitchCase,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"LoadImageFromUrl": "Load Image From URL", "LoadImageFromUrl": "Load Image From URL",
@@ -1229,4 +1315,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"NumberScaler": "Number Scaler", "NumberScaler": "Number Scaler",
"MergeModels": "Merge Models", "MergeModels": "Merge Models",
"TextRandomMultiline": "Text Random Multiline", "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] 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): def get_dict_attribute(dict_inst: dict, name_string: str, default=None):
nested_keys = name_string.split(".") nested_keys = name_string.split(".")
value = dict_inst value = dict_inst
+1 -1
View File
@@ -1,7 +1,7 @@
[project] [project]
name = "comfyui-art-venture" 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" 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" license = "LICENSE"
dependencies = ["timm==0.6.13", "transformers", "fairscale", "pycocoevalcap", "opencv-python", "qrcode[pil]", "pytorch_lightning", "kornia", "pydantic", "segment_anything", "boto3>=1.34.101"] 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 { $el } from '../../../scripts/ui.js';
import { createImageHost } from '../../../scripts/ui/imagePreview.js'; import { createImageHost } from '../../../scripts/ui/imagePreview.js';
import { chainCallback, addKVState } from './utils.js';
const style = ` const style = `
.comfy-img-preview video { .comfy-img-preview video {
object-fit: contain; object-fit: contain;
@@ -15,24 +17,6 @@ const URL_REGEX = /^((blob:)?https?:\/\/|\/view\?|\/api\/view\?|data:image\/)/
const supportedNodes = ['LoadImageFromUrl', 'LoadImageAsMaskFromUrl', 'LoadVideoFromUrl']; 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) { function injectHidden(widget) {
widget.computeSize = (target_width) => { widget.computeSize = (target_width) => {
if (widget.hidden) { 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) { function migrateWidget(nodeType, oldWidgetName, newWidgetName) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () { chainCallback(nodeType.prototype, 'onNodeCreated', function () {
if (!this.widgets) return; 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;
}
});
});
}