mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
refactor xyzgrid
This commit is contained in:
+1
-1
@@ -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(),
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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')
|
||||
|
||||
Reference in New Issue
Block a user