flux add safetensors unet load

This commit is contained in:
Vladimir Mandic
2024-08-28 21:53:09 -04:00
parent 01fc706433
commit 8cf959047b
5 changed files with 35 additions and 11 deletions
+13 -1
View File
@@ -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'
-1
View File
@@ -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())
+2 -2
View File
@@ -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:
+19 -6
View File
@@ -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():
+1 -1
Submodule wiki updated: ba1e031e30...4663accc0c