From b5f69c7a6d5fa8210e3cabf4a34b63d11dde7a81 Mon Sep 17 00:00:00 2001 From: wailovet Date: Thu, 15 Aug 2024 01:34:55 +0800 Subject: [PATCH] 'update' --- __init__.py | 12 +++++++----- mz_fluxext_core.py | 3 ++- 2 files changed, 9 insertions(+), 6 deletions(-) diff --git a/__init__.py b/__init__.py index 520ec92..d9b7ef6 100644 --- a/__init__.py +++ b/__init__.py @@ -29,11 +29,13 @@ from . import mz_fluxext_core import importlib -class MZ_Flux1VRAM_MT_Patch: +class MZ_Flux1PartialLoad_Patch: @classmethod def INPUT_TYPES(s): return {"required": { - "model": ("MODEL", ) + "model": ("MODEL", ), + "double_blocks_cuda_size": ("INT", {"min": 0, "max": 16}), + "single_blocks_cuda_size": ("INT", {"min": 0, "max": 37}), }} RETURN_TYPES = ("MODEL",) FUNCTION = "load_unet" @@ -43,8 +45,8 @@ class MZ_Flux1VRAM_MT_Patch: def load_unet(self, **kwargs): from . import mz_fluxext_core importlib.reload(mz_fluxext_core) - return mz_fluxext_core.MZ_Flux1VRAM_MT_Patch_call(kwargs) + return mz_fluxext_core.Flux1PartialLoad_Patch(kwargs) -NODE_CLASS_MAPPINGS["MZ_Flux1VRAM_MT_Patch"] = MZ_Flux1VRAM_MT_Patch -NODE_DISPLAY_NAME_MAPPINGS["MZ_Flux1VRAM_MT_Patch"] = f"{AUTHOR_NAME} - Flux1VRAM_MT_Patch" +NODE_CLASS_MAPPINGS["MZ_Flux1PartialLoad_Patch"] = MZ_Flux1PartialLoad_Patch +NODE_DISPLAY_NAME_MAPPINGS["MZ_Flux1PartialLoad_Patch"] = f"{AUTHOR_NAME} - Flux1PartialLoad_Patch" diff --git a/mz_fluxext_core.py b/mz_fluxext_core.py index d6f1eec..99c87a2 100644 --- a/mz_fluxext_core.py +++ b/mz_fluxext_core.py @@ -11,7 +11,7 @@ import safetensors from torch import Tensor, nn -def MZ_Flux1VRAM_MT_Patch_call(args={}): +def Flux1PartialLoad_Patch(args={}): model = args.get("model") def other_to_cpu(): @@ -100,6 +100,7 @@ def MZ_Flux1VRAM_MT_Patch_call(args={}): print("other to cuda") other_to_cuda() return inp + model.model.diffusion_model.register_forward_pre_hook( pre_only_model_forward_hook)