mirror of
https://github.com/vladmandic/automatic
synced 2026-09-08 22:08:42 +02:00
@@ -41,7 +41,7 @@ def load_flux_quanto(checkpoint_info):
|
||||
except Exception:
|
||||
shared.log.error(f"Load model: type=FLUX Failed to cast transformer to {devices.dtype}, set dtype to {transformer.dtype}")
|
||||
except Exception as e:
|
||||
shared.log.error(f"Load model: type=FLUX Failed to load Quanto transformer: {e}")
|
||||
shared.log.error(f"Load model: type=FLUX failed to load Quanto transformer: {e}")
|
||||
if debug:
|
||||
from modules import errors
|
||||
errors.display(e, 'FLUX Quanto:')
|
||||
@@ -68,7 +68,7 @@ def load_flux_quanto(checkpoint_info):
|
||||
except Exception:
|
||||
shared.log.error(f"Load model: type=FLUX Failed to cast text encoder to {devices.dtype}, set dtype to {text_encoder_2.dtype}")
|
||||
except Exception as e:
|
||||
shared.log.error(f"Load model: type=FLUX Failed to load Quanto text encoder: {e}")
|
||||
shared.log.error(f"Load model: type=FLUX failed to load Quanto text encoder: {e}")
|
||||
if debug:
|
||||
from modules import errors
|
||||
errors.display(e, 'FLUX Quanto:')
|
||||
@@ -100,7 +100,7 @@ def load_flux_bnb(checkpoint_info, diffusers_load_config): # pylint: disable=unu
|
||||
else:
|
||||
transformer = diffusers.FluxTransformer2DModel.from_single_file(repo_path, **diffusers_load_config)
|
||||
except Exception as e:
|
||||
shared.log.error(f"Load model: type=FLUX Failed to load BnB transformer: {e}")
|
||||
shared.log.error(f"Load model: type=FLUX failed to load BnB transformer: {e}")
|
||||
transformer, text_encoder_2 = None, None
|
||||
if debug:
|
||||
from modules import errors
|
||||
@@ -222,7 +222,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
|
||||
shared.opts.sd_unet = 'None'
|
||||
sd_unet.failed_unet.append(shared.opts.sd_unet)
|
||||
except Exception as e:
|
||||
shared.log.error(f"Load model: type=FLUX Failed to load UNet: {e}")
|
||||
shared.log.error(f"Load model: type=FLUX failed to load UNet: {e}")
|
||||
shared.opts.sd_unet = 'None'
|
||||
if debug:
|
||||
from modules import errors
|
||||
@@ -236,7 +236,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
|
||||
else:
|
||||
text_encoder_2 = load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir)
|
||||
except Exception as e:
|
||||
shared.log.error(f"Load model: type=FLUX Failed to load T5: {e}")
|
||||
shared.log.error(f"Load model: type=FLUX failed to load T5: {e}")
|
||||
shared.opts.sd_text_encoder = 'None'
|
||||
if debug:
|
||||
from modules import errors
|
||||
@@ -251,7 +251,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
|
||||
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"Load model: type=FLUX Failed to load VAE: {e}")
|
||||
shared.log.error(f"Load model: type=FLUX failed to load VAE: {e}")
|
||||
shared.opts.sd_vae = 'None'
|
||||
if debug:
|
||||
from modules import errors
|
||||
@@ -267,7 +267,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
|
||||
if _text_encoder is not None:
|
||||
text_encoder_2 = _text_encoder
|
||||
except Exception as e:
|
||||
shared.log.error(f"Load model: type=FLUX Failed to load NF4 components: {e}")
|
||||
shared.log.error(f"Load model: type=FLUX failed to load NF4 components: {e}")
|
||||
if debug:
|
||||
from modules import errors
|
||||
errors.display(e, 'FLUX NF4:')
|
||||
@@ -279,7 +279,7 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
|
||||
if _text_encoder is not None:
|
||||
text_encoder_2 = _text_encoder
|
||||
except Exception as e:
|
||||
shared.log.error(f"Load model: type=FLUX Failed to load Quanto components: {e}")
|
||||
shared.log.error(f"Load model: type=FLUX failed to load Quanto components: {e}")
|
||||
if debug:
|
||||
from modules import errors
|
||||
errors.display(e, 'FLUX Quanto:')
|
||||
|
||||
@@ -200,7 +200,7 @@ def load_flux_nf4(checkpoint_info):
|
||||
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"Load model: type=FLUX Failed to load UNET: {e}")
|
||||
shared.log.error(f"Load model: type=FLUX failed to load UNET: {e}")
|
||||
if debug:
|
||||
from modules import errors
|
||||
errors.display(e, 'FLUX:')
|
||||
|
||||
+85
-53
@@ -1,56 +1,49 @@
|
||||
import os
|
||||
import diffusers
|
||||
import transformers
|
||||
from modules import shared, devices, sd_models, sd_unet
|
||||
|
||||
|
||||
default_repo_id = 'stabilityai/stable-diffusion-3-medium'
|
||||
def load_overrides(kwargs, cache_dir):
|
||||
if shared.opts.sd_unet != 'None':
|
||||
try:
|
||||
fn = sd_unet.unet_dict[shared.opts.sd_unet]
|
||||
kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_single_file(fn, cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Load model: type=SD3 unet="{shared.opts.sd_unet}"')
|
||||
except Exception as e:
|
||||
shared.log.error(f"Load model: type=SD3 failed to load UNet: {e}")
|
||||
shared.opts.sd_unet = 'None'
|
||||
sd_unet.failed_unet.append(shared.opts.sd_unet)
|
||||
if shared.opts.sd_text_encoder != 'None':
|
||||
try:
|
||||
from modules.model_te import load_t5, load_vit_l, load_vit_g
|
||||
if 'vit-l' in shared.opts.sd_text_encoder.lower():
|
||||
kwargs['text_encoder'] = load_vit_l()
|
||||
shared.log.debug(f'Load model: type=SD3 variant="vit-l" te="{shared.opts.sd_text_encoder}"')
|
||||
elif 'vit-g' in shared.opts.sd_text_encoder.lower():
|
||||
kwargs['text_encoder_2'] = load_vit_g()
|
||||
shared.log.debug(f'Load model: type=SD3 variant="vit-g" te="{shared.opts.sd_text_encoder}"')
|
||||
else:
|
||||
kwargs['text_encoder_3'] = load_t5(name=shared.opts.sd_text_encoder, cache_dir=shared.opts.diffusers_dir)
|
||||
shared.log.debug(f'Load model: type=SD3 variant="t5" te="{shared.opts.sd_text_encoder}"')
|
||||
except Exception as e:
|
||||
shared.log.error(f"Load model: type=SD3 failed to load T5: {e}")
|
||||
shared.opts.sd_text_encoder = 'None'
|
||||
if shared.opts.sd_vae != 'None' and shared.opts.sd_vae != 'Automatic':
|
||||
try:
|
||||
from modules import sd_vae
|
||||
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')
|
||||
kwargs['vae'] = diffusers.AutoencoderKL.from_single_file(vae_file, config=vae_config, cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Load model: type=SD3 vae="{shared.opts.sd_vae}"')
|
||||
except Exception as e:
|
||||
shared.log.error(f"Load model: type=FLUX failed to load VAE: {e}")
|
||||
shared.opts.sd_vae = 'None'
|
||||
return kwargs
|
||||
|
||||
|
||||
def load_sd3(checkpoint_info, cache_dir=None, config=None):
|
||||
from modules import shared, devices, modelloader, sd_models
|
||||
repo_id = sd_models.path_to_repo(checkpoint_info.name)
|
||||
dtype = devices.dtype
|
||||
kwargs = {}
|
||||
if checkpoint_info.path is not None and checkpoint_info.path.endswith('.safetensors') and os.path.exists(checkpoint_info.path):
|
||||
loader = diffusers.StableDiffusion3Pipeline.from_single_file
|
||||
fn_size = os.path.getsize(checkpoint_info.path)
|
||||
if fn_size < 5e9:
|
||||
kwargs = {
|
||||
'text_encoder': transformers.CLIPTextModelWithProjection.from_pretrained(
|
||||
default_repo_id,
|
||||
subfolder='text_encoder',
|
||||
cache_dir=cache_dir,
|
||||
torch_dtype=dtype,
|
||||
),
|
||||
'text_encoder_2': transformers.CLIPTextModelWithProjection.from_pretrained(
|
||||
default_repo_id,
|
||||
subfolder='text_encoder_2',
|
||||
cache_dir=cache_dir,
|
||||
torch_dtype=dtype,
|
||||
),
|
||||
'tokenizer': transformers.CLIPTokenizer.from_pretrained(
|
||||
default_repo_id,
|
||||
subfolder='tokenizer',
|
||||
cache_dir=cache_dir,
|
||||
),
|
||||
'tokenizer_2': transformers.CLIPTokenizer.from_pretrained(
|
||||
default_repo_id,
|
||||
subfolder='tokenizer_2',
|
||||
cache_dir=cache_dir,
|
||||
),
|
||||
'text_encoder_3': None,
|
||||
}
|
||||
elif fn_size < 1e10: # if model is below 10gb it does not have te3
|
||||
kwargs = {
|
||||
'text_encoder_3': None,
|
||||
}
|
||||
else:
|
||||
kwargs = {}
|
||||
else:
|
||||
modelloader.hf_login()
|
||||
loader = diffusers.StableDiffusion3Pipeline.from_pretrained
|
||||
kwargs['variant'] = 'fp16'
|
||||
|
||||
def load_quants(kwargs, repo_id, cache_dir):
|
||||
if len(shared.opts.bnb_quantization) > 0:
|
||||
from modules.model_quant import load_bnb
|
||||
load_bnb('Load model: type=SD3')
|
||||
@@ -61,18 +54,57 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None):
|
||||
bnb_4bit_quant_type=shared.opts.bnb_quantization_type,
|
||||
bnb_4bit_compute_dtype=devices.dtype
|
||||
)
|
||||
if 'Model' in shared.opts.bnb_quantization:
|
||||
transformer = diffusers.SD3Transformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype)
|
||||
if 'Model' in shared.opts.bnb_quantization and 'transformer' not in kwargs:
|
||||
kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}')
|
||||
kwargs['transformer'] = transformer
|
||||
if 'Text Encoder' in shared.opts.bnb_quantization:
|
||||
te3 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype)
|
||||
if 'Text Encoder' in shared.opts.bnb_quantization and 'text_encoder_3' not in kwargs:
|
||||
kwargs['text_encoder_3'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}')
|
||||
kwargs['text_encoder_3'] = te3
|
||||
return kwargs
|
||||
|
||||
|
||||
def load_missing(kwargs, fn, cache_dir):
|
||||
keys = sd_models.get_safetensor_keys(fn)
|
||||
size = os.stat(fn).st_size // 1024 // 1024
|
||||
if size > 15000:
|
||||
repo_id = 'stabilityai/stable-diffusion-3.5-large'
|
||||
else:
|
||||
repo_id = 'stabilityai/stable-diffusion-3-medium'
|
||||
if 'text_encoder' not in kwargs and 'text_encoder' not in keys:
|
||||
kwargs['text_encoder'] = transformers.CLIPTextModelWithProjection.from_pretrained(repo_id, subfolder='text_encoder', cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Load model: type=SD3 missing=te1 repo="{repo_id}"')
|
||||
if 'text_encoder_2' not in kwargs and 'text_encoder_2' not in keys:
|
||||
kwargs['text_encoder_2'] = transformers.CLIPTextModelWithProjection.from_pretrained(repo_id, subfolder='text_encoder_2', cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Load model: type=SD3 missing=te2 repo="{repo_id}"')
|
||||
if 'text_encoder_3' not in kwargs and 'text_encoder_3' not in keys:
|
||||
kwargs['text_encoder_3'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Load model: type=SD3 missing=te3 repo="{repo_id}"')
|
||||
# if 'transformer' not in kwargs and 'transformer' not in keys:
|
||||
# kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_pretrained(default_repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
return kwargs
|
||||
|
||||
|
||||
def load_sd3(checkpoint_info, cache_dir=None, config=None):
|
||||
repo_id = sd_models.path_to_repo(checkpoint_info.name)
|
||||
fn = checkpoint_info.path
|
||||
|
||||
kwargs = {}
|
||||
kwargs = load_overrides(kwargs, cache_dir)
|
||||
kwargs = load_quants(kwargs, repo_id, cache_dir)
|
||||
|
||||
if fn is not None and fn.endswith('.safetensors') and os.path.exists(fn):
|
||||
kwargs = load_missing(kwargs, fn, cache_dir)
|
||||
loader = diffusers.StableDiffusion3Pipeline.from_single_file
|
||||
repo_id = fn
|
||||
else:
|
||||
loader = diffusers.StableDiffusion3Pipeline.from_pretrained
|
||||
kwargs['variant'] = 'fp16'
|
||||
|
||||
shared.log.debug(f'Load model: type=FLUX preloaded={list(kwargs)}')
|
||||
|
||||
pipe = loader(
|
||||
repo_id,
|
||||
torch_dtype=dtype,
|
||||
torch_dtype=devices.dtype,
|
||||
cache_dir=cache_dir,
|
||||
config=config,
|
||||
**kwargs,
|
||||
|
||||
@@ -56,7 +56,7 @@ class YoloRestorer(Detailer):
|
||||
name = os.path.splitext(os.path.basename(f))[0]
|
||||
if name not in files:
|
||||
self.list[name] = os.path.join(shared.opts.yolo_dir, f)
|
||||
shared.log.info(f'Available Yolo: path="{shared.opts.yolo_dir} items={len(list(self.list))} downloaded={downloaded}')
|
||||
shared.log.info(f'Available Yolo: path="{shared.opts.yolo_dir}" items={len(list(self.list))} downloaded={downloaded}')
|
||||
return self.list
|
||||
|
||||
def dependencies(self):
|
||||
|
||||
+11
-1
@@ -417,6 +417,16 @@ def read_state_dict(checkpoint_file, map_location=None, what:str='model'): # pyl
|
||||
return sd
|
||||
|
||||
|
||||
def get_safetensor_keys(filename):
|
||||
keys = []
|
||||
try:
|
||||
with safetensors.torch.safe_open(filename, framework="pt", device="cpu") as f:
|
||||
keys = f.keys()
|
||||
except Exception as e:
|
||||
shared.log.error(f'Load dict: path="{filename}" {e}')
|
||||
return keys
|
||||
|
||||
|
||||
def get_checkpoint_state_dict(checkpoint_info: CheckpointInfo, timer):
|
||||
if not os.path.isfile(checkpoint_info.filename):
|
||||
return None
|
||||
@@ -1088,7 +1098,7 @@ def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op='
|
||||
sd_model = load_flux(checkpoint_info, diffusers_load_config)
|
||||
elif model_type in ['Stable Diffusion 3']:
|
||||
from modules.model_sd3 import load_sd3
|
||||
shared.log.debug(f'Load {op}: model="Stable Diffusion 3" variant=medium')
|
||||
shared.log.debug(f'Load {op}: model="Stable Diffusion 3"')
|
||||
shared.opts.scheduler = 'Default'
|
||||
sd_model = load_sd3(checkpoint_info, cache_dir=shared.opts.diffusers_dir, config=diffusers_load_config.get('config', None))
|
||||
elif model_type in ['Meissonic']: # forced pipeline
|
||||
|
||||
Reference in New Issue
Block a user