diff --git a/modules/model_flux.py b/modules/model_flux.py index 6c00a47f7..0a070f856 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -26,14 +26,15 @@ def get_quant(file_path): return 'none' -def load_flux_quanto(checkpoint_info, diffusers_load_config, transformer, text_encoder_2): +def load_flux_quanto(checkpoint_info, diffusers_load_config): + transformer, text_encoder_2 = None, None from installer import install install('optimum-quanto', quiet=True) try: 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"FLUX: Failed to import optimum-quanto: {e}") + shared.log.error(f"Loading FLUX: Failed to import optimum-quanto: {e}") raise quanto.tensor.qbits.QBitsTensor.create = lambda *args, **kwargs: quanto.tensor.qbits.QBitsTensor(*args, **kwargs) @@ -42,7 +43,7 @@ def load_flux_quanto(checkpoint_info, diffusers_load_config, transformer, text_e else: repo_path = checkpoint_info.path - if transformer is None: + try: quantization_map = os.path.join(repo_path, "transformer", "quantization_map.json") if not os.path.exists(quantization_map): repo_id = checkpoint_info.name.replace('Diffusers/', '') @@ -59,10 +60,14 @@ def load_flux_quanto(checkpoint_info, diffusers_load_config, transformer, text_e try: transformer = transformer.to(dtype=devices.dtype) except Exception: - shared.log.error(f"FLUX: Failed to cast transformer to {devices.dtype}, set dtype to {transformer.dtype}") - raise + shared.log.error(f"Loading 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}") + if debug: + from modules import errors + errors.display(e, 'FLUX Quanto:') - if text_encoder_2 is None: + try: quantization_map = os.path.join(repo_path, "text_encoder_2", "quantization_map.json") if not os.path.exists(quantization_map): repo_id = checkpoint_info.name.replace('Diffusers/', '') @@ -81,11 +86,18 @@ def load_flux_quanto(checkpoint_info, diffusers_load_config, transformer, text_e try: text_encoder_2 = text_encoder_2.to(dtype=devices.dtype) except Exception: - shared.log.error(f"FLUX: Failed to cast text encoder to {devices.dtype}, set dtype to {text_encoder_2.dtype}") - raise + shared.log.error(f"Loading 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}") + if debug: + from modules import errors + errors.display(e, 'FLUX Quanto:') + + return transformer, text_encoder_2 -def load_flux_bnb(checkpoint_info, diffusers_load_config, transformer, text_encoder_2): # pylint: disable=unused-argument +def load_flux_bnb(checkpoint_info, diffusers_load_config, ): # pylint: disable=unused-argument + transformer, text_encoder_2 = None, None if isinstance(checkpoint_info, str): repo_path = checkpoint_info else: @@ -105,9 +117,11 @@ def load_flux_bnb(checkpoint_info, diffusers_load_config, transformer, text_enco else: if transformer is None: transformer = FluxTransformer2DModel.from_single_file(repo_path, **diffusers_load_config) + return transformer, text_encoder_2 -def load_transformer(file_path, transformer): # triggered by opts.sd_unet change +def load_transformer(file_path): # triggered by opts.sd_unet change + transformer = None quant = get_quant(file_path) diffusers_load_config = { "low_cpu_mem_usage": True, @@ -117,11 +131,17 @@ def load_transformer(file_path, transformer): # triggered by opts.sd_unet change shared.log.info(f'Loading UNet: type=FLUX file="{file_path}" offload={shared.opts.diffusers_offload_mode} quant={quant} dtype={devices.dtype}') if 'nf4' in quant: from modules.model_flux_nf4 import load_flux_nf4 - load_flux_nf4(file_path, diffusers_load_config, transformer, text_encoder_2='skip') + _transformer, _text_encoder_2 = load_flux_nf4(file_path, diffusers_load_config) + if _transformer is not None: + transformer = _transformer elif quant == 'qint8' or quant == 'qint4': - load_flux_quanto(file_path, diffusers_load_config, transformer, text_encoder_2='skip') + _transformer, _text_encoder_2 = load_flux_quanto(file_path, diffusers_load_config) + if _transformer is not None: + transformer = _transformer elif quant == 'fp8' or quant == 'fp4': - load_flux_bnb(file_path, diffusers_load_config, transformer, text_encoder_2='skip') + _transformer, _text_encoder_2 = load_flux_bnb(file_path, diffusers_load_config) + if _transformer is not None: + transformer = _transformer else: from diffusers import FluxTransformer2DModel transformer = FluxTransformer2DModel.from_single_file(file_path, **diffusers_load_config) @@ -141,28 +161,70 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch # load overrides if any if shared.opts.sd_unet != 'None': - debug(f'Loading FLUX: unet="{shared.opts.sd_unet}"') - from modules import sd_unet - load_transformer(sd_unet.unet_dict[shared.opts.sd_unet], transformer) + try: + debug(f'Loading FLUX: unet="{shared.opts.sd_unet}"') + from modules import sd_unet + _transformer = load_transformer(sd_unet.unet_dict[shared.opts.sd_unet]) + if _transformer is not None: + transformer = _transformer + except Exception as e: + shared.log.error(f"Loading FLUX: Failed to load UNet: {e}") + if debug: + from modules import errors + errors.display(e, 'FLUX UNet:') if shared.opts.sd_text_encoder != 'None': - debug(f'Loading FLUX: t5="{shared.opts.sd_text_encoder}"') - from modules.model_t5 import load_t5 - text_encoder_2 = load_t5(t5=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir) - if shared.opts.sd_vae != 'None': - debug(f'Loading 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] - if os.path.exists(vae_file): - vae_config = os.path.join('configs', 'flux', 'vae', 'config.json') - vae = diffusers.AutoencoderKL.from_single_file(vae_file, config=vae_config, **diffusers_load_config) + try: + debug(f'Loading FLUX: t5="{shared.opts.sd_text_encoder}"') + from modules.model_t5 import load_t5 + _text_encoder_2 = load_t5(t5=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir) + if _text_encoder_2 is not None: + text_encoder_2 = _text_encoder_2 + except Exception as e: + shared.log.error(f"Loading FLUX: Failed to load T5: {e}") + 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}"') + 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] + if os.path.exists(vae_file): + 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}") + if debug: + from modules import errors + errors.display(e, 'FLUX VAE:') # load quantized components if any if quant == 'nf4': - from modules.model_flux_nf4 import load_flux_nf4 - load_flux_nf4(checkpoint_info, diffusers_load_config, transformer, text_encoder_2) + try: + from modules.model_flux_nf4 import load_flux_nf4 + _transformer, _text_encoder = load_flux_nf4(checkpoint_info, diffusers_load_config) + if _transformer is not None: + transformer = _transformer + 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}") + if debug: + from modules import errors + errors.display(e, 'FLUX NF4:') if quant == 'qint8' or quant == 'qint4': - load_flux_quanto(checkpoint_info, diffusers_load_config, transformer, text_encoder_2) + try: + _transformer, _text_encoder = load_flux_quanto(checkpoint_info, diffusers_load_config) + if _transformer is not None: + transformer = _transformer + 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}") + if debug: + from modules import errors + errors.display(e, 'FLUX Quanto:') # initialize pipeline with pre-loaded components components = {} @@ -172,6 +234,6 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch components['text_encoder_2'] = text_encoder_2 if vae is not None: components['vae'] = vae - debug(f'Loading FLUX: preloaded={list(components)}') + shared.log.debug(f'Loading FLUX: preloaded={list(components)}') pipe = diffusers.FluxPipeline.from_pretrained(base_repo, cache_dir=shared.opts.diffusers_dir, **components, **diffusers_load_config) return pipe diff --git a/modules/model_flux_nf4.py b/modules/model_flux_nf4.py index bfe70f72a..b91e8fa28 100644 --- a/modules/model_flux_nf4.py +++ b/modules/model_flux_nf4.py @@ -26,7 +26,7 @@ def load_bnb(): global bnb # pylint: disable=global-statement bnb = bitsandbytes except Exception as e: - shared.log.error(f"FLUX: Failed to import bitsandbytes: {e}") + shared.log.error(f"Loading FLUX: Failed to import bitsandbytes: {e}") raise @@ -162,8 +162,10 @@ def create_quantized_param( module._parameters[tensor_name] = new_value # pylint: disable=protected-access -def load_flux_nf4(checkpoint_info, diffusers_load_config, transformer, text_encoder_2): +def load_flux_nf4(checkpoint_info, diffusers_load_config): load_bnb() + transformer = None + text_encoder_2 = None if isinstance(checkpoint_info, str): repo_path = checkpoint_info else: @@ -182,7 +184,7 @@ def load_flux_nf4(checkpoint_info, diffusers_load_config, transformer, text_enco try: converted_state_dict = convert_flux_transformer_checkpoint_to_diffusers(original_state_dict) except Exception as e: - shared.log.error(f"FLUX: Failed to convert UNET: {e}") + shared.log.error(f"Loading FLUX: Failed to convert UNET: {e}") if debug: from modules import errors errors.display(e, 'FLUX convert:') @@ -209,3 +211,4 @@ def load_flux_nf4(checkpoint_info, diffusers_load_config, transformer, text_enco del original_state_dict devices.torch_gc(force=True) + return transformer, text_encoder_2 diff --git a/modules/sd_unet.py b/modules/sd_unet.py index 7271ce50a..16d942a22 100644 --- a/modules/sd_unet.py +++ b/modules/sd_unet.py @@ -30,8 +30,7 @@ def load_unet(model): model.prior_pipe.text_encoder = prior_text_encoder.to(devices.device, dtype=devices.dtype) if "Flux" in model.__class__.__name__: from modules.model_flux import load_transformer - transformer = None - load_transformer(unet_dict[shared.opts.sd_unet], 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':