update model compile
This commit is contained in:
+10
-4
@@ -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"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"
|
||||
|
||||
CATEGORY = "essentials/conditioning"
|
||||
|
||||
def patch(self, model, blocks):
|
||||
if blocks == "":
|
||||
return (model, )
|
||||
|
||||
m = model.clone()
|
||||
sd = model.model_state_dict()
|
||||
patched_blocks = []
|
||||
|
||||
"""
|
||||
Also compatible with the following format:
|
||||
@@ -235,7 +240,6 @@ class FluxBlocksBuster:
|
||||
blocks = blocks.split("\n")
|
||||
blocks = [b.strip() for b in blocks if b.strip()]
|
||||
|
||||
print("Patched blocks:")
|
||||
for k in sd:
|
||||
for block in blocks:
|
||||
block = block.split("=")
|
||||
@@ -248,9 +252,11 @@ class FluxBlocksBuster:
|
||||
|
||||
if value != 1.0 and re.search(block, k):
|
||||
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 = {
|
||||
|
||||
@@ -442,7 +442,7 @@ class ModelCompile():
|
||||
def execute(self, model, fullgraph, dynamic, mode):
|
||||
work_model = model.clone()
|
||||
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, )
|
||||
|
||||
class RemoveLatentMask:
|
||||
|
||||
Reference in New Issue
Block a user