diff --git a/__init__.py b/__init__.py index d9b7ef6..c6ab4ef 100644 --- a/__init__.py +++ b/__init__.py @@ -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" diff --git a/mz_fluxext_core.py b/mz_fluxext_core.py index 99c87a2..c17d3fa 100644 --- a/mz_fluxext_core.py +++ b/mz_fluxext_core.py @@ -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: