update model compile

This commit is contained in:
cubiq
2024-09-18 15:03:26 +02:00
parent 99aad72c84
commit ef704350e7
2 changed files with 11 additions and 5 deletions
+10 -4
View File
@@ -214,14 +214,19 @@ class FluxBlocksBuster:
#**{f"double_block_{s}": ("FLOAT", { "display": "slider", "default": 1.0, "min": 0, "max": 5, "step": 0.05 }) for s in range(19)}, #**{f"double_block_{s}": ("FLOAT", { "display": "slider", "default": 1.0, "min": 0, "max": 5, "step": 0.05 }) for s in range(19)},
#**{f"single_block_{s}": ("FLOAT", { "display": "slider", "default": 1.0, "min": 0, "max": 5, "step": 0.05 }) for s in range(38)}, #**{f"single_block_{s}": ("FLOAT", { "display": "slider", "default": 1.0, "min": 0, "max": 5, "step": 0.05 }) for s in range(38)},
}} }}
RETURN_TYPES = ("MODEL",) RETURN_TYPES = ("MODEL", "STRING")
RETURN_NAMES = ("MODEL", "patched_blocks")
FUNCTION = "patch" FUNCTION = "patch"
CATEGORY = "essentials/conditioning" CATEGORY = "essentials/conditioning"
def patch(self, model, blocks): def patch(self, model, blocks):
if blocks == "":
return (model, )
m = model.clone() m = model.clone()
sd = model.model_state_dict() sd = model.model_state_dict()
patched_blocks = []
""" """
Also compatible with the following format: Also compatible with the following format:
@@ -235,7 +240,6 @@ class FluxBlocksBuster:
blocks = blocks.split("\n") blocks = blocks.split("\n")
blocks = [b.strip() for b in blocks if b.strip()] blocks = [b.strip() for b in blocks if b.strip()]
print("Patched blocks:")
for k in sd: for k in sd:
for block in blocks: for block in blocks:
block = block.split("=") block = block.split("=")
@@ -248,9 +252,11 @@ class FluxBlocksBuster:
if value != 1.0 and re.search(block, k): if value != 1.0 and re.search(block, k):
m.add_patches({k: (None,)}, 0.0, value) m.add_patches({k: (None,)}, 0.0, value)
print(f"{k}: {value}") patched_blocks.append(f"{k}: {value}")
return (m, ) patched_blocks = "\n".join(patched_blocks)
return (m, patched_blocks,)
COND_CLASS_MAPPINGS = { COND_CLASS_MAPPINGS = {
+1 -1
View File
@@ -442,7 +442,7 @@ class ModelCompile():
def execute(self, model, fullgraph, dynamic, mode): def execute(self, model, fullgraph, dynamic, mode):
work_model = model.clone() work_model = model.clone()
torch._dynamo.config.suppress_errors = True torch._dynamo.config.suppress_errors = True
work_model.model.diffusion_model = torch.compile(work_model.model.diffusion_model, dynamic=dynamic, fullgraph=fullgraph, mode=mode) work_model.add_object_patch("diffusion_model", torch.compile(model=work_model.get_model_object("diffusion_model"), dynamic=dynamic, fullgraph=fullgraph, mode=mode))
return (work_model, ) return (work_model, )
class RemoveLatentMask: class RemoveLatentMask: