feat: ImpactSwitch

This commit is contained in:
Dr.Lt.Data
2023-08-23 22:08:57 +09:00
parent 75173d0c1b
commit 2255fba8a1
6 changed files with 130 additions and 92 deletions
+9 -8
View File
@@ -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']:
+81 -4
View File
@@ -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;
}
}
}
+1 -1
View File
@@ -2,7 +2,7 @@ import configparser
import os
version = "V3.23.1"
version = "V3.24"
dependency_version = 9
+20 -62
View File
@@ -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):
+11 -17
View File
@@ -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,),
},
}
+8
View File
@@ -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("*")