Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d0cf2abdaf | ||
|
|
36ed264742 | ||
|
|
fb8bc917ba | ||
|
|
8034af6478 | ||
|
|
713f6de761 | ||
|
|
83c3732201 | ||
|
|
de986044b3 | ||
|
|
bb6994f677 | ||
|
|
4405434b9c | ||
|
|
e7e9e58c66 | ||
|
|
51dd8fcb7c | ||
|
|
4be544aa9e | ||
|
|
08fa873f2c |
@@ -67,6 +67,21 @@ Randomizes the order of lines in a multiline string.
|
|||||||
|
|
||||||

|

|
||||||
|
|
||||||
|
#### 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.
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
### 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
|
||||||
|
|
||||||

|

|
||||||
|
|||||||
@@ -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"}),
|
||||||
|
|||||||
@@ -0,0 +1,106 @@
|
|||||||
|
import os
|
||||||
|
from typing import Dict
|
||||||
|
|
||||||
|
import folder_paths
|
||||||
|
|
||||||
|
from ..utils import load_module
|
||||||
|
|
||||||
|
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
|
||||||
|
preprocessors_dir_names = ["ControlNetPreprocessors", "comfyui_controlnet_aux"]
|
||||||
|
|
||||||
|
preprocessors: list[str] = []
|
||||||
|
_preprocessors_map = {
|
||||||
|
"canny": "CannyEdgePreprocessor",
|
||||||
|
"canny_pyra": "PyraCannyPreprocessor",
|
||||||
|
"lineart": "LineArtPreprocessor",
|
||||||
|
"lineart_anime": "AnimeLineArtPreprocessor",
|
||||||
|
"lineart_manga": "Manga2Anime_LineArt_Preprocessor",
|
||||||
|
"lineart_any": "AnyLineArtPreprocessor_aux",
|
||||||
|
"scribble": "ScribblePreprocessor",
|
||||||
|
"scribble_xdog": "Scribble_XDoG_Preprocessor",
|
||||||
|
"scribble_pidi": "Scribble_PiDiNet_Preprocessor",
|
||||||
|
"scribble_hed": "FakeScribblePreprocessor",
|
||||||
|
"hed": "HEDPreprocessor",
|
||||||
|
"pidi": "PiDiNetPreprocessor",
|
||||||
|
"mlsd": "M-LSDPreprocessor",
|
||||||
|
"pose": "DWPreprocessor",
|
||||||
|
"openpose": "OpenposePreprocessor",
|
||||||
|
"dwpose": "DWPreprocessor",
|
||||||
|
"pose_dense": "DensePosePreprocessor",
|
||||||
|
"pose_animal": "AnimalPosePreprocessor",
|
||||||
|
"normalmap_bae": "BAE-NormalMapPreprocessor",
|
||||||
|
"normalmap_dsine": "DSINE-NormalMapPreprocessor",
|
||||||
|
"normalmap_midas": "MiDaS-NormalMapPreprocessor",
|
||||||
|
"depth": "DepthAnythingV2Preprocessor",
|
||||||
|
"depth_anything": "DepthAnythingPreprocessor",
|
||||||
|
"depth_anything_v2": "DepthAnythingV2Preprocessor",
|
||||||
|
"depth_anything_zoe": "Zoe_DepthAnythingPreprocessor",
|
||||||
|
"depth_zoe": "Zoe-DepthMapPreprocessor",
|
||||||
|
"depth_midas": "MiDaS-DepthMapPreprocessor",
|
||||||
|
"depth_leres": "LeReS-DepthMapPreprocessor",
|
||||||
|
"depth_metric3d": "Metric3D-DepthMapPreprocessor",
|
||||||
|
"depth_meshgraphormer": "MeshGraphormer-DepthMapPreprocessor",
|
||||||
|
"seg_ofcoco": "OneFormer-COCO-SemSegPreprocessor",
|
||||||
|
"seg_ofade20k": "OneFormer-ADE20K-SemSegPreprocessor",
|
||||||
|
"seg_ufade20k": "UniFormer-SemSegPreprocessor",
|
||||||
|
"seg_animeface": "AnimeFace_SemSegPreprocessor",
|
||||||
|
"shuffle": "ShufflePreprocessor",
|
||||||
|
"teed": "TEEDPreprocessor",
|
||||||
|
"color": "ColorPreprocessor",
|
||||||
|
"sam": "SAMPreprocessor",
|
||||||
|
"tile": "TilePreprocessor"
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def apply_preprocessor(image, preprocessor, resolution=512):
|
||||||
|
raise NotImplementedError("apply_preprocessor is not implemented")
|
||||||
|
|
||||||
|
|
||||||
|
try:
|
||||||
|
module_path = None
|
||||||
|
|
||||||
|
for custom_node in custom_nodes:
|
||||||
|
custom_node = custom_node if not os.path.islink(custom_node) else os.readlink(custom_node)
|
||||||
|
for module_dir in preprocessors_dir_names:
|
||||||
|
if module_dir in os.listdir(custom_node):
|
||||||
|
module_path = os.path.abspath(os.path.join(custom_node, module_dir))
|
||||||
|
break
|
||||||
|
|
||||||
|
if module_path is None:
|
||||||
|
raise Exception("Could not find ControlNetPreprocessors nodes")
|
||||||
|
|
||||||
|
module = load_module(module_path)
|
||||||
|
print("Loaded ControlNetPreprocessors nodes from", module_path)
|
||||||
|
|
||||||
|
nodes: Dict = getattr(module, "NODE_CLASS_MAPPINGS")
|
||||||
|
available_preprocessors: list[str] = getattr(module, "PREPROCESSOR_OPTIONS")
|
||||||
|
|
||||||
|
AIO_Preprocessor = nodes.get("AIO_Preprocessor", None)
|
||||||
|
if AIO_Preprocessor is None:
|
||||||
|
raise Exception("Could not find AIO_Preprocessor node")
|
||||||
|
|
||||||
|
for name, preprocessor in _preprocessors_map.items():
|
||||||
|
if preprocessor in available_preprocessors:
|
||||||
|
preprocessors.append(name)
|
||||||
|
|
||||||
|
aio_preprocessor = AIO_Preprocessor()
|
||||||
|
|
||||||
|
def apply_preprocessor(image, preprocessor, resolution=512):
|
||||||
|
if preprocessor == "None":
|
||||||
|
return image
|
||||||
|
|
||||||
|
if preprocessor not in preprocessors:
|
||||||
|
raise Exception(f"Preprocessor {preprocessor} is not implemented")
|
||||||
|
|
||||||
|
preprocessor_cls = _preprocessors_map[preprocessor]
|
||||||
|
args = {"preprocessor": preprocessor_cls, "image": image, "resolution": resolution}
|
||||||
|
|
||||||
|
function_name = AIO_Preprocessor.FUNCTION
|
||||||
|
res = getattr(aio_preprocessor, function_name)(**args)
|
||||||
|
if isinstance(res, dict):
|
||||||
|
res = res["result"]
|
||||||
|
|
||||||
|
return res[0]
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
print(e)
|
||||||
@@ -1,142 +0,0 @@
|
|||||||
import os
|
|
||||||
import math
|
|
||||||
from typing import Dict
|
|
||||||
|
|
||||||
import folder_paths
|
|
||||||
|
|
||||||
from ..utils import load_module
|
|
||||||
|
|
||||||
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
|
|
||||||
preprocessors_dir_names = ["ControlNetPreprocessors", "comfyui_controlnet_aux"]
|
|
||||||
|
|
||||||
control_net_preprocessors = {}
|
|
||||||
|
|
||||||
try:
|
|
||||||
module_path = None
|
|
||||||
|
|
||||||
for custom_node in custom_nodes:
|
|
||||||
custom_node = (
|
|
||||||
custom_node if not os.path.islink(custom_node) else os.readlink(custom_node)
|
|
||||||
)
|
|
||||||
for module_dir in preprocessors_dir_names:
|
|
||||||
if module_dir in os.listdir(custom_node):
|
|
||||||
module_path = os.path.abspath(os.path.join(custom_node, module_dir))
|
|
||||||
break
|
|
||||||
|
|
||||||
if module_path is None:
|
|
||||||
raise Exception("Could not find ControlNetPreprocessors nodes")
|
|
||||||
|
|
||||||
module = load_module(module_path)
|
|
||||||
print("Loaded ControlNetPreprocessors nodes from", module_path)
|
|
||||||
|
|
||||||
nodes: Dict = getattr(module, "NODE_CLASS_MAPPINGS")
|
|
||||||
|
|
||||||
if "CannyEdgePreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["canny"] = (
|
|
||||||
nodes["CannyEdgePreprocessor"],
|
|
||||||
[100, 200],
|
|
||||||
)
|
|
||||||
if "LineArtPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["lineart"] = (
|
|
||||||
nodes["LineArtPreprocessor"],
|
|
||||||
["disable"],
|
|
||||||
)
|
|
||||||
control_net_preprocessors["lineart_coarse"] = (
|
|
||||||
nodes["LineArtPreprocessor"],
|
|
||||||
["enable"],
|
|
||||||
)
|
|
||||||
if "AnimeLineArtPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["lineart_anime"] = (
|
|
||||||
nodes["AnimeLineArtPreprocessor"],
|
|
||||||
[],
|
|
||||||
)
|
|
||||||
if "Manga2Anime_LineArt_Preprocessor" in nodes:
|
|
||||||
control_net_preprocessors["lineart_manga"] = (
|
|
||||||
nodes["Manga2Anime_LineArt_Preprocessor"],
|
|
||||||
[],
|
|
||||||
)
|
|
||||||
if "ScribblePreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["scribble"] = (nodes["ScribblePreprocessor"], [])
|
|
||||||
if "FakeScribblePreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["scribble_hed"] = (
|
|
||||||
nodes["FakeScribblePreprocessor"],
|
|
||||||
["enable"],
|
|
||||||
)
|
|
||||||
if "HEDPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["hed"] = (nodes["HEDPreprocessor"], ["disable"])
|
|
||||||
control_net_preprocessors["hed_safe"] = (nodes["HEDPreprocessor"], ["enable"])
|
|
||||||
if "PiDiNetPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["pidi"] = (
|
|
||||||
nodes["PiDiNetPreprocessor"],
|
|
||||||
["disable"],
|
|
||||||
)
|
|
||||||
control_net_preprocessors["pidi_safe"] = (
|
|
||||||
nodes["PiDiNetPreprocessor"],
|
|
||||||
["enable"],
|
|
||||||
)
|
|
||||||
if "M-LSDPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["mlsd"] = (nodes["M-LSDPreprocessor"], [0.1, 0.1])
|
|
||||||
if "OpenposePreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["openpose"] = (
|
|
||||||
nodes["OpenposePreprocessor"],
|
|
||||||
["enable", "enable", "enable"],
|
|
||||||
)
|
|
||||||
control_net_preprocessors["pose"] = control_net_preprocessors["openpose"]
|
|
||||||
if "DWPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["dwpose"] = (
|
|
||||||
nodes["DWPreprocessor"],
|
|
||||||
["enable", "enable", "enable", "yolox_l.onnx", "dw-ll_ucoco_384.onnx"],
|
|
||||||
)
|
|
||||||
# use DWPreprocessor for pose by default if available
|
|
||||||
control_net_preprocessors["pose"] = control_net_preprocessors["dwpose"]
|
|
||||||
if "BAE-NormalMapPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["normalmap_bae"] = (
|
|
||||||
nodes["BAE-NormalMapPreprocessor"],
|
|
||||||
[],
|
|
||||||
)
|
|
||||||
if "MiDaS-NormalMapPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["normalmap_midas"] = (
|
|
||||||
nodes["MiDaS-NormalMapPreprocessor"],
|
|
||||||
[math.pi * 2.0, 0.1],
|
|
||||||
)
|
|
||||||
if "MiDaS-DepthMapPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["depth_midas"] = (
|
|
||||||
nodes["MiDaS-DepthMapPreprocessor"],
|
|
||||||
[math.pi * 2.0, 0.4],
|
|
||||||
)
|
|
||||||
if "Zoe-DepthMapPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["depth"] = (nodes["Zoe-DepthMapPreprocessor"], [])
|
|
||||||
control_net_preprocessors["depth_zoe"] = (nodes["Zoe-DepthMapPreprocessor"], [])
|
|
||||||
if "OneFormer-COCO-SemSegPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["seg_ofcoco"] = (
|
|
||||||
nodes["OneFormer-COCO-SemSegPreprocessor"],
|
|
||||||
[],
|
|
||||||
)
|
|
||||||
if "OneFormer-ADE20K-SemSegPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["seg_ofade20k"] = (
|
|
||||||
nodes["OneFormer-ADE20K-SemSegPreprocessor"],
|
|
||||||
[],
|
|
||||||
)
|
|
||||||
if "UniFormer-SemSegPreprocessor" in nodes:
|
|
||||||
control_net_preprocessors["seg_ufade20k"] = (
|
|
||||||
nodes["UniFormer-SemSegPreprocessor"],
|
|
||||||
[],
|
|
||||||
)
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
print(e)
|
|
||||||
|
|
||||||
|
|
||||||
class DummyPreprocessor:
|
|
||||||
@classmethod
|
|
||||||
def INPUT_TYPES(cls):
|
|
||||||
return {
|
|
||||||
"required": {
|
|
||||||
"image": ("IMAGE",)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
FUNCTION = "process"
|
|
||||||
|
|
||||||
def process(self, image):
|
|
||||||
return (image,)
|
|
||||||
@@ -6,7 +6,7 @@ import torch.nn.functional as F
|
|||||||
import folder_paths
|
import 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
|
||||||
|
|||||||
@@ -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",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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":
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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):
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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",
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
@@ -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"]
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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;
|
||||||
|
}
|
||||||
|
});
|
||||||
|
});
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user