diff --git a/backend/README.md b/backend/README.md
index ef302e7f..9e17d133 100644
--- a/backend/README.md
+++ b/backend/README.md
@@ -1 +1 @@
-# WIP Backend for Forge
+
W.I.P Backend for Forge
diff --git a/backend/args.py b/backend/args.py
index 6758bb89..8aa3cdd3 100644
--- a/backend/args.py
+++ b/backend/args.py
@@ -61,7 +61,4 @@ parser.add_argument("--disable-gpu-warning", action="store_true")
args = parser.parse_known_args()[0]
# Some dynamic args that may be changed by webui rather than cmd flags.
-dynamic_args = dict(
- embedding_dir='./embeddings',
- emphasis_name='original'
-)
+dynamic_args = dict(embedding_dir="./embeddings", emphasis_name="original")
diff --git a/backend/attention.py b/backend/attention.py
index b0bdbc7b..df898d41 100644
--- a/backend/attention.py
+++ b/backend/attention.py
@@ -1,10 +1,10 @@
import math
-import torch
+
import einops
+import torch
-from backend.args import args
from backend import memory_management
-
+from backend.args import args
BROKEN_XFORMERS = False
if memory_management.xformers_enabled():
@@ -42,7 +42,7 @@ def attention_basic(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
b, _, dim_head = q.shape
dim_head //= heads
- scale = dim_head ** -0.5
+ scale = dim_head**-0.5
h = heads
if skip_reshape:
@@ -52,26 +52,22 @@ def attention_basic(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
)
else:
q, k, v = map(
- lambda t: t.unsqueeze(3)
- .reshape(b, -1, heads, dim_head)
- .permute(0, 2, 1, 3)
- .reshape(b * heads, -1, dim_head)
- .contiguous(),
+ lambda t: t.unsqueeze(3).reshape(b, -1, heads, dim_head).permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head).contiguous(),
(q, k, v),
)
if attn_precision == torch.float32:
- sim = torch.einsum('b i d, b j d -> b i j', q.float(), k.float()) * scale
+ sim = torch.einsum("b i d, b j d -> b i j", q.float(), k.float()) * scale
else:
- sim = torch.einsum('b i d, b j d -> b i j', q, k) * scale
+ sim = torch.einsum("b i d, b j d -> b i j", q, k) * scale
del q, k
if exists(mask):
if mask.dtype == torch.bool:
- mask = einops.rearrange(mask, 'b ... -> b (...)')
+ mask = einops.rearrange(mask, "b ... -> b (...)")
max_neg_value = -torch.finfo(sim.dtype).max
- mask = einops.repeat(mask, 'b j -> (b h) () j', h=h)
+ mask = einops.repeat(mask, "b j -> (b h) () j", h=h)
sim.masked_fill_(~mask, max_neg_value)
else:
if len(mask.shape) == 2:
@@ -82,13 +78,8 @@ def attention_basic(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
sim.add_(mask)
sim = sim.softmax(dim=-1)
- out = torch.einsum('b i j, b j d -> b i d', sim.to(v.dtype), v)
- out = (
- out.unsqueeze(0)
- .reshape(b, heads, -1, dim_head)
- .permute(0, 2, 1, 3)
- .reshape(b, -1, heads * dim_head)
- )
+ out = torch.einsum("b i j, b j d -> b i d", sim.to(v.dtype), v)
+ out = out.unsqueeze(0).reshape(b, heads, -1, dim_head).permute(0, 2, 1, 3).reshape(b, -1, heads * dim_head)
return out
@@ -101,7 +92,7 @@ def attention_split(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
b, _, dim_head = q.shape
dim_head //= heads
- scale = dim_head ** -0.5
+ scale = dim_head**-0.5
h = heads
if skip_reshape:
@@ -111,11 +102,7 @@ def attention_split(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
)
else:
q, k, v = map(
- lambda t: t.unsqueeze(3)
- .reshape(b, -1, heads, dim_head)
- .permute(0, 2, 1, 3)
- .reshape(b * heads, -1, dim_head)
- .contiguous(),
+ lambda t: t.unsqueeze(3).reshape(b, -1, heads, dim_head).permute(0, 2, 1, 3).reshape(b * heads, -1, dim_head).contiguous(),
(q, k, v),
)
@@ -130,7 +117,7 @@ def attention_split(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
element_size = q.element_size()
upcast = False
- gb = 1024 ** 3
+ gb = 1024**3
tensor_size = q.shape[0] * q.shape[1] * k.shape[1] * element_size
modifier = 3
mem_required = tensor_size * modifier
@@ -143,8 +130,7 @@ def attention_split(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
if steps > 64:
max_res = math.floor(math.sqrt(math.sqrt(mem_free_total / 2.5)) / 8) * 64
- raise RuntimeError(f'Not enough memory, use lower resolution (max approx. {max_res}x{max_res}). '
- f'Need: {mem_required / 64 / gb:0.1f}GB free, Have:{mem_free_total / gb:0.1f}GB free')
+ raise RuntimeError(f"Not enough memory, use lower resolution (max approx. {max_res}x{max_res}). " f"Need: {mem_required / 64 / gb:0.1f}GB free, Have:{mem_free_total / gb:0.1f}GB free")
if mask is not None:
if len(mask.shape) == 2:
@@ -162,10 +148,10 @@ def attention_split(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
for i in range(0, q.shape[1], slice_size):
end = i + slice_size
if upcast:
- with torch.autocast(enabled=False, device_type='cuda'):
- s1 = torch.einsum('b i d, b j d -> b i j', q[:, i:end].float(), k.float()) * scale
+ with torch.autocast(enabled=False, device_type="cuda"):
+ s1 = torch.einsum("b i d, b j d -> b i j", q[:, i:end].float(), k.float()) * scale
else:
- s1 = torch.einsum('b i d, b j d -> b i j', q[:, i:end], k) * scale
+ s1 = torch.einsum("b i d, b j d -> b i j", q[:, i:end], k) * scale
if mask is not None:
if len(mask.shape) == 2:
@@ -177,7 +163,7 @@ def attention_split(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
del s1
first_op_done = True
- r1[:, i:end] = torch.einsum('b i j, b j d -> b i d', s2, v)
+ r1[:, i:end] = torch.einsum("b i j, b j d -> b i d", s2, v)
del s2
break
except memory_management.OOM_EXCEPTION as e:
@@ -196,12 +182,7 @@ def attention_split(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
del q, k, v
- r1 = (
- r1.unsqueeze(0)
- .reshape(b, heads, -1, dim_head)
- .permute(0, 2, 1, 3)
- .reshape(b, -1, heads * dim_head)
- )
+ r1 = r1.unsqueeze(0).reshape(b, heads, -1, dim_head).permute(0, 2, 1, 3).reshape(b, -1, heads * dim_head)
return r1
@@ -229,22 +210,15 @@ def attention_xformers(q, k, v, heads, mask=None, attn_precision=None, skip_resh
if mask is not None:
pad = 8 - q.shape[1] % 8
mask_out = torch.empty([q.shape[0], q.shape[1], q.shape[1] + pad], dtype=q.dtype, device=q.device)
- mask_out[:, :, :mask.shape[-1]] = mask
- mask = mask_out[:, :, :mask.shape[-1]]
+ mask_out[:, :, : mask.shape[-1]] = mask
+ mask = mask_out[:, :, : mask.shape[-1]]
out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=mask)
if skip_reshape:
- out = (
- out.unsqueeze(0)
- .reshape(b, heads, -1, dim_head)
- .permute(0, 2, 1, 3)
- .reshape(b, -1, heads * dim_head)
- )
+ out = out.unsqueeze(0).reshape(b, heads, -1, dim_head).permute(0, 2, 1, 3).reshape(b, -1, heads * dim_head)
else:
- out = (
- out.reshape(b, -1, heads * dim_head)
- )
+ out = out.reshape(b, -1, heads * dim_head)
return out
@@ -261,19 +235,17 @@ def attention_pytorch(q, k, v, heads, mask=None, attn_precision=None, skip_resha
)
out = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False)
- out = (
- out.transpose(1, 2).reshape(b, -1, heads * dim_head)
- )
+ out = out.transpose(1, 2).reshape(b, -1, heads * dim_head)
return out
def slice_attention_single_head_spatial(q, k, v):
r1 = torch.zeros_like(k, device=q.device)
- scale = (int(q.shape[-1]) ** (-0.5))
+ scale = int(q.shape[-1]) ** (-0.5)
mem_free_total = memory_management.get_free_memory(q.device)
- gb = 1024 ** 3
+ gb = 1024**3
tensor_size = q.shape[0] * q.shape[1] * k.shape[2] * q.element_size()
modifier = 3 if q.element_size() == 2 else 2.5
mem_required = tensor_size * modifier
@@ -332,8 +304,7 @@ def xformers_attention_single_head_spatial(q, k, v):
out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=None)
out = out.transpose(1, 2).reshape(B, C, H, W)
except NotImplementedError as e:
- out = slice_attention_single_head_spatial(q.view(B, -1, C), k.view(B, -1, C).transpose(1, 2),
- v.view(B, -1, C).transpose(1, 2)).reshape(B, C, H, W)
+ out = slice_attention_single_head_spatial(q.view(B, -1, C), k.view(B, -1, C).transpose(1, 2), v.view(B, -1, C).transpose(1, 2)).reshape(B, C, H, W)
return out
@@ -350,8 +321,7 @@ def pytorch_attention_single_head_spatial(q, k, v):
out = out.transpose(2, 3).reshape(B, C, H, W)
except memory_management.OOM_EXCEPTION as e:
print("scaled_dot_product_attention OOMed: switched to slice attention")
- out = slice_attention_single_head_spatial(q.view(B, -1, C), k.view(B, -1, C).transpose(1, 2),
- v.view(B, -1, C).transpose(1, 2)).reshape(B, C, H, W)
+ out = slice_attention_single_head_spatial(q.view(B, -1, C), k.view(B, -1, C).transpose(1, 2), v.view(B, -1, C).transpose(1, 2)).reshape(B, C, H, W)
return out
@@ -392,9 +362,7 @@ class AttentionProcessorForge:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
- batch_size, sequence_length, _ = (
- hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
- )
+ batch_size, sequence_length, _ = hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
diff --git a/backend/loader.py b/backend/loader.py
index f57d2058..b4934d0c 100644
--- a/backend/loader.py
+++ b/backend/loader.py
@@ -1,27 +1,28 @@
-import os
-import torch
-import logging
import importlib
+import logging
+import os
-import backend.args
import huggingface_guess
-
+import torch
from diffusers import DiffusionPipeline
from transformers import modeling_utils
+import backend.args
from backend import memory_management
-from backend.utils import read_arbitrary_config, load_torch_file, beautiful_print_gguf_state_dict_statics
-from backend.state_dict import try_filter_state_dict, load_state_dict
-from backend.operations import using_forge_operations
-from backend.nn.vae import IntegratedAutoencoderKL
-from backend.nn.clip import IntegratedCLIP
-from backend.nn.unet import IntegratedUNet2DConditionModel
-
+from backend.diffusion_engine.chroma import Chroma
+from backend.diffusion_engine.flux import Flux
from backend.diffusion_engine.sd15 import StableDiffusion
from backend.diffusion_engine.sdxl import StableDiffusionXL, StableDiffusionXLRefiner
-from backend.diffusion_engine.flux import Flux
-from backend.diffusion_engine.chroma import Chroma
-
+from backend.nn.clip import IntegratedCLIP
+from backend.nn.unet import IntegratedUNet2DConditionModel
+from backend.nn.vae import IntegratedAutoencoderKL
+from backend.operations import using_forge_operations
+from backend.state_dict import load_state_dict, try_filter_state_dict
+from backend.utils import (
+ beautiful_print_gguf_state_dict_statics,
+ load_torch_file,
+ read_arbitrary_config,
+)
possible_models = [StableDiffusion, StableDiffusionXLRefiner, StableDiffusionXL, Chroma, Flux]
@@ -33,34 +34,35 @@ dir_path = os.path.dirname(__file__)
def load_huggingface_component(guess, component_name, lib_name, cls_name, repo_path, state_dict):
config_path = os.path.join(repo_path, component_name)
- if component_name in ['feature_extractor', 'safety_checker']:
+ if component_name in ["feature_extractor", "safety_checker"]:
return None
- if lib_name in ['transformers', 'diffusers']:
- if component_name in ['scheduler']:
+ if lib_name in ["transformers", "diffusers"]:
+ if component_name in ["scheduler"]:
cls = getattr(importlib.import_module(lib_name), cls_name)
return cls.from_pretrained(os.path.join(repo_path, component_name))
- if component_name.startswith('tokenizer'):
+ if component_name.startswith("tokenizer"):
cls = getattr(importlib.import_module(lib_name), cls_name)
comp = cls.from_pretrained(os.path.join(repo_path, component_name))
comp._eventual_warn_about_too_long_sequence = lambda *args, **kwargs: None
return comp
- if cls_name in ['AutoencoderKL']:
- assert isinstance(state_dict, dict) and len(state_dict) > 16, 'You do not have VAE state dict!'
+ if cls_name in ["AutoencoderKL"]:
+ assert isinstance(state_dict, dict) and len(state_dict) > 16, "You do not have VAE state dict!"
config = IntegratedAutoencoderKL.load_config(config_path)
with using_forge_operations(device=memory_management.cpu, dtype=memory_management.vae_dtype()):
model = IntegratedAutoencoderKL.from_config(config)
- if 'decoder.up_blocks.0.resnets.0.norm1.weight' in state_dict.keys(): #diffusers format
+ if "decoder.up_blocks.0.resnets.0.norm1.weight" in state_dict.keys(): # diffusers format
state_dict = huggingface_guess.diffusers_convert.convert_vae_state_dict(state_dict)
- load_state_dict(model, state_dict, ignore_start='loss.')
+ load_state_dict(model, state_dict, ignore_start="loss.")
return model
- if component_name.startswith('text_encoder') and cls_name in ['CLIPTextModel', 'CLIPTextModelWithProjection']:
- assert isinstance(state_dict, dict) and len(state_dict) > 16, 'You do not have CLIP state dict!'
+ if component_name.startswith("text_encoder") and cls_name in ["CLIPTextModel", "CLIPTextModelWithProjection"]:
+ assert isinstance(state_dict, dict) and len(state_dict) > 16, "You do not have CLIP state dict!"
from transformers import CLIPTextConfig, CLIPTextModel
+
config = CLIPTextConfig.from_pretrained(config_path)
to_args = dict(device=memory_management.cpu, dtype=memory_management.text_encoder_dtype())
@@ -69,33 +71,30 @@ def load_huggingface_component(guess, component_name, lib_name, cls_name, repo_p
with using_forge_operations(**to_args, manual_cast_enabled=True):
model = IntegratedCLIP(CLIPTextModel, config, add_text_projection=True).to(**to_args)
- load_state_dict(model, state_dict, ignore_errors=[
- 'transformer.text_projection.weight',
- 'transformer.text_model.embeddings.position_ids',
- 'logit_scale'
- ], log_name=cls_name)
+ load_state_dict(model, state_dict, ignore_errors=["transformer.text_projection.weight", "transformer.text_model.embeddings.position_ids", "logit_scale"], log_name=cls_name)
return model
- if cls_name == 'T5EncoderModel':
- assert isinstance(state_dict, dict) and len(state_dict) > 16, 'You do not have T5 state dict!'
+ if cls_name == "T5EncoderModel":
+ assert isinstance(state_dict, dict) and len(state_dict) > 16, "You do not have T5 state dict!"
from backend.nn.t5 import IntegratedT5
+
config = read_arbitrary_config(config_path)
storage_dtype = memory_management.text_encoder_dtype()
state_dict_dtype = memory_management.state_dict_dtype(state_dict)
- if state_dict_dtype in [torch.float8_e4m3fn, torch.float8_e5m2, 'nf4', 'fp4', 'gguf']:
- print(f'Using Detected T5 Data Type: {state_dict_dtype}')
+ if state_dict_dtype in [torch.float8_e4m3fn, torch.float8_e5m2, "nf4", "fp4", "gguf"]:
+ print(f"Using Detected T5 Data Type: {state_dict_dtype}")
storage_dtype = state_dict_dtype
- if state_dict_dtype in ['nf4', 'fp4', 'gguf']:
- print(f'Using pre-quant state dict!')
- if state_dict_dtype in ['gguf']:
+ if state_dict_dtype in ["nf4", "fp4", "gguf"]:
+ print(f"Using pre-quant state dict!")
+ if state_dict_dtype in ["gguf"]:
beautiful_print_gguf_state_dict_statics(state_dict)
else:
- print(f'Using Default T5 Data Type: {storage_dtype}')
+ print(f"Using Default T5 Data Type: {storage_dtype}")
- if storage_dtype in ['nf4', 'fp4', 'gguf']:
+ if storage_dtype in ["nf4", "fp4", "gguf"]:
with modeling_utils.no_init_weights():
with using_forge_operations(device=memory_management.cpu, dtype=memory_management.text_encoder_dtype(), manual_cast_enabled=False, bnb_dtype=storage_dtype):
model = IntegratedT5(config)
@@ -104,20 +103,22 @@ def load_huggingface_component(guess, component_name, lib_name, cls_name, repo_p
with using_forge_operations(device=memory_management.cpu, dtype=storage_dtype, manual_cast_enabled=True):
model = IntegratedT5(config)
- load_state_dict(model, state_dict, log_name=cls_name, ignore_errors=['transformer.encoder.embed_tokens.weight', 'logit_scale'])
+ load_state_dict(model, state_dict, log_name=cls_name, ignore_errors=["transformer.encoder.embed_tokens.weight", "logit_scale"])
return model
- if cls_name in ['UNet2DConditionModel', 'FluxTransformer2DModel', 'SD3Transformer2DModel', 'ChromaTransformer2DModel']:
- assert isinstance(state_dict, dict) and len(state_dict) > 16, 'You do not have model state dict!'
+ if cls_name in ["UNet2DConditionModel", "FluxTransformer2DModel", "SD3Transformer2DModel", "ChromaTransformer2DModel"]:
+ assert isinstance(state_dict, dict) and len(state_dict) > 16, "You do not have model state dict!"
model_loader = None
- if cls_name == 'UNet2DConditionModel':
+ if cls_name == "UNet2DConditionModel":
model_loader = lambda c: IntegratedUNet2DConditionModel.from_config(c)
- elif cls_name == 'FluxTransformer2DModel':
+ elif cls_name == "FluxTransformer2DModel":
from backend.nn.flux import IntegratedFluxTransformer2DModel
+
model_loader = lambda c: IntegratedFluxTransformer2DModel(**c)
- elif cls_name == 'ChromaTransformer2DModel':
+ elif cls_name == "ChromaTransformer2DModel":
from backend.nn.chroma import IntegratedChromaTransformer2DModel
+
model_loader = lambda c: IntegratedChromaTransformer2DModel(**c)
# elif cls_name == 'SD3Transformer2DModel':
# from backend.nn.mmditx import MMDiTX
@@ -129,23 +130,23 @@ def load_huggingface_component(guess, component_name, lib_name, cls_name, repo_p
storage_dtype = memory_management.unet_dtype(model_params=state_dict_parameters, supported_dtypes=guess.supported_inference_dtypes)
- unet_storage_dtype_overwrite = backend.args.dynamic_args.get('forge_unet_storage_dtype')
+ unet_storage_dtype_overwrite = backend.args.dynamic_args.get("forge_unet_storage_dtype")
if unet_storage_dtype_overwrite is not None:
storage_dtype = unet_storage_dtype_overwrite
- elif state_dict_dtype in [torch.float8_e4m3fn, torch.float8_e5m2, 'nf4', 'fp4', 'gguf']:
- print(f'Using Detected UNet Type: {state_dict_dtype}')
+ elif state_dict_dtype in [torch.float8_e4m3fn, torch.float8_e5m2, "nf4", "fp4", "gguf"]:
+ print(f"Using Detected UNet Type: {state_dict_dtype}")
storage_dtype = state_dict_dtype
- if state_dict_dtype in ['nf4', 'fp4', 'gguf']:
- print(f'Using pre-quant state dict!')
- if state_dict_dtype in ['gguf']:
+ if state_dict_dtype in ["nf4", "fp4", "gguf"]:
+ print(f"Using pre-quant state dict!")
+ if state_dict_dtype in ["gguf"]:
beautiful_print_gguf_state_dict_statics(state_dict)
load_device = memory_management.get_torch_device()
computation_dtype = memory_management.get_computation_dtype(load_device, parameters=state_dict_parameters, supported_dtypes=guess.supported_inference_dtypes)
offload_device = memory_management.unet_offload_device()
- if storage_dtype in ['nf4', 'fp4', 'gguf']:
+ if storage_dtype in ["nf4", "fp4", "gguf"]:
initial_device = memory_management.unet_inital_load_device(parameters=state_dict_parameters, dtype=computation_dtype)
with using_forge_operations(device=initial_device, dtype=computation_dtype, manual_cast_enabled=False, bnb_dtype=storage_dtype):
model = model_loader(unet_config)
@@ -159,7 +160,7 @@ def load_huggingface_component(guess, component_name, lib_name, cls_name, repo_p
load_state_dict(model, state_dict)
- if hasattr(model, '_internal_dict'):
+ if hasattr(model, "_internal_dict"):
model._internal_dict = unet_config
else:
model.config = unet_config
@@ -172,7 +173,7 @@ def load_huggingface_component(guess, component_name, lib_name, cls_name, repo_p
return model
- print(f'Skipped: {component_name} = {lib_name}.{cls_name}')
+ print(f"Skipped: {component_name} = {lib_name}.{cls_name}")
return None
@@ -180,7 +181,7 @@ def replace_state_dict(sd, asd, guess):
vae_key_prefix = guess.vae_key_prefix[0]
text_encoder_key_prefix = guess.text_encoder_key_prefix[0]
- if 'enc.blk.0.attn_k.weight' in asd:
+ if "enc.blk.0.attn_k.weight" in asd:
wierd_t5_format_from_city96 = {
"enc.": "encoder.",
".blk.": ".block.",
@@ -197,7 +198,7 @@ def replace_state_dict(sd, asd, guess):
"ffn_gate": "layer.1.DenseReluDense.wi_0",
"ffn_norm": "layer.1.layer_norm",
}
- wierd_t5_pre_quant_keys_from_city96 = ['shared.weight']
+ wierd_t5_pre_quant_keys_from_city96 = ["shared.weight"]
asd_new = {}
for k, v in asd.items():
for s, d in wierd_t5_format_from_city96.items():
@@ -215,7 +216,6 @@ def replace_state_dict(sd, asd, guess):
for k, v in asd.items():
sd[vae_key_prefix + k] = v
-
## identify model type
flux_test_key = "model.diffusion_model.double_blocks.0.img_attn.norm.key_norm.scale"
sd3_test_key = "model.diffusion_model.final_layer.adaLN_modulation.1.bias"
@@ -229,7 +229,7 @@ def replace_state_dict(sd, asd, guess):
case 1024:
model_type = "sd2"
case 1280:
- model_type = "xlrf" # sdxl refiner model
+ model_type = "xlrf" # sdxl refiner model
case 2048:
model_type = "sdxl"
elif flux_test_key in sd:
@@ -239,36 +239,35 @@ def replace_state_dict(sd, asd, guess):
## prefixes used by various model types for CLIP-L
prefix_L = {
- "-" : None,
- "sd1" : "cond_stage_model.transformer.",
- "sd2" : None,
+ "-": None,
+ "sd1": "cond_stage_model.transformer.",
+ "sd2": None,
"xlrf": None,
"sdxl": "conditioner.embedders.0.transformer.",
"flux": "text_encoders.clip_l.transformer.",
- "sd3" : "text_encoders.clip_l.transformer.",
+ "sd3": "text_encoders.clip_l.transformer.",
}
## prefixes used by various model types for CLIP-G
prefix_G = {
- "-" : None,
- "sd1" : None,
- "sd2" : None,
+ "-": None,
+ "sd1": None,
+ "sd2": None,
"xlrf": "conditioner.embedders.0.model.transformer.",
"sdxl": "conditioner.embedders.1.model.transformer.",
"flux": None,
- "sd3" : "text_encoders.clip_g.transformer.",
+ "sd3": "text_encoders.clip_g.transformer.",
}
## prefixes used by various model types for CLIP-H
prefix_H = {
- "-" : None,
- "sd1" : None,
- "sd2" : "conditioner.embedders.0.model.",
+ "-": None,
+ "sd1": None,
+ "sd2": "conditioner.embedders.0.model.",
"xlrf": None,
"sdxl": None,
"flux": None,
- "sd3" : None,
+ "sd3": None,
}
-
## VAE format 0 (extracted from model, could be sd1, sd2, sdxl, sd3).
if "first_stage_model.decoder.conv_in.weight" in asd:
channels = asd["first_stage_model.decoder.conv_in.weight"].shape[1]
@@ -282,10 +281,10 @@ def replace_state_dict(sd, asd, guess):
sd[k] = v
## CLIP-H
- CLIP_H = { # key to identify source model old_prefix
- 'cond_stage_model.model.ln_final.weight' : 'cond_stage_model.model.',
-# 'text_model.encoder.layers.0.layer_norm1.bias' : 'text_model'. # would need converting
- }
+ CLIP_H = { # key to identify source model old_prefix
+ "cond_stage_model.model.ln_final.weight": "cond_stage_model.model.",
+ # 'text_model.encoder.layers.0.layer_norm1.bias' : 'text_model'. # would need converting
+ }
for CLIP_key in CLIP_H.keys():
if CLIP_key in asd and asd[CLIP_key].shape[0] == 1024:
new_prefix = prefix_H[model_type]
@@ -297,36 +296,32 @@ def replace_state_dict(sd, asd, guess):
sd[new_k] = v
## CLIP-G
- CLIP_G = { # key to identify source model old_prefix
- 'conditioner.embedders.1.model.transformer.resblocks.0.ln_1.bias' : 'conditioner.embedders.1.model.transformer.',
- 'text_encoders.clip_g.transformer.text_model.encoder.layers.0.layer_norm1.bias' : 'text_encoders.clip_g.transformer.',
- 'text_model.encoder.layers.0.layer_norm1.bias' : '',
- 'transformer.resblocks.0.ln_1.bias' : 'transformer.'
- }
+ CLIP_G = {"conditioner.embedders.1.model.transformer.resblocks.0.ln_1.bias": "conditioner.embedders.1.model.transformer.", "text_encoders.clip_g.transformer.text_model.encoder.layers.0.layer_norm1.bias": "text_encoders.clip_g.transformer.", "text_model.encoder.layers.0.layer_norm1.bias": "", "transformer.resblocks.0.ln_1.bias": "transformer."} # key to identify source model old_prefix
for CLIP_key in CLIP_G.keys():
if CLIP_key in asd and asd[CLIP_key].shape[0] == 1280:
new_prefix = prefix_G[model_type]
old_prefix = CLIP_G[CLIP_key]
if new_prefix is not None:
- if "resblocks" not in CLIP_key and model_type != "sd3": # need to convert
+ if "resblocks" not in CLIP_key and model_type != "sd3": # need to convert
+
def convert_transformers(statedict, prefix_from, prefix_to, number):
keys_to_replace = {
- "{}text_model.embeddings.position_embedding.weight" : "{}positional_embedding",
- "{}text_model.embeddings.token_embedding.weight" : "{}token_embedding.weight",
- "{}text_model.final_layer_norm.weight" : "{}ln_final.weight",
- "{}text_model.final_layer_norm.bias" : "{}ln_final.bias",
- "text_projection.weight" : "{}text_projection",
+ "{}text_model.embeddings.position_embedding.weight": "{}positional_embedding",
+ "{}text_model.embeddings.token_embedding.weight": "{}token_embedding.weight",
+ "{}text_model.final_layer_norm.weight": "{}ln_final.weight",
+ "{}text_model.final_layer_norm.bias": "{}ln_final.bias",
+ "text_projection.weight": "{}text_projection",
}
resblock_to_replace = {
- "layer_norm1" : "ln_1",
- "layer_norm2" : "ln_2",
- "mlp.fc1" : "mlp.c_fc",
- "mlp.fc2" : "mlp.c_proj",
- "self_attn.out_proj" : "attn.out_proj" ,
+ "layer_norm1": "ln_1",
+ "layer_norm2": "ln_2",
+ "mlp.fc1": "mlp.c_fc",
+ "mlp.fc2": "mlp.c_proj",
+ "self_attn.out_proj": "attn.out_proj",
}
- for x in keys_to_replace: # remove trailing 'transformer.' from new prefix
+ for x in keys_to_replace: # remove trailing 'transformer.' from new prefix
k = x.format(prefix_from)
statedict[keys_to_replace[x].format(prefix_to[:-12])] = statedict.pop(k)
@@ -363,13 +358,7 @@ def replace_state_dict(sd, asd, guess):
sd[new_k] = v
## CLIP-L
- CLIP_L = { # key to identify source model old_prefix
- 'cond_stage_model.transformer.text_model.encoder.layers.0.layer_norm1.bias' : 'cond_stage_model.transformer.',
- 'conditioner.embedders.0.transformer.text_model.encoder.layers.0.layer_norm1.bias' : 'conditioner.embedders.0.transformer.',
- 'text_encoders.clip_l.transformer.text_model.encoder.layers.0.layer_norm1.bias' : 'text_encoders.clip_l.transformer.',
- 'text_model.encoder.layers.0.layer_norm1.bias' : '',
- 'transformer.resblocks.0.ln_1.bias' : 'transformer.'
- }
+ CLIP_L = {"cond_stage_model.transformer.text_model.encoder.layers.0.layer_norm1.bias": "cond_stage_model.transformer.", "conditioner.embedders.0.transformer.text_model.encoder.layers.0.layer_norm1.bias": "conditioner.embedders.0.transformer.", "text_encoders.clip_l.transformer.text_model.encoder.layers.0.layer_norm1.bias": "text_encoders.clip_l.transformer.", "text_model.encoder.layers.0.layer_norm1.bias": "", "transformer.resblocks.0.ln_1.bias": "transformer."} # key to identify source model old_prefix
for CLIP_key in CLIP_L.keys():
if CLIP_key in asd and asd[CLIP_key].shape[0] == 768:
@@ -377,21 +366,22 @@ def replace_state_dict(sd, asd, guess):
old_prefix = CLIP_L[CLIP_key]
if new_prefix is not None:
- if "resblocks" in CLIP_key: # need to convert
+ if "resblocks" in CLIP_key: # need to convert
+
def transformers_convert(statedict, prefix_from, prefix_to, number):
keys_to_replace = {
- "positional_embedding" : "{}text_model.embeddings.position_embedding.weight",
+ "positional_embedding": "{}text_model.embeddings.position_embedding.weight",
"token_embedding.weight": "{}text_model.embeddings.token_embedding.weight",
- "ln_final.weight" : "{}text_model.final_layer_norm.weight",
- "ln_final.bias" : "{}text_model.final_layer_norm.bias",
- "text_projection" : "text_projection.weight",
+ "ln_final.weight": "{}text_model.final_layer_norm.weight",
+ "ln_final.bias": "{}text_model.final_layer_norm.bias",
+ "text_projection": "text_projection.weight",
}
resblock_to_replace = {
- "ln_1" : "layer_norm1",
- "ln_2" : "layer_norm2",
- "mlp.c_fc" : "mlp.fc1",
- "mlp.c_proj" : "mlp.fc2",
- "attn.out_proj" : "self_attn.out_proj",
+ "ln_1": "layer_norm1",
+ "ln_2": "layer_norm2",
+ "mlp.c_fc": "mlp.fc1",
+ "mlp.c_proj": "mlp.fc2",
+ "attn.out_proj": "self_attn.out_proj",
}
for k in keys_to_replace:
@@ -410,7 +400,7 @@ def replace_state_dict(sd, asd, guess):
for x in range(3):
p = ["self_attn.q_proj", "self_attn.k_proj", "self_attn.v_proj"]
k_to = "{}text_model.encoder.layers.{}.{}.{}".format(prefix_to, resblock, p[x], y)
- statedict[k_to] = weights[shape_from*x:shape_from*(x + 1)]
+ statedict[k_to] = weights[shape_from * x : shape_from * (x + 1)]
return statedict
asd = transformers_convert(asd, old_prefix, new_prefix, 12)
@@ -426,8 +416,7 @@ def replace_state_dict(sd, asd, guess):
new_k = k.replace(old_prefix, new_prefix)
sd[new_k] = v
-
- if 'encoder.block.0.layer.0.SelfAttention.k.weight' in asd:
+ if "encoder.block.0.layer.0.SelfAttention.k.weight" in asd:
keys_to_delete = [k for k in sd if k.startswith(f"{text_encoder_key_prefix}t5xxl.")]
for k in keys_to_delete:
del sd[k]
@@ -457,26 +446,23 @@ def split_state_dict(sd, additional_state_dicts: list = None):
guess.clip_target = guess.clip_target(sd)
guess.model_type = guess.model_type(sd)
- guess.ztsnr = 'ztsnr' in sd
+ guess.ztsnr = "ztsnr" in sd
sd = guess.process_vae_state_dict(sd)
- state_dict = {
- guess.unet_target: try_filter_state_dict(sd, guess.unet_key_prefix),
- guess.vae_target: try_filter_state_dict(sd, guess.vae_key_prefix)
- }
+ state_dict = {guess.unet_target: try_filter_state_dict(sd, guess.unet_key_prefix), guess.vae_target: try_filter_state_dict(sd, guess.vae_key_prefix)}
sd = guess.process_clip_state_dict(sd)
for k, v in guess.clip_target.items():
- state_dict[v] = try_filter_state_dict(sd, [k + '.'])
+ state_dict[v] = try_filter_state_dict(sd, [k + "."])
- state_dict['ignore'] = sd
+ state_dict["ignore"] = sd
print_dict = {k: len(v) for k, v in state_dict.items()}
- print(f'StateDict Keys: {print_dict}')
+ print(f"StateDict Keys: {print_dict}")
- del state_dict['ignore']
+ del state_dict["ignore"]
return state_dict, guess
@@ -486,11 +472,11 @@ def forge_loader(sd, additional_state_dicts=None):
try:
state_dicts, estimated_config = split_state_dict(sd, additional_state_dicts=additional_state_dicts)
except:
- raise ValueError('Failed to recognize model type!')
+ raise ValueError("Failed to recognize model type!")
repo_name = estimated_config.huggingface_repo
- local_path = os.path.join(dir_path, 'huggingface', repo_name)
+ local_path = os.path.join(dir_path, "huggingface", repo_name)
config: dict = DiffusionPipeline.load_config(local_path)
huggingface_components = {}
for component_name, v in config.items():
@@ -507,44 +493,43 @@ def forge_loader(sd, additional_state_dicts=None):
yaml_config_prediction_type = None
try:
- import yaml
from pathlib import Path
- config_filename = os.path.splitext(sd)[0] + '.yaml'
+
+ import yaml
+
+ config_filename = os.path.splitext(sd)[0] + ".yaml"
if Path(config_filename).is_file():
- with open(config_filename, 'r') as stream:
+ with open(config_filename, "r") as stream:
yaml_config = yaml.safe_load(stream)
except ImportError:
pass
# Fix Huggingface prediction type using .yaml config or estimated config detection
prediction_types = {
- 'EPS': 'epsilon',
- 'V_PREDICTION': 'v_prediction',
- 'EDM': 'edm',
+ "EPS": "epsilon",
+ "V_PREDICTION": "v_prediction",
+ "EDM": "edm",
}
- has_prediction_type = 'scheduler' in huggingface_components and hasattr(huggingface_components['scheduler'], 'config') and 'prediction_type' in huggingface_components['scheduler'].config
+ has_prediction_type = "scheduler" in huggingface_components and hasattr(huggingface_components["scheduler"], "config") and "prediction_type" in huggingface_components["scheduler"].config
if yaml_config is not None:
- yaml_config_prediction_type: str = (
- yaml_config.get('model', {}).get('params', {}).get('parameterization', '')
- or yaml_config.get('model', {}).get('params', {}).get('denoiser_config', {}).get('params', {}).get('scaling_config', {}).get('target', '')
- )
- if yaml_config_prediction_type == 'v' or yaml_config_prediction_type.endswith(".VScaling"):
- yaml_config_prediction_type = 'v_prediction'
+ yaml_config_prediction_type: str = yaml_config.get("model", {}).get("params", {}).get("parameterization", "") or yaml_config.get("model", {}).get("params", {}).get("denoiser_config", {}).get("params", {}).get("scaling_config", {}).get("target", "")
+ if yaml_config_prediction_type == "v" or yaml_config_prediction_type.endswith(".VScaling"):
+ yaml_config_prediction_type = "v_prediction"
else:
# Use estimated prediction config if no suitable prediction type found
- yaml_config_prediction_type = ''
+ yaml_config_prediction_type = ""
if has_prediction_type:
if yaml_config_prediction_type:
- huggingface_components['scheduler'].config.prediction_type = yaml_config_prediction_type
+ huggingface_components["scheduler"].config.prediction_type = yaml_config_prediction_type
else:
- huggingface_components['scheduler'].config.prediction_type = prediction_types.get(estimated_config.model_type.name, huggingface_components['scheduler'].config.prediction_type)
+ huggingface_components["scheduler"].config.prediction_type = prediction_types.get(estimated_config.model_type.name, huggingface_components["scheduler"].config.prediction_type)
for M in possible_models:
if any(isinstance(estimated_config, x) for x in M.matched_guesses):
return M(estimated_config=estimated_config, huggingface_components=huggingface_components)
- print('Failed to recognize model type!')
+ print("Failed to recognize model type!")
return None
diff --git a/backend/memory_management.py b/backend/memory_management.py
index 5f0c8312..7c205066 100644
--- a/backend/memory_management.py
+++ b/backend/memory_management.py
@@ -1,17 +1,17 @@
# Cherry-picked some good parts from ComfyUI with some bad parts fixed
+import platform
import sys
import time
+from enum import Enum
+
import psutil
import torch
-import platform
-from enum import Enum
from backend import stream, utils
from backend.args import args
-
-cpu = torch.device('cpu')
+cpu = torch.device("cpu")
class VRAMState(Enum):
@@ -105,7 +105,7 @@ def get_total_memory(dev=None, torch_total_too=False):
if dev is None:
dev = get_torch_device()
- if hasattr(dev, 'type') and (dev.type == 'cpu' or dev.type == 'mps'):
+ if hasattr(dev, "type") and (dev.type == "cpu" or dev.type == "mps"):
mem_total = psutil.virtual_memory().total
mem_total_torch = mem_total
else:
@@ -114,12 +114,12 @@ def get_total_memory(dev=None, torch_total_too=False):
mem_total_torch = mem_total
elif is_intel_xpu():
stats = torch.xpu.memory_stats(dev)
- mem_reserved = stats['reserved_bytes.all.current']
+ mem_reserved = stats["reserved_bytes.all.current"]
mem_total_torch = mem_reserved
mem_total = torch.xpu.get_device_properties(dev).total_memory
else:
stats = torch.cuda.memory_stats(dev)
- mem_reserved = stats['reserved_bytes.all.current']
+ mem_reserved = stats["reserved_bytes.all.current"]
_, mem_total_cuda = torch.cuda.mem_get_info(dev)
mem_total_torch = mem_reserved
mem_total = mem_total_cuda
@@ -258,7 +258,7 @@ if PIN_SHARED_MEMORY:
def get_torch_device_name(device):
- if hasattr(device, 'type'):
+ if hasattr(device, "type"):
if device.type == "cuda":
try:
allocator_backend = torch.cuda.get_allocator_backend()
@@ -277,12 +277,12 @@ try:
torch_device_name = get_torch_device_name(get_torch_device())
print("Device: {}".format(torch_device_name))
except:
- torch_device_name = ''
+ torch_device_name = ""
print("Could not pick default device.")
-if 'rtx' in torch_device_name.lower():
+if "rtx" in torch_device_name.lower():
if not args.cuda_malloc:
- print('Hint: your device supports --cuda-malloc for potential speed improvements.')
+ print("Hint: your device supports --cuda-malloc for potential speed improvements.")
current_loaded_models = []
@@ -310,12 +310,12 @@ def state_dict_parameters(sd):
def state_dict_dtype(state_dict):
for k, v in state_dict.items():
- if hasattr(v, 'gguf_cls'):
- return 'gguf'
- if 'bitsandbytes__nf4' in k:
- return 'nf4'
- if 'bitsandbytes__fp4' in k:
- return 'fp4'
+ if hasattr(v, "gguf_cls"):
+ return "gguf"
+ if "bitsandbytes__nf4" in k:
+ return "nf4"
+ if "bitsandbytes__fp4" in k:
+ return "fp4"
dtype_counts = {}
@@ -338,11 +338,11 @@ def state_dict_dtype(state_dict):
def bake_gguf_model(model):
- if getattr(model, 'gguf_baked', False):
+ if getattr(model, "gguf_baked", False):
return
for p in model.parameters():
- gguf_cls = getattr(p, 'gguf_cls', None)
+ gguf_cls = getattr(p, "gguf_cls", None)
if gguf_cls is not None:
gguf_cls.bake(p)
@@ -356,7 +356,7 @@ def bake_gguf_model(model):
def module_size(module, exclude_device=None, include_device=None, return_split=False):
module_mem = 0
weight_mem = 0
- weight_patterns = ['weight']
+ weight_patterns = ["weight"]
for k, p in module.named_parameters():
t = p.data
@@ -371,7 +371,7 @@ def module_size(module, exclude_device=None, include_device=None, return_split=F
element_size = t.element_size()
- if getattr(p, 'quant_type', None) in ['fp4', 'nf4']:
+ if getattr(p, "quant_type", None) in ["fp4", "nf4"]:
if element_size > 1:
# not quanted yet
element_size = 0.55 # a bit more than 0.5 because of quant state parameters
@@ -472,7 +472,7 @@ class LoadedModel:
raise e
if do_not_need_cpu_swap:
- print('All loaded to GPU.')
+ print("All loaded to GPU.")
else:
gpu_modules, gpu_modules_only_extras, cpu_modules = build_module_profile(self.real_model, model_gpu_memory_when_using_cpu_swap)
pin_memory = PIN_SHARED_MEMORY and is_device_cpu(self.model.offload_device)
@@ -495,8 +495,8 @@ class LoadedModel:
for m in gpu_modules_only_extras:
m.prev_parameters_manual_cast = m.parameters_manual_cast
m.parameters_manual_cast = True
- module_move(m, device=self.device, recursive=False, excluded_pattens=['weight'])
- if hasattr(m, 'weight') and m.weight is not None:
+ module_move(m, device=self.device, recursive=False, excluded_pattens=["weight"])
+ if hasattr(m, "weight") and m.weight is not None:
if pin_memory:
m.weight = utils.tensor2parameter(m.weight.to(self.model.offload_device).pin_memory())
else:
@@ -504,8 +504,8 @@ class LoadedModel:
mem_counter += m.extra_mem
swap_counter += m.weight_mem
- swap_flag = 'Shared' if PIN_SHARED_MEMORY else 'CPU'
- method_flag = 'asynchronous' if stream.should_use_stream() else 'blocked'
+ swap_flag = "Shared" if PIN_SHARED_MEMORY else "CPU"
+ method_flag = "asynchronous" if stream.should_use_stream() else "blocked"
print(f"{swap_flag} Swap Loaded ({method_flag} method): {swap_counter / (1024 * 1024):.2f} MB, GPU Loaded: {mem_counter / (1024 * 1024):.2f} MB")
self.model_accelerated = True
@@ -596,17 +596,14 @@ def free_memory(memory_required, device, keep_loaded=[], free_all=False):
if mem_free_torch > mem_free_total * 0.25:
soft_empty_cache()
- print('Done.')
+ print("Done.")
return
def compute_model_gpu_memory_when_using_cpu_swap(current_free_mem, inference_memory):
maximum_memory_available = current_free_mem - inference_memory
- suggestion = max(
- maximum_memory_available / 1.3,
- maximum_memory_available - 1024 * 1024 * 1024 * 1.25
- )
+ suggestion = max(maximum_memory_available / 1.3, maximum_memory_available - 1024 * 1024 * 1024 * 1.25)
return int(max(0, suggestion))
@@ -638,7 +635,7 @@ def load_models_gpu(models, memory_required=0, hard_memory_preservation=0):
moving_time = time.perf_counter() - execution_start_time
if moving_time > 0.1:
- print(f'Memory cleanup has taken {moving_time:.2f} seconds')
+ print(f"Memory cleanup has taken {moving_time:.2f} seconds")
return
@@ -685,7 +682,7 @@ def load_models_gpu(models, memory_required=0, hard_memory_preservation=0):
current_loaded_models.insert(0, loaded_model)
moving_time = time.perf_counter() - execution_start_time
- print(f'Moving model(s) has taken {moving_time:.2f} seconds')
+ print(f"Moving model(s) has taken {moving_time:.2f} seconds")
return
@@ -859,7 +856,7 @@ print(f"VAE dtype preferences: {VAE_DTYPES} -> {vae_dtype()}")
def get_autocast_device(dev):
- if hasattr(dev, 'type'):
+ if hasattr(dev, "type"):
return dev.type
return "cuda"
@@ -938,7 +935,7 @@ def cast_to_device(tensor, device, dtype, copy=False):
if tensor.dtype == torch.float32 or tensor.dtype == torch.float16:
device_supports_cast = True
elif tensor.dtype == torch.bfloat16:
- if hasattr(device, 'type') and device.type.startswith("cuda"):
+ if hasattr(device, "type") and device.type.startswith("cuda"):
device_supports_cast = True
elif is_intel_xpu():
device_supports_cast = True
@@ -995,7 +992,7 @@ def pytorch_attention_flash_attention():
def force_upcast_attention_dtype():
upcast = args.force_upcast_attention
try:
- if platform.mac_ver()[0] in ['14.5']: # black image bug on OSX Sonoma 14.5
+ if platform.mac_ver()[0] in ["14.5"]: # black image bug on OSX Sonoma 14.5
upcast = True
except:
pass
@@ -1010,7 +1007,7 @@ def get_free_memory(dev=None, torch_free_too=False):
if dev is None:
dev = get_torch_device()
- if hasattr(dev, 'type') and (dev.type == 'cpu' or dev.type == 'mps'):
+ if hasattr(dev, "type") and (dev.type == "cpu" or dev.type == "mps"):
mem_free_total = psutil.virtual_memory().available
mem_free_torch = mem_free_total
else:
@@ -1019,15 +1016,15 @@ def get_free_memory(dev=None, torch_free_too=False):
mem_free_torch = mem_free_total
elif is_intel_xpu():
stats = torch.xpu.memory_stats(dev)
- mem_active = stats['active_bytes.all.current']
- mem_reserved = stats['reserved_bytes.all.current']
+ mem_active = stats["active_bytes.all.current"]
+ mem_reserved = stats["reserved_bytes.all.current"]
mem_free_torch = mem_reserved - mem_active
mem_free_xpu = torch.xpu.get_device_properties(dev).total_memory - mem_reserved
mem_free_total = mem_free_xpu + mem_free_torch
else:
stats = torch.cuda.memory_stats(dev)
- mem_active = stats['active_bytes.all.current']
- mem_reserved = stats['reserved_bytes.all.current']
+ mem_active = stats["active_bytes.all.current"]
+ mem_reserved = stats["reserved_bytes.all.current"]
mem_free_cuda, _ = torch.cuda.mem_get_info(dev)
mem_free_torch = mem_reserved - mem_active
mem_free_total = mem_free_cuda + mem_free_torch
@@ -1049,22 +1046,22 @@ def mps_mode():
def is_device_type(device, type):
- if hasattr(device, 'type'):
- if (device.type == type):
+ if hasattr(device, "type"):
+ if device.type == type:
return True
return False
def is_device_cpu(device):
- return is_device_type(device, 'cpu')
+ return is_device_type(device, "cpu")
def is_device_mps(device):
- return is_device_type(device, 'mps')
+ return is_device_type(device, "mps")
def is_device_cuda(device):
- return is_device_type(device, 'cuda')
+ return is_device_type(device, "cuda")
def should_use_fp16(device=None, model_params=0, prioritize_performance=True, manual_cast=False):
@@ -1111,7 +1108,7 @@ def should_use_fp16(device=None, model_params=0, prioritize_performance=True, ma
if x in props.name.lower():
if manual_cast:
# For storage dtype
- free_model_memory = (get_free_memory() * 0.9 - minimum_inference_memory())
+ free_model_memory = get_free_memory() * 0.9 - minimum_inference_memory()
if (not prioritize_performance) or model_params * 4 > free_model_memory:
return True
else:
@@ -1166,7 +1163,7 @@ def should_use_bf16(device=None, model_params=0, prioritize_performance=True, ma
# So in this case bf16 should only be used as storge dtype
if manual_cast:
# For storage dtype
- free_model_memory = (get_free_memory() * 0.9 - minimum_inference_memory())
+ free_model_memory = get_free_memory() * 0.9 - minimum_inference_memory()
if (not prioritize_performance) or model_params * 4 > free_model_memory:
return True
@@ -1178,7 +1175,7 @@ def can_install_bnb():
if not torch.cuda.is_available():
return False
- cuda_version = tuple(int(x) for x in torch.version.cuda.split('.'))
+ cuda_version = tuple(int(x) for x in torch.version.cuda.split("."))
if cuda_version >= (11, 7):
return True
diff --git a/backend/nn/cnets/cldm.py b/backend/nn/cnets/cldm.py
index c177bce5..d695f928 100644
--- a/backend/nn/cnets/cldm.py
+++ b/backend/nn/cnets/cldm.py
@@ -1,6 +1,14 @@
import torch.nn as nn
-from backend.nn.unet import Downsample, ResBlock, SpatialTransformer, TimestepEmbedSequential, conv_nd, exists, timestep_embedding
+from backend.nn.unet import (
+ Downsample,
+ ResBlock,
+ SpatialTransformer,
+ TimestepEmbedSequential,
+ conv_nd,
+ exists,
+ timestep_embedding,
+)
class ControlNet(nn.Module):
diff --git a/backend/operations.py b/backend/operations.py
index 7760756a..092e90e9 100644
--- a/backend/operations.py
+++ b/backend/operations.py
@@ -1,40 +1,40 @@
# Copyright Forge 2024
-import time
-import torch
import contextlib
+import time
-from backend import stream, memory_management, utils
+import torch
+
+from backend import memory_management, stream, utils
from backend.patcher.lora import merge_lora_to_weight
-
stash = {}
def get_weight_and_bias(layer, weight_args=None, bias_args=None, weight_fn=None, bias_fn=None):
- scale_weight = getattr(layer, 'scale_weight', None)
- patches = getattr(layer, 'forge_online_loras', None)
+ scale_weight = getattr(layer, "scale_weight", None)
+ patches = getattr(layer, "forge_online_loras", None)
weight_patches, bias_patches = None, None
if patches is not None:
- weight_patches = patches.get('weight', None)
+ weight_patches = patches.get("weight", None)
if patches is not None:
- bias_patches = patches.get('bias', None)
+ bias_patches = patches.get("bias", None)
weight = None
if layer.weight is not None:
weight = layer.weight
if weight_fn is not None:
if weight_args is not None:
- fn_device = weight_args.get('device', None)
+ fn_device = weight_args.get("device", None)
if fn_device is not None:
weight = weight.to(device=fn_device)
weight = weight_fn(weight)
if weight_args is not None:
weight = weight.to(**weight_args)
if scale_weight is not None:
- weight = weight*scale_weight.to(device=weight.device, dtype=weight.dtype)
+ weight = weight * scale_weight.to(device=weight.device, dtype=weight.dtype)
if weight_patches is not None:
weight = merge_lora_to_weight(patches=weight_patches, weight=weight, key="online weight lora", computation_dtype=weight.dtype)
@@ -43,7 +43,7 @@ def get_weight_and_bias(layer, weight_args=None, bias_args=None, weight_fn=None,
bias = layer.bias
if bias_fn is not None:
if bias_args is not None:
- fn_device = bias_args.get('device', None)
+ fn_device = bias_args.get("device", None)
if fn_device is not None:
bias = bias.to(device=fn_device)
bias = bias_fn(bias)
@@ -58,7 +58,7 @@ def weights_manual_cast(layer, x, skip_weight_dtype=False, skip_bias_dtype=False
weight, bias, signal = None, None, None
non_blocking = True
- if getattr(x.device, 'type', None) == 'mps':
+ if getattr(x.device, "type", None) == "mps":
non_blocking = False
target_dtype = x.dtype
@@ -135,13 +135,13 @@ class ForgeOperations:
self.parameters_manual_cast = current_manual_cast_enabled
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):
- if hasattr(self, 'dummy'):
- if prefix + 'weight' in state_dict:
- self.weight = torch.nn.Parameter(state_dict[prefix + 'weight'].to(self.dummy))
- if prefix + 'scale_weight' in state_dict:
- self.scale_weight = torch.nn.Parameter(state_dict[prefix + 'scale_weight'])
- if prefix + 'bias' in state_dict:
- self.bias = torch.nn.Parameter(state_dict[prefix + 'bias'].to(self.dummy))
+ if hasattr(self, "dummy"):
+ if prefix + "weight" in state_dict:
+ self.weight = torch.nn.Parameter(state_dict[prefix + "weight"].to(self.dummy))
+ if prefix + "scale_weight" in state_dict:
+ self.scale_weight = torch.nn.Parameter(state_dict[prefix + "scale_weight"])
+ if prefix + "bias" in state_dict:
+ self.bias = torch.nn.Parameter(state_dict[prefix + "bias"].to(self.dummy))
del self.dummy
else:
super()._load_from_state_dict(state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs)
@@ -158,8 +158,8 @@ class ForgeOperations:
class Conv2d(torch.nn.Conv2d):
def __init__(self, *args, **kwargs):
- kwargs['device'] = current_device
- kwargs['dtype'] = current_dtype
+ kwargs["device"] = current_device
+ kwargs["dtype"] = current_dtype
super().__init__(*args, **kwargs)
self.parameters_manual_cast = current_manual_cast_enabled
@@ -178,8 +178,8 @@ class ForgeOperations:
class Conv3d(torch.nn.Conv3d):
def __init__(self, *args, **kwargs):
- kwargs['device'] = current_device
- kwargs['dtype'] = current_dtype
+ kwargs["device"] = current_device
+ kwargs["dtype"] = current_dtype
super().__init__(*args, **kwargs)
self.parameters_manual_cast = current_manual_cast_enabled
@@ -198,8 +198,8 @@ class ForgeOperations:
class Conv1d(torch.nn.Conv1d):
def __init__(self, *args, **kwargs):
- kwargs['device'] = current_device
- kwargs['dtype'] = current_dtype
+ kwargs["device"] = current_device
+ kwargs["dtype"] = current_dtype
super().__init__(*args, **kwargs)
self.parameters_manual_cast = current_manual_cast_enabled
@@ -218,8 +218,8 @@ class ForgeOperations:
class ConvTranspose2d(torch.nn.ConvTranspose2d):
def __init__(self, *args, **kwargs):
- kwargs['device'] = current_device
- kwargs['dtype'] = current_dtype
+ kwargs["device"] = current_device
+ kwargs["dtype"] = current_dtype
super().__init__(*args, **kwargs)
self.parameters_manual_cast = current_manual_cast_enabled
@@ -243,8 +243,8 @@ class ForgeOperations:
class ConvTranspose1d(torch.nn.ConvTranspose1d):
def __init__(self, *args, **kwargs):
- kwargs['device'] = current_device
- kwargs['dtype'] = current_dtype
+ kwargs["device"] = current_device
+ kwargs["dtype"] = current_dtype
super().__init__(*args, **kwargs)
self.parameters_manual_cast = current_manual_cast_enabled
@@ -268,8 +268,8 @@ class ForgeOperations:
class ConvTranspose3d(torch.nn.ConvTranspose3d):
def __init__(self, *args, **kwargs):
- kwargs['device'] = current_device
- kwargs['dtype'] = current_dtype
+ kwargs["device"] = current_device
+ kwargs["dtype"] = current_dtype
super().__init__(*args, **kwargs)
self.parameters_manual_cast = current_manual_cast_enabled
@@ -293,8 +293,8 @@ class ForgeOperations:
class GroupNorm(torch.nn.GroupNorm):
def __init__(self, *args, **kwargs):
- kwargs['device'] = current_device
- kwargs['dtype'] = current_dtype
+ kwargs["device"] = current_device
+ kwargs["dtype"] = current_dtype
super().__init__(*args, **kwargs)
self.parameters_manual_cast = current_manual_cast_enabled
@@ -312,8 +312,8 @@ class ForgeOperations:
class LayerNorm(torch.nn.LayerNorm):
def __init__(self, *args, **kwargs):
- kwargs['device'] = current_device
- kwargs['dtype'] = current_dtype
+ kwargs["device"] = current_device
+ kwargs["dtype"] = current_dtype
super().__init__(*args, **kwargs)
self.parameters_manual_cast = current_manual_cast_enabled
@@ -331,7 +331,7 @@ class ForgeOperations:
class Embedding(torch.nn.Embedding):
def __init__(self, *args, **kwargs):
- kwargs['device'] = current_device
+ kwargs["device"] = current_device
super().__init__(*args, **kwargs)
self.parameters_manual_cast = current_manual_cast_enabled
self.bias = None
@@ -350,7 +350,12 @@ class ForgeOperations:
try:
- from backend.operations_bnb import ForgeLoader4Bit, ForgeParams4bit, functional_linear_4bits, functional_dequantize_4bit
+ from backend.operations_bnb import (
+ ForgeLoader4Bit,
+ ForgeParams4bit,
+ functional_dequantize_4bit,
+ functional_linear_4bits,
+ )
class ForgeOperationsBNB4bits(ForgeOperations):
class Linear(ForgeLoader4Bit):
@@ -364,7 +369,7 @@ try:
# And it only invokes one time, and most linear does not have bias
self.bias = utils.tensor2parameter(self.bias.to(x.dtype))
- if hasattr(self, 'forge_online_loras'):
+ if hasattr(self, "forge_online_loras"):
weight, bias, signal = weights_manual_cast(self, x, weight_fn=functional_dequantize_4bit, bias_fn=None, skip_bias_dtype=True)
with main_stream_worker(weight, bias, signal):
return torch.nn.functional.linear(x, weight, bias)
@@ -372,7 +377,7 @@ try:
if not self.parameters_manual_cast:
return functional_linear_4bits(x, self.weight, self.bias)
elif not self.weight.bnb_quantized:
- assert x.device.type == 'cuda', 'BNB Must Use CUDA as Computation Device!'
+ assert x.device.type == "cuda", "BNB Must Use CUDA as Computation Device!"
layer_original_device = self.weight.device
self.weight = self.weight._quantize(x.device)
bias = self.bias.to(x.device) if self.bias is not None else None
@@ -402,23 +407,23 @@ class ForgeOperationsGGUF(ForgeOperations):
self.parameters_manual_cast = current_manual_cast_enabled
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):
- if hasattr(self, 'dummy'):
+ if hasattr(self, "dummy"):
computation_dtype = self.dummy.dtype
if computation_dtype not in [torch.float16, torch.bfloat16]:
# GGUF cast only supports 16bits otherwise super slow
computation_dtype = torch.float16
- if prefix + 'weight' in state_dict:
- self.weight = state_dict[prefix + 'weight'].to(device=self.dummy.device)
+ if prefix + "weight" in state_dict:
+ self.weight = state_dict[prefix + "weight"].to(device=self.dummy.device)
self.weight.computation_dtype = computation_dtype
- if prefix + 'bias' in state_dict:
- self.bias = state_dict[prefix + 'bias'].to(device=self.dummy.device)
+ if prefix + "bias" in state_dict:
+ self.bias = state_dict[prefix + "bias"].to(device=self.dummy.device)
self.bias.computation_dtype = computation_dtype
del self.dummy
else:
- if prefix + 'weight' in state_dict:
- self.weight = state_dict[prefix + 'weight']
- if prefix + 'bias' in state_dict:
- self.bias = state_dict[prefix + 'bias']
+ if prefix + "weight" in state_dict:
+ self.weight = state_dict[prefix + "weight"]
+ if prefix + "bias" in state_dict:
+ self.bias = state_dict[prefix + "bias"]
return
def _apply(self, fn, recurse=True):
@@ -430,7 +435,7 @@ class ForgeOperationsGGUF(ForgeOperations):
if self.bias is not None and self.bias.dtype != x.dtype:
self.bias = utils.tensor2parameter(dequantize_tensor(self.bias).to(x.dtype))
- if self.weight is not None and self.weight.dtype != x.dtype and getattr(self.weight, 'gguf_cls', None) is None:
+ if self.weight is not None and self.weight.dtype != x.dtype and getattr(self.weight, "gguf_cls", None) is None:
self.weight = utils.tensor2parameter(self.weight.to(x.dtype))
weight, bias, signal = weights_manual_cast(self, x, weight_fn=dequantize_tensor, bias_fn=None, skip_bias_dtype=True)
@@ -445,14 +450,14 @@ def using_forge_operations(operations=None, device=None, dtype=None, manual_cast
current_device, current_dtype, current_manual_cast_enabled, current_bnb_dtype = device, dtype, manual_cast_enabled, bnb_dtype
if operations is None:
- if bnb_dtype in ['gguf']:
+ if bnb_dtype in ["gguf"]:
operations = ForgeOperationsGGUF
- elif bnb_avaliable and bnb_dtype in ['nf4', 'fp4']:
+ elif bnb_avaliable and bnb_dtype in ["nf4", "fp4"]:
operations = ForgeOperationsBNB4bits
else:
operations = ForgeOperations
- op_names = ['Linear', 'Conv1d', 'Conv2d', 'Conv3d', 'ConvTranspose1d', 'ConvTranspose2d', 'ConvTranspose3d', 'GroupNorm', 'LayerNorm', 'Embedding']
+ op_names = ["Linear", "Conv1d", "Conv2d", "Conv3d", "ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d", "GroupNorm", "LayerNorm", "Embedding"]
backups = {op_name: getattr(torch.nn, op_name) for op_name in op_names}
try:
@@ -469,17 +474,14 @@ def using_forge_operations(operations=None, device=None, dtype=None, manual_cast
def shift_manual_cast(model, enabled):
for m in model.modules():
- if hasattr(m, 'parameters_manual_cast'):
+ if hasattr(m, "parameters_manual_cast"):
m.parameters_manual_cast = enabled
return
@contextlib.contextmanager
def automatic_memory_management():
- memory_management.free_memory(
- memory_required=3 * 1024 * 1024 * 1024,
- device=memory_management.get_torch_device()
- )
+ memory_management.free_memory(memory_required=3 * 1024 * 1024 * 1024, device=memory_management.get_torch_device())
module_list = []
@@ -511,7 +513,7 @@ def automatic_memory_management():
memory_management.soft_empty_cache()
end = time.perf_counter()
- print(f'Automatic Memory Management: {len(module_list)} Modules in {(end - start):.2f} seconds.')
+ print(f"Automatic Memory Management: {len(module_list)} Modules in {(end - start):.2f} seconds.")
return
@@ -519,11 +521,11 @@ class DynamicSwapInstaller:
@staticmethod
def _install_module(module: torch.nn.Module, target_device: torch.device):
original_class = module.__class__
- module.__dict__['forge_backup_original_class'] = original_class
+ module.__dict__["forge_backup_original_class"] = original_class
def hacked_get_attr(self, name: str):
- if '_parameters' in self.__dict__:
- _parameters = self.__dict__['_parameters']
+ if "_parameters" in self.__dict__:
+ _parameters = self.__dict__["_parameters"]
if name in _parameters:
p = _parameters[name]
if p is None:
@@ -532,22 +534,26 @@ class DynamicSwapInstaller:
return torch.nn.Parameter(p.to(target_device), requires_grad=p.requires_grad)
else:
return p.to(target_device)
- if '_buffers' in self.__dict__:
- _buffers = self.__dict__['_buffers']
+ if "_buffers" in self.__dict__:
+ _buffers = self.__dict__["_buffers"]
if name in _buffers:
return _buffers[name].to(target_device)
return super(original_class, self).__getattr__(name)
- module.__class__ = type('DynamicSwap_' + original_class.__name__, (original_class,), {
- '__getattr__': hacked_get_attr,
- })
+ module.__class__ = type(
+ "DynamicSwap_" + original_class.__name__,
+ (original_class,),
+ {
+ "__getattr__": hacked_get_attr,
+ },
+ )
return
@staticmethod
def _uninstall_module(module: torch.nn.Module):
- if 'forge_backup_original_class' in module.__dict__:
- module.__class__ = module.__dict__.pop('forge_backup_original_class')
+ if "forge_backup_original_class" in module.__dict__:
+ module.__class__ = module.__dict__.pop("forge_backup_original_class")
return
@staticmethod
@@ -561,4 +567,3 @@ class DynamicSwapInstaller:
for m in model.modules():
DynamicSwapInstaller._uninstall_module(m)
return
-
diff --git a/backend/operations_bnb.py b/backend/operations_bnb.py
index b45ab542..e7293ed5 100644
--- a/backend/operations_bnb.py
+++ b/backend/operations_bnb.py
@@ -1,11 +1,11 @@
# Copyright Forge 2024
-import torch
import bitsandbytes as bnb
-
-from backend import utils, memory_management
-from bitsandbytes.nn.modules import Params4bit, QuantState
+import torch
from bitsandbytes.functional import dequantize_4bit
+from bitsandbytes.nn.modules import Params4bit, QuantState
+
+from backend import memory_management, utils
def functional_linear_4bits(x, weight, bias):
@@ -20,12 +20,12 @@ def functional_dequantize_4bit(weight):
weight_original_device = weight.device
- if weight_original_device.type != 'cuda':
+ if weight_original_device.type != "cuda":
weight = weight.cuda()
weight = dequantize_4bit(weight, quant_state=weight.quant_state, blocksize=weight.blocksize, quant_type=weight.quant_type)
- if weight_original_device.type != 'cuda':
+ if weight_original_device.type != "cuda":
weight = weight.to(device=weight_original_device)
return weight
@@ -118,26 +118,26 @@ class ForgeLoader4Bit(torch.nn.Module):
return
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):
- quant_state_keys = {k[len(prefix + "weight."):] for k in state_dict.keys() if k.startswith(prefix + "weight.")}
+ quant_state_keys = {k[len(prefix + "weight.") :] for k in state_dict.keys() if k.startswith(prefix + "weight.")}
- if any('bitsandbytes' in k for k in quant_state_keys):
+ if any("bitsandbytes" in k for k in quant_state_keys):
quant_state_dict = {k: state_dict[prefix + "weight." + k] for k in quant_state_keys}
self.weight = ForgeParams4bit.from_prequantized(
- data=state_dict[prefix + 'weight'],
+ data=state_dict[prefix + "weight"],
quantized_stats=quant_state_dict,
requires_grad=False,
device=self.dummy.device,
)
- if prefix + 'bias' in state_dict:
- self.bias = torch.nn.Parameter(state_dict[prefix + 'bias'].to(self.dummy))
+ if prefix + "bias" in state_dict:
+ self.bias = torch.nn.Parameter(state_dict[prefix + "bias"].to(self.dummy))
del self.dummy
- elif hasattr(self, 'dummy'):
- if prefix + 'weight' in state_dict:
+ elif hasattr(self, "dummy"):
+ if prefix + "weight" in state_dict:
self.weight = ForgeParams4bit(
- state_dict[prefix + 'weight'].to(self.dummy),
+ state_dict[prefix + "weight"].to(self.dummy),
requires_grad=False,
compress_statistics=False,
blocksize=64,
@@ -145,8 +145,8 @@ class ForgeLoader4Bit(torch.nn.Module):
quant_storage=torch.uint8,
)
- if prefix + 'bias' in state_dict:
- self.bias = torch.nn.Parameter(state_dict[prefix + 'bias'].to(self.dummy))
+ if prefix + "bias" in state_dict:
+ self.bias = torch.nn.Parameter(state_dict[prefix + "bias"].to(self.dummy))
del self.dummy
else:
@@ -154,16 +154,8 @@ class ForgeLoader4Bit(torch.nn.Module):
def reload_weight(self, weight):
weight_original_device = weight.device
- weight = ForgeParams4bit(
- weight,
- requires_grad=False,
- compress_statistics=self.weight.compress_statistics,
- blocksize=self.weight.blocksize,
- quant_type=self.weight.quant_type,
- quant_storage=self.weight.quant_storage,
- bnb_quantized=False
- )
- if weight_original_device.type == 'cuda':
+ weight = ForgeParams4bit(weight, requires_grad=False, compress_statistics=self.weight.compress_statistics, blocksize=self.weight.blocksize, quant_type=self.weight.quant_type, quant_storage=self.weight.quant_storage, bnb_quantized=False)
+ if weight_original_device.type == "cuda":
weight = weight.to(weight_original_device)
else:
weight = weight.cuda().to(weight_original_device)
diff --git a/backend/operations_gguf.py b/backend/operations_gguf.py
index 468e4991..e010b792 100644
--- a/backend/operations_gguf.py
+++ b/backend/operations_gguf.py
@@ -1,7 +1,6 @@
import gguf
import torch
-
quants_mapping = {
gguf.GGMLQuantizationType.Q2_K: gguf.Q2_K,
gguf.GGMLQuantizationType.Q3_K: gguf.Q3_K,
@@ -60,7 +59,7 @@ def dequantize_tensor(tensor):
if tensor is None:
return None
- if not hasattr(tensor, 'gguf_cls'):
+ if not hasattr(tensor, "gguf_cls"):
return tensor
gguf_cls = tensor.gguf_cls
diff --git a/backend/patcher/controlnet.py b/backend/patcher/controlnet.py
index 32b8150d..772a393d 100644
--- a/backend/patcher/controlnet.py
+++ b/backend/patcher/controlnet.py
@@ -5,7 +5,12 @@ import torch
from backend import memory_management, state_dict, utils
from backend.misc import image_resize
from backend.nn.cnets import cldm, t2i_adapter
-from backend.operations import ForgeOperations, main_stream_worker, using_forge_operations, weights_manual_cast
+from backend.operations import (
+ ForgeOperations,
+ main_stream_worker,
+ using_forge_operations,
+ weights_manual_cast,
+)
from backend.patcher.base import ModelPatcher
diff --git a/backend/sampling/sampling_function.py b/backend/sampling/sampling_function.py
index 03091f5d..0c44f19a 100644
--- a/backend/sampling/sampling_function.py
+++ b/backend/sampling/sampling_function.py
@@ -11,7 +11,11 @@ import torch
from backend import memory_management, utils
from backend.args import args, dynamic_args
from backend.operations import cleanup_cache
-from backend.sampling.condition import Condition, compile_conditions, compile_weighted_conditions
+from backend.sampling.condition import (
+ Condition,
+ compile_conditions,
+ compile_weighted_conditions,
+)
def get_area_and_mult(conds, x_in, timestep_in):
diff --git a/backend/state_dict.py b/backend/state_dict.py
index 317bfce5..789cb389 100644
--- a/backend/state_dict.py
+++ b/backend/state_dict.py
@@ -1,6 +1,3 @@
-import torch
-
-
def load_state_dict(model, sd, ignore_errors=[], log_name=None, ignore_start=None):
missing, unexpected = model.load_state_dict(sd, strict=False)
missing = [x for x in missing if x not in ignore_errors]
@@ -12,9 +9,9 @@ def load_state_dict(model, sd, ignore_errors=[], log_name=None, ignore_start=Non
log_name = log_name or type(model).__name__
if len(missing) > 0:
- print(f'{log_name} Missing: {missing}')
+ print(f"{log_name} Missing: {missing}")
if len(unexpected) > 0:
- print(f'{log_name} Unexpected: {unexpected}')
+ print(f"{log_name} Unexpected: {unexpected}")
return
@@ -22,18 +19,18 @@ def state_dict_has(sd, prefix):
return any(x.startswith(prefix) for x in sd.keys())
-def filter_state_dict_with_prefix(sd, prefix, new_prefix=''):
+def filter_state_dict_with_prefix(sd, prefix, new_prefix=""):
new_sd = {}
for k, v in list(sd.items()):
if k.startswith(prefix):
- new_sd[new_prefix + k[len(prefix):]] = v
+ new_sd[new_prefix + k[len(prefix) :]] = v
del sd[k]
return new_sd
-def try_filter_state_dict(sd, prefix_list, new_prefix=''):
+def try_filter_state_dict(sd, prefix_list, new_prefix=""):
for prefix in prefix_list:
if state_dict_has(sd, prefix):
return filter_state_dict_with_prefix(sd, prefix, new_prefix)
@@ -77,7 +74,7 @@ def transformers_convert(sd, prefix_from, prefix_to, number):
for x in range(3):
p = ["self_attn.q_proj", "self_attn.k_proj", "self_attn.v_proj"]
k_to = "{}encoder.layers.{}.{}.{}".format(prefix_to, resblock, p[x], y)
- sd[k_to] = weights[shape_from*x:shape_from*(x + 1)]
+ sd[k_to] = weights[shape_from * x : shape_from * (x + 1)]
return sd
@@ -94,7 +91,7 @@ def state_dict_prefix_replace(state_dict, replace_prefix, filter_keys=False):
else:
out = state_dict
for rp in replace_prefix:
- replace = list(map(lambda a: (a, "{}{}".format(replace_prefix[rp], a[len(rp):])), filter(lambda a: a.startswith(rp), state_dict.keys())))
+ replace = list(map(lambda a: (a, "{}{}".format(replace_prefix[rp], a[len(rp) :])), filter(lambda a: a.startswith(rp), state_dict.keys())))
for x in replace:
w = state_dict.pop(x[0])
out[x[1]] = w
diff --git a/backend/stream.py b/backend/stream.py
index f3fcd7bc..d9343787 100644
--- a/backend/stream.py
+++ b/backend/stream.py
@@ -1,4 +1,5 @@
import torch
+
from backend.args import args
diff --git a/backend/utils.py b/backend/utils.py
index c88fceae..3703fc5b 100644
--- a/backend/utils.py
+++ b/backend/utils.py
@@ -1,19 +1,21 @@
-import gguf
-import torch
-import os
import json
+import os
+
+import gguf
import safetensors.torch
+import torch
+
import backend.misc.checkpoint_pickle
from backend.operations_gguf import ParameterGGUF
def read_arbitrary_config(directory):
- config_path = os.path.join(directory, 'config.json')
+ config_path = os.path.join(directory, "config.json")
if not os.path.exists(config_path):
raise FileNotFoundError(f"No config.json file found in the directory: {directory}")
- with open(config_path, 'rt', encoding='utf-8') as file:
+ with open(config_path, "rt", encoding="utf-8") as file:
config_data = json.load(file)
return config_data
@@ -31,7 +33,7 @@ def load_torch_file(ckpt, safe_load=False, device=None):
sd[str(tensor.name)] = ParameterGGUF(tensor)
else:
if safe_load:
- if not 'weights_only' in torch.load.__code__.co_varnames:
+ if not "weights_only" in torch.load.__code__.co_varnames:
print("Warning torch.load doesn't support weights_only on this pytorch version, loading unsafely.")
safe_load = False
if safe_load:
@@ -147,9 +149,9 @@ def nested_move_to_device(obj, **kwargs):
return obj
-def get_state_dict_after_quant(model, prefix=''):
+def get_state_dict_after_quant(model, prefix=""):
for m in model.modules():
- if hasattr(m, 'weight') and hasattr(m.weight, 'bnb_quantized'):
+ if hasattr(m, "weight") and hasattr(m.weight, "bnb_quantized"):
if not m.weight.bnb_quantized:
original_device = m.weight.device
m.cuda()
@@ -162,14 +164,15 @@ def get_state_dict_after_quant(model, prefix=''):
def beautiful_print_gguf_state_dict_statics(state_dict):
from gguf.constants import GGMLQuantizationType
+
type_counts = {}
for k, v in state_dict.items():
- gguf_cls = getattr(v, 'gguf_cls', None)
+ gguf_cls = getattr(v, "gguf_cls", None)
if gguf_cls is not None:
type_name = gguf_cls.__name__
if type_name in type_counts:
type_counts[type_name] += 1
else:
type_counts[type_name] = 1
- print(f'GGUF state dict: {type_counts}')
+ print(f"GGUF state dict: {type_counts}")
return