From 8cf959047b881880a5b601f63b74aa84f67eac83 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Wed, 28 Aug 2024 21:53:09 -0400 Subject: [PATCH] flux add safetensors unet load --- modules/model_flux.py | 14 +++++++++++++- modules/model_flux_nf4.py | 1 - modules/model_stablecascade.py | 4 ++-- modules/sd_unet.py | 25 +++++++++++++++++++------ wiki | 2 +- 5 files changed, 35 insertions(+), 11 deletions(-) diff --git a/modules/model_flux.py b/modules/model_flux.py index 899e75af8..82a7a6456 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -1,4 +1,3 @@ -import os import json import torch import diffusers @@ -36,6 +35,19 @@ def load_quanto_text_encoder_2(repo_path): return text_encoder_2 +def load_transformer(file_path): + diffusers_load_config = { + "low_cpu_mem_usage": True, + "torch_dtype": devices.dtype, + "cache_dir": shared.opts.hfcache_dir, + } + from diffusers import FluxTransformer2DModel + transformer = FluxTransformer2DModel.from_single_file(file_path, **diffusers_load_config) + if transformer is None: + shared.log.error('Failed to load UNet model') + return transformer + + def load_flux(checkpoint_info, diffusers_load_config): if "qint8" in checkpoint_info.path.lower(): quant = 'qint8' diff --git a/modules/model_flux_nf4.py b/modules/model_flux_nf4.py index fe74d6330..1320e94a7 100644 --- a/modules/model_flux_nf4.py +++ b/modules/model_flux_nf4.py @@ -166,7 +166,6 @@ def load_flux_nf4(checkpoint_info, diffusers_load_config): raise with init_empty_weights(): - # config = FluxTransformer2DModel.load_config(checkpoint_info.path) config = FluxTransformer2DModel.load_config("black-forest-labs/flux.1-dev", subfolder="transformer") model = FluxTransformer2DModel.from_config(config).to(devices.dtype) expected_state_dict_keys = list(model.state_dict().keys()) diff --git a/modules/model_stablecascade.py b/modules/model_stablecascade.py index 7ff3569ec..831c313fc 100644 --- a/modules/model_stablecascade.py +++ b/modules/model_stablecascade.py @@ -271,7 +271,7 @@ class StableCascadeDecoderPipelineFixed(diffusers.StableCascadeDecoderPipeline): else: alphas_cumprod = [] - self._num_timesteps = len(timesteps) + self._num_timesteps = len(timesteps) # pylint: disable=attribute-defined-outside-init for i, t in enumerate(self.progress_bar(timesteps)): if not isinstance(self.scheduler, diffusers.DDPMWuerstchenScheduler): if len(alphas_cumprod) > 0: @@ -321,7 +321,7 @@ class StableCascadeDecoderPipelineFixed(diffusers.StableCascadeDecoderPipeline): f"Only the output types `pt`, `np`, `pil` and `latent` are supported not output_type={output_type}" ) - if not output_type == "latent": + if output_type != "latent": if shared.opts.diffusers_offload_mode == "balanced": shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) else: diff --git a/modules/sd_unet.py b/modules/sd_unet.py index a4aeaba5f..c948c223b 100644 --- a/modules/sd_unet.py +++ b/modules/sd_unet.py @@ -11,9 +11,6 @@ def load_unet(model): if shared.opts.sd_unet not in list(unet_dict): shared.log.error(f'UNet model not found: {shared.opts.sd_unet}') return - if (not hasattr(model, 'unet') or model.unet is None) and not (hasattr(model, 'prior_pipe') and hasattr(model.prior_pipe, "prior")): - shared.log.error('UNet not found in current model') - return config_file = os.path.splitext(unet_dict[shared.opts.sd_unet])[0] + '.json' if os.path.exists(config_file): config = shared.readfile(config_file) @@ -24,12 +21,28 @@ def load_unet(model): if "StableCascade" in model.__class__.__name__: from modules.model_stablecascade import load_prior prior_unet, prior_text_encoder = load_prior(unet_dict[shared.opts.sd_unet], config_file=config_file) - model.prior_pipe.prior = None # Prevent OOM - model.prior_pipe.prior = prior_unet.to(devices.device, dtype=devices.dtype_unet) + if prior_unet is not None: + model.prior_pipe.prior = None # Prevent OOM + model.prior_pipe.prior = prior_unet.to(devices.device, dtype=devices.dtype_unet) if prior_text_encoder is not None: model.prior_pipe.text_encoder = None # Prevent OOM model.prior_pipe.text_encoder = prior_text_encoder.to(devices.device, dtype=devices.dtype) + if "Flux" in model.__class__.__name__: + shared.log.info(f'Loading UNet: name="{shared.opts.sd_unet}" file="{unet_dict[shared.opts.sd_unet]}" offload={shared.opts.diffusers_offload_mode}') + from modules.model_flux import load_transformer + transformer = load_transformer(unet_dict[shared.opts.sd_unet]) + if transformer is not None: + model.transformer = None + if shared.opts.diffusers_offload_mode == 'none': + model.transformer = transformer.to(devices.device, devices.dtype) + else: + model.transformer = transformer + from modules.sd_models import set_diffuser_offload + set_diffuser_offload(model, 'model') else: + if not hasattr(model, 'unet') or model.unet is None: + shared.log.error('UNet not found in current model') + return shared.log.info(f'Loading UNet: name="{shared.opts.sd_unet}" file="{unet_dict[shared.opts.sd_unet]}" config="{config_file}"') from diffusers import UNet2DConditionModel from safetensors.torch import load_file @@ -38,9 +51,9 @@ def load_unet(model): unet.load_state_dict(state_dict) model.unet = unet.to(devices.device, devices.dtype_unet) except Exception as e: - unet = None shared.log.error(f'Failed to load UNet model: {e}') return + devices.torch_gc() def refresh_unet_list(): diff --git a/wiki b/wiki index ba1e031e3..4663accc0 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit ba1e031e30b08b0dda42c6fc84054a0acdd81163 +Subproject commit 4663accc0cfde6aef3aad86a29a249a4d4fcb902