diff --git a/__init__.py b/__init__.py index 4f63b6a..f4b498e 100644 --- a/__init__.py +++ b/__init__.py @@ -96,6 +96,7 @@ from impact.impact_pack import * from impact.detectors import * from impact.pipe import * from impact.logics import * +from impact.util_nodes import * impact.wildcards.read_wildcard_dict(wildcards_path) impact.wildcards.read_wildcard_dict(custom_wildcards_path) @@ -180,8 +181,9 @@ NODE_CLASS_MAPPINGS = { "LatentSender": LatentSender, "LatentReceiver": LatentReceiver, "ImageMaskSwitch": ImageMaskSwitch, - "LatentSwitch": LatentSwitch, - "SEGSSwitch": SEGSSwitch, + "LatentSwitch": GeneralSwitch, + "SEGSSwitch": GeneralSwitch, + "ImpactSwitch": GeneralSwitch, # "SaveConditioning": SaveConditioning, # "LoadConditioning": LoadConditioning, @@ -196,9 +198,6 @@ NODE_CLASS_MAPPINGS = { "ImpactSEGSToMaskList": SEGSToMaskList, "ImpactSEGSConcat": SEGSConcat, - # "SEGPick": SEGPick, - # "SEGEdit": SEGEdit, - "ImpactKSamplerBasicPipe": KSamplerBasicPipe, "ImpactKSamplerAdvancedBasicPipe": KSamplerAdvancedBasicPipe, @@ -285,13 +284,15 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ImageSender": "Image Sender", "ImageReceiver": "Image Receiver", "ImageMaskSwitch": "Switch (images, mask)", - "LatentSwitch": "Switch (latent)", - "SEGSSwitch": "Switch (SEGS)", + "ImpactSwitch": "Switch (Any)", "MasksToMaskList": "Masks to Mask List", "ImpactImageBatchToImageList": "Image batch to Image List", "ImpactMakeImageList": "Make Image List", - "ImpactStringSelector": "String Selector" + "ImpactStringSelector": "String Selector", + + "LatentSwitch": "Switch (latent/legacy)", + "SEGSSwitch": "Switch (SEGS/legacy)" } if not impact.config.get_config()['mmdet_skip']: diff --git a/js/impact-pack.js b/js/impact-pack.js index 20caa2e..7ad8903 100644 --- a/js/impact-pack.js +++ b/js/impact-pack.js @@ -191,10 +191,80 @@ app.registerExtension({ impactProgressBadge.addStatusHandler(nodeType); } - if (nodeData.name === 'ImpactMakeImageList') { + if (nodeData.name === 'ImpactMakeImageList' || nodeData.name === 'ImpactSwitch' || nodeData.name === 'LatentSwitch' || nodeData.name == 'SEGSSwitch') { + var input_name = "input"; + + switch(nodeData.name) { + case 'ImpactMakeImageList': + input_name = "image"; + break; + + case 'LatentSwitch': + input_name = "input"; + break; + + case 'SEGSSwitch': + input_name = "input"; + break; + + case 'ImpactSwitch': + input_name = "input"; + } + const onConnectionsChange = nodeType.prototype.onConnectionsChange nodeType.prototype.onConnectionsChange = function (type, index, connected, link_info) { + if(!link_info) + return; + + if(type == 2) { + // connect output + if(connected){ + if(this.outputs[0].type == '*'){ + if(link_info.type == '*') { + this.disconnectOutput(link_info.origin_slot); + } + else { + // propagate type + this.outputs[0].type = link_info.type; + this.outputs[0].label = link_info.type; + this.outputs[0].name = link_info.type; + + for(let i in this.inputs) { + this.inputs[i].type = link_info.type; + } + } + } + } + + return; + } + else { + // connect input + if(this.inputs[0].type == '*'){ + const node = app.graph.getNodeById(link_info.origin_id); + let origin_type = node.outputs[link_info.origin_slot].type; + + if(origin_type == '*') { + this.disconnectInput(link_info.target_slot); + return; + } + + for(let i in this.inputs) { + this.inputs[i].type = origin_type; + } + + this.outputs[0].type = origin_type; + this.outputs[0].label = origin_type; + this.outputs[0].name = origin_type; + } + } + if (!connected && this.inputs.length > 1) { + const stackTrace = new Error().stack; + if(stackTrace.includes('LGraphNode.connect')) { + return; // replace connection: don't remove connection + } + if (this.widgets) { const w = this.widgets.find((w) => w.name === this.inputs[index].name) if (w) { @@ -206,12 +276,19 @@ app.registerExtension({ } for (let i = 0; i < this.inputs.length; i++) { - this.inputs[i].label = `image${i + 1}` - this.inputs[i].name = `image${i + 1}` + this.inputs[i].label = `${input_name}${i + 1}` + this.inputs[i].name = `${input_name}${i + 1}` } if (this.inputs[this.inputs.length - 1].link != undefined) { - this.addInput(`image${this.inputs.length + 1}`, 'IMAGE'); + this.addInput(`${input_name}${this.inputs.length + 1}`, this.inputs[0].type); + } + + if(this.widgets) { + this.widgets[0].options.max = this.inputs.length-1; + this.widgets[0].value = Math.min(this.widgets[0].value, this.widgets[0].options.max); + if(this.widgets[0].options.max > 0 && this.widgets[0].value == 0) + this.widgets[0].value = 1; } } } diff --git a/modules/impact/config.py b/modules/impact/config.py index 33630dc..d548444 100644 --- a/modules/impact/config.py +++ b/modules/impact/config.py @@ -2,7 +2,7 @@ import configparser import os -version = "V3.23.1" +version = "V3.24" dependency_version = 9 diff --git a/modules/impact/impact_pack.py b/modules/impact/impact_pack.py index b7c0507..b0b3126 100644 --- a/modules/impact/impact_pack.py +++ b/modules/impact/impact_pack.py @@ -2357,15 +2357,9 @@ class LatentSwitch: @classmethod def INPUT_TYPES(s): return {"required": { - "select": ("INT", {"default": 1, "min": 1, "max": 4, "step": 1}), + "select": ("INT", {"default": 1, "min": 1, "max": 99999, "step": 1}), "latent1": ("LATENT",), }, - - "optional": { - "latent2_opt": ("LATENT",), - "latent3_opt": ("LATENT",), - "latent4_opt": ("LATENT",), - }, } RETURN_TYPES = ("LATENT", ) @@ -2376,29 +2370,22 @@ class LatentSwitch: CATEGORY = "ImpactPack/Util" - def doit(self, select, latent1, latent2_opt=None, latent3_opt=None, latent4_opt=None): - if select == 1: - return (latent1,) - elif select == 2: - return (latent2_opt,) - elif select == 3: - return (latent3_opt,) + def doit(self, *args, **kwargs): + input_name = f"latent{int(kwargs['select'])}" + + if input_name in kwargs: + return (kwargs[input_name],) else: - return (latent4_opt,) + print(f"LatentSwitch: invalid select index ('latent1' is selected)") + return (kwargs['latent1'],) class SEGSSwitch: @classmethod def INPUT_TYPES(s): return {"required": { - "select": ("INT", {"default": 1, "min": 1, "max": 4, "step": 1}), - "segs": ("SEGS",), - }, - - "optional": { - "segs2_opt": ("SEGS",), - "segs3_opt": ("SEGS",), - "segs4_opt": ("SEGS",), + "select": ("INT", {"default": 1, "min": 1, "max": 99999, "step": 1}), + "segs1": ("SEGS",), }, } @@ -2410,43 +2397,14 @@ class SEGSSwitch: CATEGORY = "ImpactPack/Util" - def doit(self, select, segs, segs2_opt=None, segs3_opt=None, segs4_opt=None): - if select == 1: - return (segs,) - elif select == 2: - return (segs2_opt,) - elif select == 3: - return (segs3_opt,) + def doit(self, *args, **kwargs): + input_name = f"segs{int(kwargs['select'])}" + + if input_name in kwargs: + return (kwargs[input_name],) else: - return (segs4_opt,) - - -# class SEGPick: -# @classmethod -# def INPUT_TYPES(s): -# return {"required": { -# "select": ("INT", {"default": 1, "min": 1, "max": 99999, "step": 1}), -# "segs": ("SEGS",), -# }, -# } -# -# RETURN_TYPES = ("SEGS", ) -# -# OUTPUT_NODE = True -# -# FUNCTION = "doit" -# -# CATEGORY = "ImpactPack/Util" -# -# def doit(self, select, segs): -# if select == 1: -# return (segs,) -# elif select == 2: -# return (segs2_opt,) -# elif select == 3: -# return (segs3_opt,) -# else: -# return (segs4_opt,) + print(f"SEGSSwitch: invalid select index ('segs1' is selected)") + return (kwargs['segs1'],) class SaveConditioning: @@ -2770,13 +2728,13 @@ class StringSelector: return (selected, ) -from impact.logics import AnyType +from impact.utils import any_typ class ImpactLogger: @classmethod def INPUT_TYPES(s): return {"required": { - "data": (AnyType("*"), ""), + "data": (any_typ, ""), }, "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, } @@ -2815,7 +2773,7 @@ class ImpactDummyInput: CATEGORY = "ImpactPack/Debug" - RETURN_TYPES = (AnyType("*"),) + RETURN_TYPES = (any_typ,) FUNCTION = "doit" def doit(self): diff --git a/modules/impact/logics.py b/modules/impact/logics.py index aaab6a3..5fd0f24 100644 --- a/modules/impact/logics.py +++ b/modules/impact/logics.py @@ -1,12 +1,6 @@ import sys from server import PromptServer - -# wildcard trick is taken from pythongossss's -class AnyType(str): - def __ne__(self, __value: object) -> bool: - return False - -any = AnyType("*") +from impact.utils import any_typ class ImpactCompare: @@ -15,8 +9,8 @@ class ImpactCompare: return { "required": { "cmp": (['a = b', 'a <> b', 'a > b', 'a < b', 'a >= b', 'a <= b', 'tt', 'ff'],), - "a": (any, ), - "b": (any, ), + "a": (any_typ, ), + "b": (any_typ, ), }, } @@ -50,15 +44,15 @@ class ImpactConditionalBranch: return { "required": { "cond": ("BOOLEAN", {"forceInput": True}), - "tt_value": (any,), - "ff_value": (any,), + "tt_value": (any_typ,), + "ff_value": (any_typ,), }, } FUNCTION = "doit" CATEGORY = "ImpactPack/Logic" - RETURN_TYPES = (any, ) + RETURN_TYPES = (any_typ, ) def doit(self, cond, tt_value, ff_value): if cond: @@ -143,7 +137,7 @@ class ImpactValueSender: @classmethod def INPUT_TYPES(cls): return {"required": { - "value": (any, ), + "value": (any_typ, ), "link_id": ("INT", {"default": 0, "min": 0, "max": sys.maxsize, "step": 1}), }, } @@ -165,7 +159,7 @@ class ImpactIntConstSender: @classmethod def INPUT_TYPES(cls): return {"required": { - "signal": (any, ), + "signal": (any_typ, ), "value": ("INT", {"default": 0, "min": 0, "max": sys.maxsize, "step": 1}), "link_id": ("INT", {"default": 0, "min": 0, "max": sys.maxsize, "step": 1}), }, @@ -198,7 +192,7 @@ class ImpactValueReceiver: CATEGORY = "ImpactPack/Logic" - RETURN_TYPES = (any, ) + RETURN_TYPES = (any_typ, ) def doit(self, typ, value, link_id=0): if typ == "INT": @@ -235,8 +229,8 @@ class ImpactMinMax: def INPUT_TYPES(cls): return {"required": { "mode": ("BOOLEAN", {"default": True, "label_on": "max", "label_off": "min"}), - "a": (any,), - "b": (any,), + "a": (any_typ,), + "b": (any_typ,), }, } diff --git a/modules/impact/utils.py b/modules/impact/utils.py index d06ceb4..eabccce 100644 --- a/modules/impact/utils.py +++ b/modules/impact/utils.py @@ -238,3 +238,11 @@ class NonListIterable: def __getitem__(self, index): return self.data[index] + + +# wildcard trick is taken from pythongossss's +class AnyType(str): + def __ne__(self, __value: object) -> bool: + return False + +any_typ = AnyType("*")