mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
flux add safetensors unet load
This commit is contained in:
+13
-1
@@ -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'
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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
@@ -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
Reference in New Issue
Block a user