refactor xyzgrid

This commit is contained in:
Vladimir Mandic
2024-09-13 19:22:02 -04:00
parent bf870c7644
commit fe93ad6929
14 changed files with 626 additions and 984 deletions
+1 -1
View File
@@ -67,7 +67,7 @@ def get_gpu_info():
}
elif torch.version.cuda:
return {
'device': f'{torch.cuda.get_device_name(torch.cuda.current_device())} n={torch.cuda.device_count()} arch={torch.cuda.get_arch_list()[-1]} cap={torch.cuda.get_device_capability(device)}',
'device': f'{torch.cuda.get_device_name(torch.cuda.current_device())} n={torch.cuda.device_count()} arch={torch.cuda.get_arch_list()[-1]} capability={torch.cuda.get_device_capability(device)}',
'cuda': torch.version.cuda,
'cudnn': torch.backends.cudnn.version(),
'driver': get_driver(),
+5
View File
@@ -39,6 +39,9 @@ timer.startup.record("torch")
import transformers # pylint: disable=W0611,C0411
timer.startup.record("transformers")
import accelerate # pylint: disable=W0611,C0411
timer.startup.record("accelerate")
import onnxruntime # pylint: disable=W0611,C0411
onnxruntime.set_default_logger_severity(3)
timer.startup.record("onnx")
@@ -86,6 +89,8 @@ def get_packages():
"torch": getattr(torch, "__long_version__", torch.__version__),
"diffusers": diffusers.__version__,
"gradio": gradio.__version__,
"transformers": transformers.__version__,
"accelerate": accelerate.__version__,
}
errors.log.info(f'Load packages: {get_packages()}')
+28 -15
View File
@@ -107,17 +107,27 @@ def load_flux_bnb(checkpoint_info, diffusers_load_config): # pylint: disable=unu
install('bitsandbytes', quiet=True)
from diffusers import FluxTransformer2DModel
quant = get_quant(repo_path)
if quant == 'fp8':
quantization_config = transformers.BitsAndBytesConfig(load_in_8bit=True)
if transformer is None:
try:
if quant == 'fp8':
quantization_config = transformers.BitsAndBytesConfig(load_in_8bit=True, bnb_4bit_compute_dtype=devices.dtype)
debug(f'Quantization: {quantization_config}')
transformer = FluxTransformer2DModel.from_single_file(repo_path, **diffusers_load_config, quantization_config=quantization_config)
elif quant == 'fp4':
quantization_config = transformers.BitsAndBytesConfig(load_in_4bit=True)
if transformer is None:
elif quant == 'fp4':
quantization_config = transformers.BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=devices.dtype, bnb_4bit_quant_type= 'fp4')
debug(f'Quantization: {quantization_config}')
transformer = FluxTransformer2DModel.from_single_file(repo_path, **diffusers_load_config, quantization_config=quantization_config)
else:
if transformer is None:
elif quant == 'nf4':
quantization_config = transformers.BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=devices.dtype, bnb_4bit_quant_type= 'nf4')
debug(f'Quantization: {quantization_config}')
transformer = FluxTransformer2DModel.from_single_file(repo_path, **diffusers_load_config, quantization_config=quantization_config)
else:
transformer = 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}")
transformer, text_encoder_2 = None, None
if debug:
from modules import errors
errors.display(e, 'FLUX:')
return transformer, text_encoder_2
@@ -130,19 +140,19 @@ def load_transformer(file_path): # triggered by opts.sd_unet change
"cache_dir": shared.opts.hfcache_dir,
}
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
_transformer, _text_encoder_2 = load_flux_nf4(file_path)
if _transformer is not None:
transformer = _transformer
elif quant == 'qint8' or quant == 'qint4':
if quant == 'qint8' or quant == 'qint4':
_transformer, _text_encoder_2 = load_flux_quanto(file_path)
if _transformer is not None:
transformer = _transformer
elif quant == 'fp8' or quant == 'fp4':
elif quant == 'fp8' or quant == 'fp4' or quant == 'nf4':
_transformer, _text_encoder_2 = load_flux_bnb(file_path, diffusers_load_config)
if _transformer is not None:
transformer = _transformer
elif 'nf4' in quant: # TODO right now this is not working for civitai published nf4 models
from modules.model_flux_nf4 import load_flux_nf4
_transformer, _text_encoder_2 = load_flux_nf4(file_path)
if _transformer is not None:
transformer = _transformer
else:
from diffusers import FluxTransformer2DModel
transformer = FluxTransformer2DModel.from_single_file(file_path, **diffusers_load_config)
@@ -169,7 +179,10 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
from modules import sd_unet
_transformer = load_transformer(sd_unet.unet_dict[shared.opts.sd_unet])
if _transformer is not None:
sd_unet.loaded_unet = shared.opts.sd_unet
transformer = _transformer
else:
sd_unet.failed_unet.append(shared.opts.sd_unet)
except Exception as e:
shared.log.error(f"Loading FLUX: Failed to load UNet: {e}")
if debug:
+17 -10
View File
@@ -198,16 +198,23 @@ def load_flux_nf4(checkpoint_info):
_replace_with_bnb_linear(transformer, "nf4")
for param_name, param in converted_state_dict.items():
if param_name not in expected_state_dict_keys:
continue
is_param_float8_e4m3fn = hasattr(torch, "float8_e4m3fn") and param.dtype == torch.float8_e4m3fn
if torch.is_floating_point(param) and not is_param_float8_e4m3fn:
param = param.to(devices.dtype)
if not check_quantized_param(transformer, param_name):
set_module_tensor_to_device(transformer, param_name, device=0, value=param)
else:
create_quantized_param(transformer, param, param_name, target_device=0, state_dict=original_state_dict, pre_quantized=True)
try:
for param_name, param in converted_state_dict.items():
if param_name not in expected_state_dict_keys:
continue
is_param_float8_e4m3fn = hasattr(torch, "float8_e4m3fn") and param.dtype == torch.float8_e4m3fn
if torch.is_floating_point(param) and not is_param_float8_e4m3fn:
param = param.to(devices.dtype)
if not check_quantized_param(transformer, param_name):
set_module_tensor_to_device(transformer, param_name, device=0, value=param)
else:
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}")
if debug:
from modules import errors
errors.display(e, 'FLUX:')
del original_state_dict
devices.torch_gc(force=True)
+12 -2
View File
@@ -3,10 +3,13 @@ from modules import shared, devices, files_cache, sd_models
unet_dict = {}
loaded_unet = None
failed_unet = []
debug = os.environ.get('SD_LOAD_DEBUG', None) is not None
def load_unet(model):
global loaded_unet # pylint: disable=global-statement
if shared.opts.sd_unet == 'None':
return
if shared.opts.sd_unet not in list(unet_dict):
@@ -19,16 +22,21 @@ def load_unet(model):
config = None
config_file = 'default'
try:
if "StableCascade" in model.__class__.__name__:
if shared.opts.sd_unet == loaded_unet or shared.opts.sd_unet in failed_unet:
pass
elif "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)
loaded_unet = shared.opts.sd_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__:
elif "Flux" in model.__class__.__name__:
sd_models.load_diffuser() # TODO forcing reloading entire flux as loading transformers only leads to massive memory usage
"""
from modules.model_flux import load_transformer
transformer = load_transformer(unet_dict[shared.opts.sd_unet])
if transformer is not None:
@@ -36,8 +44,10 @@ def load_unet(model):
if shared.opts.diffusers_offload_mode == 'none':
sd_models.move_model(transformer, devices.device)
model.transformer = transformer
loaded_unet = shared.opts.sd_unet
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')