add DisplayAny node
This commit is contained in:
+3
-1
@@ -7,6 +7,8 @@ from .misc import MISC_CLASS_MAPPINGS, MISC_NAME_MAPPINGS
|
||||
from .conditioning import COND_CLASS_MAPPINGS, COND_NAME_MAPPINGS
|
||||
from .text import TEXT_CLASS_MAPPINGS, TEXT_NAME_MAPPINGS
|
||||
|
||||
WEB_DIRECTORY = "./js"
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
@@ -31,4 +33,4 @@ NODE_DISPLAY_NAME_MAPPINGS.update(TEXT_NAME_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(MISC_CLASS_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(MISC_NAME_MAPPINGS)
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', "WEB_DIRECTORY"]
|
||||
|
||||
+19
-12
@@ -205,28 +205,33 @@ class SD3AttentionSeekerT5:
|
||||
|
||||
return (m, )
|
||||
|
||||
"""
|
||||
class FluxXAttentionSeeker:
|
||||
class FluxModelAttentionSeeker:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "model": ("MODEL",),
|
||||
"q": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"k": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"v": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"out": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
}}
|
||||
return {"required": {
|
||||
"model": ("MODEL",),
|
||||
**{f"block_{s}": ("FLOAT", { "display": "slider", "default": 1.0, "min": 0, "max": 5, "step": 0.05 }) for s in range(38)},
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
|
||||
CATEGORY = "_for_testing/attention_experiments"
|
||||
|
||||
def patch(self, model, q, k, v, out):
|
||||
def patch(self, model, **blocks):
|
||||
m = model.clone()
|
||||
sd = model.model_state_dict()
|
||||
print(sd.keys())
|
||||
|
||||
for k in sd:
|
||||
block = re.search(r"(\d+)\.(img|txt)_(mod|attn|mlp)\.(lin|qkv|proj|0|2)\.(weight|bias)", k)
|
||||
block = int(block.group(1)) if block else None
|
||||
|
||||
if block is not None and blocks[f"block_{block}"] != 1.0:
|
||||
print(k, blocks[f"block_{block}"])
|
||||
m.add_patches({k: (None,)}, 0.0, blocks[f"block_{block}"])
|
||||
|
||||
return (m, )
|
||||
"""
|
||||
|
||||
|
||||
|
||||
COND_CLASS_MAPPINGS = {
|
||||
"CLIPTextEncodeSDXL+": CLIPTextEncodeSDXLSimplified,
|
||||
"ConditioningCombineMultiple+": ConditioningCombineMultiple,
|
||||
@@ -234,6 +239,7 @@ COND_CLASS_MAPPINGS = {
|
||||
"FluxAttentionSeeker+": FluxAttentionSeeker,
|
||||
"SD3AttentionSeekerLG+": SD3AttentionSeekerLG,
|
||||
"SD3AttentionSeekerT5+": SD3AttentionSeekerT5,
|
||||
#"FluxModelAttentionSeeker+": FluxModelAttentionSeeker,
|
||||
}
|
||||
|
||||
COND_NAME_MAPPINGS = {
|
||||
@@ -243,4 +249,5 @@ COND_NAME_MAPPINGS = {
|
||||
"FluxAttentionSeeker+": "🔧 Flux Attention Seeker",
|
||||
"SD3AttentionSeekerLG+": "🔧 SD3 Attention Seeker L/G",
|
||||
"SD3AttentionSeekerT5+": "🔧 SD3 Attention Seeker T5",
|
||||
#"FluxModelAttentionSeeker+": "🔧 Flux Model Attention Seeker",
|
||||
}
|
||||
@@ -487,10 +487,38 @@ class SDXLEmptyLatentSizePicker:
|
||||
|
||||
return ({"samples":latent}, width, height,)
|
||||
|
||||
class DisplayAny:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input": (("*",{})),
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(s, input_types):
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "execute"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
CATEGORY = "essentials/utilities"
|
||||
|
||||
def execute(self, input):
|
||||
text = str(input)
|
||||
|
||||
return {"ui": {"text": text}, "result": ()}
|
||||
|
||||
MISC_CLASS_MAPPINGS = {
|
||||
"BatchCount+": BatchCount,
|
||||
"ConsoleDebug+": ConsoleDebug,
|
||||
"DebugTensorShape+": DebugTensorShape,
|
||||
"DisplayAny": DisplayAny,
|
||||
"ModelCompile+": ModelCompile,
|
||||
"RemoveLatentMask+": RemoveLatentMask,
|
||||
"SDXLEmptyLatentSizePicker+": SDXLEmptyLatentSizePicker,
|
||||
@@ -511,6 +539,7 @@ MISC_NAME_MAPPINGS = {
|
||||
"BatchCount+": "🔧 Batch Count",
|
||||
"ConsoleDebug+": "🔧 Console Debug",
|
||||
"DebugTensorShape+": "🔧 Debug Tensor Shape",
|
||||
"DisplayAny": "🔧 Display Any",
|
||||
"ModelCompile+": "🔧 Model Compile",
|
||||
"RemoveLatentMask+": "🔧 Remove Latent Mask",
|
||||
"SDXLEmptyLatentSizePicker+": "🔧 Empty Latent Size Picker",
|
||||
|
||||
Reference in New Issue
Block a user