mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
flux fix quant preloads
This commit is contained in:
+93
-31
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
+1
-2
@@ -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':
|
||||
|
||||
Reference in New Issue
Block a user