mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
post release dev merge, see changelog for details
This commit is contained in:
@@ -1,5 +1,24 @@
|
||||
# Change Log for SD.Next
|
||||
|
||||
## Update for 2024-09-01
|
||||
|
||||
- flux improve logging, warn when attempting to load unet as base model
|
||||
- flux unet support fp8/fp4 quantization
|
||||
- flux vae support fp16
|
||||
- flux lora support additional training tools (*1)
|
||||
- flux model support loading all-in-one safetensors (*1)
|
||||
not recommended due to massive duplication of components, but added due to popular demand
|
||||
- taesd configurable number of layers
|
||||
can be used to speed-up taesd decoding by reducing number of ops
|
||||
e.g. if generating 1024px image, reducing layers by 1 will result in preview being 512px
|
||||
set via *settings -> live preview -> taesd decode layers*
|
||||
- xhinker prompt parser handle offloaded models
|
||||
- t5 enum manually downloaded models (*2)
|
||||
|
||||
*notes*:
|
||||
- (*1) requires `diffusers==0.31.0.dev0`
|
||||
- (*2) work-in-progress
|
||||
|
||||
## Update for 2024-08-31
|
||||
|
||||
### Highlights for 2024-08-31
|
||||
|
||||
+93
-58
@@ -5,16 +5,46 @@ import diffusers
|
||||
import transformers
|
||||
from safetensors.torch import load_file
|
||||
from huggingface_hub import hf_hub_download
|
||||
from accelerate.utils import compute_module_sizes
|
||||
from modules import shared, devices
|
||||
|
||||
|
||||
def load_quanto_transformer(checkpoint_info):
|
||||
from optimum.quanto import requantize # pylint: disable=no-name-in-module
|
||||
repo_path = checkpoint_info.path
|
||||
debug = os.environ.get('SD_LOAD_DEBUG', None) is not None
|
||||
|
||||
|
||||
def get_quant(file_path):
|
||||
if "qint8" in file_path.lower():
|
||||
return 'qint8'
|
||||
if "qint4" in file_path.lower():
|
||||
return 'qint4'
|
||||
if "fp8" in file_path.lower():
|
||||
return 'fp8'
|
||||
if "fp4" in file_path.lower():
|
||||
return 'fp4'
|
||||
if "nf4" in file_path.lower():
|
||||
return 'nf4'
|
||||
return 'none'
|
||||
|
||||
|
||||
def load_flux_quanto(checkpoint_info, diffusers_load_config, transformer_only=False):
|
||||
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}")
|
||||
raise
|
||||
quanto.tensor.qbits.QBitsTensor.create = lambda *args, **kwargs: quanto.tensor.qbits.QBitsTensor(*args, **kwargs)
|
||||
|
||||
if isinstance(checkpoint_info, str):
|
||||
repo_path = checkpoint_info
|
||||
else:
|
||||
repo_path = checkpoint_info.path
|
||||
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/', '')
|
||||
quantization_map = hf_hub_download(repo_id, subfolder='transformer', filename='quantization_map.json', cache_dir=shared.opts.diffusers_dir)
|
||||
quantization_map = hf_hub_download(repo_id, subfolder='transformer', filename='quantization_map.json', **diffusers_load_config)
|
||||
with open(quantization_map, "r", encoding='utf8') as f:
|
||||
quantization_map = json.load(f)
|
||||
state_dict = load_file(os.path.join(repo_path, "transformer", "diffusion_pytorch_model.safetensors"))
|
||||
@@ -23,16 +53,19 @@ def load_quanto_transformer(checkpoint_info):
|
||||
transformer = diffusers.FluxTransformer2DModel.from_config(os.path.join(repo_path, "transformer", "config.json")).to(dtype=dtype)
|
||||
requantize(transformer, state_dict, quantization_map, device=torch.device("cpu"))
|
||||
transformer.eval()
|
||||
return transformer
|
||||
if transformer.dtype != devices.dtype:
|
||||
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
|
||||
if transformer_only:
|
||||
return transformer, None
|
||||
|
||||
|
||||
def load_quanto_text_encoder_2(checkpoint_info):
|
||||
from optimum.quanto import requantize # pylint: disable=no-name-in-module
|
||||
repo_path = checkpoint_info.path
|
||||
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/', '')
|
||||
quantization_map = hf_hub_download(repo_id, subfolder='text_encoder_2', filename='quantization_map.json', cache_dir=shared.opts.diffusers_dir)
|
||||
quantization_map = hf_hub_download(repo_id, subfolder='text_encoder_2', filename='quantization_map.json', **diffusers_load_config)
|
||||
with open(quantization_map, "r", encoding='utf8') as f:
|
||||
quantization_map = json.load(f)
|
||||
with open(os.path.join(repo_path, "text_encoder_2", "config.json"), encoding='utf8') as f:
|
||||
@@ -43,71 +76,73 @@ def load_quanto_text_encoder_2(checkpoint_info):
|
||||
text_encoder_2 = transformers.T5EncoderModel(t5_config).to(dtype=dtype)
|
||||
requantize(text_encoder_2, state_dict, quantization_map, device=torch.device("cpu"))
|
||||
text_encoder_2.eval()
|
||||
return text_encoder_2
|
||||
if text_encoder_2.dtype != devices.dtype:
|
||||
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
|
||||
return transformer, text_encoder_2
|
||||
|
||||
|
||||
def load_transformer(file_path):
|
||||
def load_flux_bnb(checkpoint_info, diffusers_load_config, transformer_only=False):
|
||||
if isinstance(checkpoint_info, str):
|
||||
repo_path = checkpoint_info
|
||||
else:
|
||||
repo_path = checkpoint_info.path
|
||||
from installer import install
|
||||
install('bitsandbytes', quiet=True)
|
||||
from diffusers import FluxTransformer2DModel
|
||||
quant = get_quant(repo_path)
|
||||
if quant == 'fp8':
|
||||
quantization_config = transformers.BitsAndBytesConfig(load_in_8bit=True)
|
||||
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)
|
||||
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)
|
||||
if transformer_only:
|
||||
return transformer, None
|
||||
# TODO load text_encoder_2
|
||||
|
||||
|
||||
def load_transformer(file_path): # triggered by opts.sd_unet change
|
||||
quant = get_quant(file_path)
|
||||
diffusers_load_config = {
|
||||
"low_cpu_mem_usage": True,
|
||||
"torch_dtype": devices.dtype,
|
||||
"cache_dir": shared.opts.hfcache_dir,
|
||||
}
|
||||
from diffusers import FluxTransformer2DModel
|
||||
transformer = FluxTransformer2DModel.from_single_file(file_path, **diffusers_load_config)
|
||||
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 = load_flux_nf4(file_path, diffusers_load_config, transformer_only=True)
|
||||
elif quant == 'qint8' or quant == 'qint4':
|
||||
transformer, _ = load_flux_quanto(file_path, diffusers_load_config, transformer_only=True)
|
||||
elif quant == 'fp8' or quant == 'fp4':
|
||||
transformer, _ = load_flux_bnb(file_path, diffusers_load_config, transformer_only=True)
|
||||
else:
|
||||
from diffusers import FluxTransformer2DModel
|
||||
transformer = FluxTransformer2DModel.from_single_file(file_path, **diffusers_load_config)
|
||||
if transformer is None:
|
||||
shared.log.error('Failed to load UNet model')
|
||||
if debug:
|
||||
shared.log.debug(f'FLUX transformer: size={round(compute_module_sizes(transformer)[""] / 1024 / 1204)}')
|
||||
return transformer
|
||||
|
||||
|
||||
def load_flux(checkpoint_info, diffusers_load_config):
|
||||
if "qint8" in checkpoint_info.path.lower():
|
||||
quant = 'qint8'
|
||||
elif "qint4" in checkpoint_info.path.lower():
|
||||
quant = 'qint4'
|
||||
elif "nf4" in checkpoint_info.path.lower():
|
||||
quant = 'nf4'
|
||||
else:
|
||||
quant = None
|
||||
shared.log.debug(f'Loading FLUX: model="{checkpoint_info.name}" quant={quant}')
|
||||
def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_checkpoint change
|
||||
quant = get_quant(checkpoint_info.path)
|
||||
shared.log.debug(f'Loading FLUX: model="{checkpoint_info.name}" quant={quant} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}')
|
||||
if quant == 'nf4':
|
||||
from installer import install
|
||||
install('bitsandbytes', quiet=True)
|
||||
try:
|
||||
import bitsandbytes # pylint: disable=unused-import
|
||||
except Exception as e:
|
||||
shared.log.error(f"FLUX: Failed to import bitsandbytes: {e}")
|
||||
raise
|
||||
from modules.model_flux_nf4 import load_flux_nf4
|
||||
pipe = load_flux_nf4(checkpoint_info, diffusers_load_config)
|
||||
elif quant == 'qint8' or quant == 'qint4':
|
||||
from installer import install
|
||||
install('optimum-quanto', quiet=True)
|
||||
try:
|
||||
from optimum import quanto # pylint: disable=no-name-in-module
|
||||
except Exception as e:
|
||||
shared.log.error(f"FLUX: Failed to import optimum-quanto: {e}")
|
||||
raise
|
||||
quanto.tensor.qbits.QBitsTensor.create = lambda *args, **kwargs: quanto.tensor.qbits.QBitsTensor(*args, **kwargs)
|
||||
pipe = diffusers.FluxPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, transformer=None, text_encoder_2=None, **diffusers_load_config)
|
||||
pipe.transformer = load_quanto_transformer(checkpoint_info)
|
||||
pipe.text_encoder_2 = load_quanto_text_encoder_2(checkpoint_info)
|
||||
if pipe.transformer.dtype != devices.dtype:
|
||||
try:
|
||||
pipe.transformer = pipe.transformer.to(dtype=devices.dtype)
|
||||
except Exception:
|
||||
shared.log.error(f"FLUX: Failed to cast transformer to {devices.dtype}, set dtype to {pipe.transformer.dtype}")
|
||||
raise
|
||||
if pipe.text_encoder_2.dtype != devices.dtype:
|
||||
try:
|
||||
pipe.text_encoder_2 = pipe.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 {pipe.text_encoder_2.dtype}")
|
||||
raise
|
||||
pipe.transformer, pipe.text_encoder_2 = load_flux_quanto(checkpoint_info, diffusers_load_config)
|
||||
else:
|
||||
pipe = diffusers.FluxPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config)
|
||||
if devices.dtype == torch.float16 and not shared.opts.no_half_vae:
|
||||
shared.log.warning("FLUX: does not support FP16 VAE, enabling no-half-vae")
|
||||
shared.opts.no_half_vae = True
|
||||
# from accelerate.utils import compute_module_sizes
|
||||
# shared.log.debug(f'FLUX computed size: {round(compute_module_sizes(pipe.transformer)[""] / 1024 / 1204)}')
|
||||
if debug:
|
||||
shared.log.debug(f'FLUX transformer: size={round(compute_module_sizes(pipe.transformer)[""] / 1024 / 1204)}')
|
||||
return pipe
|
||||
|
||||
+34
-10
@@ -5,7 +5,6 @@ Copied from: https://github.com/huggingface/diffusers/issues/9165
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import bitsandbytes as bnb
|
||||
from transformers.quantizers.quantizers_utils import get_module_from_name
|
||||
from huggingface_hub import hf_hub_download
|
||||
from accelerate import init_empty_weights
|
||||
@@ -16,6 +15,22 @@ import safetensors.torch
|
||||
from modules import shared, devices
|
||||
|
||||
|
||||
bnb = None
|
||||
debug = os.environ.get('SD_LOAD_DEBUG', None) is not None
|
||||
|
||||
|
||||
def load_bnb():
|
||||
from installer import install
|
||||
install('bitsandbytes', quiet=True)
|
||||
try:
|
||||
import bitsandbytes
|
||||
global bnb # pylint: disable=global-statement
|
||||
bnb = bitsandbytes
|
||||
except Exception as e:
|
||||
shared.log.error(f"FLUX: Failed to import bitsandbytes: {e}")
|
||||
raise
|
||||
|
||||
|
||||
def _replace_with_bnb_linear(
|
||||
model,
|
||||
method="nf4",
|
||||
@@ -148,25 +163,31 @@ def create_quantized_param(
|
||||
module._parameters[tensor_name] = new_value # pylint: disable=protected-access
|
||||
|
||||
|
||||
def load_flux_nf4(checkpoint_info, diffusers_load_config):
|
||||
repo_path = checkpoint_info.path
|
||||
def load_flux_nf4(checkpoint_info, diffusers_load_config, transformer_only=False):
|
||||
load_bnb()
|
||||
if isinstance(checkpoint_info, str):
|
||||
repo_path = checkpoint_info
|
||||
else:
|
||||
repo_path = checkpoint_info.path
|
||||
if os.path.exists(repo_path) and os.path.isfile(repo_path):
|
||||
ckpt_path = repo_path
|
||||
if os.path.exists(repo_path) and os.path.isdir(repo_path) and os.path.exists(os.path.join(repo_path, "diffusion_pytorch_model.safetensors")):
|
||||
elif os.path.exists(repo_path) and os.path.isdir(repo_path) and os.path.exists(os.path.join(repo_path, "diffusion_pytorch_model.safetensors")):
|
||||
ckpt_path = os.path.join(repo_path, "diffusion_pytorch_model.safetensors")
|
||||
else:
|
||||
ckpt_path = hf_hub_download(repo_path, filename="diffusion_pytorch_model.safetensors", cache_dir=shared.opts.diffusers_dir)
|
||||
original_state_dict = safetensors.torch.load_file(ckpt_path)
|
||||
|
||||
if 'sayakpaul' in checkpoint_info.path:
|
||||
if 'sayakpaul' in repo_path:
|
||||
converted_state_dict = original_state_dict # already converted
|
||||
else:
|
||||
try:
|
||||
converted_state_dict = convert_flux_transformer_checkpoint_to_diffusers(original_state_dict)
|
||||
except Exception as e:
|
||||
from modules import errors
|
||||
errors.display(e, 'FLUX convert:')
|
||||
raise
|
||||
shared.log.error(f"FLUX: Failed to convert UNET: {e}")
|
||||
if debug:
|
||||
from modules import errors
|
||||
errors.display(e, 'FLUX convert:')
|
||||
converted_state_dict = original_state_dict
|
||||
|
||||
with init_empty_weights():
|
||||
config = FluxTransformer2DModel.load_config("black-forest-labs/flux.1-dev", subfolder="transformer")
|
||||
@@ -187,6 +208,9 @@ def load_flux_nf4(checkpoint_info, diffusers_load_config):
|
||||
create_quantized_param(model, param, param_name, target_device=0, state_dict=original_state_dict, pre_quantized=True)
|
||||
|
||||
del original_state_dict
|
||||
pipe = FluxPipeline.from_pretrained("black-forest-labs/flux.1-dev", transformer=model, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config)
|
||||
devices.torch_gc(force=True)
|
||||
return pipe
|
||||
if transformer_only:
|
||||
return model
|
||||
else:
|
||||
pipe = FluxPipeline.from_pretrained("black-forest-labs/flux.1-dev", transformer=model, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config)
|
||||
return pipe
|
||||
|
||||
+26
-36
@@ -1,54 +1,40 @@
|
||||
import os
|
||||
import torch
|
||||
import transformers
|
||||
from modules import shared, devices, files_cache
|
||||
|
||||
|
||||
t5_dict = {}
|
||||
|
||||
|
||||
def load_t5(t5=None, cache_dir=None):
|
||||
from modules import devices, modelloader
|
||||
from modules import modelloader
|
||||
repo_id = 'stabilityai/stable-diffusion-3-medium-diffusers'
|
||||
if 'fp16' in t5.lower():
|
||||
fn = t5_dict.get(t5) if t5 in t5_dict else None
|
||||
if fn is not None:
|
||||
shared.log.error(f'Loading T5: file="{fn}" unsupported')
|
||||
t5 = None
|
||||
elif 'fp16' in t5.lower():
|
||||
modelloader.hf_login()
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(
|
||||
repo_id,
|
||||
subfolder='text_encoder_3',
|
||||
# torch_dtype=dtype,
|
||||
cache_dir=cache_dir,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
elif 'fp4' in t5.lower():
|
||||
modelloader.hf_login()
|
||||
from installer import install
|
||||
install('bitsandbytes', quiet=True)
|
||||
quantization_config = transformers.BitsAndBytesConfig(load_in_4bit=True)
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(
|
||||
repo_id,
|
||||
subfolder='text_encoder_3',
|
||||
quantization_config=quantization_config,
|
||||
cache_dir=cache_dir,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', quantization_config=quantization_config, cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
elif 'fp8' in t5.lower():
|
||||
modelloader.hf_login()
|
||||
from installer import install
|
||||
install('bitsandbytes', quiet=True)
|
||||
quantization_config = transformers.BitsAndBytesConfig(load_in_8bit=True)
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(
|
||||
repo_id,
|
||||
subfolder='text_encoder_3',
|
||||
quantization_config=quantization_config,
|
||||
cache_dir=cache_dir,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', quantization_config=quantization_config, cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
elif 'qint8' in t5.lower():
|
||||
modelloader.hf_login()
|
||||
from installer import install
|
||||
install('optimum-quanto', quiet=True)
|
||||
from modules.sd_models_compile import optimum_quanto_model
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(
|
||||
repo_id,
|
||||
subfolder='text_encoder_3',
|
||||
cache_dir=cache_dir,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
t5 = optimum_quanto_model(t5, weights="qint8", activations="none")
|
||||
elif 'int8' in t5.lower():
|
||||
modelloader.hf_login()
|
||||
@@ -56,12 +42,7 @@ def load_t5(t5=None, cache_dir=None):
|
||||
install('nncf==2.7.0', quiet=True)
|
||||
from modules.sd_models_compile import nncf_compress_model
|
||||
from modules.sd_hijack import NNCF_T5DenseGatedActDense
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(
|
||||
repo_id,
|
||||
subfolder='text_encoder_3',
|
||||
cache_dir=cache_dir,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
for i in range(len(t5.encoder.block)):
|
||||
t5.encoder.block[i].layer[1].DenseReluDense = NNCF_T5DenseGatedActDense(
|
||||
t5.encoder.block[i].layer[1].DenseReluDense,
|
||||
@@ -74,10 +55,11 @@ def load_t5(t5=None, cache_dir=None):
|
||||
|
||||
|
||||
def set_t5(pipe, module, t5=None, cache_dir=None):
|
||||
from modules import devices, shared
|
||||
if pipe is None or not hasattr(pipe, module):
|
||||
return pipe
|
||||
t5 = load_t5(t5=t5, cache_dir=cache_dir)
|
||||
if module == "text_encoder_2" and t5 is None: # do not unload te2
|
||||
return
|
||||
setattr(pipe, module, t5)
|
||||
if shared.opts.diffusers_offload_mode == "sequential":
|
||||
from accelerate import cpu_offload
|
||||
@@ -90,3 +72,11 @@ def set_t5(pipe, module, t5=None, cache_dir=None):
|
||||
pipe.maybe_free_model_hooks()
|
||||
devices.torch_gc()
|
||||
return pipe
|
||||
|
||||
|
||||
def refresh_t5_list():
|
||||
t5_dict.clear()
|
||||
for file in files_cache.list_files(shared.opts.t5_dir, ext_filter=[".safetensors"]):
|
||||
name = os.path.splitext(os.path.basename(file))[0]
|
||||
t5_dict[name] = file
|
||||
shared.log.debug(f'Available T5s: path="{shared.opts.t5_dir}" items={len(t5_dict)}')
|
||||
|
||||
@@ -101,6 +101,7 @@ def create_paths(opts):
|
||||
create_path(fix_path('diffusers_dir'))
|
||||
create_path(fix_path('vae_dir'))
|
||||
create_path(fix_path('unet_dir'))
|
||||
create_path(fix_path('t5_dir'))
|
||||
create_path(fix_path('lora_dir'))
|
||||
create_path(fix_path('embeddings_dir'))
|
||||
create_path(fix_path('hypernetwork_dir'))
|
||||
|
||||
@@ -99,7 +99,7 @@ def taesd_vae_decode(latents):
|
||||
debug(f'VAE decode: name=TAESD images={len(latents)} latents={latents.shape} slicing={shared.opts.diffusers_vae_slicing}')
|
||||
if len(latents) == 0:
|
||||
return []
|
||||
if shared.opts.diffusers_vae_slicing:
|
||||
if shared.opts.diffusers_vae_slicing and len(latents) > 1:
|
||||
decoded = torch.zeros((len(latents), 3, latents.shape[2] * 8, latents.shape[3] * 8), dtype=devices.dtype_vae, device=devices.device)
|
||||
for i in range(latents.shape[0]):
|
||||
decoded[i] = sd_vae_taesd.decode(latents[i])
|
||||
|
||||
@@ -463,13 +463,13 @@ def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", cl
|
||||
te1_device, te2_device, te3_device = None, None, None
|
||||
if hasattr(pipe, "text_encoder") and pipe.text_encoder.device != devices.device:
|
||||
te1_device = pipe.text_encoder.device
|
||||
pipe.text_encoder = pipe.text_encoder.to(devices.device)
|
||||
sd_models.move_model(pipe.text_encoder, devices.device)
|
||||
if hasattr(pipe, "text_encoder_2") and pipe.text_encoder_2.device != devices.device:
|
||||
te2_device = pipe.text_encoder_2.device
|
||||
pipe.text_encoder_2 = pipe.text_encoder_2.to(devices.device)
|
||||
sd_models.move_model(pipe.text_encoder_2, devices.device)
|
||||
if hasattr(pipe, "text_encoder_3") and pipe.text_encoder_3.device != devices.device:
|
||||
te3_device = pipe.text_encoder_3.device
|
||||
pipe.text_encoder_3 = pipe.text_encoder_3.to(devices.device)
|
||||
sd_models.move_model(pipe.text_encoder_3, devices.device)
|
||||
|
||||
if is_sd3:
|
||||
prompt_embed, negative_embed, positive_pooled, negative_pooled = get_weighted_text_embeddings_sd3(pipe=pipe, prompt=prompt, neg_prompt=neg_prompt, use_t5_encoder=bool(pipe.text_encoder_3))
|
||||
@@ -481,10 +481,10 @@ def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", cl
|
||||
prompt_embed, negative_embed = get_weighted_text_embeddings_sd15(pipe=pipe, prompt=prompt, neg_prompt=neg_prompt, clip_skip=clip_skip)
|
||||
|
||||
if te1_device is not None:
|
||||
pipe.text_encoder = pipe.text_encoder.to(te1_device)
|
||||
sd_models.move_model(pipe.text_encoder, te1_device)
|
||||
if te2_device is not None:
|
||||
pipe.text_encoder_2 = pipe.text_encoder_2.to(te2_device)
|
||||
sd_models.move_model(pipe.text_encoder_2, te1_device)
|
||||
if te3_device is not None:
|
||||
pipe.text_encoder_3 = pipe.text_encoder_3.to(te3_device)
|
||||
sd_models.move_model(pipe.text_encoder_3, te1_device)
|
||||
|
||||
return prompt_embed, positive_pooled, negative_embed, negative_pooled
|
||||
|
||||
@@ -269,12 +269,12 @@ def get_weighted_text_embeddings_sd15(
|
||||
# get positive prompt embeddings with weights
|
||||
token_tensor = torch.tensor(
|
||||
[prompt_token_groups[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder.device
|
||||
)
|
||||
weight_tensor = torch.tensor(
|
||||
prompt_weight_groups[i]
|
||||
, dtype=torch.float16
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder.device
|
||||
)
|
||||
|
||||
token_embedding = pipe.text_encoder(token_tensor)[0].squeeze(0)
|
||||
@@ -286,12 +286,12 @@ def get_weighted_text_embeddings_sd15(
|
||||
# get negative prompt embeddings with weights
|
||||
neg_token_tensor = torch.tensor(
|
||||
[neg_prompt_token_groups[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder.device
|
||||
)
|
||||
neg_weight_tensor = torch.tensor(
|
||||
neg_prompt_weight_groups[i]
|
||||
, dtype=torch.float16
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder.device
|
||||
)
|
||||
neg_token_embedding = pipe.text_encoder(neg_token_tensor)[0].squeeze(0)
|
||||
for z in range(len(neg_weight_tensor)):
|
||||
@@ -449,36 +449,36 @@ def get_weighted_text_embeddings_sdxl(
|
||||
# get positive prompt embeddings with weights
|
||||
token_tensor = torch.tensor(
|
||||
[prompt_token_groups[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder.device
|
||||
)
|
||||
weight_tensor = torch.tensor(
|
||||
prompt_weight_groups[i]
|
||||
, dtype=torch.float16
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder.device
|
||||
)
|
||||
|
||||
token_tensor_2 = torch.tensor(
|
||||
[prompt_token_groups_2[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder_2.device
|
||||
)
|
||||
|
||||
# use first text encoder
|
||||
prompt_embeds_1 = pipe.text_encoder(
|
||||
token_tensor.to(pipe.device)
|
||||
token_tensor.to(pipe.text_encoder.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
prompt_embeds_1_hidden_states = prompt_embeds_1.hidden_states[-2]
|
||||
|
||||
# use second text encoder
|
||||
prompt_embeds_2 = pipe.text_encoder_2(
|
||||
token_tensor_2.to(pipe.device)
|
||||
token_tensor_2.to(pipe.text_encoder_2.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
prompt_embeds_2_hidden_states = prompt_embeds_2.hidden_states[-2]
|
||||
pooled_prompt_embeds = prompt_embeds_2[0]
|
||||
|
||||
prompt_embeds_list = [prompt_embeds_1_hidden_states, prompt_embeds_2_hidden_states]
|
||||
token_embedding = torch.concat(prompt_embeds_list, dim=-1).squeeze(0).to(pipe.device)
|
||||
token_embedding = torch.concat(prompt_embeds_list, dim=-1).squeeze(0).to(pipe.text_encoder.device)
|
||||
|
||||
for j in range(len(weight_tensor)):
|
||||
if weight_tensor[j] != 1.0:
|
||||
@@ -509,35 +509,35 @@ def get_weighted_text_embeddings_sdxl(
|
||||
# get negative prompt embeddings with weights
|
||||
neg_token_tensor = torch.tensor(
|
||||
[neg_prompt_token_groups[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder.device
|
||||
)
|
||||
neg_token_tensor_2 = torch.tensor(
|
||||
[neg_prompt_token_groups_2[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder_2.device
|
||||
)
|
||||
neg_weight_tensor = torch.tensor(
|
||||
neg_prompt_weight_groups[i]
|
||||
, dtype=torch.float16
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder.device
|
||||
)
|
||||
|
||||
# use first text encoder
|
||||
neg_prompt_embeds_1 = pipe.text_encoder(
|
||||
neg_token_tensor.to(pipe.device)
|
||||
neg_token_tensor.to(pipe.text_encoder.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
neg_prompt_embeds_1_hidden_states = neg_prompt_embeds_1.hidden_states[-2]
|
||||
|
||||
# use second text encoder
|
||||
neg_prompt_embeds_2 = pipe.text_encoder_2(
|
||||
neg_token_tensor_2.to(pipe.device)
|
||||
neg_token_tensor_2.to(pipe.text_encoder_2.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
neg_prompt_embeds_2_hidden_states = neg_prompt_embeds_2.hidden_states[-2]
|
||||
negative_pooled_prompt_embeds = neg_prompt_embeds_2[0]
|
||||
|
||||
neg_prompt_embeds_list = [neg_prompt_embeds_1_hidden_states, neg_prompt_embeds_2_hidden_states]
|
||||
neg_token_embedding = torch.concat(neg_prompt_embeds_list, dim=-1).squeeze(0).to(pipe.device)
|
||||
neg_token_embedding = torch.concat(neg_prompt_embeds_list, dim=-1).squeeze(0).to(pipe.text_encoder.device)
|
||||
|
||||
for z in range(len(neg_weight_tensor)):
|
||||
if neg_weight_tensor[z] != 1.0:
|
||||
@@ -657,18 +657,18 @@ def get_weighted_text_embeddings_sdxl_refiner(
|
||||
# get positive prompt embeddings with weights
|
||||
token_tensor_2 = torch.tensor(
|
||||
[prompt_token_groups_2[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder_2.device
|
||||
)
|
||||
|
||||
weight_tensor_2 = torch.tensor(
|
||||
prompt_weight_groups_2[i]
|
||||
, dtype=torch.float16
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder_2.device
|
||||
)
|
||||
|
||||
# use second text encoder
|
||||
prompt_embeds_2 = pipe.text_encoder_2(
|
||||
token_tensor_2.to(pipe.device)
|
||||
token_tensor_2.to(pipe.text_encoder_2.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
prompt_embeds_2_hidden_states = prompt_embeds_2.hidden_states[-2]
|
||||
@@ -703,17 +703,17 @@ def get_weighted_text_embeddings_sdxl_refiner(
|
||||
# get negative prompt embeddings with weights
|
||||
neg_token_tensor_2 = torch.tensor(
|
||||
[neg_prompt_token_groups_2[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder_2.device
|
||||
)
|
||||
neg_weight_tensor_2 = torch.tensor(
|
||||
neg_prompt_weight_groups_2[i]
|
||||
, dtype=torch.float16
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder_2.device
|
||||
)
|
||||
|
||||
# use second text encoder
|
||||
neg_prompt_embeds_2 = pipe.text_encoder_2(
|
||||
neg_token_tensor_2.to(pipe.device)
|
||||
neg_token_tensor_2.to(pipe.text_encoder_2.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
neg_prompt_embeds_2_hidden_states = neg_prompt_embeds_2.hidden_states[-2]
|
||||
@@ -787,8 +787,6 @@ def get_weighted_text_embeddings_sdxl_2p(
|
||||
"""
|
||||
prompt_2 = prompt_2 or prompt
|
||||
neg_prompt_2 = neg_prompt_2 or neg_prompt
|
||||
|
||||
import math
|
||||
eos = pipe.tokenizer.eos_token_id
|
||||
|
||||
# tokenizer 1
|
||||
@@ -907,33 +905,33 @@ def get_weighted_text_embeddings_sdxl_2p(
|
||||
# get positive prompt embeddings with weights
|
||||
token_tensor = torch.tensor(
|
||||
[prompt_token_groups[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder.device
|
||||
)
|
||||
weight_tensor = torch.tensor(
|
||||
prompt_weight_groups[i]
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder.device
|
||||
)
|
||||
|
||||
token_tensor_2 = torch.tensor(
|
||||
[prompt_token_groups_2[i]]
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder_2.device
|
||||
)
|
||||
|
||||
weight_tensor_2 = torch.tensor(
|
||||
prompt_weight_groups_2[i]
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder_2.device
|
||||
)
|
||||
|
||||
# use first text encoder
|
||||
prompt_embeds_1 = pipe.text_encoder(
|
||||
token_tensor.to(pipe.device)
|
||||
token_tensor.to(pipe.text_encoder.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
prompt_embeds_1_hidden_states = prompt_embeds_1.hidden_states[-2]
|
||||
|
||||
# use second text encoder
|
||||
prompt_embeds_2 = pipe.text_encoder_2(
|
||||
token_tensor_2.to(pipe.device)
|
||||
token_tensor_2.to(pipe.text_encoder_2.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
prompt_embeds_2_hidden_states = prompt_embeds_2.hidden_states[-2]
|
||||
@@ -966,31 +964,31 @@ def get_weighted_text_embeddings_sdxl_2p(
|
||||
# get negative prompt embeddings with weights
|
||||
neg_token_tensor = torch.tensor(
|
||||
[neg_prompt_token_groups[i]]
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder.device
|
||||
)
|
||||
neg_token_tensor_2 = torch.tensor(
|
||||
[neg_prompt_token_groups_2[i]]
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder_2.device
|
||||
)
|
||||
neg_weight_tensor = torch.tensor(
|
||||
neg_prompt_weight_groups[i]
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder.device
|
||||
)
|
||||
neg_weight_tensor_2 = torch.tensor(
|
||||
neg_prompt_weight_groups_2[i]
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder_2.device
|
||||
)
|
||||
|
||||
# use first text encoder
|
||||
neg_prompt_embeds_1 = pipe.text_encoder(
|
||||
neg_token_tensor.to(pipe.device)
|
||||
neg_token_tensor.to(pipe.text_encoder.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
neg_prompt_embeds_1_hidden_states = neg_prompt_embeds_1.hidden_states[-2]
|
||||
|
||||
# use second text encoder
|
||||
neg_prompt_embeds_2 = pipe.text_encoder_2(
|
||||
neg_token_tensor_2.to(pipe.device)
|
||||
neg_token_tensor_2.to(pipe.text_encoder_2.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
neg_prompt_embeds_2_hidden_states = neg_prompt_embeds_2.hidden_states[-2]
|
||||
@@ -1049,7 +1047,6 @@ def get_weighted_text_embeddings_sd3(
|
||||
pooled_prompt_embeds (torch.Tensor)
|
||||
negative_pooled_prompt_embeds (torch.Tensor)
|
||||
"""
|
||||
import math
|
||||
eos = pipe.tokenizer.eos_token_id
|
||||
|
||||
# tokenizer 1
|
||||
@@ -1161,22 +1158,22 @@ def get_weighted_text_embeddings_sd3(
|
||||
# get positive prompt embeddings with weights
|
||||
token_tensor = torch.tensor(
|
||||
[prompt_token_groups[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder.device
|
||||
)
|
||||
weight_tensor = torch.tensor(
|
||||
prompt_weight_groups[i]
|
||||
, dtype=torch.float16
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder.device
|
||||
)
|
||||
|
||||
token_tensor_2 = torch.tensor(
|
||||
[prompt_token_groups_2[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder_2.device
|
||||
)
|
||||
|
||||
# use first text encoder
|
||||
prompt_embeds_1 = pipe.text_encoder(
|
||||
token_tensor.to(pipe.device)
|
||||
token_tensor.to(pipe.text_encoder.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
prompt_embeds_1_hidden_states = prompt_embeds_1.hidden_states[-2]
|
||||
@@ -1184,14 +1181,14 @@ def get_weighted_text_embeddings_sd3(
|
||||
|
||||
# use second text encoder
|
||||
prompt_embeds_2 = pipe.text_encoder_2(
|
||||
token_tensor_2.to(pipe.device)
|
||||
token_tensor_2.to(pipe.text_encoder_2.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
prompt_embeds_2_hidden_states = prompt_embeds_2.hidden_states[-2]
|
||||
pooled_prompt_embeds_2 = prompt_embeds_2[0]
|
||||
|
||||
prompt_embeds_list = [prompt_embeds_1_hidden_states, prompt_embeds_2_hidden_states]
|
||||
token_embedding = torch.concat(prompt_embeds_list, dim=-1).squeeze(0).to(pipe.device)
|
||||
token_embedding = torch.concat(prompt_embeds_list, dim=-1).squeeze(0).to(pipe.text_encoder.device)
|
||||
|
||||
for j in range(len(weight_tensor)):
|
||||
if weight_tensor[j] != 1.0:
|
||||
@@ -1222,21 +1219,21 @@ def get_weighted_text_embeddings_sd3(
|
||||
# get negative prompt embeddings with weights
|
||||
neg_token_tensor = torch.tensor(
|
||||
[neg_prompt_token_groups[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder.device
|
||||
)
|
||||
neg_token_tensor_2 = torch.tensor(
|
||||
[neg_prompt_token_groups_2[i]]
|
||||
, dtype=torch.long, device=pipe.device
|
||||
, dtype=torch.long, device=pipe.text_encoder_2.device
|
||||
)
|
||||
neg_weight_tensor = torch.tensor(
|
||||
neg_prompt_weight_groups[i]
|
||||
, dtype=torch.float16
|
||||
, device=pipe.device
|
||||
, device=pipe.text_encoder.device
|
||||
)
|
||||
|
||||
# use first text encoder
|
||||
neg_prompt_embeds_1 = pipe.text_encoder(
|
||||
neg_token_tensor.to(pipe.device)
|
||||
neg_token_tensor.to(pipe.text_encoder.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
neg_prompt_embeds_1_hidden_states = neg_prompt_embeds_1.hidden_states[-2]
|
||||
@@ -1244,14 +1241,14 @@ def get_weighted_text_embeddings_sd3(
|
||||
|
||||
# use second text encoder
|
||||
neg_prompt_embeds_2 = pipe.text_encoder_2(
|
||||
neg_token_tensor_2.to(pipe.device)
|
||||
neg_token_tensor_2.to(pipe.text_encoder_2.device)
|
||||
, output_hidden_states=True
|
||||
)
|
||||
neg_prompt_embeds_2_hidden_states = neg_prompt_embeds_2.hidden_states[-2]
|
||||
negative_pooled_prompt_embeds_2 = neg_prompt_embeds_2[0]
|
||||
|
||||
neg_prompt_embeds_list = [neg_prompt_embeds_1_hidden_states, neg_prompt_embeds_2_hidden_states]
|
||||
neg_token_embedding = torch.concat(neg_prompt_embeds_list, dim=-1).squeeze(0).to(pipe.device)
|
||||
neg_token_embedding = torch.concat(neg_prompt_embeds_list, dim=-1).squeeze(0).to(pipe.text_encoder.device)
|
||||
|
||||
for z in range(len(neg_weight_tensor)):
|
||||
if neg_weight_tensor[z] != 1.0:
|
||||
@@ -1286,8 +1283,8 @@ def get_weighted_text_embeddings_sd3(
|
||||
# ----------------- generate positive t5 embeddings --------------------
|
||||
prompt_tokens_3 = torch.tensor([prompt_tokens_3], dtype=torch.long)
|
||||
|
||||
t5_prompt_embeds = pipe.text_encoder_3(prompt_tokens_3.to(pipe.device))[0].squeeze(0)
|
||||
t5_prompt_embeds = t5_prompt_embeds.to(device=pipe.device)
|
||||
t5_prompt_embeds = pipe.text_encoder_3(prompt_tokens_3.to(pipe.text_encoder_3.device))[0].squeeze(0)
|
||||
t5_prompt_embeds = t5_prompt_embeds.to(device=pipe.text_encoder_3.device)
|
||||
|
||||
# add weight to t5 prompt
|
||||
for z in range(len(prompt_weights_3)):
|
||||
@@ -1296,7 +1293,7 @@ def get_weighted_text_embeddings_sd3(
|
||||
t5_prompt_embeds = t5_prompt_embeds.unsqueeze(0)
|
||||
else:
|
||||
t5_prompt_embeds = torch.zeros(1, 4096, dtype=prompt_embeds.dtype).unsqueeze(0)
|
||||
t5_prompt_embeds = t5_prompt_embeds.to(device=pipe.device)
|
||||
t5_prompt_embeds = t5_prompt_embeds.to(device=pipe.text_encoder_3.device)
|
||||
|
||||
# merge with the clip embedding 1 and clip embedding 2
|
||||
clip_prompt_embeds = torch.nn.functional.pad(
|
||||
@@ -1308,8 +1305,8 @@ def get_weighted_text_embeddings_sd3(
|
||||
# ---------------------- get neg t5 embeddings -------------------------
|
||||
neg_prompt_tokens_3 = torch.tensor([neg_prompt_tokens_3], dtype=torch.long)
|
||||
|
||||
t5_neg_prompt_embeds = pipe.text_encoder_3(neg_prompt_tokens_3.to(pipe.device))[0].squeeze(0)
|
||||
t5_neg_prompt_embeds = t5_neg_prompt_embeds.to(device=pipe.device)
|
||||
t5_neg_prompt_embeds = pipe.text_encoder_3(neg_prompt_tokens_3.to(pipe.pipe.text_encoder_3.device))[0].squeeze(0)
|
||||
t5_neg_prompt_embeds = t5_neg_prompt_embeds.to(device=pipe.text_encoder_3.device)
|
||||
|
||||
# add weight to neg t5 embeddings
|
||||
for z in range(len(neg_prompt_weights_3)):
|
||||
@@ -1318,7 +1315,7 @@ def get_weighted_text_embeddings_sd3(
|
||||
t5_neg_prompt_embeds = t5_neg_prompt_embeds.unsqueeze(0)
|
||||
else:
|
||||
t5_neg_prompt_embeds = torch.zeros(1, 4096, dtype=prompt_embeds.dtype).unsqueeze(0)
|
||||
t5_neg_prompt_embeds = t5_prompt_embeds.to(device=pipe.device)
|
||||
t5_neg_prompt_embeds = t5_prompt_embeds.to(device=pipe.text_encoder_3.device)
|
||||
|
||||
clip_neg_prompt_embeds = torch.nn.functional.pad(
|
||||
negative_prompt_embeds, (0, t5_neg_prompt_embeds.shape[-1] - negative_prompt_embeds.shape[-1])
|
||||
@@ -1359,7 +1356,7 @@ def get_weighted_text_embeddings_flux1(
|
||||
"""
|
||||
prompt2 = prompt if prompt2 is None else prompt2
|
||||
if device is None:
|
||||
device = pipe.device
|
||||
device = pipe.text_encoder.device
|
||||
|
||||
# tokenizer 1 - openai/clip-vit-large-patch14
|
||||
prompt_tokens, prompt_weights = get_prompts_tokens_with_weights(
|
||||
|
||||
@@ -566,7 +566,7 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False):
|
||||
# guess by size
|
||||
if os.path.isfile(f) and f.endswith('.safetensors'):
|
||||
size = round(os.path.getsize(f) / 1024 / 1024)
|
||||
if size < 128:
|
||||
if (size > 0 and size < 128):
|
||||
warn(f'Model size smaller than expected: {f} size={size} MB')
|
||||
elif (size >= 316 and size <= 324) or (size >= 156 and size <= 164): # 320 or 160
|
||||
warn(f'Model detected as VAE model, but attempting to load as model: {op}={f} size={size} MB')
|
||||
@@ -591,6 +591,8 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False):
|
||||
guess = 'Stable Diffusion XL'
|
||||
elif (size > 5692 and size < 5698) or (size > 4134 and size < 4138) or (size > 10362 and size < 10366) or (size > 15028 and size < 15228):
|
||||
guess = 'Stable Diffusion 3'
|
||||
elif (size > 20000 and size < 40000):
|
||||
guess = 'FLUX'
|
||||
# guess by name
|
||||
"""
|
||||
if 'LCM_' in f.upper() or 'LCM-' in f.upper() or '_LCM' in f.upper() or '-LCM' in f.upper():
|
||||
@@ -620,8 +622,10 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False):
|
||||
guess = 'Kolors'
|
||||
if 'auraflow' in f.lower():
|
||||
guess = 'AuraFlow'
|
||||
if 'flux.1' in f.lower() or 'flux1' in f.lower():
|
||||
if 'flux' in f.lower():
|
||||
guess = 'FLUX'
|
||||
if size > 11000 and size < 20000:
|
||||
warn(f'Model detected as FLUX UNET model, but attempting to load a base model: {op}={f} size={size} MB')
|
||||
# switch for specific variant
|
||||
if guess == 'Stable Diffusion' and 'inpaint' in f.lower():
|
||||
guess = 'Stable Diffusion Inpaint'
|
||||
|
||||
@@ -71,7 +71,7 @@ def create_sampler(name, model):
|
||||
sampler = config.constructor(model)
|
||||
if shared.sd_model_type == 'f1':
|
||||
if 'base_image_seq_len' not in sampler.sampler.config or 'max_image_seq_len' not in sampler.sampler.config or 'base_shift' not in sampler.sampler.config or 'max_shift' not in sampler.sampler.config:
|
||||
shared.log.warning('FLUX sampler: attempting to use a non compatible scheduler')
|
||||
shared.log.warning(f'FLUX: sampler="{name}" non compatible')
|
||||
return None
|
||||
if not hasattr(model, 'scheduler_config'):
|
||||
model.scheduler_config = sampler.sampler.config.copy()
|
||||
|
||||
+7
-5
@@ -1,8 +1,9 @@
|
||||
import os
|
||||
from modules import shared, devices, files_cache
|
||||
from modules import shared, devices, files_cache, sd_models
|
||||
|
||||
|
||||
unet_dict = {}
|
||||
debug = os.environ.get('SD_LOAD_DEBUG', None) is not None
|
||||
|
||||
|
||||
def load_unet(model):
|
||||
@@ -28,15 +29,13 @@ def load_unet(model):
|
||||
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__:
|
||||
shared.log.info(f'Loading UNet: name="{shared.opts.sd_unet}" file="{unet_dict[shared.opts.sd_unet]}" offload={shared.opts.diffusers_offload_mode}')
|
||||
from modules.model_flux import load_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':
|
||||
model.transformer = transformer.to(devices.device, devices.dtype)
|
||||
else:
|
||||
model.transformer = transformer
|
||||
sd_models.move_model(transformer, devices.device)
|
||||
model.transformer = transformer
|
||||
from modules.sd_models import set_diffuser_offload
|
||||
set_diffuser_offload(model, 'model')
|
||||
else:
|
||||
@@ -52,6 +51,9 @@ def load_unet(model):
|
||||
model.unet = unet.to(devices.device, devices.dtype_unet)
|
||||
except Exception as e:
|
||||
shared.log.error(f'Failed to load UNet model: {e}')
|
||||
if debug:
|
||||
from modules import errors
|
||||
errors.display(e, 'UNet load:')
|
||||
return
|
||||
devices.torch_gc()
|
||||
|
||||
|
||||
@@ -233,6 +233,8 @@ def load_vae_diffusers(model_file, vae_file=None, vae_source="unknown-source"):
|
||||
global loaded_vae_file # pylint: disable=global-statement
|
||||
loaded_vae_file = os.path.basename(vae_file)
|
||||
# shared.log.debug(f'Diffusers VAE config: {vae.config}')
|
||||
if shared.opts.diffusers_offload_mode == 'none':
|
||||
sd_models.move_model(vae, devices.device)
|
||||
return vae
|
||||
except Exception as e:
|
||||
shared.log.error(f"Loading VAE failed: model={vae_file} {e}")
|
||||
|
||||
@@ -55,6 +55,8 @@ def Decoder(latent_channels=4):
|
||||
return nn.Sequential(
|
||||
Clamp(), conv(latent_channels, 64), nn.ReLU(),
|
||||
Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False),
|
||||
Block(64, 64), Block(64, 64), Block(64, 64), nn.Identity(), conv(64, 64, bias=False),
|
||||
Block(64, 64), Block(64, 64), Block(64, 64), nn.Identity(), conv(64, 64, bias=False),
|
||||
Block(64, 64), conv(64, 3),
|
||||
)
|
||||
elif shared.opts.live_preview_taesd_layers == 2:
|
||||
@@ -62,6 +64,7 @@ def Decoder(latent_channels=4):
|
||||
Clamp(), conv(latent_channels, 64), nn.ReLU(),
|
||||
Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False),
|
||||
Block(64, 64), Block(64, 64), Block(64, 64), nn.Upsample(scale_factor=2), conv(64, 64, bias=False),
|
||||
Block(64, 64), Block(64, 64), Block(64, 64), nn.Identity(), conv(64, 64, bias=False),
|
||||
Block(64, 64), conv(64, 3),
|
||||
)
|
||||
else:
|
||||
@@ -86,9 +89,9 @@ class TAESD(nn.Module): # pylint: disable=abstract-method
|
||||
self.encoder = Encoder(latent_channels)
|
||||
self.decoder = Decoder(latent_channels)
|
||||
if encoder_path is not None:
|
||||
self.encoder.load_state_dict(torch.load(encoder_path, map_location="cpu"))
|
||||
self.encoder.load_state_dict(torch.load(encoder_path, map_location="cpu"), strict=False)
|
||||
if decoder_path is not None:
|
||||
self.decoder.load_state_dict(torch.load(decoder_path, map_location="cpu"))
|
||||
self.decoder.load_state_dict(torch.load(decoder_path, map_location="cpu"), strict=False)
|
||||
|
||||
def guess_latent_channels(self, decoder_path, encoder_path):
|
||||
"""guess latent channel count based on encoder filename"""
|
||||
|
||||
+3
-1
@@ -406,7 +406,8 @@ options_templates.update(options_section(('sd', "Execution & Models"), {
|
||||
"sd_model_refiner": OptionInfo('None', "Refiner model", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_tiles()}, refresh=refresh_checkpoints),
|
||||
"sd_vae": OptionInfo("Automatic", "VAE model", gr.Dropdown, lambda: {"choices": shared_items.sd_vae_items()}, refresh=shared_items.refresh_vae_list),
|
||||
"sd_unet": OptionInfo("None", "UNET model", gr.Dropdown, lambda: {"choices": shared_items.sd_unet_items()}, refresh=shared_items.refresh_unet_list),
|
||||
"sd_text_encoder": OptionInfo('None', "Text encoder model", gr.Dropdown, lambda: {"choices": ['None', 'T5 FP4', 'T5 FP8', 'T5 INT8', 'T5 QINT8', 'T5 FP16']}),
|
||||
# "sd_text_encoder": OptionInfo('None', "Text encoder model", gr.Dropdown, lambda: {"choices": ['None', 'T5 FP4', 'T5 FP8', 'T5 INT8', 'T5 QINT8', 'T5 FP16']}),
|
||||
"sd_text_encoder": OptionInfo('None', "Text encoder model", gr.Dropdown, lambda: {"choices": shared_items.sd_t5_items()}, refresh=shared_items.refresh_t5_list),
|
||||
"sd_model_dict": OptionInfo('None', "Use separate base dict", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_tiles()}, refresh=refresh_checkpoints),
|
||||
"sd_checkpoint_autoload": OptionInfo(True, "Model autoload on start"),
|
||||
"sd_textencoder_cache": OptionInfo(True, "Cache text encoder results"),
|
||||
@@ -575,6 +576,7 @@ options_templates.update(options_section(('system-paths', "System Paths"), {
|
||||
"hfcache_dir": OptionInfo(os.path.join(os.path.expanduser('~'), '.cache', 'huggingface', 'hub'), "Folder for Huggingface cache", folder=True),
|
||||
"vae_dir": OptionInfo(os.path.join(paths.models_path, 'VAE'), "Folder with VAE files", folder=True),
|
||||
"unet_dir": OptionInfo(os.path.join(paths.models_path, 'UNET'), "Folder with UNET files", folder=True),
|
||||
"t5_dir": OptionInfo(os.path.join(paths.models_path, 'T5'), "Folder with T5 files", folder=True),
|
||||
"sd_lora": OptionInfo("", "Add LoRA to prompt", gr.Textbox, {"visible": False}),
|
||||
"lora_dir": OptionInfo(os.path.join(paths.models_path, 'Lora'), "Folder with LoRA network(s)", folder=True),
|
||||
"lyco_dir": OptionInfo(os.path.join(paths.models_path, 'LyCORIS'), "Folder with LyCORIS network(s)", gr.Text, {"visible": False}),
|
||||
|
||||
@@ -23,6 +23,17 @@ def refresh_unet_list():
|
||||
modules.sd_unet.refresh_unet_list()
|
||||
|
||||
|
||||
def sd_t5_items():
|
||||
import modules.model_t5
|
||||
predefined = ['None', 'T5 FP4', 'T5 FP8', 'T5 INT8', 'T5 QINT8', 'T5 FP16']
|
||||
return predefined + list(modules.model_t5.t5_dict)
|
||||
|
||||
|
||||
def refresh_t5_list():
|
||||
import modules.model_t5
|
||||
modules.model_t5.refresh_t5_list()
|
||||
|
||||
|
||||
def list_crossattention(diffusers=False):
|
||||
if diffusers:
|
||||
return [
|
||||
|
||||
@@ -24,6 +24,7 @@ import modules.scripts
|
||||
import modules.sd_models
|
||||
import modules.sd_vae
|
||||
import modules.sd_unet
|
||||
import modules.model_t5
|
||||
import modules.progress
|
||||
import modules.ui
|
||||
import modules.txt2img
|
||||
@@ -90,6 +91,9 @@ def initialize():
|
||||
modules.sd_unet.refresh_unet_list()
|
||||
timer.startup.record("unet")
|
||||
|
||||
modules.model_t5.refresh_t5_list()
|
||||
timer.startup.record("unet")
|
||||
|
||||
extensions.list_extensions()
|
||||
timer.startup.record("extensions")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user