From adc82313ebf4ade36c71a71ddcc1fb487fdcb599 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 22 Mar 2024 00:37:21 +0200 Subject: [PATCH] Make model loader node without sdxl_model selection for clarity --- __init__.py | 12 ++-- nodes_v2.py | 178 +++++++++++++++++++++++++++++++++++++++++++++------- 2 files changed, 164 insertions(+), 26 deletions(-) diff --git a/__init__.py b/__init__.py index c9bcb99..d372fdc 100644 --- a/__init__.py +++ b/__init__.py @@ -1,5 +1,5 @@ from .nodes import SUPIR_Upscale -from .nodes_v2 import SUPIR_sample, SUPIR_model_loader, SUPIR_first_stage, SUPIR_encode, SUPIR_decode, SUPIR_conditioner, SUPIR_tiles +from .nodes_v2 import SUPIR_sample, SUPIR_model_loader, SUPIR_first_stage, SUPIR_encode, SUPIR_decode, SUPIR_conditioner, SUPIR_tiles, SUPIR_model_loader_v2 NODE_CLASS_MAPPINGS = { "SUPIR_Upscale": SUPIR_Upscale, @@ -9,16 +9,18 @@ NODE_CLASS_MAPPINGS = { "SUPIR_encode": SUPIR_encode, "SUPIR_decode": SUPIR_decode, "SUPIR_conditioner": SUPIR_conditioner, - "SUPIR_tiles": SUPIR_tiles + "SUPIR_tiles": SUPIR_tiles, + "SUPIR_model_loader_v2": SUPIR_model_loader_v2 } NODE_DISPLAY_NAME_MAPPINGS = { - "SUPIR_Upscale": "SUPIR Upscale", + "SUPIR_Upscale": "SUPIR Upscale (Legacy)", "SUPIR_sample": "SUPIR Sampler", - "SUPIR_model_loader": "SUPIR Model Loader", + "SUPIR_model_loader": "SUPIR Model Loader (Legacy)", "SUPIR_first_stage": "SUPIR First Stage (Denoiser)", "SUPIR_encode": "SUPIR Encode", "SUPIR_decode": "SUPIR Decode", "SUPIR_conditioner": "SUPIR Conditioner", - "SUPIR_tiles": "SUPIR Tiles" + "SUPIR_tiles": "SUPIR Tiles", + "SUPIR_model_loader_v2": "SUPIR Model Loader (v2)" } __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/nodes_v2.py b/nodes_v2.py index 8764378..add5b0b 100644 --- a/nodes_v2.py +++ b/nodes_v2.py @@ -585,11 +585,6 @@ class SUPIR_model_loader: "default": 'auto' }), }, - "optional":{ - "model" :("MODEL",), - "clip": ("CLIP",), - "vae": ("VAE",), - } } RETURN_TYPES = ("SUPIRMODEL", "SUPIRVAE") @@ -597,7 +592,7 @@ class SUPIR_model_loader: FUNCTION = "process" CATEGORY = "SUPIR" - def process(self, supir_model, sdxl_model, diffusion_dtype, fp8_unet, model=None, clip=None, vae=None): + def process(self, supir_model, sdxl_model, diffusion_dtype, fp8_unet): device = mm.get_torch_device() mm.unload_all_models() @@ -613,6 +608,152 @@ class SUPIR_model_loader: 'diffusion_dtype': diffusion_dtype, 'supir_model': supir_model, 'fp8_unet': fp8_unet, + } + + if diffusion_dtype == 'auto': + try: + if mm.should_use_fp16(): + print("Diffusion using fp16") + dtype = torch.float16 + model_dtype = 'fp16' + elif mm.should_use_bf16(): + print("Diffusion using bf16") + dtype = torch.bfloat16 + model_dtype = 'bf16' + else: + print("Diffusion using using fp32") + dtype = torch.float32 + model_dtype = 'fp32' + except: + raise AttributeError("ComfyUI version too old, can't autodecet properly. Set your dtypes manually.") + else: + print(f"Diffusion using using {diffusion_dtype}") + dtype = convert_dtype(diffusion_dtype) + model_dtype = diffusion_dtype + + + if not hasattr(self, "model") or self.model is None or self.current_config != custom_config: + self.current_config = custom_config + self.model = None + + mm.soft_empty_cache() + + config = OmegaConf.load(config_path) + + if XFORMERS_IS_AVAILABLE: + config.model.params.control_stage_config.params.spatial_transformer_attn_type = "softmax-xformers" + config.model.params.network_config.params.spatial_transformer_attn_type = "softmax-xformers" + config.model.params.first_stage_config.params.ddconfig.attn_type = "vanilla-xformers" + + config.model.params.diffusion_dtype = model_dtype + config.model.target = ".SUPIR.models.SUPIR_model_v2.SUPIRModel" + pbar = comfy.utils.ProgressBar(7) + + self.model = instantiate_from_config(config.model).cpu() + pbar.update(1) + try: + print(f'Attempting to load SUPIR model: [{SUPIR_MODEL_PATH}]') + supir_state_dict = load_state_dict(SUPIR_MODEL_PATH) + pbar.update(1) + except: + raise Exception("Failed to load SUPIR model") + try: + print(f"Attempting to load SDXL model: [{SDXL_MODEL_PATH}]") + sdxl_state_dict = load_state_dict(SDXL_MODEL_PATH) + pbar.update(1) + except: + raise Exception("Failed to load SDXL model") + self.model.load_state_dict(supir_state_dict, strict=False) + pbar.update(1) + self.model.load_state_dict(sdxl_state_dict, strict=False) + pbar.update(1) + + del supir_state_dict + + #first clip model from SDXL checkpoint + try: + print("Loading first clip model from SDXL checkpoint") + + replace_prefix = {} + replace_prefix["conditioner.embedders.0.transformer."] = "" + + sd = comfy.utils.state_dict_prefix_replace(sdxl_state_dict, replace_prefix, filter_keys=False) + clip_text_config = CLIPTextConfig.from_pretrained(clip_config_path) + self.model.conditioner.embedders[0].tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path) + self.model.conditioner.embedders[0].transformer = CLIPTextModel(clip_text_config) + self.model.conditioner.embedders[0].transformer.load_state_dict(sd, strict=False) + self.model.conditioner.embedders[0].eval() + for param in self.model.conditioner.embedders[0].parameters(): + param.requires_grad = False + pbar.update(1) + except: + raise Exception("Failed to load first clip model from SDXL checkpoint") + + del sdxl_state_dict + + #second clip model from SDXL checkpoint + try: + print("Loading second clip model from SDXL checkpoint") + replace_prefix2 = {} + replace_prefix2["conditioner.embedders.1.model."] = "" + sd = comfy.utils.state_dict_prefix_replace(sd, replace_prefix2, filter_keys=True) + clip_g = build_text_model_from_openai_state_dict(sd, cast_dtype=dtype) + self.model.conditioner.embedders[1].model = clip_g + pbar.update(1) + except: + raise Exception("Failed to load second clip model from SDXL checkpoint") + + del sd, clip_g + mm.soft_empty_cache() + + self.model.to(dtype) + + #only unets and/or vae to fp8 + if fp8_unet: + self.model.model.to(torch.float8_e4m3fn) + + return (self.model, self.model.first_stage_model,) + +class SUPIR_model_loader_v2: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "model" :("MODEL",), + "clip": ("CLIP",), + "vae": ("VAE",), + "supir_model": (folder_paths.get_filename_list("checkpoints"),), + "fp8_unet": ("BOOLEAN", {"default": False}), + "diffusion_dtype": ( + [ + 'fp16', + 'bf16', + 'fp32', + 'auto' + ], { + "default": 'auto' + }), + }, + } + + RETURN_TYPES = ("SUPIRMODEL", "SUPIRVAE") + RETURN_NAMES = ("SUPIR_model","SUPIR_VAE",) + FUNCTION = "process" + CATEGORY = "SUPIR" + + def process(self, supir_model, diffusion_dtype, fp8_unet, model, clip, vae): + device = mm.get_torch_device() + mm.unload_all_models() + + SUPIR_MODEL_PATH = folder_paths.get_full_path("checkpoints", supir_model) + + config_path = os.path.join(script_directory, "options/SUPIR_v0.yaml") + clip_config_path = os.path.join(script_directory, "configs/clip_vit_config.json") + tokenizer_path = os.path.join(script_directory, "configs/tokenizer") + + custom_config = { + 'diffusion_dtype': diffusion_dtype, + 'supir_model': supir_model, + 'fp8_unet': fp8_unet, 'model': model, "clip": clip, "vae": vae @@ -666,21 +807,16 @@ class SUPIR_model_loader: except: raise Exception("Failed to load SUPIR model") try: - if model is None or clip is None or vae is None: - print(f"Attempting to load SDXL model: [{SDXL_MODEL_PATH}]") - sdxl_state_dict = load_state_dict(SDXL_MODEL_PATH) - else: - assert model is not None and clip is not None and vae is not None, "Need to pass model, clip, and vae" - print(f"Attempting to load SDXL model from node inputs") - clip_sd = None - load_models = [model] - load_models.append(clip.load_model()) - clip_sd = clip.get_sd() + print(f"Attempting to load SDXL model from node inputs") + clip_sd = None + load_models = [model] + load_models.append(clip.load_model()) + clip_sd = clip.get_sd() - mm.load_models_gpu(load_models) - - sd = model.model.state_dict_for_saving(clip_sd, vae.get_sd(), None) - sdxl_state_dict = sd + mm.load_models_gpu(load_models) + + sd = model.model.state_dict_for_saving(clip_sd, vae.get_sd(), None) + sdxl_state_dict = sd pbar.update(1) except: raise Exception("Failed to load SDXL model") @@ -734,7 +870,7 @@ class SUPIR_model_loader: self.model.model.to(torch.float8_e4m3fn) return (self.model, self.model.first_stage_model,) - + class SUPIR_tiles: @classmethod def INPUT_TYPES(s):