This commit is contained in:
wailovet
2024-08-15 01:36:26 +08:00
parent b5f69c7a6d
commit abfd0976c0
2 changed files with 6 additions and 4 deletions
+2 -2
View File
@@ -34,8 +34,8 @@ class MZ_Flux1PartialLoad_Patch:
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL", ),
"double_blocks_cuda_size": ("INT", {"min": 0, "max": 16}),
"single_blocks_cuda_size": ("INT", {"min": 0, "max": 37}),
"double_blocks_cuda_size": ("INT", {"min": 0, "max": 16, "default": 7}),
"single_blocks_cuda_size": ("INT", {"min": 0, "max": 37, "default": 7}),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "load_unet"
+4 -2
View File
@@ -13,6 +13,8 @@ from torch import Tensor, nn
def Flux1PartialLoad_Patch(args={}):
model = args.get("model")
double_blocks_cuda_size = args.get("double_blocks_cuda_size")
single_blocks_cuda_size = args.get("single_blocks_cuda_size")
def other_to_cpu():
model.model.diffusion_model.img_in.to("cpu")
@@ -105,7 +107,7 @@ def Flux1PartialLoad_Patch(args={}):
pre_only_model_forward_hook)
double_blocks_depth = len(model.model.diffusion_model.double_blocks)
steps = 7
steps = double_blocks_cuda_size
for i in range(0, double_blocks_depth, steps):
s = steps
if i + s > double_blocks_depth:
@@ -114,7 +116,7 @@ def Flux1PartialLoad_Patch(args={}):
generate_double_blocks_forward_hook(i, s))
single_blocks_depth = len(model.model.diffusion_model.single_blocks)
steps = 7
steps = single_blocks_cuda_size
for i in range(0, single_blocks_depth, steps):
s = steps
if i + s > single_blocks_depth: