auto-set upcast if first decode fails, fix flux

This commit is contained in:
Vladimir Mandic
2024-09-21 13:32:58 -04:00
parent ae221f883e
commit f059313333
6 changed files with 105 additions and 84 deletions
+4
View File
@@ -8,6 +8,7 @@
- fix diffusers local model name parsing
- full prompt parser will auto-select `xhinker` for flux models
- controlnet support for img2img and inpaint (in addition to previous txt2img controlnet)
- allow separate vae load
- **xyz grid** full refactor
- multi-mode: *selectable-script* and *alwayson-script*
- allow usage combined with other scripts
@@ -57,6 +58,8 @@
- hide token counter until tokens are known
- minor ui optimizations
- massive log cleanup
- **experimental**
- flux t5 load from gguf: requires transformers pr
## Update for 2024-09-13
@@ -163,6 +166,7 @@ Examples:
- **prompt enhance**: improve quality and/or verbosity of your prompts
simply select in *scripts -> prompt enhance*
uses [gokaygokay/Flux-Prompt-Enhance](https://huggingface.co/gokaygokay/Flux-Prompt-Enhance) model
- **decode** auto-set upcast if first decode fails
- **taesd** configurable number of layers
can be used to speed-up taesd decoding by reducing number of ops
e.g. if generating 1024px image, reducing layers by 1 will result in preview being 512px
+23 -22
View File
@@ -5,7 +5,7 @@ import diffusers
import transformers
from safetensors.torch import load_file
from huggingface_hub import hf_hub_download
from modules import shared, devices, modelloader, sd_models
from modules import shared, devices, modelloader, sd_models, sd_unet, model_te
debug = shared.log.trace if os.environ.get('SD_LOAD_DEBUG', None) is not None else lambda *args, **kwargs: None
@@ -33,7 +33,7 @@ def load_flux_quanto(checkpoint_info):
from optimum import quanto # pylint: disable=no-name-in-module
from optimum.quanto import requantize # pylint: disable=no-name-in-module
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to import optimum-quanto: {e}")
shared.log.error(f"Load model: type=FLUX Failed to import optimum-quanto: {e}")
raise
quanto.tensor.qbits.QBitsTensor.create = lambda *args, **kwargs: quanto.tensor.qbits.QBitsTensor(*args, **kwargs)
@@ -44,7 +44,7 @@ def load_flux_quanto(checkpoint_info):
try:
quantization_map = os.path.join(repo_path, "transformer", "quantization_map.json")
debug(f'Loading FLUX: quantization map="{quantization_map}" repo="{checkpoint_info.name}" component="transformer"')
debug(f'Load model: type=FLUX quantization map="{quantization_map}" repo="{checkpoint_info.name}" component="transformer"')
if not os.path.exists(quantization_map):
repo_id = sd_models.path_to_repo(checkpoint_info.name)
quantization_map = hf_hub_download(repo_id, subfolder='transformer', filename='quantization_map.json', cache_dir=shared.opts.diffusers_dir)
@@ -60,16 +60,16 @@ def load_flux_quanto(checkpoint_info):
try:
transformer = transformer.to(dtype=devices.dtype)
except Exception:
shared.log.error(f"Loading FLUX: Failed to cast transformer to {devices.dtype}, set dtype to {transformer.dtype}")
shared.log.error(f"Load model: type=FLUX Failed to cast transformer to {devices.dtype}, set dtype to {transformer.dtype}")
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to load Quanto transformer: {e}")
shared.log.error(f"Load model: type=FLUX Failed to load Quanto transformer: {e}")
if debug:
from modules import errors
errors.display(e, 'FLUX Quanto:')
try:
quantization_map = os.path.join(repo_path, "text_encoder_2", "quantization_map.json")
debug(f'Loading FLUX: quantization map="{quantization_map}" repo="{checkpoint_info.name}" component="text_encoder_2"')
debug(f'Load model: type=FLUX quantization map="{quantization_map}" repo="{checkpoint_info.name}" component="text_encoder_2"')
if not os.path.exists(quantization_map):
repo_id = sd_models.path_to_repo(checkpoint_info.name)
quantization_map = hf_hub_download(repo_id, subfolder='text_encoder_2', filename='quantization_map.json', cache_dir=shared.opts.diffusers_dir)
@@ -87,9 +87,9 @@ def load_flux_quanto(checkpoint_info):
try:
text_encoder_2 = text_encoder_2.to(dtype=devices.dtype)
except Exception:
shared.log.error(f"Loading FLUX: Failed to cast text encoder to {devices.dtype}, set dtype to {text_encoder_2.dtype}")
shared.log.error(f"Load model: type=FLUX Failed to cast text encoder to {devices.dtype}, set dtype to {text_encoder_2.dtype}")
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to load Quanto text encoder: {e}")
shared.log.error(f"Load model: type=FLUX Failed to load Quanto text encoder: {e}")
if debug:
from modules import errors
errors.display(e, 'FLUX Quanto:')
@@ -122,7 +122,7 @@ def load_flux_bnb(checkpoint_info, diffusers_load_config): # pylint: disable=unu
else:
transformer = diffusers.FluxTransformer2DModel.from_single_file(repo_path, **diffusers_load_config)
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to load BnB transformer: {e}")
shared.log.error(f"Load model: type=FLUX Failed to load BnB transformer: {e}")
transformer, text_encoder_2 = None, None
if debug:
from modules import errors
@@ -131,7 +131,7 @@ def load_flux_bnb(checkpoint_info, diffusers_load_config): # pylint: disable=unu
def load_flux_gguf(file_path): # TODO add support for GGUF flux models
shared.log.error(f"Loading FLUX: GGUF UNET is not supported: {file_path}")
shared.log.error(f"Load model: type=FLUX GGUF UNET is not supported: {file_path}")
"""
with torch.device("meta"):
transformer = diffusers.FluxTransformer2DModel.from_config(os.path.join("configs", "flux", "transformer", "config.json")).to(dtype=devices.dtype)
@@ -182,8 +182,8 @@ def load_transformer(file_path): # triggered by opts.sd_unet change
def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_checkpoint change
quant = get_quant(checkpoint_info.path)
repo_id = sd_models.path_to_repo(checkpoint_info.name)
shared.log.debug(f'Loading FLUX: model="{checkpoint_info.name}" repo="{repo_id}" unet="{shared.opts.sd_unet}" t5="{shared.opts.sd_text_encoder}" vae="{shared.opts.sd_vae}" quant={quant} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}')
debug(f'Loading FLUX: config={diffusers_load_config}')
shared.log.debug(f'Load model: type=FLUX model="{checkpoint_info.name}" repo="{repo_id}" unet="{shared.opts.sd_unet}" t5="{shared.opts.sd_text_encoder}" vae="{shared.opts.sd_vae}" quant={quant} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}')
debug(f'Load model: type=FLUX config={diffusers_load_config}')
modelloader.hf_login()
transformer = None
@@ -193,8 +193,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
# load overrides if any
if shared.opts.sd_unet != 'None':
try:
debug(f'Loading FLUX: unet="{shared.opts.sd_unet}"')
from modules import sd_unet
debug(f'Load model: type=FLUX unet="{shared.opts.sd_unet}"')
_transformer = load_transformer(sd_unet.unet_dict[shared.opts.sd_unet])
if _transformer is not None:
transformer = _transformer
@@ -202,14 +201,14 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
shared.opts.sd_unet = 'None'
sd_unet.failed_unet.append(shared.opts.sd_unet)
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to load UNet: {e}")
shared.log.error(f"Load model: type=FLUX Failed to load UNet: {e}")
shared.opts.sd_unet = 'None'
if debug:
from modules import errors
errors.display(e, 'FLUX UNet:')
if shared.opts.sd_text_encoder != 'None':
try:
debug(f'Loading FLUX: t5="{shared.opts.sd_text_encoder}"')
debug(f'Load model: type=FLUX t5="{shared.opts.sd_text_encoder}"')
from modules.model_te import load_t5
_text_encoder_2 = load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir)
if _text_encoder_2 is not None:
@@ -217,14 +216,14 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
else:
shared.opts.sd_text_encoder = 'None'
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to load T5: {e}")
shared.log.error(f"Load model: type=FLUX Failed to load T5: {e}")
shared.opts.sd_text_encoder = 'None'
if debug:
from modules import errors
errors.display(e, 'FLUX T5:')
if shared.opts.sd_vae != 'None' and shared.opts.sd_vae != 'Automatic':
try:
debug(f'Loading FLUX: vae="{shared.opts.sd_vae}"')
debug(f'Load model: type=FLUX vae="{shared.opts.sd_vae}"')
from modules import sd_vae
# vae = sd_vae.load_vae_diffusers(None, sd_vae.vae_dict[shared.opts.sd_vae], 'override')
vae_file = sd_vae.vae_dict[shared.opts.sd_vae]
@@ -232,7 +231,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
vae_config = os.path.join('configs', 'flux', 'vae', 'config.json')
vae = diffusers.AutoencoderKL.from_single_file(vae_file, config=vae_config, **diffusers_load_config)
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to load VAE: {e}")
shared.log.error(f"Load model: type=FLUX Failed to load VAE: {e}")
shared.opts.sd_vae = 'None'
if debug:
from modules import errors
@@ -248,7 +247,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
if _text_encoder is not None:
text_encoder_2 = _text_encoder
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to load NF4 components: {e}")
shared.log.error(f"Load model: type=FLUX Failed to load NF4 components: {e}")
if debug:
from modules import errors
errors.display(e, 'FLUX NF4:')
@@ -260,7 +259,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
if _text_encoder is not None:
text_encoder_2 = _text_encoder
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to load Quanto components: {e}")
shared.log.error(f"Load model: type=FLUX Failed to load Quanto components: {e}")
if debug:
from modules import errors
errors.display(e, 'FLUX Quanto:')
@@ -269,11 +268,13 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
components = {}
if transformer is not None:
components['transformer'] = transformer
sd_unet.loaded_unet = shared.opts.sd_unet
if text_encoder_2 is not None:
components['text_encoder_2'] = text_encoder_2
model_te.loaded_te = shared.opts.sd_text_encoder
if vae is not None:
components['vae'] = vae
shared.log.debug(f'Loading FLUX: preloaded={list(components)}')
shared.log.debug(f'Load model: type=FLUX preloaded={list(components)}')
if repo_id == 'sayakpaul/flux.1-dev-nf4':
repo_id = 'black-forest-labs/FLUX.1-dev' # workaround since sayakpaul model is missing model_index.json
pipe = diffusers.FluxPipeline.from_pretrained(repo_id, cache_dir=shared.opts.diffusers_dir, **components, **diffusers_load_config)
+3 -3
View File
@@ -26,7 +26,7 @@ def load_bnb():
global bnb # pylint: disable=global-statement
bnb = bitsandbytes
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to import bitsandbytes: {e}")
shared.log.error(f"Load model: type=FLUX Failed to import bitsandbytes: {e}")
raise
@@ -184,7 +184,7 @@ def load_flux_nf4(checkpoint_info):
try:
converted_state_dict = convert_flux_transformer_checkpoint_to_diffusers(original_state_dict)
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to convert UNET: {e}")
shared.log.error(f"Load model: type=FLUX Failed to convert UNET: {e}")
if debug:
from modules import errors
errors.display(e, 'FLUX convert:')
@@ -211,7 +211,7 @@ def load_flux_nf4(checkpoint_info):
create_quantized_param(transformer, param, param_name, target_device=0, state_dict=original_state_dict, pre_quantized=True)
except Exception as e:
transformer, text_encoder_2 = None, None
shared.log.error(f"Loading FLUX: Failed to load UNET: {e}")
shared.log.error(f"Load model: type=FLUX Failed to load UNET: {e}")
if debug:
from modules import errors
errors.display(e, 'FLUX:')
+10 -4
View File
@@ -339,6 +339,7 @@ def img2img_image_conditioning(p, source_image, latent_image, image_mask=None):
def validate_sample(tensor):
if not isinstance(tensor, np.ndarray) and not isinstance(tensor, torch.Tensor):
return tensor
dtype = tensor.dtype
if tensor.dtype == torch.bfloat16: # numpy does not support bf16
tensor = tensor.to(torch.float16)
if isinstance(tensor, torch.Tensor) and hasattr(tensor, 'detach'):
@@ -346,16 +347,21 @@ def validate_sample(tensor):
elif isinstance(tensor, np.ndarray):
sample = tensor
else:
shared.log.warning(f'Unknown sample type: {type(tensor)}')
shared.log.warning(f'Decode: type={type(tensor)} unknown sample')
sample = 255.0 * np.moveaxis(sample, 0, 2) if not shared.native else 255.0 * sample
with warnings.catch_warnings(record=True) as w:
cast = sample.astype(np.uint8)
if len(w) > 0:
minimum, maximum, mean = np.min(cast), np.max(cast), np.mean(cast)
if len(w) > 0 or minimum == maximum:
nans = np.isnan(sample).sum()
cast = np.nan_to_num(sample)
minimum, maximum, mean = np.min(cast), np.max(cast), np.mean(cast)
cast = cast.astype(np.uint8)
shared.log.error(f'Failed to validate samples: sample={sample.shape} min={minimum:.2f} max={maximum:.2f} mean={mean:.2f} invalid={nans}')
vae = shared.sd_model.vae.dtype if hasattr(shared.sd_model, 'vae') else None
upcast = getattr(shared.sd_model.vae.config, 'force_upcast', None) if hasattr(shared.sd_model, 'vae') and hasattr(shared.sd_model.vae, 'config') else None
shared.log.error(f'Decode: sample={sample.shape} invalid={nans} mean={mean} dtype={dtype} vae={vae} upcast={upcast} failed to validate')
if upcast is not None and not upcast:
setattr(shared.sd_model.vae.config, 'force_upcast', True) # noqa: B010
shared.log.warning('Decode: upcast=True set, retry operation')
return cast
+13 -5
View File
@@ -36,21 +36,28 @@ def create_latents(image, p, dtype=None, device=None):
def full_vae_decode(latents, model):
t0 = time.time()
if not hasattr(model, 'vae'):
shared.log.error('VAE not found in model')
return []
if debug:
devices.torch_gc(force=True)
shared.mem_mon.reset()
base_device = None
if shared.opts.diffusers_move_unet and not getattr(model, 'has_accelerate', False):
base_device = sd_models.move_base(model, devices.cpu)
if shared.opts.diffusers_offload_mode == "balanced":
shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model)
elif not shared.opts.diffusers_offload_mode == "sequential" and hasattr(model, 'vae'):
elif shared.opts.diffusers_offload_mode != "sequential":
sd_models.move_model(model.vae, devices.device)
latents.to(model.vae.device)
upcast = (model.vae.dtype == torch.float16) and getattr(model.vae.config, 'force_upcast', False) and hasattr(model, 'upcast_vae')
if upcast: # this is done by diffusers automatically if output_type != 'latent'
model.upcast_vae()
upcast = (model.vae.dtype == torch.float16) and getattr(model.vae.config, 'force_upcast', False)
if upcast:
if hasattr(model, 'upcast_vae'): # this is done by diffusers automatically if output_type != 'latent'
model.upcast_vae()
model.vae = model.vae.to(dtype=torch.float32)
latents = latents.to(torch.float32)
if getattr(model.vae, "post_quant_conv", None) is not None:
latents = latents.to(next(iter(model.vae.post_quant_conv.parameters())).dtype)
@@ -70,6 +77,7 @@ def full_vae_decode(latents, model):
latents = latents / scaling_factor
if shift_factor:
latents = latents + shift_factor
latents = latents.to(model.vae.device)
decoded = model.vae.decode(latents, return_dict=False)[0]
# delete vae after OpenVINO compile
+52 -50
View File
@@ -161,14 +161,14 @@ def list_models():
if shared.cmd_opts.ckpt is not None:
if not os.path.exists(shared.cmd_opts.ckpt) and not shared.native:
if shared.cmd_opts.ckpt.lower() != "none":
shared.log.warning(f"Requested model not found: {shared.cmd_opts.ckpt}")
shared.log.warning(f'Load model: path="{shared.cmd_opts.ckpt}" not found')
else:
checkpoint_info = CheckpointInfo(shared.cmd_opts.ckpt)
if checkpoint_info.name is not None:
checkpoint_info.register()
shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title
elif shared.cmd_opts.ckpt != shared.default_sd_model_file and shared.cmd_opts.ckpt is not None:
shared.log.warning(f"Model not found: {shared.cmd_opts.ckpt}")
shared.log.warning(f'Load model: path="{shared.cmd_opts.ckpt}" not found')
shared.log.info(f'Available Models: path="{shared.opts.ckpt_dir}" items={len(checkpoints_list)} time={time.time()-t0:.2f}')
checkpoints_list = dict(sorted(checkpoints_list.items(), key=lambda cp: cp[1].filename))
@@ -250,7 +250,7 @@ def select_checkpoint(op='model'):
return None
checkpoint_info = get_closet_checkpoint_match(model_checkpoint)
if checkpoint_info is not None:
shared.log.info(f'Select: {op}="{checkpoint_info.title if checkpoint_info is not None else None}"')
shared.log.info(f'Load {op}: select="{checkpoint_info.title if checkpoint_info is not None else None}"')
return checkpoint_info
if len(checkpoints_list) == 0:
shared.log.warning("Cannot generate without a checkpoint")
@@ -262,13 +262,13 @@ def select_checkpoint(op='model'):
# checkpoint_info = next(iter(checkpoints_list.values()))
if model_checkpoint is not None:
if model_checkpoint != 'model.ckpt' and model_checkpoint != 'stabilityai/stable-diffusion-xl-base-1.0':
shared.log.warning(f'Selected: {op}="{model_checkpoint}" not found')
shared.log.warning(f'Load {op}: select="{model_checkpoint}" not found')
else:
shared.log.info("Selecting first available checkpoint")
# shared.log.warning(f"Loading fallback checkpoint: {checkpoint_info.title}")
# shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title
else:
shared.log.info(f'Select: {op}="{checkpoint_info.title if checkpoint_info is not None else None}"')
shared.log.info(f'Load {op}: select="{checkpoint_info.title if checkpoint_info is not None else None}"')
return checkpoint_info
@@ -384,7 +384,7 @@ def read_metadata_from_safetensors(filename):
def read_state_dict(checkpoint_file, map_location=None): # pylint: disable=unused-argument
if not os.path.isfile(checkpoint_file):
shared.log.error(f"Model is not a file: {checkpoint_file}")
shared.log.error(f'Load dict: path="{checkpoint_file}" not a file')
return None
try:
pl_sd = None
@@ -421,7 +421,7 @@ def get_checkpoint_state_dict(checkpoint_info: CheckpointInfo, timer):
if not os.path.isfile(checkpoint_info.filename):
return None
if checkpoint_info in checkpoints_loaded:
shared.log.info("Model weights loading: from cache")
shared.log.info("Load model: cache")
checkpoints_loaded.move_to_end(checkpoint_info, last=True) # FIFO -> LRU cache
return checkpoints_loaded[checkpoint_info]
res = read_state_dict(checkpoint_info.filename)
@@ -437,7 +437,7 @@ def get_checkpoint_state_dict(checkpoint_info: CheckpointInfo, timer):
def load_model_weights(model: torch.nn.Module, checkpoint_info: CheckpointInfo, state_dict, timer):
_pipeline, _model_type = detect_pipeline(checkpoint_info.path, 'model')
shared.log.debug(f'Model weights loading: {memory_stats()}')
shared.log.debug(f'Load model: memory={memory_stats()}')
timer.record("hash")
if model_data.sd_dict == 'None':
shared.opts.data["sd_model_checkpoint"] = checkpoint_info.title
@@ -446,7 +446,7 @@ def load_model_weights(model: torch.nn.Module, checkpoint_info: CheckpointInfo,
try:
model.load_state_dict(state_dict, strict=False)
except Exception as e:
shared.log.error(f'Error loading model weights: {checkpoint_info.filename}')
shared.log.error(f'Load model: path="{checkpoint_info.filename}"')
shared.log.error(' '.join(str(e).splitlines()[:2]))
return False
del state_dict
@@ -638,21 +638,21 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False):
# get actual pipeline
pipeline = shared_items.get_pipelines().get(guess, None)
if not quiet:
shared.log.info(f'Autodetect: {op}="{guess}" class={pipeline.__name__} file="{f}" size={size}MB')
shared.log.info(f'Autodetect {op}: detect="{guess}" class={pipeline.__name__} file="{f}" size={size}MB')
except Exception as e:
shared.log.error(f'Error detecting diffusers pipeline: model={f} {e}')
shared.log.error(f'Autodetect {op}: file="{f}" {e}')
return None, None
else:
try:
size = round(os.path.getsize(f) / 1024 / 1024)
pipeline = shared_items.get_pipelines().get(guess, None)
if not quiet:
shared.log.info(f'Diffusers: {op}="{guess}" class={pipeline.__name__} file="{f}" size={size}MB')
shared.log.info(f'Load {op}: detect="{guess}" class={pipeline.__name__} file="{f}" size={size}MB')
except Exception as e:
shared.log.error(f'Error loading diffusers pipeline: model={f} {e}')
shared.log.error(f'Load {op}: detect="{guess}" file="{f}" {e}')
if pipeline is None:
shared.log.warning(f'Autodetect: pipeline not recognized: {guess}: {op}={f} size={size}')
shared.log.warning(f'Load {op}: detect="{guess}" file="{f}" size={size} not recognized')
pipeline = diffusers.StableDiffusionPipeline
return pipeline, guess
@@ -713,13 +713,13 @@ def set_diffuser_options(sd_model, vae = None, op: str = 'model', offload=True):
sd_model.fuse_qkv_projections()
shared.log.debug(f'Setting {op}: fused-qkv=True')
except Exception as e:
shared.log.error(f'Error enabling fused projections: {e}')
shared.log.error(f'Setting {op}: fused-qkv=True {e}')
if shared.opts.diffusers_fuse_projections and hasattr(sd_model, 'transformer') and hasattr(sd_model.transformer, 'fuse_qkv_projections'):
try:
sd_model.transformer.fuse_qkv_projections()
shared.log.debug(f'Setting {op}: fused-qkv=True')
except Exception as e:
shared.log.error(f'Error enabling fused projections: {e}')
shared.log.error(f'Setting {op}: fused-qkv=True {e}')
if shared.opts.diffusers_eval:
def eval_model(model, op=None, sd_model=None): # pylint: disable=unused-argument
if hasattr(model, "requires_grad_"):
@@ -761,7 +761,7 @@ def set_diffuser_offload(sd_model, op: str = 'model'):
sd_model.maybe_free_model_hooks()
sd_model.has_accelerate = True
except Exception as e:
shared.log.error(f'Model offload error: mode={shared.opts.diffusers_offload_mode} {e}')
shared.log.error(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} {e}')
if hasattr(sd_model, "enable_sequential_cpu_offload"):
if shared.opts.diffusers_offload_mode == "sequential":
try:
@@ -782,13 +782,13 @@ def set_diffuser_offload(sd_model, op: str = 'model'):
sd_model.enable_sequential_cpu_offload(device=devices.device)
sd_model.has_accelerate = True
except Exception as e:
shared.log.error(f'Model offload error: mode={shared.opts.diffusers_offload_mode} {e}')
shared.log.error(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} {e}')
if shared.opts.diffusers_offload_mode == "balanced":
try:
shared.log.debug(f'Setting {op}: offload={shared.opts.diffusers_offload_mode}')
sd_model = apply_balanced_offload(sd_model)
except Exception as e:
shared.log.error(f'Model offload error: mode={shared.opts.diffusers_offload_mode} {e}')
shared.log.error(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} {e}')
def apply_balanced_offload(sd_model):
@@ -929,12 +929,14 @@ def move_model(model, device=None, force=False):
def move_base(model, device):
key = 'unet'
if isinstance(model, diffusers.FluxPipeline):
if hasattr(model, 'transformer'):
key = 'transformer'
if not hasattr(model, key):
elif hasattr(model, 'unet'):
key = 'unet'
else:
shared.log.warning(f'Model move: model={model.__class__} device={device} key=unknown')
return None
shared.log.debug(f'Moving to CPU: model={key}')
shared.log.debug(f'Model move: module={key} device={device}')
model = getattr(model, key)
R = model.device
move_model(model, device)
@@ -1035,7 +1037,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
if shared.opts.diffusers_model_load_variant != 'default':
diffusers_load_config['variant'] = shared.opts.diffusers_model_load_variant
if shared.opts.diffusers_pipeline == 'Custom Diffusers Pipeline' and len(shared.opts.custom_diffusers_pipeline) > 0:
shared.log.debug(f'Diffusers custom pipeline: {shared.opts.custom_diffusers_pipeline}')
shared.log.debug(f'Model pipeline: pipeline="{shared.opts.custom_diffusers_pipeline}"')
diffusers_load_config['custom_pipeline'] = shared.opts.custom_diffusers_pipeline
# if 'LCM' in checkpoint_info.path:
# diffusers_load_config['custom_pipeline'] = 'latent_consistency_txt2img'
@@ -1055,14 +1057,14 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
ckpt_basename = os.path.basename(shared.cmd_opts.ckpt)
model_name = modelloader.find_diffuser(ckpt_basename)
if model_name is not None:
shared.log.info(f'Load model {op}: {model_name}')
shared.log.info(f'Load model {op}: path="{model_name}"')
model_file = modelloader.download_diffusers_model(hub_id=model_name, variant=diffusers_load_config.get('variant', None))
try:
shared.log.debug(f'Model load {op} config: {diffusers_load_config}')
shared.log.debug(f'Load {op}: config={diffusers_load_config}')
sd_model = diffusers.DiffusionPipeline.from_pretrained(model_file, **diffusers_load_config)
except Exception as e:
shared.log.error(f'Failed loading model: {model_file} {e}')
errors.display(e, f'Load model: {model_file}')
errors.display(e, f'Load model: path="{model_file}"')
list_models() # rescan for downloaded model
checkpoint_info = CheckpointInfo(model_name)
@@ -1071,7 +1073,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
unload_model_weights(op=op)
return
shared.log.debug(f'Diffusers loading: path="{checkpoint_info.path}"')
shared.log.debug(f'Load {op}: path="{checkpoint_info.path}"')
pipeline, model_type = detect_pipeline(checkpoint_info.path, op)
vae = None
@@ -1091,7 +1093,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
from modules.model_stablecascade import load_cascade_combined
sd_model = load_cascade_combined(checkpoint_info, diffusers_load_config)
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}')
if debug_load:
errors.display(e, 'Load')
return
@@ -1100,7 +1102,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
pipeline = diffusers.utils.get_class_from_dynamic_module('instaflow_one_step', module_file='pipeline.py')
sd_model = pipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config)
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}')
if debug_load:
errors.display(e, 'Load')
return
@@ -1110,7 +1112,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
sd_model = SegMoEPipeline(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config)
sd_model = sd_model.pipe # segmoe pipe does its stuff in __init__ and __call__ is the original pipeline
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}')
if debug_load:
errors.display(e, 'Load')
return
@@ -1119,7 +1121,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
from modules.model_pixart import load_pixart
sd_model = load_pixart(checkpoint_info, diffusers_load_config)
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}')
if debug_load:
errors.display(e, 'Load')
return
@@ -1128,7 +1130,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
from modules.model_lumina import load_lumina
sd_model = load_lumina(checkpoint_info, diffusers_load_config)
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}')
if debug_load:
errors.display(e, 'Load')
return
@@ -1137,7 +1139,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
from modules.model_kolors import load_kolors
sd_model = load_kolors(checkpoint_info, diffusers_load_config)
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}')
if debug_load:
errors.display(e, 'Load')
return
@@ -1146,7 +1148,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
from modules.model_auraflow import load_auraflow
sd_model = load_auraflow(checkpoint_info, diffusers_load_config)
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}')
if debug_load:
errors.display(e, 'Load')
return
@@ -1155,18 +1157,18 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
from modules.model_flux import load_flux
sd_model = load_flux(checkpoint_info, diffusers_load_config)
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}')
if debug_load:
errors.display(e, 'Load')
return
elif model_type in ['Stable Diffusion 3']:
try:
from modules.model_sd3 import load_sd3
shared.log.debug('Loading: model="Stable Diffusion 3" variant=medium type=diffusers')
shared.log.debuc(f'Load {op}: model="Stable Diffusion 3" variant=medium')
shared.opts.scheduler = 'Default'
sd_model = load_sd3(cache_dir=shared.opts.diffusers_dir, config=diffusers_load_config.get('config', None))
except Exception as e:
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}')
if debug_load:
errors.display(e, 'Load')
return
@@ -1174,7 +1176,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
try:
sd_model = pipeline.from_pretrained(checkpoint_info.path)
except Exception as e:
shared.log.error(f'ONNX Failed loading {op}: {checkpoint_info.path} {e}')
shared.log.error(f'Load {op}: type=ONNX path="{checkpoint_info.path}" {e}')
if debug_load:
errors.display(e, 'Load')
return
@@ -1182,14 +1184,14 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
err1, err2, err3 = None, None, None
# diffusers_load_config['use_safetensors'] = True
if debug_load:
shared.log.debug(f'Diffusers load args: {diffusers_load_config}')
shared.log.debug(f'Load {op}: args={diffusers_load_config}')
try: # 1 - autopipeline, best choice but not all pipelines are available
try:
sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config)
sd_model.model_type = sd_model.__class__.__name__
except ValueError as e:
if 'no variant default' in str(e):
shared.log.warning(f'Load: variant={diffusers_load_config["variant"]} model={checkpoint_info.path} using default variant')
shared.log.warning(f'Load {op}: variant={diffusers_load_config["variant"]} model="{checkpoint_info.path}" using default variant')
diffusers_load_config.pop('variant', None)
sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config)
sd_model.model_type = sd_model.__class__.__name__
@@ -1219,13 +1221,13 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
if debug_load:
errors.display(e, "Load StableDiffusionPipeline")
if err3 is not None:
shared.log.error(f'Failed loading {op}: {checkpoint_info.path} auto={err1} diffusion={err2}')
shared.log.error(f'Load {op}: {checkpoint_info.path} auto={err1} diffusion={err2}')
return
elif os.path.isfile(checkpoint_info.path) and checkpoint_info.path.lower().endswith('.safetensors'):
diffusers_load_config["local_files_only"] = diffusers_version < 28 # must be true for old diffusers, otherwise false but we override config for sd15/sdxl
diffusers_load_config["extract_ema"] = shared.opts.diffusers_extract_ema
if pipeline is None:
shared.log.error(f'Diffusers {op} pipeline not initialized: {shared.opts.diffusers_pipeline}')
shared.log.error(f'Load {op}: pipeline={shared.opts.diffusers_pipeline} not initialized')
return
try:
if model_type.startswith('Stable Diffusion'):
@@ -1235,7 +1237,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
model_config = get_load_config(checkpoint_info.path, model_type, config_type='json')
if model_config is not None:
if debug_load:
shared.log.debug(f'Model config: path="{model_config}"')
shared.log.debug(f'Load {op}: config="{model_config}"')
diffusers_load_config['config'] = model_config
if model_type.startswith('Stable Diffusion 3'):
from modules.model_sd3 import load_sd3
@@ -1276,7 +1278,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
errors.display(e, f'loading {op}={checkpoint_info.path} pipeline={shared.opts.diffusers_pipeline}/{sd_model.__class__.__name__}')
return
else:
shared.log.error(f'Diffusers cannot load: {op}={checkpoint_info.path}')
shared.log.error(f'Load {op}: path="{checkpoint_info.path}" failed')
return
if "StableDiffusion" in sd_model.__class__.__name__:
@@ -1351,7 +1353,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
timer.record("compile")
except Exception as e:
shared.log.error("Failed to load model")
shared.log.error(f"Load {op}: {e}")
errors.display(e, "Model")
devices.torch_gc(force=True)
@@ -1653,9 +1655,9 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None,
state_dict = get_checkpoint_state_dict(checkpoint_info, timer)
checkpoint_config = sd_models_config.find_checkpoint_config(state_dict, checkpoint_info)
if state_dict is None or checkpoint_config is None:
shared.log.error(f"Failed to load checkpooint: {checkpoint_info.filename}")
shared.log.error(f'Load {op}: path="{checkpoint_info.filename}"')
if current_checkpoint_info is not None:
shared.log.info(f"Restoring previous checkpoint: {current_checkpoint_info.filename}")
shared.log.info(f'Load {op}: previous="{current_checkpoint_info.filename}" restore')
load_model(current_checkpoint_info, None)
return
shared.log.debug(f'Model dict loaded: {memory_stats()}')
@@ -1729,11 +1731,11 @@ def reload_text_encoder(initial=False):
t5 = [k for k, v in signature.items() if 'T5EncoderModel' in str(v)]
if len(t5) > 0:
from modules.model_te import set_t5
shared.log.debug(f'Load: t5={shared.opts.sd_text_encoder} module="{t5[0]}"')
shared.log.debug(f'Load module: type=t5 path="{shared.opts.sd_text_encoder}" module="{t5[0]}"')
set_t5(pipe=shared.sd_model, module=t5[0], t5=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir)
elif hasattr(shared.sd_model, 'text_encoder_3'):
from modules.model_te import set_t5
shared.log.debug(f'Load: t5={shared.opts.sd_text_encoder} module="text_encoder_3"')
shared.log.debug(f'Load module: type=t5 path="{shared.opts.sd_text_encoder}" module="text_encoder_3"')
set_t5(pipe=shared.sd_model, module='text_encoder_3', t5=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir)
elif hasattr(shared.sd_model, 'text_encoder') and 'vit' in shared.opts.sd_text_encoder.lower():
from modules.model_te import set_clip