flux fix quant preloads

This commit is contained in:
Vladimir Mandic
2024-09-02 16:53:26 -04:00
parent 3bebc69ad7
commit fa136103b2
3 changed files with 100 additions and 36 deletions
+93 -31
View File
@@ -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
+6 -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"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
View File
@@ -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':