post release dev merge, see changelog for details

This commit is contained in:
Vladimir Mandic
2024-09-01 12:20:08 -04:00
parent 66f06fba55
commit 0d9ce663e4
16 changed files with 272 additions and 178 deletions
+19
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)}')
+1
View File
@@ -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'))
+1 -1
View File
@@ -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])
+6 -6
View File
@@ -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
+53 -56
View File
@@ -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(
+6 -2
View File
@@ -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'
+1 -1
View File
@@ -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
View File
@@ -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()
+2
View File
@@ -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}")
+5 -2
View File
@@ -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
View File
@@ -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}),
+11
View File
@@ -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 [
+4
View File
@@ -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")