Add files via upload

This commit is contained in:
TripleHeadedMonkey
2026-02-12 13:28:20 +00:00
committed by GitHub
parent 4da08fe219
commit 1d11ce1f07
40 changed files with 14757 additions and 0 deletions
+69
View File
@@ -0,0 +1,69 @@
import torch
class ZImageConditioningContrast:
"""
Applies Non-Linear Contrast (Power Scaling) to Z-Image/Qwen embeddings.
Why this works where Multiplication fails:
Z-Image applies Layer Normalization to inputs, which mathematically cancels out
any linear multiplication (1.2 * x -> normalized -> x).
This node applies a Power Function (x ^ exponent), which changes the
distribution shape (Kurtosis) of the embedding. Normalization cannot
revert this, allowing you to effectively 'sharpen' or 'soften' the prompt.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"conditioning": ("CONDITIONING", ),
"contrast": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.05}),
}
}
RETURN_TYPES = ("CONDITIONING",)
RETURN_NAMES = ("conditioning",)
FUNCTION = "apply_contrast"
CATEGORY = "RES4LYF/conditioning"
def apply_contrast(self, conditioning, contrast):
c = []
for t in conditioning:
# 1. Handle Standard CLIP Embeddings (Less important for Z-Image, but good practice)
# Formula: sign(x) * |x|^contrast
# We use sign() to preserve negative values, as x^2 would lose them.
original_emb = t[0]
new_emb = original_emb.sign() * original_emb.abs().pow(contrast)
# 2. Handle Metadata Dictionary (The Critical Qwen Part)
original_dict = t[1]
new_dict = original_dict.copy()
# Keys known to hold Qwen/Llama embeddings
target_keys = ["conditioning_llama3", "llama_embeds", "pooled_output"]
for key, value in new_dict.items():
if isinstance(value, torch.Tensor):
# Skip masks to prevent crashes
if "mask" in key or "ids" in key or "size" in key:
continue
# Target specific embedding keys OR keys that look like high-dim embeddings
if key in target_keys or (value.dim() >= 3 and value.shape[-1] > 64):
# Apply Power Scaling
# .contiguous() is strictly required for Flash Attention in Turbo models
new_dict[key] = (value.sign() * value.abs().pow(contrast)).contiguous()
c.append([new_emb, new_dict])
return (c,)
# Register the node
NODE_CLASS_MAPPINGS = {
"ZImageConditioningContrast": ZImageConditioningContrast
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZImageConditioningContrast": "Z-Image Conditioning Contrast"
}
+173
View File
@@ -0,0 +1,173 @@
import torch
import folder_paths
import comfy.sd
import comfy.ops
import comfy.model_management
from .loader import gguf_sd_loader, gguf_clip_loader
from .ops import GGMLOps, manual_resize_tensor
from .nodes2 import GGUFModelPatcher
from .Universal_LoRA import ZImageRawWrapper
# ==============================================================================
# 1. STANDALONE GGUF LOADER (Reader + Fixer)
# ==============================================================================
class ZImageGGUFStandaloneLoader:
@classmethod
def INPUT_TYPES(s):
models = folder_paths.get_filename_list("unet_gguf")
clips = folder_paths.get_filename_list("clip_gguf")
return {
"required": {
"transformer_name": (["None"] + models, ),
"text_encoder_name": (["None"] + clips, ),
"type": (
[
"stable_diffusion",
"stable_diffusion_xl",
"sd3",
"flux",
"qwen_image",
"gemma",
"hunyuan_di_t"
],
),
}
}
RETURN_TYPES = ("RAW_MODEL", "RAW_CLIP")
FUNCTION = "load_gguf_raw"
CATEGORY = "Z-Image/Loaders"
def _fix_tensor_padding(self, sd):
"""
Scans state dict for specific GGUF padding sizes and trims them.
2720 -> 2560 (Qwen/LLM Hidden Size)
4352 -> 4096 (Gemma/LLM Hidden Size)
"""
new_sd = {}
for k, v in sd.items():
current_shape = list(v.shape)
needs_resize = False
# Iterate through all dimensions to find padding
for i, dim in enumerate(current_shape):
if dim == 2720:
current_shape[i] = 2560
needs_resize = True
elif dim == 4352:
current_shape[i] = 4096
needs_resize = True
if needs_resize:
new_v = manual_resize_tensor(v, torch.Size(current_shape))
new_sd[k] = new_v
else:
new_sd[k] = v
return new_sd
def load_gguf_raw(self, transformer_name, text_encoder_name, type="stable_diffusion"):
raw_model = None
raw_clip = None
# --- 1. LOAD MODEL (UNET) & FIX SHAPES ---
if transformer_name != "None":
path = folder_paths.get_full_path("unet_gguf", transformer_name)
sd = gguf_sd_loader(path)
# Apply padding fix if the type warrants it
if type in ["qwen_image", "gemma", "stable_diffusion"]:
sd = self._fix_tensor_padding(sd)
raw_model = ZImageRawWrapper(sd, path)
# --- 2. LOAD CLIP & FIX SHAPES ---
if text_encoder_name != "None":
path = folder_paths.get_full_path("clip_gguf", text_encoder_name)
sd = gguf_clip_loader(path)
# Apply padding fix (handles 2720 AND 4352)
if type in ["qwen_image", "gemma", "stable_diffusion"]:
sd = self._fix_tensor_padding(sd)
raw_clip = ZImageRawWrapper(sd, path)
return (raw_model, raw_clip)
# ==============================================================================
# 2. GGUF INJECTOR (Builder)
# ==============================================================================
class ZImageGGUFInjector:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"raw_model": ("RAW_MODEL",),
"raw_clip": ("RAW_CLIP",),
"clip_type": (
[
"stable_diffusion",
"stable_diffusion_xl",
"sd3",
"flux",
"qwen_image",
"gemma",
"hunyuan_di_t"
],
),
}
}
RETURN_TYPES = ("MODEL", "CLIP")
FUNCTION = "inject_gguf"
CATEGORY = "Z-Image/Injectors"
def inject_gguf(self, raw_model, raw_clip, clip_type):
model = None
clip = None
if raw_model is not None:
sd = raw_model.sd
# We still keep GGMLOps here for the specialized Linear layers
ops = GGMLOps()
model_obj = comfy.sd.load_diffusion_model_state_dict(
sd,
model_options={"custom_operations": ops}
)
model = GGUFModelPatcher.clone(model_obj)
if raw_clip is not None:
try:
clip_type_enum = getattr(comfy.sd.CLIPType, clip_type.upper())
except AttributeError:
clip_type_enum = comfy.sd.CLIPType.STABLE_DIFFUSION
clip = comfy.sd.load_text_encoder_state_dicts(
[raw_clip.sd],
embedding_directory=folder_paths.get_folder_paths("embeddings"),
clip_type=clip_type_enum,
model_options = {
"custom_operations": GGMLOps,
"initial_device": comfy.model_management.text_encoder_offload_device()
}
)
clip.patcher = GGUFModelPatcher.clone(clip.patcher)
return (model, clip)
# ==============================================================================
# MAPPINGS
# ==============================================================================
NODE_CLASS_MAPPINGS = {
"ZImageGGUFStandaloneLoader": ZImageGGUFStandaloneLoader,
"ZImageGGUFInjector": ZImageGGUFInjector
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZImageGGUFStandaloneLoader": "Z-Image GGUF Standalone Loader",
"ZImageGGUFInjector": "Z-Image GGUF Injector"
}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+4
View File
@@ -0,0 +1,4 @@
MIT License
Copyright (c) 2026 TripleHeadedMonkey
...
+153
View File
@@ -0,0 +1,153 @@
import os
import folder_paths
import comfy.utils
import comfy.sd
import comfy.lora
import logging
def _split_lora_state_dict(sd: dict):
te_sd = {}
rest_sd = {}
for k, v in sd.items():
if k.startswith("lora_te."):
te_sd[k] = v
else:
rest_sd[k] = v
return te_sd, rest_sd
def _get_cond_stage_model(clip):
# In your build, clip has cond_stage_model and that is where the Qwen weights live.
if hasattr(clip, "cond_stage_model"):
return clip.cond_stage_model
return clip
def _infer_qwen_prefix_from_state_dict_keys(keys):
"""
We saw keys like:
qwen3_4b.transformer.model.layers.0.self_attn.q_proj.weight
We want to discover the prefix:
qwen3_4b.transformer.
Return something like "qwen3_4b.transformer."
"""
for k in keys:
if k.endswith(".weight") and ".transformer.model.layers." in k:
# everything before "model.layers"
# e.g. "qwen3_4b.transformer.model.layers..." -> "qwen3_4b.transformer."
idx = k.find("model.layers.")
return k[:idx] # includes trailing dot
return None
def _build_qwen_te_key_map(cond_stage_model):
"""
Map ai-toolkit TE LoRA prefixes to actual TE weight keys.
ai-toolkit LoRA prefix: lora_te.<path_without_.weight>
e.g. lora_te.model.layers.0.self_attn.q_proj
actual weight key: <QWEN_PREFIX>model.layers.0.self_attn.q_proj.weight
e.g. qwen3_4b.transformer.model.layers.0.self_attn.q_proj.weight
"""
sd = cond_stage_model.state_dict()
keys = list(sd.keys())
qwen_prefix = _infer_qwen_prefix_from_state_dict_keys(keys)
if qwen_prefix is None:
raise RuntimeError("Could not infer Qwen prefix. Expected keys containing '.transformer.model.layers.'")
key_map = {}
# Build mapping for every .weight key under qwen_prefix + "model."
for k in keys:
if not k.endswith(".weight"):
continue
if not k.startswith(qwen_prefix):
continue
# Strip ".weight" and strip leading prefix => path like "model.layers.0.self_attn.q_proj"
bare = k[:-len(".weight")]
if not bare.startswith(qwen_prefix):
continue
bare_no_prefix = bare[len(qwen_prefix):] # "model.layers...."
lora_key = "lora_te." + bare_no_prefix
key_map[lora_key] = k
return key_map
class ZImageQwenTELoRALoader:
"""
Like Load LoRA, but applies ai-toolkit TE LoRA keys (lora_te.model.layers...)
onto ComfyUI's Qwen TE keyspace (qwen3_4b.transformer.model.layers...).
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"clip": ("CLIP",),
"lora_name": (folder_paths.get_filename_list("loras"),),
"strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
"strength_clip": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
}
}
RETURN_TYPES = ("MODEL", "CLIP")
FUNCTION = "load_lora"
CATEGORY = "loaders"
def load_lora(self, model, clip, lora_name, strength_model, strength_clip):
lora_path = folder_paths.get_full_path("loras", lora_name)
if lora_path is None or not os.path.exists(lora_path):
raise FileNotFoundError(f"LoRA file not found: {lora_name}")
lora_sd = comfy.utils.load_torch_file(lora_path, safe_load=True)
te_sd, rest_sd = _split_lora_state_dict(lora_sd)
# 1) Apply non-TE LoRA weights normally (diffusion/transformer part)
model_out, clip_out = comfy.sd.load_lora_for_models(model, clip, rest_sd, strength_model, 0.0)
# 2) Apply TE LoRA directly to the Qwen cond_stage_model
if te_sd and abs(strength_clip) > 1e-8:
try:
clip_out = clip_out.clone() if hasattr(clip_out, "clone") else clip_out
cond = _get_cond_stage_model(clip_out)
key_map = _build_qwen_te_key_map(cond)
# NOTE: lm_head won't exist in your cond_stage_model state_dict, so it will be skipped naturally.
patches = comfy.lora.load_lora(te_sd, key_map)
# Add patches to the CLIP object (patcher exists on clip wrapper)
# FIXED: strength_model must be 1.0 to preserve the original base weights.
# Only strength_patch should use the slider value (strength_clip).
if hasattr(clip_out, "add_patches"):
clip_out.add_patches(patches, strength_patch=strength_clip, strength_model=1.0)
elif hasattr(clip_out, "patcher") and hasattr(clip_out.patcher, "add_patches"):
clip_out.patcher.add_patches(patches, strength_patch=strength_clip, strength_model=1.0)
else:
raise RuntimeError("Could not find add_patches on CLIP wrapper")
except Exception:
logging.exception("Failed to apply Z-Image Qwen TE LoRA")
# keep diffusion LoRA applied even if TE fails
pass
return (model_out, clip_out)
NODE_CLASS_MAPPINGS = {
"ZImageQwenTELoRALoader": ZImageQwenTELoRALoader
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZImageQwenTELoRALoader": "Load LoRA (Z-Image Qwen TE)"
}
+123
View File
@@ -0,0 +1,123 @@
import torch
import copy
from comfy.model_patcher import ModelPatcher
class ZImageTIESMerge:
"""
Implements a simplified TIES Merging (Trim-Only for Single Pair):
1. Calculate Delta = Turbo - Base
2. Trim: Zero out the bottom (1 - density)% of values by magnitude.
3. Merge: Base + Strength * (Trimmed_Delta)
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_base": ("MODEL",),
"model_turbo": ("MODEL",),
"density": ("FLOAT", {
"default": 0.2,
"min": 0.01,
"max": 1.0,
"step": 0.05,
"display": "number",
"tooltip": "Fraction of weights to keep (0.2 = keep top 20% of changes)"
}),
"strength": ("FLOAT", {
"default": 1.0,
"min": 0.0,
"max": 5.0,
"step": 0.1,
"display": "number"
}),
}
}
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("merged_model",)
FUNCTION = "apply_ties_merge"
CATEGORY = "Experimental"
def apply_ties_merge(self, model_base, model_turbo, density, strength):
print(f"Applying TIES Merge (Density: {density}, Strength: {strength})")
# 1. Clone Base Model (Target)
if isinstance(model_base, ModelPatcher):
new_model_patcher = model_base.clone()
base_model_obj = new_model_patcher.model
else:
base_model_obj = model_base
new_model_patcher = copy.deepcopy(model_base)
base_sd = base_model_obj.diffusion_model.state_dict()
# 2. Get Turbo State Dict
if isinstance(model_turbo, ModelPatcher):
turbo_sd = model_turbo.model.diffusion_model.state_dict()
else:
turbo_sd = model_turbo.diffusion_model.state_dict()
merged_sd = {}
keys_processed = 0
# We perform operations on CPU to avoid "Tensor on different device" errors
# and to avoid OOM on GPU during the heavy sort/top-k operations.
calculation_device = torch.device("cpu")
for key in base_sd.keys():
if key in turbo_sd:
# 3. Load tensors and force to same device/dtype
w_base = base_sd[key].to(device=calculation_device, dtype=torch.float32)
w_turbo = turbo_sd[key].to(device=calculation_device, dtype=torch.float32)
if w_base.shape != w_turbo.shape:
print(f"Skipping {key}: Shape mismatch.")
merged_sd[key] = base_sd[key]
continue
# 4. Calculate Task Vector (Delta)
delta = w_turbo - w_base
# 5. TIES-TRIM: Filter out small values (noise)
# We only want the top 'density' (e.g. 20%) of changes
if density < 1.0:
# Flatten to find the global threshold for this layer
flat_delta = delta.abs().view(-1)
k = int(flat_delta.numel() * density)
if k > 0:
# Find the k-th largest value
top_k_value, _ = torch.kthvalue(flat_delta, flat_delta.numel() - k + 1)
threshold = top_k_value.item()
# Zero out elements below threshold
mask = delta.abs() >= threshold
delta = delta * mask
else:
delta.zero_()
# 6. Apply Scaled Delta
merged_weight = w_base + (strength * delta)
# Cast back to original dtype/device of the base model logic (handled by Comfy loading)
# We store it back to CPU dict to be safe
merged_sd[key] = merged_weight.to(base_sd[key].dtype)
keys_processed += 1
else:
merged_sd[key] = base_sd[key]
print(f"TIES Merge complete. Processed {keys_processed} keys.")
# Load weights back
base_model_obj.diffusion_model.load_state_dict(merged_sd, strict=False)
return (new_model_patcher,)
NODE_CLASS_MAPPINGS = {
"ZImageTIESMerge": ZImageTIESMerge
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZImageTIESMerge": "Z-Image TIES Merge (Method 2)"
}
+242
View File
@@ -0,0 +1,242 @@
import torch
import folder_paths
from safetensors.torch import load_file
import logging
import os
import re
# Setup logging
file_handler = logging.FileHandler('z_image_universal.log', mode='w')
file_handler.setFormatter(logging.Formatter('%(message)s'))
logger = logging.getLogger("ZImageUniversal")
logger.setLevel(logging.INFO)
logger.addHandler(file_handler)
logger.addHandler(logging.StreamHandler())
# ==============================================================================
# 1. HELPER: RAW IMAGE WRAPPER
# ==============================================================================
class ZImageRawWrapper:
def __init__(self, sd, path):
self.sd = sd
self.path = path
# ==============================================================================
# 2. MATH HELPERS
# ==============================================================================
def robust_matmul(a, b):
if len(a.shape) == 4 and a.shape[2] == 1: a = a.squeeze(3).squeeze(2)
if len(b.shape) == 4 and b.shape[2] == 1: b = b.squeeze(3).squeeze(2)
if a.shape[-1] == b.shape[0]: return a @ b
elif b.shape[-1] == a.shape[0]: return b @ a
elif a.shape[0] == b.shape[0]: return a.T @ b
elif a.shape[-1] == b.shape[-1]: return a @ b.T
return None
def make_lora(wa, wb, scale):
res = robust_matmul(wa, wb)
if res is None: return None
return res * scale
# ==============================================================================
# 3. UNIVERSAL PATCHING LOGIC
# ==============================================================================
def normalize_key(k):
"""
Standardizes keys for comparison.
"""
k = k.replace("model.diffusion_model.", "")
k = k.replace("diffusion_model.", "").replace("transformer.", "")
k = k.replace("text_model.", "").replace("model.", "")
k = k.replace("lora_unet_", "").replace("lora_te_", "")
k = k.replace("lycoris_all_", "").replace("lycoris_", "")
# Clean Artifacts
k = k.replace("_2-1", "").replace(".2-1", "")
k = re.sub(r'(\d+)_\d+', r'\1', k) # Fix "blocks_1_2" -> "blocks_1"
# FIX: Standardize Attention Output
# Converts "to_out_0", "to_out.0", "to_out" all to "out"
k = k.replace("to_out_0", "out").replace("to_out.0", "out").replace("to_out", "out")
# Standardize QKV
k = k.replace("to_k", "k").replace("to_q", "q").replace("to_v", "v")
k = k.replace("k_proj", "k").replace("q_proj", "q").replace("v_proj", "v")
k = k.replace("qkv_proj", "qkv")
# Strip separators
k = k.replace(".", "").replace("_", "").replace("-", "")
return k.lower()
def apply_universal_lora(target_dict, lora_path, strength):
if strength == 0 or target_dict is None: return 0
filename = os.path.basename(lora_path)
logger.info(f"Applying Universal Patch: {filename}")
lora_sd = load_file(lora_path)
# 1. Group Keys
modules = {}
for k, v in lora_sd.items():
if "lora_up" in k:
base = k.replace(".lora_up.weight", "")
param = "up"
elif "lora_down" in k:
base = k.replace(".lora_down.weight", "")
param = "down"
elif "lora_A" in k:
base = k.replace(".lora_A.weight", "")
param = "down"
elif "lora_B" in k:
base = k.replace(".lora_B.weight", "")
param = "up"
elif "alpha" in k:
base = k.replace(".alpha", "")
param = "alpha"
else:
continue
if base not in modules: modules[base] = {}
modules[base][param] = v
# 2. Map Target Dict
normalized_target_map = {}
for k in target_dict.keys():
if not k.endswith(".weight"): continue
bare = k[:-7]
norm = normalize_key(bare)
normalized_target_map[norm] = k
patch_count = 0
skipped_modules = []
# 3. Patching Loop
for lora_prefix, params in modules.items():
if "up" not in params or "down" not in params: continue
up, down = params["up"].float(), params["down"].float()
alpha = params.get("alpha", None)
dim = down.shape[0]
scale = (float(alpha) / dim) * strength if alpha is not None else strength
diff = make_lora(up, down, scale)
if diff is None: continue
# --- MATCHING STRATEGY ---
norm_lora = normalize_key(lora_prefix)
target_key = normalized_target_map.get(norm_lora)
# Strategy B: QKV Fusion Patching
slice_idx = None
if not target_key:
if norm_lora.endswith("q") or norm_lora.endswith("k") or norm_lora.endswith("v"):
base_qkv = norm_lora[:-1] + "qkv"
parent_key = normalized_target_map.get(base_qkv)
if parent_key:
target_key = parent_key
if norm_lora.endswith("q"): slice_idx = 0
elif norm_lora.endswith("k"): slice_idx = 1
elif norm_lora.endswith("v"): slice_idx = 2
if not target_key:
skipped_modules.append(lora_prefix)
continue
# --- APPLY PATCH ---
w = target_dict[target_key]
try:
# Case 1: Standard Patch
if w.shape == diff.shape:
target_dict[target_key] = w + diff.to(w.dtype).to(w.device)
patch_count += 1
# Case 2: QKV Sliced Patch
elif slice_idx is not None:
chunk_size = w.shape[0] // 3
start = slice_idx * chunk_size
end = start + chunk_size
if diff.shape[0] == chunk_size and diff.shape[1] == w.shape[1]:
target_dict[target_key][start:end, :] += diff.to(w.dtype).to(w.device)
patch_count += 1
# Case 3: Reshape
elif len(w.shape) == 4 and diff.dim() == 2:
diff_reshaped = diff.reshape(w.shape)
target_dict[target_key] = w + diff_reshaped.to(w.dtype).to(w.device)
patch_count += 1
except Exception as e:
logger.error(f"Patch error {target_key}: {e}")
logger.info(f"Total Patched: {patch_count} / {len(modules)}")
if len(skipped_modules) > 0:
logger.info("--- Skipped Modules ---")
for s in skipped_modules[:5]:
logger.info(f"Skipped: {s} (Norm: {normalize_key(s)})")
return patch_count
# ==============================================================================
# 4. NODE DEFINITION
# ==============================================================================
class ZImageUniversalLoRALoader:
@classmethod
def INPUT_TYPES(s):
models = ["None"] + folder_paths.get_filename_list("diffusion_models")
tes = ["None"] + folder_paths.get_filename_list("text_encoders")
loras = ["None"] + folder_paths.get_filename_list("loras")
return {
"required": {
"transformer_name": (models, ),
"text_encoder_name": (tes, ),
"lora_name": (loras, ),
"strength_model": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step": 0.01}),
"strength_clip": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step": 0.01}),
}
}
RETURN_TYPES = ("RAW_MODEL", "RAW_CLIP")
FUNCTION = "load_universal"
CATEGORY = "Z-Image/Loaders"
def load_universal(self, transformer_name, text_encoder_name, lora_name, strength_model, strength_clip):
unet_sd = None
clip_sd = None
unet_path = None
clip_path = None
if transformer_name != "None":
unet_path = folder_paths.get_full_path("diffusion_models", transformer_name)
logger.info(f"Loading Model: {os.path.basename(unet_path)}")
unet_sd = load_file(unet_path)
if text_encoder_name != "None":
clip_path = folder_paths.get_full_path("text_encoders", text_encoder_name)
if not clip_path: clip_path = folder_paths.get_full_path("checkpoints", text_encoder_name)
logger.info(f"Loading CLIP: {os.path.basename(clip_path)}")
clip_sd = load_file(clip_path)
if lora_name != "None":
lora_path = folder_paths.get_full_path("loras", lora_name)
if lora_path:
if strength_model != 0 and unet_sd:
c = apply_universal_lora(unet_sd, lora_path, strength_model)
logging.info(f"Universal Model Patch: {c}")
if strength_clip != 0 and clip_sd:
c = apply_universal_lora(clip_sd, lora_path, strength_clip)
logging.info(f"Universal CLIP Patch: {c}")
raw_model = ZImageRawWrapper(unet_sd, unet_path) if unet_sd else None
raw_clip = ZImageRawWrapper(clip_sd, clip_path) if clip_sd else None
return (raw_model, raw_clip)
NODE_CLASS_MAPPINGS = {
"ZImageUniversalLoRALoader": ZImageUniversalLoRALoader
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZImageUniversalLoRALoader": "Z-Image Universal LoRA Loader"
}
+124
View File
@@ -0,0 +1,124 @@
import torch
class ZImageAdvancedConditioning:
"""
FINAL PRODUCTION VERSION.
- Fixes 'Drift' bug (No more in-place modification of cached inputs).
- Removes debug prints for speed.
- Safely handles Z-Image/Qwen/Llama architectures.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"conditioning_to": ("CONDITIONING", ),
"conditioning_from": ("CONDITIONING", ),
"operation": (["mix_slerp", "purge_ortho", "add_perpendicular"], ),
"strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
}
}
RETURN_TYPES = ("CONDITIONING",)
RETURN_NAMES = ("conditioning",)
FUNCTION = "process"
CATEGORY = "RES4LYF/conditioning"
def process(self, conditioning_to, conditioning_from, operation, strength):
results = []
# Match list lengths
min_len = min(len(conditioning_to), len(conditioning_from))
for i in range(min_len):
# --- 1. PROCESS MAIN TENSOR (t[0]) ---
t0_input = conditioning_to[i][0]
t0_ref = conditioning_from[i][0]
# Create a SAFE COPY of the main tensor to avoid modifying the input cache
t0_output = t0_input.clone()
# Handle Shape Mismatch (Slice to common length)
common_len = min(t0_input.shape[1], t0_ref.shape[1])
# Extract working slices
v0 = t0_input[:, :common_len, :].clone()
v1 = t0_ref[:, :common_len, :].clone()
# Apply Math
v_processed = self.apply_math(v0, v1, operation, strength)
# Write into the NEW output tensor (not the input!)
t0_output[:, :common_len, :] = v_processed
# --- 2. PROCESS DICTIONARY ---
# Create a shallow copy of the dict, but we will replace values with new tensors
new_dict = conditioning_to[i][1].copy()
ref_dict = conditioning_from[i][1]
target_keys = ["conditioning_llama3", "llama_embeds", "pooled_output"]
for key in new_dict.keys():
if key in target_keys and key in ref_dict:
val_to = new_dict[key]
val_from = ref_dict[key]
if val_to is not None and val_from is not None and isinstance(val_to, torch.Tensor):
# Ensure we only process if shapes align
if val_to.shape == val_from.shape:
# Apply Math
# apply_math returns a new tensor, so this is safe
new_dict[key] = self.apply_math(val_to, val_from, operation, strength)
results.append([t0_output, new_dict])
return (results,)
def apply_math(self, v0, v1, operation, strength):
"""Helper for vector operations"""
# Epsilon for stability
eps = 1e-8
if operation == "mix_slerp":
# Normalize
v0_n = v0 / (v0.norm(dim=-1, keepdim=True) + eps)
v1_n = v1 / (v1.norm(dim=-1, keepdim=True) + eps)
dot = (v0_n * v1_n).sum(dim=-1, keepdim=True)
dot = torch.clamp(dot, -0.9995, 0.9995)
theta = torch.acos(dot)
sin_theta = torch.sin(theta) + eps
w0 = torch.sin((1.0 - strength) * theta) / sin_theta
w1 = torch.sin(strength * theta) / sin_theta
return (w0 * v0 + w1 * v1).contiguous()
elif operation == "purge_ortho":
# Project v0 onto v1
dot_v0_v1 = (v0 * v1).sum(dim=-1, keepdim=True)
dot_v1_v1 = (v1 * v1).sum(dim=-1, keepdim=True)
proj = (dot_v0_v1 / (dot_v1_v1 + eps)) * v1
return (v0 - (proj * strength)).contiguous()
elif operation == "add_perpendicular":
# Find part of v1 orthogonal to v0
dot_v1_v0 = (v1 * v0).sum(dim=-1, keepdim=True)
dot_v0_v0 = (v0 * v0).sum(dim=-1, keepdim=True)
proj = (dot_v1_v0 / (dot_v0_v0 + eps)) * v0
ortho_v1 = v1 - proj
return (v0 + (ortho_v1 * strength)).contiguous()
return v0.contiguous()
# Register
NODE_CLASS_MAPPINGS = {
"ZImageAdvancedConditioning": ZImageAdvancedConditioning
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZImageAdvancedConditioning": "Z-Image Advanced Mixing"
}
+102
View File
@@ -0,0 +1,102 @@
# =====================================================
# Z-Image Core Nodes
# =====================================================
from .nodes import NODE_CLASS_MAPPINGS as NODES_CLASS_MAPPINGS
from .nodes import NODE_DISPLAY_NAME_MAPPINGS as NODES_DISPLAY_NAME_MAPPINGS
from .ZCondAdv import NODE_CLASS_MAPPINGS as ZCOND_CLASS_MAPPINGS
from .ZCondAdv import NODE_DISPLAY_NAME_MAPPINGS as ZCOND_DISPLAY_NAME_MAPPINGS
from .CondMul import NODE_CLASS_MAPPINGS as CONDMUL_CLASS_MAPPINGS
from .CondMul import NODE_DISPLAY_NAME_MAPPINGS as CONDMUL_DISPLAY_NAME_MAPPINGS
from .diffz import NODE_CLASS_MAPPINGS as DIFFZ_CLASS_MAPPINGS
from .diffz import NODE_DISPLAY_NAME_MAPPINGS as DIFFZ_DISPLAY_NAME_MAPPINGS
from .Universal_LoRA import NODE_CLASS_MAPPINGS as UNIVERSAL_CLASS_MAPPINGS
from .Universal_LoRA import NODE_DISPLAY_NAME_MAPPINGS as UNIVERSAL_DISPLAY_NAME_MAPPINGS
from .Qwen_Lora import NODE_CLASS_MAPPINGS as QWEN_CLASS_MAPPINGS
from .Qwen_Lora import NODE_DISPLAY_NAME_MAPPINGS as QWEN_DISPLAY_NAME_MAPPINGS
# =====================================================
# Experimental Merge Nodes
# =====================================================
from .z_image_vector_merge import NODE_CLASS_MAPPINGS as VECTOR_CLASS_MAPPINGS
from .z_image_vector_merge import NODE_DISPLAY_NAME_MAPPINGS as VECTOR_DISPLAY_NAME_MAPPINGS
from .TIES import NODE_CLASS_MAPPINGS as TIES_CLASS_MAPPINGS
from .TIES import NODE_DISPLAY_NAME_MAPPINGS as TIES_DISPLAY_NAME_MAPPINGS
# =====================================================
# Utility / Logic Nodes
# =====================================================
from .getimagesizeplus import NODE_CLASS_MAPPINGS as SIZE_CLASS_MAPPINGS
from .getimagesizeplus import NODE_DISPLAY_NAME_MAPPINGS as SIZE_DISPLAY_NAME_MAPPINGS
from .keyword_match_gate import NODE_CLASS_MAPPINGS as GATE_CLASS_MAPPINGS
from .keyword_match_gate import NODE_DISPLAY_NAME_MAPPINGS as GATE_DISPLAY_NAME_MAPPINGS
from .primitive_widget_to_string import NODE_CLASS_MAPPINGS as PRIMITIVE_CLASS_MAPPINGS
from .primitive_widget_to_string import NODE_DISPLAY_NAME_MAPPINGS as PRIMITIVE_DISPLAY_NAME_MAPPINGS
# =====================================================
# GGUF RAW System (Root + Components)
# =====================================================
# Root GGUF initializer
from .GGUF_RAW import NODE_CLASS_MAPPINGS as GGUF_ROOT_CLASS_MAPPINGS
from .GGUF_RAW import NODE_DISPLAY_NAME_MAPPINGS as GGUF_ROOT_DISPLAY_NAME_MAPPINGS
# =====================================================
# Merge All Mappings
# =====================================================
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
ALL_CLASS_MAPPINGS = [
NODES_CLASS_MAPPINGS,
ZCOND_CLASS_MAPPINGS,
CONDMUL_CLASS_MAPPINGS,
DIFFZ_CLASS_MAPPINGS,
UNIVERSAL_CLASS_MAPPINGS,
QWEN_CLASS_MAPPINGS,
VECTOR_CLASS_MAPPINGS,
TIES_CLASS_MAPPINGS,
SIZE_CLASS_MAPPINGS,
GATE_CLASS_MAPPINGS,
PRIMITIVE_CLASS_MAPPINGS, # <-- ADDED
GGUF_ROOT_CLASS_MAPPINGS,
]
ALL_DISPLAY_MAPPINGS = [
NODES_DISPLAY_NAME_MAPPINGS,
ZCOND_DISPLAY_NAME_MAPPINGS,
CONDMUL_DISPLAY_NAME_MAPPINGS,
DIFFZ_DISPLAY_NAME_MAPPINGS,
UNIVERSAL_DISPLAY_NAME_MAPPINGS,
QWEN_DISPLAY_NAME_MAPPINGS,
VECTOR_DISPLAY_NAME_MAPPINGS,
TIES_DISPLAY_NAME_MAPPINGS,
SIZE_DISPLAY_NAME_MAPPINGS,
GATE_DISPLAY_NAME_MAPPINGS,
PRIMITIVE_DISPLAY_NAME_MAPPINGS, # <-- ADDED
GGUF_ROOT_DISPLAY_NAME_MAPPINGS,
]
for mapping in ALL_CLASS_MAPPINGS:
NODE_CLASS_MAPPINGS.update(mapping)
for mapping in ALL_DISPLAY_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS.update(mapping)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+301
View File
@@ -0,0 +1,301 @@
# (c) City96 || Apache-2.0 (apache.org/licenses/LICENSE-2.0)
import gguf
import torch
from tqdm import tqdm
TORCH_COMPATIBLE_QTYPES = (None, gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16)
def is_torch_compatible(tensor):
return tensor is None or getattr(tensor, "tensor_type", None) in TORCH_COMPATIBLE_QTYPES
def is_quantized(tensor):
return not is_torch_compatible(tensor)
def dequantize_tensor(tensor, dtype=None, dequant_dtype=None):
qtype = getattr(tensor, "tensor_type", None)
oshape = getattr(tensor, "tensor_shape", tensor.shape)
if qtype in TORCH_COMPATIBLE_QTYPES:
return tensor.to(dtype)
elif qtype in dequantize_functions:
dequant_dtype = dtype if dequant_dtype == "target" else dequant_dtype
return dequantize(tensor.data, qtype, oshape, dtype=dequant_dtype).to(dtype)
else:
# this is incredibly slow
tqdm.write(f"Falling back to numpy dequant for qtype: {getattr(qtype, 'name', repr(qtype))}")
new = gguf.quants.dequantize(tensor.cpu().numpy(), qtype)
return torch.from_numpy(new).to(tensor.device, dtype=dtype)
def dequantize(data, qtype, oshape, dtype=None):
"""
Dequantize tensor back to usable shape/dtype
"""
block_size, type_size = gguf.GGML_QUANT_SIZES[qtype]
dequantize_blocks = dequantize_functions[qtype]
rows = data.reshape(
(-1, data.shape[-1])
).view(torch.uint8)
n_blocks = rows.numel() // type_size
blocks = rows.reshape((n_blocks, type_size))
blocks = dequantize_blocks(blocks, block_size, type_size, dtype)
return blocks.reshape(oshape)
def to_uint32(x):
# no uint32 :(
x = x.view(torch.uint8).to(torch.int32)
return (x[:, 0] | x[:, 1] << 8 | x[:, 2] << 16 | x[:, 3] << 24).unsqueeze(1)
def to_uint16(x):
x = x.view(torch.uint8).to(torch.int32)
return (x[:, 0] | x[:, 1] << 8).unsqueeze(1)
def split_block_dims(blocks, *args):
n_max = blocks.shape[1]
dims = list(args) + [n_max - sum(args)]
return torch.split(blocks, dims, dim=1)
# Full weights #
def dequantize_blocks_BF16(blocks, block_size, type_size, dtype=None):
return (blocks.view(torch.int16).to(torch.int32) << 16).view(torch.float32)
# Legacy Quants #
def dequantize_blocks_Q8_0(blocks, block_size, type_size, dtype=None):
d, x = split_block_dims(blocks, 2)
d = d.view(torch.float16).to(dtype)
x = x.view(torch.int8)
return (d * x)
def dequantize_blocks_Q5_1(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
d, m, qh, qs = split_block_dims(blocks, 2, 2, 4)
d = d.view(torch.float16).to(dtype)
m = m.view(torch.float16).to(dtype)
qh = to_uint32(qh)
qh = qh.reshape((n_blocks, 1)) >> torch.arange(32, device=d.device, dtype=torch.int32).reshape(1, 32)
ql = qs.reshape((n_blocks, -1, 1, block_size // 2)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape(1, 1, 2, 1)
qh = (qh & 1).to(torch.uint8)
ql = (ql & 0x0F).reshape((n_blocks, -1))
qs = (ql | (qh << 4))
return (d * qs) + m
def dequantize_blocks_Q5_0(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
d, qh, qs = split_block_dims(blocks, 2, 4)
d = d.view(torch.float16).to(dtype)
qh = to_uint32(qh)
qh = qh.reshape(n_blocks, 1) >> torch.arange(32, device=d.device, dtype=torch.int32).reshape(1, 32)
ql = qs.reshape(n_blocks, -1, 1, block_size // 2) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape(1, 1, 2, 1)
qh = (qh & 1).to(torch.uint8)
ql = (ql & 0x0F).reshape(n_blocks, -1)
qs = (ql | (qh << 4)).to(torch.int8) - 16
return (d * qs)
def dequantize_blocks_Q4_1(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
d, m, qs = split_block_dims(blocks, 2, 2)
d = d.view(torch.float16).to(dtype)
m = m.view(torch.float16).to(dtype)
qs = qs.reshape((n_blocks, -1, 1, block_size // 2)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape(1, 1, 2, 1)
qs = (qs & 0x0F).reshape(n_blocks, -1)
return (d * qs) + m
def dequantize_blocks_Q4_0(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
d, qs = split_block_dims(blocks, 2)
d = d.view(torch.float16).to(dtype)
qs = qs.reshape((n_blocks, -1, 1, block_size // 2)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape((1, 1, 2, 1))
qs = (qs & 0x0F).reshape((n_blocks, -1)).to(torch.int8) - 8
return (d * qs)
# K Quants #
QK_K = 256
K_SCALE_SIZE = 12
def get_scale_min(scales):
n_blocks = scales.shape[0]
scales = scales.view(torch.uint8)
scales = scales.reshape((n_blocks, 3, 4))
d, m, m_d = torch.split(scales, scales.shape[-2] // 3, dim=-2)
sc = torch.cat([d & 0x3F, (m_d & 0x0F) | ((d >> 2) & 0x30)], dim=-1)
min = torch.cat([m & 0x3F, (m_d >> 4) | ((m >> 2) & 0x30)], dim=-1)
return (sc.reshape((n_blocks, 8)), min.reshape((n_blocks, 8)))
def dequantize_blocks_Q6_K(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
ql, qh, scales, d, = split_block_dims(blocks, QK_K // 2, QK_K // 4, QK_K // 16)
scales = scales.view(torch.int8).to(dtype)
d = d.view(torch.float16).to(dtype)
d = (d * scales).reshape((n_blocks, QK_K // 16, 1))
ql = ql.reshape((n_blocks, -1, 1, 64)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape((1, 1, 2, 1))
ql = (ql & 0x0F).reshape((n_blocks, -1, 32))
qh = qh.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 2, 4, 6], device=d.device, dtype=torch.uint8).reshape((1, 1, 4, 1))
qh = (qh & 0x03).reshape((n_blocks, -1, 32))
q = (ql | (qh << 4)).to(torch.int8) - 32
q = q.reshape((n_blocks, QK_K // 16, -1))
return (d * q).reshape((n_blocks, QK_K))
def dequantize_blocks_Q5_K(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
d, dmin, scales, qh, qs = split_block_dims(blocks, 2, 2, K_SCALE_SIZE, QK_K // 8)
d = d.view(torch.float16).to(dtype)
dmin = dmin.view(torch.float16).to(dtype)
sc, m = get_scale_min(scales)
d = (d * sc).reshape((n_blocks, -1, 1))
dm = (dmin * m).reshape((n_blocks, -1, 1))
ql = qs.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape((1, 1, 2, 1))
qh = qh.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([i for i in range(8)], device=d.device, dtype=torch.uint8).reshape((1, 1, 8, 1))
ql = (ql & 0x0F).reshape((n_blocks, -1, 32))
qh = (qh & 0x01).reshape((n_blocks, -1, 32))
q = (ql | (qh << 4))
return (d * q - dm).reshape((n_blocks, QK_K))
def dequantize_blocks_Q4_K(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
d, dmin, scales, qs = split_block_dims(blocks, 2, 2, K_SCALE_SIZE)
d = d.view(torch.float16).to(dtype)
dmin = dmin.view(torch.float16).to(dtype)
sc, m = get_scale_min(scales)
d = (d * sc).reshape((n_blocks, -1, 1))
dm = (dmin * m).reshape((n_blocks, -1, 1))
qs = qs.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape((1, 1, 2, 1))
qs = (qs & 0x0F).reshape((n_blocks, -1, 32))
return (d * qs - dm).reshape((n_blocks, QK_K))
def dequantize_blocks_Q3_K(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
hmask, qs, scales, d = split_block_dims(blocks, QK_K // 8, QK_K // 4, 12)
d = d.view(torch.float16).to(dtype)
lscales, hscales = scales[:, :8], scales[:, 8:]
lscales = lscales.reshape((n_blocks, 1, 8)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape((1, 2, 1))
lscales = lscales.reshape((n_blocks, 16))
hscales = hscales.reshape((n_blocks, 1, 4)) >> torch.tensor([0, 2, 4, 6], device=d.device, dtype=torch.uint8).reshape((1, 4, 1))
hscales = hscales.reshape((n_blocks, 16))
scales = (lscales & 0x0F) | ((hscales & 0x03) << 4)
scales = (scales.to(torch.int8) - 32)
dl = (d * scales).reshape((n_blocks, 16, 1))
ql = qs.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 2, 4, 6], device=d.device, dtype=torch.uint8).reshape((1, 1, 4, 1))
qh = hmask.reshape(n_blocks, -1, 1, 32) >> torch.tensor([i for i in range(8)], device=d.device, dtype=torch.uint8).reshape((1, 1, 8, 1))
ql = ql.reshape((n_blocks, 16, QK_K // 16)) & 3
qh = (qh.reshape((n_blocks, 16, QK_K // 16)) & 1) ^ 1
q = (ql.to(torch.int8) - (qh << 2).to(torch.int8))
return (dl * q).reshape((n_blocks, QK_K))
def dequantize_blocks_Q2_K(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
scales, qs, d, dmin = split_block_dims(blocks, QK_K // 16, QK_K // 4, 2)
d = d.view(torch.float16).to(dtype)
dmin = dmin.view(torch.float16).to(dtype)
# (n_blocks, 16, 1)
dl = (d * (scales & 0xF)).reshape((n_blocks, QK_K // 16, 1))
ml = (dmin * (scales >> 4)).reshape((n_blocks, QK_K // 16, 1))
shift = torch.tensor([0, 2, 4, 6], device=d.device, dtype=torch.uint8).reshape((1, 1, 4, 1))
qs = (qs.reshape((n_blocks, -1, 1, 32)) >> shift) & 3
qs = qs.reshape((n_blocks, QK_K // 16, 16))
qs = dl * qs - ml
return qs.reshape((n_blocks, -1))
# IQ quants
KVALUES = torch.tensor([-127, -104, -83, -65, -49, -35, -22, -10, 1, 13, 25, 38, 53, 69, 89, 113], dtype=torch.int8)
def dequantize_blocks_IQ4_NL(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
d, qs = split_block_dims(blocks, 2)
d = d.view(torch.float16).to(dtype)
qs = qs.reshape((n_blocks, -1, 1, block_size//2)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape((1, 1, 2, 1))
qs = (qs & 0x0F).reshape((n_blocks, -1, 1)).to(torch.int32)
kvalues = KVALUES.to(qs.device).expand(*qs.shape[:-1], 16)
qs = torch.gather(kvalues, dim=-1, index=qs).reshape((n_blocks, -1))
del kvalues # should still be view, but just to be safe
return (d * qs)
def dequantize_blocks_IQ4_XS(blocks, block_size, type_size, dtype=None):
n_blocks = blocks.shape[0]
d, scales_h, scales_l, qs = split_block_dims(blocks, 2, 2, QK_K // 64)
d = d.view(torch.float16).to(dtype)
scales_h = to_uint16(scales_h)
shift_a = torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape((1, 1, 2))
shift_b = torch.tensor([2 * i for i in range(QK_K // 32)], device=d.device, dtype=torch.uint8).reshape((1, -1, 1))
scales_l = scales_l.reshape((n_blocks, -1, 1)) >> shift_a.reshape((1, 1, 2))
scales_h = scales_h.reshape((n_blocks, -1, 1)) >> shift_b.reshape((1, -1, 1))
scales_l = scales_l.reshape((n_blocks, -1)) & 0x0F
scales_h = scales_h.reshape((n_blocks, -1)).to(torch.uint8) & 0x03
scales = (scales_l | (scales_h << 4)).to(torch.int8) - 32
dl = (d * scales.to(dtype)).reshape((n_blocks, -1, 1))
qs = qs.reshape((n_blocks, -1, 1, 16)) >> shift_a.reshape((1, 1, 2, 1))
qs = qs.reshape((n_blocks, -1, 32, 1)) & 0x0F
kvalues = KVALUES.to(qs.device).expand(*qs.shape[:-1], 16)
qs = torch.gather(kvalues, dim=-1, index=qs.to(torch.int32)).reshape((n_blocks, -1, 32))
del kvalues # see IQ4_NL
del shift_a
del shift_b
return (dl * qs).reshape((n_blocks, -1))
dequantize_functions = {
gguf.GGMLQuantizationType.BF16: dequantize_blocks_BF16,
gguf.GGMLQuantizationType.Q8_0: dequantize_blocks_Q8_0,
gguf.GGMLQuantizationType.Q5_1: dequantize_blocks_Q5_1,
gguf.GGMLQuantizationType.Q5_0: dequantize_blocks_Q5_0,
gguf.GGMLQuantizationType.Q4_1: dequantize_blocks_Q4_1,
gguf.GGMLQuantizationType.Q4_0: dequantize_blocks_Q4_0,
gguf.GGMLQuantizationType.Q6_K: dequantize_blocks_Q6_K,
gguf.GGMLQuantizationType.Q5_K: dequantize_blocks_Q5_K,
gguf.GGMLQuantizationType.Q4_K: dequantize_blocks_Q4_K,
gguf.GGMLQuantizationType.Q3_K: dequantize_blocks_Q3_K,
gguf.GGMLQuantizationType.Q2_K: dequantize_blocks_Q2_K,
gguf.GGMLQuantizationType.IQ4_NL: dequantize_blocks_IQ4_NL,
gguf.GGMLQuantizationType.IQ4_XS: dequantize_blocks_IQ4_XS,
}
+73
View File
@@ -0,0 +1,73 @@
import torch
import folder_paths
import comfy.utils
import comfy.lora
import logging
class ZImageDiffSynthLoader:
"""
A specialized LoRA loader for Z-Image Turbo that automatically fixes
key mismatches from DiffSynth-trained LoRAs.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"clip": ("CLIP",),
"lora_name": (folder_paths.get_filename_list("loras"), ),
"strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
"strength_clip": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
}
}
RETURN_TYPES = ("MODEL", "CLIP")
FUNCTION = "load_lora"
CATEGORY = "Z-Image/Loaders"
def load_lora(self, model, clip, lora_name, strength_model, strength_clip):
lora_path = folder_paths.get_full_path("loras", lora_name)
lora_sd = comfy.utils.load_torch_file(lora_path, safe_load=True)
# --- KEY MAPPING LOGIC ---
# DiffSynth keys are "flat" (e.g. "layers.0.attention...")
# ComfyUI Z-Image expects "diffusion_model." prefix for these.
fixed_sd = {}
for k, v in lora_sd.items():
new_key = k
# 1. Handle specialized Z-Image Refiners
if k.startswith("context_refiner."):
new_key = f"diffusion_model.{k}"
elif k.startswith("noise_refiner."):
new_key = f"diffusion_model.{k}"
# 2. Handle standard Llama-3 Backbone Layers
# If it starts with "layers." it belongs to the DiT backbone
elif k.startswith("layers."):
new_key = f"diffusion_model.{k}"
# 3. Handle Input/Output blocks if they are bare
elif k.startswith("final_layer.") or k.startswith("label_emb."):
new_key = f"diffusion_model.{k}"
fixed_sd[new_key] = v
# --- APPLY PATCHES ---
# We pass the fixed dictionary to ComfyUI's standard loader
model_lora, clip_lora = comfy.sd.load_lora_for_models(
model, clip, fixed_sd, strength_model, strength_clip
)
return (model_lora, clip_lora)
# Node Registration
NODE_CLASS_MAPPINGS = {
"ZImageDiffSynthLoader": ZImageDiffSynthLoader
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZImageDiffSynthLoader": "Z-Image DiffSynth LoRA Loader"
}
+61
View File
@@ -0,0 +1,61 @@
import math
import torch
MAX_RESOLUTION = 8192
class GetImageSizePlus:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"target_width": ("INT", {"default": 1024, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"target_height": ("INT", {"default": 1024, "min": 0, "max": MAX_RESOLUTION, "step": 1}),
"calculate_from_target": ("BOOLEAN", {"default": False}),
"multiple_of": (["8", "16", "32", "64", "None"], {"default": "8"}),
}
}
RETURN_TYPES = ("INT", "INT", "INT",)
RETURN_NAMES = ("width", "height", "count")
FUNCTION = "execute"
CATEGORY = "essentials_mb/image utils"
def execute(self, image, target_width, target_height, calculate_from_target, multiple_of):
# 1. Get current dimensions
orig_width = image.shape[2]
orig_height = image.shape[1]
count = image.shape[0]
# 2. If toggle is off, return original dimensions
if not calculate_from_target:
return (orig_width, orig_height, count)
# 3. Calculate Target Pixel Count (Area)
target_area = target_width * target_height
# 4. Calculate Aspect Ratio
aspect_ratio = orig_width / orig_height
# 5. Calculate new dimensions
new_height = math.sqrt(target_area / aspect_ratio)
new_width = new_height * aspect_ratio
# 6. Snap to nearest multiple (e.g., 8)
if multiple_of == "None":
final_width = int(round(new_width))
final_height = int(round(new_height))
else:
divisor = int(multiple_of)
final_width = int(round(new_width / divisor) * divisor)
final_height = int(round(new_height / divisor) * divisor)
return (final_width, final_height, count)
NODE_CLASS_MAPPINGS = {
"GetImageSizePlus+": GetImageSizePlus
}
NODE_DISPLAY_NAME_MAPPINGS = {
"GetImageSizePlus+": "🔧 Get Image Size Plus"
}
+5
View File
@@ -0,0 +1,5 @@
software_meta = {
"name": "AI Toolkit",
"version": "1.0.0",
"url": "https://github.com/ostris/ai-toolkit"
}
+107
View File
@@ -0,0 +1,107 @@
import re
class KeywordMatchGate:
"""
Keyword Match Gate Node
Acts as a string-based AND gate:
- If the keyword matches (with the chosen options), returns a string.
- If not, returns the literal string "no match".
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"input_text": ("STRING", {
"multiline": True,
"default": "",
}),
"match_keyword": ("STRING", {
"multiline": False,
"default": "",
}),
"case_sensitive": ("BOOLEAN", {
"default": False,
}),
"match_whole_word": ("BOOLEAN", {
"default": True,
}),
"echo_input": ("BOOLEAN", {
"default": False,
}),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("gated_text",)
FUNCTION = "execute"
CATEGORY = "Logic"
@classmethod
def IS_CHANGED(cls, **kwargs):
# Pure function: change when inputs change
return float("nan")
def _match(
self,
input_text: str,
match_keyword: str,
case_sensitive: bool,
match_whole_word: bool,
) -> bool:
# Empty values are treated as "no match"
if not input_text or not match_keyword:
return False
# Case normalization
if not case_sensitive:
norm_input = input_text.lower()
norm_keyword = match_keyword.lower()
else:
norm_input = input_text
norm_keyword = match_keyword
try:
if match_whole_word:
# Whole-word regex using word boundaries
escaped = re.escape(norm_keyword)
pattern = rf"\b{escaped}\b"
return re.search(pattern, norm_input) is not None
else:
# Simple substring check
return norm_keyword in norm_input
except re.error:
# Fail-safe: treat as no match
return False
def execute(
self,
input_text,
match_keyword,
case_sensitive,
match_whole_word,
echo_input,
):
match_found = self._match(
input_text=input_text,
match_keyword=match_keyword,
case_sensitive=case_sensitive,
match_whole_word=match_whole_word,
)
if match_found:
# Gate open: return either the keyword or the original input
return (input_text if echo_input else match_keyword,)
else:
# Gate closed: explicit "no match" so it can be mapped to 0 later
return ("no match",)
NODE_CLASS_MAPPINGS = {
"KeywordMatchGate": KeywordMatchGate,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"KeywordMatchGate": "🔹 Keyword Match Gate",
}
+356
View File
@@ -0,0 +1,356 @@
# (c) City96 || Apache-2.0 (apache.org/licenses/LICENSE-2.0)
import warnings
import logging
import torch
import gguf
import re
import os
from .ops import GGMLTensor
from .dequant import is_quantized, dequantize_tensor
IMG_ARCH_LIST = {"flux", "sd1", "sdxl", "sd3", "aura", "hidream", "cosmos", "ltxv", "hyvid", "wan", "lumina2", "qwen_image"}
TXT_ARCH_LIST = {"t5", "t5encoder", "llama", "qwen2vl", "qwen3"}
VIS_TYPE_LIST = {"clip-vision", "mmproj"}
def get_orig_shape(reader, tensor_name):
field_key = f"comfy.gguf.orig_shape.{tensor_name}"
field = reader.get_field(field_key)
if field is None:
return None
# Has original shape metadata, so we try to decode it.
if len(field.types) != 2 or field.types[0] != gguf.GGUFValueType.ARRAY or field.types[1] != gguf.GGUFValueType.INT32:
raise TypeError(f"Bad original shape metadata for {field_key}: Expected ARRAY of INT32, got {field.types}")
return torch.Size(tuple(int(field.parts[part_idx][0]) for part_idx in field.data))
def get_field(reader, field_name, field_type):
field = reader.get_field(field_name)
if field is None:
return None
elif field_type == str:
# extra check here as this is used for checking arch string
if len(field.types) != 1 or field.types[0] != gguf.GGUFValueType.STRING:
raise TypeError(f"Bad type for GGUF {field_name} key: expected string, got {field.types!r}")
return str(field.parts[field.data[-1]], encoding="utf-8")
elif field_type in [int, float, bool]:
return field_type(field.parts[field.data[-1]])
else:
raise TypeError(f"Unknown field type {field_type}")
def get_list_field(reader, field_name, field_type):
field = reader.get_field(field_name)
if field is None:
return None
elif field_type == str:
return tuple(str(field.parts[part_idx], encoding="utf-8") for part_idx in field.data)
elif field_type in [int, float, bool]:
return tuple(field_type(field.parts[part_idx][0]) for part_idx in field.data)
else:
raise TypeError(f"Unknown field type {field_type}")
def gguf_sd_loader(path, handle_prefix="model.diffusion_model.", return_arch=False, is_text_model=False):
"""
Read state dict as fake tensors
"""
reader = gguf.GGUFReader(path)
# filter and strip prefix
has_prefix = False
if handle_prefix is not None:
prefix_len = len(handle_prefix)
tensor_names = set(tensor.name for tensor in reader.tensors)
has_prefix = any(s.startswith(handle_prefix) for s in tensor_names)
tensors = []
for tensor in reader.tensors:
sd_key = tensor_name = tensor.name
if has_prefix:
if not tensor_name.startswith(handle_prefix):
continue
sd_key = tensor_name[prefix_len:]
tensors.append((sd_key, tensor))
# detect and verify architecture
compat = None
arch_str = get_field(reader, "general.architecture", str)
type_str = get_field(reader, "general.type", str)
if arch_str in [None, "pig"]:
if is_text_model:
raise ValueError(f"This text model is incompatible with llama.cpp!\nConsider using the safetensors version\n({path})")
compat = "sd.cpp" if arch_str is None else arch_str
# import here to avoid changes to convert.py breaking regular models
from .tools.convert import detect_arch
try:
arch_str = detect_arch(set(val[0] for val in tensors)).arch
except Exception as e:
raise ValueError(f"This model is not currently supported - ({e})")
elif arch_str not in TXT_ARCH_LIST and is_text_model:
if type_str not in VIS_TYPE_LIST:
raise ValueError(f"Unexpected text model architecture type in GGUF file: {arch_str!r}")
elif arch_str not in IMG_ARCH_LIST and not is_text_model:
raise ValueError(f"Unexpected architecture type in GGUF file: {arch_str!r}")
if compat:
logging.warning(f"Warning: This gguf model file is loaded in compatibility mode '{compat}' [arch:{arch_str}]")
# main loading loop
state_dict = {}
qtype_dict = {}
for sd_key, tensor in tensors:
tensor_name = tensor.name
# torch_tensor = torch.from_numpy(tensor.data) # mmap
# NOTE: line above replaced with this block to avoid persistent numpy warning about mmap
with warnings.catch_warnings():
warnings.filterwarnings("ignore", message="The given NumPy array is not writable")
torch_tensor = torch.from_numpy(tensor.data) # mmap
shape = get_orig_shape(reader, tensor_name)
if shape is None:
shape = torch.Size(tuple(int(v) for v in reversed(tensor.shape)))
# Workaround for stable-diffusion.cpp SDXL detection.
if compat == "sd.cpp" and arch_str == "sdxl":
if any([tensor_name.endswith(x) for x in (".proj_in.weight", ".proj_out.weight")]):
while len(shape) > 2 and shape[-1] == 1:
shape = shape[:-1]
# add to state dict
if tensor.tensor_type in {gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16}:
torch_tensor = torch_tensor.view(*shape)
state_dict[sd_key] = GGMLTensor(torch_tensor, tensor_type=tensor.tensor_type, tensor_shape=shape)
# keep track of loaded tensor types
tensor_type_str = getattr(tensor.tensor_type, "name", repr(tensor.tensor_type))
qtype_dict[tensor_type_str] = qtype_dict.get(tensor_type_str, 0) + 1
# print loaded tensor type counts
logging.info("gguf qtypes: " + ", ".join(f"{k} ({v})" for k, v in qtype_dict.items()))
# mark largest tensor for vram estimation
qsd = {k:v for k,v in state_dict.items() if is_quantized(v)}
if len(qsd) > 0:
max_key = max(qsd.keys(), key=lambda k: qsd[k].numel())
state_dict[max_key].is_largest_weight = True
if return_arch:
return (state_dict, arch_str)
return state_dict
# for remapping llama.cpp -> original key names
T5_SD_MAP = {
"enc.": "encoder.",
".blk.": ".block.",
"token_embd": "shared",
"output_norm": "final_layer_norm",
"attn_q": "layer.0.SelfAttention.q",
"attn_k": "layer.0.SelfAttention.k",
"attn_v": "layer.0.SelfAttention.v",
"attn_o": "layer.0.SelfAttention.o",
"attn_norm": "layer.0.layer_norm",
"attn_rel_b": "layer.0.SelfAttention.relative_attention_bias",
"ffn_up": "layer.1.DenseReluDense.wi_1",
"ffn_down": "layer.1.DenseReluDense.wo",
"ffn_gate": "layer.1.DenseReluDense.wi_0",
"ffn_norm": "layer.1.layer_norm",
}
LLAMA_SD_MAP = {
"blk.": "model.layers.",
"attn_norm": "input_layernorm",
"attn_q_norm.": "self_attn.q_norm.",
"attn_k_norm.": "self_attn.k_norm.",
"attn_v_norm.": "self_attn.v_norm.",
"attn_q": "self_attn.q_proj",
"attn_k": "self_attn.k_proj",
"attn_v": "self_attn.v_proj",
"attn_output": "self_attn.o_proj",
"ffn_up": "mlp.up_proj",
"ffn_down": "mlp.down_proj",
"ffn_gate": "mlp.gate_proj",
"ffn_norm": "post_attention_layernorm",
"token_embd": "model.embed_tokens",
"output_norm": "model.norm",
"output.weight": "lm_head.weight",
}
CLIP_VISION_SD_MAP = {
"mm.": "visual.merger.mlp.",
"v.post_ln.": "visual.merger.ln_q.",
"v.patch_embd": "visual.patch_embed.proj",
"v.blk.": "visual.blocks.",
"ffn_up": "mlp.up_proj",
"ffn_down": "mlp.down_proj",
"ffn_gate": "mlp.gate_proj",
"attn_out.": "attn.proj.",
"ln1.": "norm1.",
"ln2.": "norm2.",
}
def sd_map_replace(raw_sd, key_map):
sd = {}
for k,v in raw_sd.items():
for s,d in key_map.items():
k = k.replace(s,d)
sd[k] = v
return sd
def llama_permute(raw_sd, n_head, n_head_kv):
# Reverse version of LlamaModel.permute in llama.cpp convert script
sd = {}
permute = lambda x,h: x.reshape(h, x.shape[0] // h // 2, 2, *x.shape[1:]).swapaxes(1, 2).reshape(x.shape)
for k,v in raw_sd.items():
if k.endswith(("q_proj.weight", "q_proj.bias")):
v.data = permute(v.data, n_head)
if k.endswith(("k_proj.weight", "k_proj.bias")):
v.data = permute(v.data, n_head_kv)
sd[k] = v
return sd
def strip_quant_suffix(name):
pattern = r"[-_]?(?:ud-)?i?q[0-9]_[a-z0-9_\-]{1,8}$"
match = re.search(pattern, name, re.IGNORECASE)
if match:
name = name[:match.start()]
return name
def gguf_mmproj_loader(path):
# Reverse version of Qwen2VLVisionModel.modify_tensors
logging.info("Attenpting to find mmproj file for text encoder...")
# get name to match w/o quant suffix
tenc_fname = os.path.basename(path)
tenc = os.path.splitext(tenc_fname)[0].lower()
tenc = strip_quant_suffix(tenc)
# try and find matching mmproj
target = []
root = os.path.dirname(path)
for fname in os.listdir(root):
name, ext = os.path.splitext(fname)
if ext.lower() != ".gguf":
continue
if "mmproj" not in name.lower():
continue
if tenc in name.lower():
target.append(fname)
if len(target) == 0:
logging.error(f"Error: Can't find mmproj file for '{tenc_fname}' (matching:'{tenc}')! Qwen-Image-Edit will be broken!")
return {}
if len(target) > 1:
logging.error(f"Ambiguous mmproj for text encoder '{tenc_fname}', will use first match.")
logging.info(f"Using mmproj '{target[0]}' for text encoder '{tenc_fname}'.")
target = os.path.join(root, target[0])
vsd = gguf_sd_loader(target, is_text_model=True)
# concat 4D to 5D
if "v.patch_embd.weight.1" in vsd:
w1 = dequantize_tensor(vsd.pop("v.patch_embd.weight"), dtype=torch.float32)
w2 = dequantize_tensor(vsd.pop("v.patch_embd.weight.1"), dtype=torch.float32)
vsd["v.patch_embd.weight"] = torch.stack([w1, w2], dim=2)
# run main replacement
vsd = sd_map_replace(vsd, CLIP_VISION_SD_MAP)
# handle split Q/K/V
if "visual.blocks.0.attn_q.weight" in vsd:
attns = {}
# filter out attentions + group
for k,v in vsd.items():
if any(x in k for x in ["attn_q", "attn_k", "attn_v"]):
k_attn, k_name = k.rsplit(".attn_", 1)
k_attn += ".attn.qkv." + k_name.split(".")[-1]
if k_attn not in attns:
attns[k_attn] = {}
attns[k_attn][k_name] = dequantize_tensor(
v, dtype=(torch.bfloat16 if is_quantized(v) else torch.float16)
)
# recombine
for k,v in attns.items():
suffix = k.split(".")[-1]
vsd[k] = torch.cat([
v[f"q.{suffix}"],
v[f"k.{suffix}"],
v[f"v.{suffix}"],
], dim=0)
del attns
return vsd
def gguf_tokenizer_loader(path, temb_shape):
# convert gguf tokenizer to spiece
logging.info("Attempting to recreate sentencepiece tokenizer from GGUF file metadata...")
try:
from sentencepiece import sentencepiece_model_pb2 as model
except ImportError:
raise ImportError("Please make sure sentencepiece and protobuf are installed.\npip install sentencepiece protobuf")
spm = model.ModelProto()
reader = gguf.GGUFReader(path)
if get_field(reader, "tokenizer.ggml.model", str) == "t5":
if temb_shape == (256384, 4096): # probably UMT5
spm.trainer_spec.model_type == 1 # Unigram (do we have a T5 w/ BPE?)
else:
raise NotImplementedError("Unknown model, can't set tokenizer!")
else:
raise NotImplementedError("Unknown model, can't set tokenizer!")
spm.normalizer_spec.add_dummy_prefix = get_field(reader, "tokenizer.ggml.add_space_prefix", bool)
spm.normalizer_spec.remove_extra_whitespaces = get_field(reader, "tokenizer.ggml.remove_extra_whitespaces", bool)
tokens = get_list_field(reader, "tokenizer.ggml.tokens", str)
scores = get_list_field(reader, "tokenizer.ggml.scores", float)
toktypes = get_list_field(reader, "tokenizer.ggml.token_type", int)
for idx, (token, score, toktype) in enumerate(zip(tokens, scores, toktypes)):
# # These aren't present in the original?
# if toktype == 5 and idx >= temb_shape[0]%1000):
# continue
piece = spm.SentencePiece()
piece.piece = token
piece.score = score
piece.type = toktype
spm.pieces.append(piece)
# unsure if any of these are correct
spm.trainer_spec.byte_fallback = True
spm.trainer_spec.vocab_size = len(tokens) # split off unused?
spm.trainer_spec.max_sentence_length = 4096
spm.trainer_spec.eos_id = get_field(reader, "tokenizer.ggml.eos_token_id", int)
spm.trainer_spec.pad_id = get_field(reader, "tokenizer.ggml.padding_token_id", int)
logging.info(f"Created tokenizer with vocab size of {len(spm.pieces)}")
del reader
return torch.ByteTensor(list(spm.SerializeToString()))
def gguf_clip_loader(path):
sd, arch = gguf_sd_loader(path, return_arch=True, is_text_model=True)
if arch in {"t5", "t5encoder"}:
temb_key = "token_embd.weight"
if temb_key in sd and sd[temb_key].shape == (256384, 4096):
# non-standard Comfy-Org tokenizer
sd["spiece_model"] = gguf_tokenizer_loader(path, sd[temb_key].shape)
# TODO: dequantizing token embed here is janky but otherwise we OOM due to tensor being massive.
logging.warning(f"Dequantizing {temb_key} to prevent runtime OOM.")
sd[temb_key] = dequantize_tensor(sd[temb_key], dtype=torch.float16)
sd = sd_map_replace(sd, T5_SD_MAP)
elif arch in {"llama", "qwen2vl", "qwen3"}:
# TODO: pass model_options["vocab_size"] to loader somehow
temb_key = "token_embd.weight"
if temb_key in sd and sd[temb_key].shape[0] >= (64 * 1024):
# See note above for T5.
logging.warning(f"Dequantizing {temb_key} to prevent runtime OOM.")
sd[temb_key] = dequantize_tensor(sd[temb_key], dtype=torch.float16)
sd = sd_map_replace(sd, LLAMA_SD_MAP)
if arch == "llama":
sd = llama_permute(sd, 32, 8) # L3
if arch == "qwen2vl":
vsd = gguf_mmproj_loader(path)
sd.update(vsd)
else:
pass
return sd
+831
View File
@@ -0,0 +1,831 @@
import torch
import folder_paths
import comfy.sd
import comfy.utils
from safetensors.torch import load_file, save_file
import logging
import os
import sys
import re
import glob
from unittest.mock import patch
# ==============================================================================
# 1. MATH HELPERS
# ==============================================================================
def robust_matmul(a, b):
if len(a.shape) == 4 and a.shape[2] == 1 and a.shape[3] == 1: a = a.squeeze(3).squeeze(2)
if len(b.shape) == 4 and b.shape[2] == 1 and b.shape[3] == 1: b = b.squeeze(3).squeeze(2)
if a.shape[-1] == b.shape[0]: return a @ b
elif b.shape[-1] == a.shape[0]: return b @ a
elif a.shape[0] == b.shape[0]: return a.T @ b
elif a.shape[-1] == b.shape[-1]: return a @ b.T
else: return None
def make_kron(w1, w2, scale):
if len(w2.shape) == 4: w1 = w1.unsqueeze(2).unsqueeze(2)
w2 = w2.contiguous()
return torch.kron(w1, w2) * scale
def make_hada(w1a, w1b, w2a, w2b, scale):
w1 = robust_matmul(w1a, w1b)
w2 = robust_matmul(w2a, w2b)
if w1 is None or w2 is None: return None
return (w1 * w2) * scale
def make_lora(wa, wb, scale):
res = robust_matmul(wa, wb)
if res is None: return None
return res * scale
# ==============================================================================
# 2. PATCHING ENGINES
# ==============================================================================
def natural_sort_key(s):
return [int(c) if c.isdigit() else c for c in re.split(r'(\d+)', s)]
def group_lora_keys(lora_sd):
modules = {}
suffixes = [
".lora_A.weight", ".lora_B.weight",
".lora_up.weight", ".lora_down.weight",
".hada_w1_a", ".hada_w1_b", ".hada_w2_a", ".hada_w2_b",
".lokr_w1", ".lokr_w2",
".alpha"
]
for key, value in lora_sd.items():
base = key
param_name = None
for s in suffixes:
if key.endswith(s):
base = key[:-len(s)]
param_name = s.strip(".")
break
if param_name is None:
if key.endswith(".weight"):
base = key[:-7]
param_name = "weight"
else:
continue
if base not in modules:
modules[base] = {}
modules[base][param_name] = value
return modules
def apply_lycoris_to_dict(target_dict, lora_path, strength, is_clip=False):
if strength == 0 or target_dict is None: return 0
filename = os.path.basename(lora_path)
logging.info(f"Applying LyCORIS Patch: {filename}")
lora_sd = load_file(lora_path)
modules = group_lora_keys(lora_sd)
sorted_prefixes = sorted(modules.keys(), key=natural_sort_key)
target_linear_keys = []
for k in sorted(target_dict.keys(), key=natural_sort_key):
v = target_dict[k]
if k.endswith(".weight") and len(v.shape) >= 2:
target_linear_keys.append(k)
patch_count = 0
used_targets = set()
for i, prefix in enumerate(sorted_prefixes):
params = modules[prefix]
diff = None
alpha = params.get("alpha", None)
try:
if "hada_w1_a" in params:
w1a, w1b = params["hada_w1_a"].float(), params["hada_w1_b"].float()
w2a, w2b = params["hada_w2_a"].float(), params["hada_w2_b"].float()
alpha_val = float(alpha) if alpha else float(min(w1a.shape))
base_scale = alpha_val / min(w1a.shape)
diff = make_hada(w1a, w1b, w2a, w2b, base_scale * strength)
elif "lokr_w1" in params:
w1, w2 = params["lokr_w1"].float(), params["lokr_w2"].float()
dim = w1.shape[0] * w2.shape[0]
alpha_val = float(alpha) if alpha else float(dim)
base_scale = alpha_val / dim
diff = make_kron(w1, w2, base_scale * strength)
elif "lora_up.weight" in params or "lora_A.weight" in params:
up_key = "lora_up.weight" if "lora_up.weight" in params else "lora_A.weight"
down_key = "lora_down.weight" if "lora_down.weight" in params else "lora_B.weight"
if "lora_A.weight" in params and "lora_B.weight" in params:
down = params["lora_A.weight"].float()
up = params["lora_B.weight"].float()
else:
up = params[up_key].float()
down = params[down_key].float()
alpha_val = float(alpha) if alpha else float(down.shape[0])
base_scale = alpha_val / down.shape[0]
diff = make_lora(up, down, base_scale * strength)
except Exception:
continue
if diff is None: continue
applied = False
clean_prefix = prefix.replace("lora_unet_", "").replace("lora_te_", "").replace("transformer.", "")
candidates = []
candidates.append(f"{clean_prefix}.weight")
if is_clip:
candidates.append(f"text_encoder.{clean_prefix}.weight")
candidates.append(f"model.{clean_prefix}.weight")
candidates.append(f"transformer.{clean_prefix}.weight")
if "text_model" in clean_prefix:
candidates.append(f"{clean_prefix.replace('text_model.', '')}.weight")
else:
candidates.append(f"diffusion_model.{clean_prefix}.weight")
candidates.append(f"model.{clean_prefix}.weight")
candidates.append(f"transformer.{clean_prefix}.weight")
candidates.append(f"model.diffusion_model.{clean_prefix}.weight")
if "blocks" in clean_prefix:
swapped = clean_prefix.replace("blocks", "layers")
candidates.append(f"{swapped}.weight")
candidates.append(f"model.{swapped}.weight")
candidates.append(f"transformer.{swapped}.weight")
for target_key in candidates:
if target_key in target_dict:
target_param = target_dict[target_key]
if target_param.shape == diff.shape:
target_dict[target_key] = target_param + diff.to(target_param.dtype)
patch_count += 1
applied = True
used_targets.add(target_key)
break
if not applied and ("modules" in prefix or "blocks" in prefix):
for t_key in target_linear_keys:
if t_key in used_targets: continue
target_param = target_dict[t_key]
if target_param.shape == diff.shape or target_param.shape == diff.T.shape:
if target_param.shape == diff.T.shape:
diff = diff.T
target_dict[t_key] = target_param + diff.to(target_param.dtype)
patch_count += 1
applied = True
used_targets.add(t_key)
break
return patch_count
def apply_aitk_to_dict(target_dict, lora_path, strength, is_clip=False):
if strength == 0 or target_dict is None: return 0
filename = os.path.basename(lora_path)
logging.info(f"Applying AITK Patch: {filename}")
lora_sd = load_file(lora_path)
modules = group_lora_keys(lora_sd)
patch_count = 0
for prefix, params in modules.items():
try:
if "lora_A.weight" in params and "lora_B.weight" in params:
down = params["lora_A.weight"].float()
up = params["lora_B.weight"].float()
elif "lora_up.weight" in params and "lora_down.weight" in params:
up = params["lora_up.weight"].float()
down = params["lora_down.weight"].float()
else:
continue
alpha = params.get("alpha", None)
alpha_val = float(alpha) if alpha else float(down.shape[0])
scale = (alpha_val / down.shape[0]) * strength
diff = make_lora(up, down, scale)
if diff is None: continue
except Exception:
continue
target_keys = []
if is_clip:
if prefix.startswith("lora_te."):
bare = prefix.replace("lora_te.", "")
target_keys.append(f"transformer.{bare}.weight")
target_keys.append(f"{bare}.weight")
else:
if "lora_unet" in prefix:
bare = prefix.replace("lora_unet_", "").replace("lora_unet.", "")
target_keys.append(f"{bare}.weight")
target_keys.append(f"diffusion_model.{bare}.weight")
if "blocks" in bare:
swapped = bare.replace("blocks", "layers")
target_keys.append(f"{swapped}.weight")
target_keys.append(f"diffusion_model.{swapped}.weight")
for t_key in target_keys:
if t_key in target_dict:
w = target_dict[t_key]
if w.shape == diff.shape:
target_dict[t_key] = w + diff.to(w.dtype)
patch_count += 1
break
return patch_count
def merge_raw_state_dicts(sd_a, sd_b, strength):
m = {}
keys = set(sd_a.keys()) | set(sd_b.keys())
for k in keys:
if k in sd_a and k in sd_b:
wa = sd_a[k]
wb = sd_b[k]
if wa.shape == wb.shape:
res = wa.to(dtype=torch.float32) * (1.0 - strength) + wb.to(dtype=torch.float32) * strength
m[k] = res.to(dtype=wa.dtype)
else:
logging.warning(f"Merge shape mismatch for {k}. Keeping A.")
m[k] = wa
elif k in sd_a:
m[k] = sd_a[k]
else:
m[k] = sd_b[k]
return m
# ==============================================================================
# 3. NODES
# ==============================================================================
class ZImageRawWrapper:
def __init__(self, sd, path):
self.sd = sd
self.path = path
class ZImageAITKLoRALoader:
@classmethod
def INPUT_TYPES(s):
models = ["None"] + folder_paths.get_filename_list("diffusion_models")
tes = ["None"] + folder_paths.get_filename_list("text_encoders")
loras = ["None"] + folder_paths.get_filename_list("loras")
return {
"required": {
"transformer_name": (models, ),
"text_encoder_name": (tes, ),
"lora_name": (loras, ),
"strength_model": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step": 0.01}),
"strength_clip": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step": 0.01}),
},
"optional": {
"lora_stack": ("LYCORIS_STACK", )
}
}
RETURN_TYPES = ("RAW_MODEL", "RAW_CLIP")
FUNCTION = "load_and_patch_aitk"
CATEGORY = "Z-Image/Loaders"
def load_and_patch_aitk(self, transformer_name, text_encoder_name, lora_name, strength_model, strength_clip, lora_stack=None):
unet_sd = None
clip_sd = None
unet_path = None
clip_path = None
# -------------------------------------------------
# Load Base Model / CLIP
# -------------------------------------------------
if transformer_name != "None":
unet_path = folder_paths.get_full_path("diffusion_models", transformer_name)
logging.info(f"Loading Model: {os.path.basename(unet_path)}")
unet_sd = load_file(unet_path)
if text_encoder_name != "None":
clip_path = folder_paths.get_full_path("text_encoders", text_encoder_name)
if not clip_path:
clip_path = folder_paths.get_full_path("checkpoints", text_encoder_name)
logging.info(f"Loading CLIP: {os.path.basename(clip_path)}")
clip_sd = load_file(clip_path)
# -------------------------------------------------
# Build Patch Job List (STACK SUPPORT ADDED HERE)
# -------------------------------------------------
jobs = []
# Add stacked LoRAs first (if provided)
if lora_stack:
jobs.extend(lora_stack)
# Add single LoRA input
if lora_name != "None":
lora_path = folder_paths.get_full_path("loras", lora_name)
if lora_path:
jobs.append({
"path": lora_path,
"str_model": strength_model,
"str_clip": strength_clip
})
# -------------------------------------------------
# Apply All AITK LoRAs Sequentially
# -------------------------------------------------
if jobs:
logging.info(f"Applying {len(jobs)} AITK LoRA Patches...")
for job in jobs:
lora_path = job["path"]
if job["str_model"] != 0 and unet_sd:
c = apply_aitk_to_dict(unet_sd, lora_path, job["str_model"], is_clip=False)
logging.info(f"AITK Model Patched ({os.path.basename(lora_path)}): {c} layers")
if job["str_clip"] != 0 and clip_sd:
c = apply_aitk_to_dict(clip_sd, lora_path, job["str_clip"], is_clip=True)
logging.info(f"AITK CLIP Patched ({os.path.basename(lora_path)}): {c} layers")
raw_model = ZImageRawWrapper(unet_sd, unet_path) if unet_sd else None
raw_clip = ZImageRawWrapper(clip_sd, clip_path) if clip_sd else None
return (raw_model, raw_clip)
class ZImageDiffusersLoader:
"""
Robust loader for Diffusers folders with automatic key conversion.
Matches z_image_convert_original_to_comfy.py logic internally.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"folder_path": ("STRING", {"default": "", "multiline": False}),
"load_transformer": ("BOOLEAN", {"default": True}),
"load_text_encoder": ("BOOLEAN", {"default": True}),
"load_vae": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("RAW_MODEL", "RAW_CLIP", "VAE")
FUNCTION = "load_diffusers"
CATEGORY = "Z-Image/Loaders"
def _load_sharded_or_single(self, folder):
index_files = glob.glob(os.path.join(folder, "*.index.json"))
if index_files:
logging.info(f"Z-Image: Detected sharded model in {folder}")
combined_sd = {}
shards = glob.glob(os.path.join(folder, "*.safetensors"))
if not shards: shards = glob.glob(os.path.join(folder, "*.bin"))
for shard in shards:
if shard.endswith(".safetensors"):
sd = load_file(shard)
else:
sd = torch.load(shard, map_location="cpu")
combined_sd.update(sd)
return combined_sd, shards[0]
candidates = [
"model.safetensors", "model.fp16.safetensors",
"diffusion_pytorch_model.safetensors", "diffusion_pytorch_model.fp16.safetensors",
"diffusion_pytorch_model.bin", "model.bin"
]
for fn in candidates:
p = os.path.join(folder, fn)
if os.path.exists(p):
logging.info(f"Z-Image: Found weights {fn}")
if fn.endswith(".safetensors"):
return load_file(p), p
else:
return torch.load(p, map_location="cpu"), p
return None, None
def _convert_diffusers_to_comfy(self, sd):
"""
Converts keys using a robust two-pass method to ensure fusion works
regardless of dictionary key order.
"""
new_sd = {}
handled_keys = set()
# Mappings from user's script + standard maps
replace_keys = {
"all_final_layer.2-1.": "final_layer.",
"all_x_embedder.2-1.": "x_embedder.",
".attention.to_out.0.bias": ".attention.out.bias",
".attention.norm_k.weight": ".attention.k_norm.weight",
".attention.norm_q.weight": ".attention.q_norm.weight",
".attention.to_out.0.weight": ".attention.out.weight"
}
# --- Pass 1: QKV Fusion (Scan for to_q) ---
for k in sd.keys():
if ".attention.to_q.weight" in k:
k_q = k
k_k = k.replace(".attention.to_q.weight", ".attention.to_k.weight")
k_v = k.replace(".attention.to_q.weight", ".attention.to_v.weight")
if k_k in sd and k_v in sd:
try:
# Weights
w_q, w_k, w_v = sd[k_q], sd[k_k], sd[k_v]
fused_w = torch.cat([w_q, w_k, w_v], dim=0)
# Generate output key
# 1. Apply replacements to base key (k_q)
clean_k = k_q
for r, rr in replace_keys.items():
clean_k = clean_k.replace(r, rr)
# 2. Strip prefixes
for p in ["transformer.", "unet.", "diffusion_model."]:
if clean_k.startswith(p): clean_k = clean_k[len(p):]
# 3. Swap suffix
final_key_w = clean_k.replace(".attention.to_q.weight", ".attention.qkv.weight")
new_sd[final_key_w] = fused_w
handled_keys.update([k_q, k_k, k_v])
# Biases (Optional)
k_q_b = k_q.replace("weight", "bias")
k_k_b = k_k.replace("weight", "bias")
k_v_b = k_v.replace("weight", "bias")
if k_q_b in sd and k_k_b in sd and k_v_b in sd:
b_q, b_k, b_v = sd[k_q_b], sd[k_k_b], sd[k_v_b]
fused_b = torch.cat([b_q, b_k, b_v], dim=0)
final_key_b = final_key_w.replace("weight", "bias")
new_sd[final_key_b] = fused_b
handled_keys.update([k_q_b, k_k_b, k_v_b])
except Exception as e:
logging.warning(f"Z-Image: QKV Fusion failed for {k}: {e}")
# --- Pass 2: Process Remaining Keys ---
for k in sd.keys():
if k in handled_keys: continue
# Apply replacements
clean_k = k
for r, rr in replace_keys.items():
clean_k = clean_k.replace(r, rr)
# Regex Cleaning (artifacts)
clean_k = re.sub(r'^all_', '', clean_k)
clean_k = re.sub(r'\.\d+-\d+', '', clean_k) # .2-1
# Prefix Stripping
for p in ["transformer.", "unet.", "diffusion_model."]:
if clean_k.startswith(p):
clean_k = clean_k[len(p):]
break
# Text Encoder
if clean_k.startswith("text_model."): clean_k = clean_k.replace("text_model.", "")
new_sd[clean_k] = sd[k]
return new_sd
def load_diffusers(self, folder_path, load_transformer, load_text_encoder, load_vae):
folder_path = folder_path.strip().strip('"').strip("'")
raw_model = None
raw_clip = None
vae = None
if not os.path.isdir(folder_path):
raise FileNotFoundError(f"Diffusers folder not found: {folder_path}")
# --- 1. Transformer / UNet ---
if load_transformer:
sd = None
path = None
search_dirs = [os.path.join(folder_path, "transformer"), os.path.join(folder_path, "unet"), folder_path]
for d in search_dirs:
if os.path.isdir(d):
sd, path = self._load_sharded_or_single(d)
if sd: break
if sd:
sd = self._convert_diffusers_to_comfy(sd)
raw_model = ZImageRawWrapper(sd, path)
else:
raise FileNotFoundError(f"Z-Image: Could not find Transformer/UNet weights in {folder_path}. Checked: {search_dirs}")
# --- 2. Text Encoder ---
if load_text_encoder:
sd = None
path = None
search_dirs = [os.path.join(folder_path, "text_encoder"), os.path.join(folder_path, "text_encoder_2"), folder_path]
for d in search_dirs:
if os.path.isdir(d):
sd, path = self._load_sharded_or_single(d)
if sd: break
if sd:
new_sd = {}
for k, v in sd.items():
new_k = k
if new_k.startswith("text_model."): new_k = new_k.replace("text_model.", "")
new_sd[new_k] = v
raw_clip = ZImageRawWrapper(new_sd, path)
else:
raise FileNotFoundError(f"Z-Image: Could not find Text Encoder weights in {folder_path}. Checked: {search_dirs}")
# --- 3. VAE ---
if load_vae:
path = None
search_dirs = [os.path.join(folder_path, "vae"), folder_path]
for d in search_dirs:
if os.path.isdir(d):
_, path = self._load_sharded_or_single(d)
if path: break
if path:
vae = comfy.sd.load_vae(path)
else:
raise FileNotFoundError(f"Z-Image: Could not find VAE weights in {folder_path}")
return (raw_model, raw_clip, vae)
class ZImageLoaderAndPatcher:
@classmethod
def INPUT_TYPES(s):
models = ["None"] + folder_paths.get_filename_list("diffusion_models")
tes = ["None"] + folder_paths.get_filename_list("text_encoders")
loras = ["None"] + folder_paths.get_filename_list("loras")
return {
"required": {
"transformer_name": (models, ),
"text_encoder_name": (tes, ),
"lora_name": (loras, ),
"strength_model": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step": 0.01}),
"strength_clip": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step": 0.01}),
},
"optional": {"lora_stack": ("LYCORIS_STACK", )}
}
RETURN_TYPES = ("RAW_MODEL", "RAW_CLIP")
FUNCTION = "load_and_patch_raw"
CATEGORY = "Z-Image/Loaders"
def load_and_patch_raw(self, transformer_name, text_encoder_name, lora_name, strength_model, strength_clip, lora_stack=None):
unet_sd = None
clip_sd = None
unet_path = None
clip_path = None
if transformer_name != "None":
unet_path = folder_paths.get_full_path("diffusion_models", transformer_name)
logging.info(f"Loading Model: {os.path.basename(unet_path)}")
unet_sd = load_file(unet_path)
if text_encoder_name != "None":
clip_path = folder_paths.get_full_path("text_encoders", text_encoder_name)
if not clip_path: clip_path = folder_paths.get_full_path("checkpoints", text_encoder_name)
logging.info(f"Loading CLIP: {os.path.basename(clip_path)}")
clip_sd = load_file(clip_path)
jobs = []
if lora_stack: jobs.extend(lora_stack)
if lora_name != "None":
path = folder_paths.get_full_path("loras", lora_name)
if path: jobs.append({"path": path, "str_model": strength_model, "str_clip": strength_clip})
if jobs:
logging.info(f"Applying {len(jobs)} Manual Patches...")
for job in jobs:
if job["str_model"] != 0 and unet_sd:
c = apply_lycoris_to_dict(unet_sd, job["path"], job["str_model"], is_clip=False)
logging.info(f"Model Layers Patched: {c}")
if job["str_clip"] != 0 and clip_sd:
c = apply_lycoris_to_dict(clip_sd, job["path"], job["str_clip"], is_clip=True)
logging.info(f"CLIP Layers Patched: {c}")
raw_model = ZImageRawWrapper(unet_sd, unet_path) if unet_sd else None
raw_clip = ZImageRawWrapper(clip_sd, clip_path) if clip_sd else None
return (raw_model, raw_clip)
class ZImageLycorisStacker:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"lora_name": (folder_paths.get_filename_list("loras"), ),
"strength_model": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step": 0.01}),
"strength_clip": ("FLOAT", {"default": 1.0, "min": -100.0, "max": 100.0, "step": 0.01}),
},
"optional": {
"input_stack": ("LYCORIS_STACK", ),
}
}
RETURN_TYPES = ("LYCORIS_STACK",)
FUNCTION = "stack_lora"
CATEGORY = "Z-Image/Loaders"
def stack_lora(self, lora_name, strength_model, strength_clip, input_stack=None):
stack = []
if input_stack: stack.extend(input_stack)
lora_path = folder_paths.get_full_path("loras", lora_name)
if lora_path:
stack.append({"path": lora_path, "str_model": strength_model, "str_clip": strength_clip})
return (stack,)
class ZImageRawModelMerge:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_a": ("RAW_MODEL",),
"model_b": ("RAW_MODEL",),
"strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
}
}
RETURN_TYPES = ("RAW_MODEL",)
FUNCTION = "merge"
CATEGORY = "Z-Image/Loaders"
def merge(self, model_a, model_b, strength):
if not model_a or not model_b:
return (model_a if model_a else model_b,)
logging.info(f"Merging models with strength {strength}")
merged_sd = merge_raw_state_dicts(model_a.sd, model_b.sd, strength)
return (ZImageRawWrapper(merged_sd, "merged_model.safetensors"),)
class ZImageRawClipMerge:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"clip_a": ("RAW_CLIP",),
"clip_b": ("RAW_CLIP",),
"strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
}
}
RETURN_TYPES = ("RAW_CLIP",)
FUNCTION = "merge"
CATEGORY = "Z-Image/Loaders"
def merge(self, clip_a, clip_b, strength):
if not clip_a or not clip_b:
return (clip_a if clip_a else clip_b,)
logging.info(f"Merging CLIPs with strength {strength}")
merged_sd = merge_raw_state_dicts(clip_a.sd, clip_b.sd, strength)
return (ZImageRawWrapper(merged_sd, "merged_clip.safetensors"),)
class ZImageComfyInjector:
@classmethod
def INPUT_TYPES(s):
return {
"required": {},
"optional": {
"raw_model": ("RAW_MODEL",),
"raw_clip": ("RAW_CLIP",),
}
}
RETURN_TYPES = ("MODEL", "CLIP")
FUNCTION = "inject"
CATEGORY = "Z-Image/Loaders"
def inject(self, raw_model=None, raw_clip=None):
m, c = None, None
if raw_model:
logging.info("Injecting Diffusion Model...")
with patch('comfy.utils.load_torch_file', return_value=(raw_model.sd, None)):
m = comfy.sd.load_diffusion_model(raw_model.path, model_options={})
if raw_clip:
logging.info("Injecting CLIP...")
with patch('comfy.utils.load_torch_file', return_value=(raw_clip.sd, None)):
c = comfy.sd.load_clip(ckpt_paths=[raw_clip.path], embedding_directory=folder_paths.get_folder_paths("embeddings"))
return (m, c)
class ZImageComfyUninjector:
@classmethod
def INPUT_TYPES(s):
return {
"required": {},
"optional": {
"model": ("MODEL",),
"clip": ("CLIP",),
}
}
RETURN_TYPES = ("RAW_MODEL", "RAW_CLIP")
FUNCTION = "uninject"
CATEGORY = "Z-Image/Loaders"
def bake_patches(self, patcher, sd):
out_sd = {}
for key, weight in sd.items():
if key in patcher.patches:
patches = patcher.patches[key]
w = weight.to("cpu", copy=True)
try:
for p in patches:
patch_weights = p[1]
patch_strength = p[2]
if hasattr(patch_weights, "shape") and patch_weights.shape == w.shape:
w += patch_weights.to("cpu") * patch_strength
elif isinstance(patch_weights, dict) and "diff" in patch_weights:
w += patch_weights["diff"].to("cpu") * patch_strength
except Exception:
pass
out_sd[key] = w
else:
out_sd[key] = weight
return out_sd
def sanitize_keys(self, sd, is_clip=False):
new_sd = {}
for k, v in sd.items():
new_k = k
if is_clip:
new_k = new_k.replace("cond_stage_model.", "") # Unwrap
else:
new_k = new_k.replace("model.diffusion_model.", "")
new_sd[new_k] = v
return new_sd
def uninject(self, model=None, clip=None):
raw_model = None
raw_clip = None
if model:
try:
sd = model.model.state_dict()
if hasattr(model, "patcher"): sd = self.bake_patches(model.patcher, sd)
sanitized_sd = self.sanitize_keys(sd, is_clip=False)
raw_model = ZImageRawWrapper(sanitized_sd, "uninject_model.safetensors")
except Exception as e:
logging.error(f"Z-Image Uninjector: Failed Model: {e}")
if clip:
try:
sd = clip.cond_stage_model.state_dict() if hasattr(clip, "cond_stage_model") else clip.state_dict()
if hasattr(clip, "patcher"): sd = self.bake_patches(clip.patcher, sd)
sanitized_sd = self.sanitize_keys(sd, is_clip=True)
raw_clip = ZImageRawWrapper(sanitized_sd, "uninject_clip.safetensors")
except Exception as e:
logging.error(f"Z-Image Uninjector: Failed CLIP: {e}")
return (raw_model, raw_clip)
class ZImageSaveTransformer:
@classmethod
def INPUT_TYPES(s):
return {"required": {"raw_model": ("RAW_MODEL",), "filename_prefix": ("STRING", {"default": "z_image_patched"})}}
RETURN_TYPES = ()
OUTPUT_NODE = True
FUNCTION = "save"
CATEGORY = "Z-Image/Saving"
def save(self, raw_model, filename_prefix):
if not raw_model: return {}
path = os.path.join(folder_paths.get_output_directory(), "diffusion_models", f"{filename_prefix}.safetensors")
os.makedirs(os.path.dirname(path), exist_ok=True)
save_file(raw_model.sd, path)
return {}
class ZImageSaveTextEncoder:
@classmethod
def INPUT_TYPES(s):
return {"required": {"raw_clip": ("RAW_CLIP",), "filename_prefix": ("STRING", {"default": "qwen_patched"})}}
RETURN_TYPES = ()
OUTPUT_NODE = True
FUNCTION = "save"
CATEGORY = "Z-Image/Saving"
def save(self, raw_clip, filename_prefix):
if not raw_clip: return {}
path = os.path.join(folder_paths.get_output_directory(), "text_encoders", f"{filename_prefix}.safetensors")
os.makedirs(os.path.dirname(path), exist_ok=True)
save_file(raw_clip.sd, path)
return {}
NODE_CLASS_MAPPINGS = {
"ZImageLoaderAndPatcher": ZImageLoaderAndPatcher,
"ZImageAITKLoRALoader": ZImageAITKLoRALoader,
"ZImageComfyInjector": ZImageComfyInjector,
"ZImageComfyUninjector": ZImageComfyUninjector,
"ZImageLycorisStacker": ZImageLycorisStacker,
"ZImageDiffusersLoader": ZImageDiffusersLoader,
"ZImageRawModelMerge": ZImageRawModelMerge,
"ZImageRawClipMerge": ZImageRawClipMerge,
"ZImageSaveTransformer": ZImageSaveTransformer,
"ZImageSaveTextEncoder": ZImageSaveTextEncoder
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZImageLoaderAndPatcher": "Z-Image Loader & Patcher (LyCORIS/LoHA)",
"ZImageAITKLoRALoader": "Z-Image AITK LoRA Loader (Standard)",
"ZImageComfyInjector": "Z-Image Comfy Injector",
"ZImageComfyUninjector": "Z-Image Comfy Uninjector",
"ZImageLycorisStacker": "Z-Image Stacker",
"ZImageDiffusersLoader": "Z-Image Diffusers Loader",
"ZImageRawModelMerge": "Z-Image Raw Model Merge",
"ZImageRawClipMerge": "Z-Image Raw CLIP Merge",
"ZImageSaveTransformer": "Z-Image Save Transformer",
"ZImageSaveTextEncoder": "Z-Image Save Text Encoder"
}
+321
View File
@@ -0,0 +1,321 @@
# (c) City96 || Apache-2.0 (apache.org/licenses/LICENSE-2.0)
import torch
import logging
import collections
import nodes
import comfy.sd
import comfy.lora
import comfy.float
import comfy.utils
import comfy.model_patcher
import comfy.model_management
import folder_paths
from .ops import GGMLOps, move_patch_to_device
from .loader import gguf_sd_loader, gguf_clip_loader
from .dequant import is_quantized, is_torch_compatible
def update_folder_names_and_paths(key, targets=[]):
# check for existing key
base = folder_paths.folder_names_and_paths.get(key, ([], {}))
base = base[0] if isinstance(base[0], (list, set, tuple)) else []
# find base key & add w/ fallback, sanity check + warning
target = next((x for x in targets if x in folder_paths.folder_names_and_paths), targets[0])
orig, _ = folder_paths.folder_names_and_paths.get(target, ([], {}))
folder_paths.folder_names_and_paths[key] = (orig or base, {".gguf"})
if base and base != orig:
logging.warning(f"Unknown file list already present on key {key}: {base}")
# Add a custom keys for files ending in .gguf
update_folder_names_and_paths("unet_gguf", ["diffusion_models", "unet"])
update_folder_names_and_paths("clip_gguf", ["text_encoders", "clip"])
class GGUFModelPatcher(comfy.model_patcher.ModelPatcher):
patch_on_device = False
def patch_weight_to_device(self, key, device_to=None, inplace_update=False):
if key not in self.patches:
return
weight = comfy.utils.get_attr(self.model, key)
patches = self.patches[key]
if is_quantized(weight):
out_weight = weight.to(device_to)
patches = move_patch_to_device(patches, self.load_device if self.patch_on_device else self.offload_device)
# TODO: do we ever have legitimate duplicate patches? (i.e. patch on top of patched weight)
out_weight.patches = [(patches, key)]
else:
inplace_update = self.weight_inplace_update or inplace_update
if key not in self.backup:
self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(
weight.to(device=self.offload_device, copy=inplace_update), inplace_update
)
if device_to is not None:
temp_weight = comfy.model_management.cast_to_device(weight, device_to, torch.float32, copy=True)
else:
temp_weight = weight.to(torch.float32, copy=True)
out_weight = comfy.lora.calculate_weight(patches, temp_weight, key)
out_weight = comfy.float.stochastic_rounding(out_weight, weight.dtype)
if inplace_update:
comfy.utils.copy_to_param(self.model, key, out_weight)
else:
comfy.utils.set_attr_param(self.model, key, out_weight)
def unpatch_model(self, device_to=None, unpatch_weights=True):
if unpatch_weights:
for p in self.model.parameters():
if is_torch_compatible(p):
continue
patches = getattr(p, "patches", [])
if len(patches) > 0:
p.patches = []
# TODO: Find another way to not unload after patches
return super().unpatch_model(device_to=device_to, unpatch_weights=unpatch_weights)
def pin_weight_to_device(self, key):
op_key = key.rsplit('.', 1)[0]
if not self.mmap_released and op_key in self.named_modules_to_munmap:
# TODO: possible to OOM, find better way to detach
self.named_modules_to_munmap[op_key].to(self.load_device).to(self.offload_device)
del self.named_modules_to_munmap[op_key]
super().pin_weight_to_device(key)
mmap_released = False
named_modules_to_munmap = {}
def load(self, *args, force_patch_weights=False, **kwargs):
if not self.mmap_released:
self.named_modules_to_munmap = dict(self.model.named_modules())
# always call `patch_weight_to_device` even for lowvram
super().load(*args, force_patch_weights=True, **kwargs)
# make sure nothing stays linked to mmap after first load
if not self.mmap_released:
linked = []
if kwargs.get("lowvram_model_memory", 0) > 0:
for n, m in self.named_modules_to_munmap.items():
if hasattr(m, "weight"):
device = getattr(m.weight, "device", None)
if device == self.offload_device:
linked.append((n, m))
continue
if hasattr(m, "bias"):
device = getattr(m.bias, "device", None)
if device == self.offload_device:
linked.append((n, m))
continue
if linked and self.load_device != self.offload_device:
logging.info(f"Attempting to release mmap ({len(linked)})")
for n, m in linked:
# TODO: possible to OOM, find better way to detach
m.to(self.load_device).to(self.offload_device)
self.mmap_released = True
self.named_modules_to_munmap = {}
def clone(self, *args, **kwargs):
src_cls = self.__class__
self.__class__ = GGUFModelPatcher
n = super().clone(*args, **kwargs)
n.__class__ = GGUFModelPatcher
self.__class__ = src_cls
# GGUF specific clone values below
n.patch_on_device = getattr(self, "patch_on_device", False)
n.mmap_released = getattr(self, "mmap_released", False)
if src_cls != GGUFModelPatcher:
n.size = 0 # force recalc
return n
class UnetLoaderGGUF:
@classmethod
def INPUT_TYPES(s):
unet_names = [x for x in folder_paths.get_filename_list("unet_gguf")]
return {
"required": {
"unet_name": (unet_names,),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "load_unet"
CATEGORY = "bootleg"
TITLE = "Unet Loader (GGUF)"
def load_unet(self, unet_name, dequant_dtype=None, patch_dtype=None, patch_on_device=None):
ops = GGMLOps()
if dequant_dtype in ("default", None):
ops.Linear.dequant_dtype = None
elif dequant_dtype in ["target"]:
ops.Linear.dequant_dtype = dequant_dtype
else:
ops.Linear.dequant_dtype = getattr(torch, dequant_dtype)
if patch_dtype in ("default", None):
ops.Linear.patch_dtype = None
elif patch_dtype in ["target"]:
ops.Linear.patch_dtype = patch_dtype
else:
ops.Linear.patch_dtype = getattr(torch, patch_dtype)
# init model
unet_path = folder_paths.get_full_path("unet", unet_name)
sd = gguf_sd_loader(unet_path)
model = comfy.sd.load_diffusion_model_state_dict(
sd, model_options={"custom_operations": ops}
)
if model is None:
logging.error("ERROR UNSUPPORTED UNET {}".format(unet_path))
raise RuntimeError("ERROR: Could not detect model type of: {}".format(unet_path))
model = GGUFModelPatcher.clone(model)
model.patch_on_device = patch_on_device
return (model,)
class UnetLoaderGGUFAdvanced(UnetLoaderGGUF):
@classmethod
def INPUT_TYPES(s):
unet_names = [x for x in folder_paths.get_filename_list("unet_gguf")]
return {
"required": {
"unet_name": (unet_names,),
"dequant_dtype": (["default", "target", "float32", "float16", "bfloat16"], {"default": "default"}),
"patch_dtype": (["default", "target", "float32", "float16", "bfloat16"], {"default": "default"}),
"patch_on_device": ("BOOLEAN", {"default": False}),
}
}
TITLE = "Unet Loader (GGUF/Advanced)"
class CLIPLoaderGGUF:
@classmethod
def INPUT_TYPES(s):
base = nodes.CLIPLoader.INPUT_TYPES()
return {
"required": {
"clip_name": (s.get_filename_list(),),
"type": base["required"]["type"],
}
}
RETURN_TYPES = ("CLIP",)
FUNCTION = "load_clip"
CATEGORY = "bootleg"
TITLE = "CLIPLoader (GGUF)"
@classmethod
def get_filename_list(s):
files = []
files += folder_paths.get_filename_list("clip")
files += folder_paths.get_filename_list("clip_gguf")
return sorted(files)
def load_data(self, ckpt_paths):
clip_data = []
for p in ckpt_paths:
if p.endswith(".gguf"):
sd = gguf_clip_loader(p)
else:
sd = comfy.utils.load_torch_file(p, safe_load=True)
if "scaled_fp8" in sd: # NOTE: Scaled FP8 would require different custom ops, but only one can be active
raise NotImplementedError(f"Mixing scaled FP8 with GGUF is not supported! Use regular CLIP loader or switch model(s)\n({p})")
clip_data.append(sd)
return clip_data
def load_patcher(self, clip_paths, clip_type, clip_data):
clip = comfy.sd.load_text_encoder_state_dicts(
clip_type = clip_type,
state_dicts = clip_data,
model_options = {
"custom_operations": GGMLOps,
"initial_device": comfy.model_management.text_encoder_offload_device()
},
embedding_directory = folder_paths.get_folder_paths("embeddings"),
)
clip.patcher = GGUFModelPatcher.clone(clip.patcher)
return clip
def load_clip(self, clip_name, type="stable_diffusion"):
clip_path = folder_paths.get_full_path("clip", clip_name)
clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION)
return (self.load_patcher([clip_path], clip_type, self.load_data([clip_path])),)
class DualCLIPLoaderGGUF(CLIPLoaderGGUF):
@classmethod
def INPUT_TYPES(s):
base = nodes.DualCLIPLoader.INPUT_TYPES()
file_options = (s.get_filename_list(), )
return {
"required": {
"clip_name1": file_options,
"clip_name2": file_options,
"type": base["required"]["type"],
}
}
TITLE = "DualCLIPLoader (GGUF)"
def load_clip(self, clip_name1, clip_name2, type):
clip_path1 = folder_paths.get_full_path("clip", clip_name1)
clip_path2 = folder_paths.get_full_path("clip", clip_name2)
clip_paths = (clip_path1, clip_path2)
clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION)
return (self.load_patcher(clip_paths, clip_type, self.load_data(clip_paths)),)
class TripleCLIPLoaderGGUF(CLIPLoaderGGUF):
@classmethod
def INPUT_TYPES(s):
file_options = (s.get_filename_list(), )
return {
"required": {
"clip_name1": file_options,
"clip_name2": file_options,
"clip_name3": file_options,
}
}
TITLE = "TripleCLIPLoader (GGUF)"
def load_clip(self, clip_name1, clip_name2, clip_name3, type="sd3"):
clip_path1 = folder_paths.get_full_path("clip", clip_name1)
clip_path2 = folder_paths.get_full_path("clip", clip_name2)
clip_path3 = folder_paths.get_full_path("clip", clip_name3)
clip_paths = (clip_path1, clip_path2, clip_path3)
clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION)
return (self.load_patcher(clip_paths, clip_type, self.load_data(clip_paths)),)
class QuadrupleCLIPLoaderGGUF(CLIPLoaderGGUF):
@classmethod
def INPUT_TYPES(s):
file_options = (s.get_filename_list(), )
return {
"required": {
"clip_name1": file_options,
"clip_name2": file_options,
"clip_name3": file_options,
"clip_name4": file_options,
}
}
TITLE = "QuadrupleCLIPLoader (GGUF)"
def load_clip(self, clip_name1, clip_name2, clip_name3, clip_name4, type="stable_diffusion"):
clip_path1 = folder_paths.get_full_path("clip", clip_name1)
clip_path2 = folder_paths.get_full_path("clip", clip_name2)
clip_path3 = folder_paths.get_full_path("clip", clip_name3)
clip_path4 = folder_paths.get_full_path("clip", clip_name4)
clip_paths = (clip_path1, clip_path2, clip_path3, clip_path4)
clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION)
return (self.load_patcher(clip_paths, clip_type, self.load_data(clip_paths)),)
NODE_CLASS_MAPPINGS = {
"UnetLoaderGGUF": UnetLoaderGGUF,
"CLIPLoaderGGUF": CLIPLoaderGGUF,
"DualCLIPLoaderGGUF": DualCLIPLoaderGGUF,
"TripleCLIPLoaderGGUF": TripleCLIPLoaderGGUF,
"QuadrupleCLIPLoaderGGUF": QuadrupleCLIPLoaderGGUF,
"UnetLoaderGGUFAdvanced": UnetLoaderGGUFAdvanced,
}
+298
View File
@@ -0,0 +1,298 @@
# (c) City96 || Apache-2.0 (apache.org/licenses/LICENSE-2.0)
import gguf
import torch
import logging
import comfy.ops
import comfy.lora
import comfy.model_management
from .dequant import dequantize_tensor, is_quantized
def chained_hasattr(obj, chained_attr):
probe = obj
for attr in chained_attr.split('.'):
if hasattr(probe, attr):
probe = getattr(probe, attr)
else:
return False
return True
# A backward and forward compatible way to get `torch.compiler.disable`.
def get_torch_compiler_disable_decorator():
def dummy_decorator(*args, **kwargs):
def noop(x):
return x
return noop
from packaging import version
if not chained_hasattr(torch, "compiler.disable"):
logging.info("ComfyUI-GGUF: Torch too old for torch.compile - bypassing")
return dummy_decorator # torch too old
elif version.parse(torch.__version__) >= version.parse("2.8"):
logging.info("ComfyUI-GGUF: Allowing full torch compile")
return dummy_decorator # torch compile works
if chained_hasattr(torch, "_dynamo.config.nontraceable_tensor_subclasses"):
logging.info("ComfyUI-GGUF: Allowing full torch compile (nightly)")
return dummy_decorator # torch compile works, nightly before 2.8 release
else:
logging.info("ComfyUI-GGUF: Partial torch compile only, consider updating pytorch")
return torch.compiler.disable
torch_compiler_disable = get_torch_compiler_disable_decorator()
class GGMLTensor(torch.Tensor):
"""
Main tensor-like class for storing quantized weights
"""
def __init__(self, *args, tensor_type, tensor_shape, patches=[], **kwargs):
super().__init__()
self.tensor_type = tensor_type
self.tensor_shape = tensor_shape
self.patches = patches
def __new__(cls, *args, tensor_type, tensor_shape, patches=[], **kwargs):
return super().__new__(cls, *args, **kwargs)
def to(self, *args, **kwargs):
new = super().to(*args, **kwargs)
new.tensor_type = getattr(self, "tensor_type", None)
new.tensor_shape = getattr(self, "tensor_shape", new.data.shape)
new.patches = getattr(self, "patches", []).copy()
return new
def clone(self, *args, **kwargs):
return self
def detach(self, *args, **kwargs):
return self
def copy_(self, *args, **kwargs):
# fixes .weight.copy_ in comfy/clip_model/CLIPTextModel
try:
return super().copy_(*args, **kwargs)
except Exception as e:
logging.warning(f"ignoring 'copy_' on tensor: {e}")
def new_empty(self, size, *args, **kwargs):
# Intel Arc fix, ref#50
new_tensor = super().new_empty(size, *args, **kwargs)
return GGMLTensor(
new_tensor,
tensor_type = getattr(self, "tensor_type", None),
tensor_shape = size,
patches = getattr(self, "patches", []).copy()
)
@property
def shape(self):
if not hasattr(self, "tensor_shape"):
self.tensor_shape = self.size()
return self.tensor_shape
class GGMLLayer(torch.nn.Module):
"""
This (should) be responsible for de-quantizing on the fly
"""
comfy_cast_weights = True
dequant_dtype = None
patch_dtype = None
largest_layer = False
torch_compatible_tensor_types = {None, gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16}
def is_ggml_quantized(self, *, weight=None, bias=None):
if weight is None:
weight = self.weight
if bias is None:
bias = self.bias
return is_quantized(weight) or is_quantized(bias)
def _load_from_state_dict(self, state_dict, prefix, *args, **kwargs):
weight, bias = state_dict.get(f"{prefix}weight"), state_dict.get(f"{prefix}bias")
# NOTE: using modified load for linear due to not initializing on creation, see GGMLOps todo
if self.is_ggml_quantized(weight=weight, bias=bias) or isinstance(self, torch.nn.Linear):
return self.ggml_load_from_state_dict(state_dict, prefix, *args, **kwargs)
# Not strictly required, but fixes embedding shape mismatch. Threshold set in loader.py
if isinstance(self, torch.nn.Embedding) and self.weight.shape[0] >= (64 * 1024):
return self.ggml_load_from_state_dict(state_dict, prefix, *args, **kwargs)
return super()._load_from_state_dict(state_dict, prefix, *args, **kwargs)
def ggml_load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):
prefix_len = len(prefix)
for k,v in state_dict.items():
if k[prefix_len:] == "weight":
self.weight = torch.nn.Parameter(v, requires_grad=False)
elif k[prefix_len:] == "bias" and v is not None:
self.bias = torch.nn.Parameter(v, requires_grad=False)
else:
unexpected_keys.append(k)
# For Linear layer with missing weight
if self.weight is None and isinstance(self, torch.nn.Linear):
v = torch.zeros(self.in_features, self.out_features)
self.weight = torch.nn.Parameter(v, requires_grad=False)
missing_keys.append(prefix+"weight")
# for vram estimation (TODO: less fragile logic?)
if getattr(self.weight, "is_largest_weight", False):
self.largest_layer = True
def _save_to_state_dict(self, *args, **kwargs):
if self.is_ggml_quantized():
return self.ggml_save_to_state_dict(*args, **kwargs)
return super()._save_to_state_dict(*args, **kwargs)
def ggml_save_to_state_dict(self, destination, prefix, keep_vars):
# This is a fake state dict for vram estimation
weight = torch.zeros_like(self.weight, device=torch.device("meta"))
destination[prefix + "weight"] = weight
if self.bias is not None:
bias = torch.zeros_like(self.bias, device=torch.device("meta"))
destination[prefix + "bias"] = bias
# Take into account space required for dequantizing the largest tensor
if self.largest_layer:
shape = getattr(self.weight, "tensor_shape", self.weight.shape)
dtype = self.dequant_dtype if self.dequant_dtype and self.dequant_dtype != "target" else torch.float16
temp = torch.empty(*shape, device=torch.device("meta"), dtype=dtype)
destination[prefix + "temp.weight"] = temp
return
# This would return the dequantized state dict
destination[prefix + "weight"] = self.get_weight(self.weight)
if bias is not None:
destination[prefix + "bias"] = self.get_weight(self.bias)
def get_weight(self, tensor, dtype):
if tensor is None:
return
# consolidate and load patches to GPU in async
patch_list = []
device = tensor.device
for patches, key in getattr(tensor, "patches", []):
patch_list += move_patch_to_device(patches, device)
# dequantize tensor while patches load
weight = dequantize_tensor(tensor, dtype, self.dequant_dtype)
# prevent propagating custom tensor class
if isinstance(weight, GGMLTensor):
weight = torch.Tensor(weight)
# apply patches
if len(patch_list) > 0:
if self.patch_dtype is None:
weight = comfy.lora.calculate_weight(patch_list, weight, key)
else:
# for testing, may degrade image quality
patch_dtype = dtype if self.patch_dtype == "target" else self.patch_dtype
weight = comfy.lora.calculate_weight(patch_list, weight, key, patch_dtype)
return weight
@torch_compiler_disable()
def cast_bias_weight(s, input=None, dtype=None, device=None, bias_dtype=None):
if input is not None:
if dtype is None:
dtype = getattr(input, "dtype", torch.float32)
if bias_dtype is None:
bias_dtype = dtype
if device is None:
device = input.device
bias = None
non_blocking = comfy.model_management.device_supports_non_blocking(device)
if s.bias is not None:
bias = s.get_weight(s.bias.to(device), dtype)
bias = comfy.ops.cast_to(bias, bias_dtype, device, non_blocking=non_blocking, copy=False)
weight = s.get_weight(s.weight.to(device), dtype)
weight = comfy.ops.cast_to(weight, dtype, device, non_blocking=non_blocking, copy=False)
return weight, bias
def forward_comfy_cast_weights(self, input, *args, **kwargs):
if self.is_ggml_quantized():
out = self.forward_ggml_cast_weights(input, *args, **kwargs)
else:
out = super().forward_comfy_cast_weights(input, *args, **kwargs)
# non-ggml forward might still propagate custom tensor class
if isinstance(out, GGMLTensor):
out = torch.Tensor(out)
return out
def forward_ggml_cast_weights(self, input):
raise NotImplementedError
class GGMLOps(comfy.ops.manual_cast):
"""
Dequantize weights on the fly before doing the compute
"""
class Linear(GGMLLayer, comfy.ops.manual_cast.Linear):
def __init__(self, in_features, out_features, bias=True, device=None, dtype=None):
torch.nn.Module.__init__(self)
# TODO: better workaround for reserved memory spike on windows
# Issue is with `torch.empty` still reserving the full memory for the layer
# Windows doesn't over-commit memory so without this 24GB+ of pagefile is used
self.in_features = in_features
self.out_features = out_features
self.weight = None
self.bias = None
def forward_ggml_cast_weights(self, input):
weight, bias = self.cast_bias_weight(input)
return torch.nn.functional.linear(input, weight, bias)
class Conv2d(GGMLLayer, comfy.ops.manual_cast.Conv2d):
def forward_ggml_cast_weights(self, input):
weight, bias = self.cast_bias_weight(input)
return self._conv_forward(input, weight, bias)
class Embedding(GGMLLayer, comfy.ops.manual_cast.Embedding):
def forward_ggml_cast_weights(self, input, out_dtype=None):
output_dtype = out_dtype
if self.weight.dtype == torch.float16 or self.weight.dtype == torch.bfloat16:
out_dtype = None
weight, _bias = self.cast_bias_weight(self, device=input.device, dtype=out_dtype)
return torch.nn.functional.embedding(
input, weight, self.padding_idx, self.max_norm, self.norm_type, self.scale_grad_by_freq, self.sparse
).to(dtype=output_dtype)
class LayerNorm(GGMLLayer, comfy.ops.manual_cast.LayerNorm):
def forward_ggml_cast_weights(self, input):
if self.weight is None:
return super().forward_comfy_cast_weights(input)
weight, bias = self.cast_bias_weight(input)
return torch.nn.functional.layer_norm(input, self.normalized_shape, weight, bias, self.eps)
class GroupNorm(GGMLLayer, comfy.ops.manual_cast.GroupNorm):
def forward_ggml_cast_weights(self, input):
weight, bias = self.cast_bias_weight(input)
return torch.nn.functional.group_norm(input, self.num_groups, weight, bias, self.eps)
def move_patch_to_device(item, device):
if isinstance(item, torch.Tensor):
return item.to(device, non_blocking=True)
elif isinstance(item, tuple):
return tuple(move_patch_to_device(x, device) for x in item)
elif isinstance(item, list):
return [move_patch_to_device(x, device) for x in item]
else:
return item
# === IMPROVED RESIZE HELPER ===
def manual_resize_tensor(tensor, target_shape):
"""
Manually trims a tensor to match a target shape.
Handles n-dimensional slicing automatically.
"""
if tensor.shape == target_shape:
return tensor
# Build a list of slices: [0:target_dim, 0:target_dim, ...]
# This automatically trims any dimension that is too large
slices = []
for i, dim in enumerate(target_shape):
slices.append(slice(0, dim))
return tensor[tuple(slices)]
+37
View File
@@ -0,0 +1,37 @@
# primitive_widget_to_string.py
class PrimitiveWidgetToString:
"""
Takes a widget-driven value (e.g. Era Styler's 'era' or a 'folder' socket)
and simply outputs that value as a string for use elsewhere in the graph.
Uses STRING input to avoid ALL "Value not in list" validation errors.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"folder": ("STRING", {
"multiline": False,
"default": "",
}),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("text",)
FUNCTION = "passthrough"
CATEGORY = "utils"
def passthrough(self, folder: str):
# Just return the input value unchanged
return (folder,)
NODE_CLASS_MAPPINGS = {
"PrimitiveWidgetToString": PrimitiveWidgetToString,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"PrimitiveWidgetToString": "Primitive Widget → String",
}
+34
View File
@@ -0,0 +1,34 @@
# Z-Image Toolkit (LyCORIS / LoRA / GGUF / Merge Tools)
Advanced raw patching toolkit for ComfyUI.
## Features
- LyCORIS / LoHA / LoKR Loader
- AITK LoRA Loader (with stacking)
- GGUF Raw Loader & Injector
- Vector Merge
- TIES Merge
- Raw Model Merge
- Diffusers Loader
- CLIP & Model Inject / Uninject
- Utility Nodes
## Categories
- Z-Image/Loaders
- Z-Image/Injectors
- Z-Image/Saving
- Experimental
- utils
## Installation
Install via ComfyUI Manager or clone manually:
cd ComfyUI/custom_nodes
git clone https://github.com/TripleHeadedMonkey/ComfyUI-Zlycoris.git
Restart ComfyUI.
+15
View File
@@ -0,0 +1,15 @@
torch
numpy
safetensors>=0.4.0
gguf
packaging
transformers
diffusers
huggingface_hub
accelerate
optimum
sentencepiece
tqdm
pyyaml
torchaudio
lycoris
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+156
View File
@@ -0,0 +1,156 @@
import torch
import torch.nn as nn
from toolkit.network_mixins import ToolkitNetworkMixin
from toolkit.models.loha import LoHaModule
import re
import weakref
class LoHaNetwork(ToolkitNetworkMixin, nn.Module):
def __init__(
self,
text_encoder,
unet,
lora_dim=4,
alpha=1,
dropout=None,
rank_dropout=None,
module_dropout=None,
multiplier=1.0,
train_unet=True,
train_text_encoder=True,
**kwargs
):
nn.Module.__init__(self)
super().__init__(
text_encoder=text_encoder,
unet=unet,
train_text_encoder=train_text_encoder,
train_unet=train_unet,
**kwargs
)
self.network_type = "loha"
self.peft_format = None
# Force None to bypass incompatible base mixin save hooks
self.base_model_ref = None
# FIX: Explicitly set is_pixart.
# The base Mixin uses this in load_weights but doesn't initialize it itself.
self.is_pixart = kwargs.get('is_pixart', False)
self.lora_dim = lora_dim
self.alpha = alpha
self.dropout = dropout
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
self._multiplier = multiplier
# Use ModuleList so PyTorch registers the parameters for saving
self.loha_modules = nn.ModuleList()
if train_unet:
self.create_modules(unet, "lora_unet")
if train_text_encoder:
if isinstance(text_encoder, list):
for i, te in enumerate(text_encoder):
self.create_modules(te, f"lora_te_{i+1}")
else:
self.create_modules(text_encoder, "lora_te")
def create_modules(self, root_module, prefix):
for name, module in root_module.named_modules():
if module.__class__.__name__ in ["Linear", "Conv2d", "LoRACompatibleLinear", "LoRACompatibleConv"]:
if not module.weight.requires_grad and "refiner" not in prefix:
pass
lora_name = f"{prefix}_{name}".replace('.', '_')
loha_mod = LoHaModule(
lora_name=lora_name,
network=self,
org_module=module,
multiplier=self._multiplier,
lora_dim=self.lora_dim,
alpha=self.alpha,
dropout=self.dropout,
rank_dropout=self.rank_dropout,
module_dropout=self.module_dropout,
)
# Append works the same way with ModuleList
self.loha_modules.append(loha_mod)
# Direct injection
module.forward = loha_mod.forward
module.loha_module = loha_mod
def apply_to(self, text_encoder, unet, train_text_encoder, train_unet):
pass
def prepare_grad_etc(self, text_encoder, unet):
for module in self.get_all_modules():
for param in module.parameters():
param.requires_grad = True
def get_all_modules(self):
return self.loha_modules
def save_weights(self, file, dtype=torch.float16, metadata=None, extra_state_dict=None):
if metadata is None:
metadata = {}
metadata["ss_network_module"] = "lycoris.kohya"
metadata["ss_network_dim"] = str(self.lora_dim)
metadata["ss_network_alpha"] = str(self.alpha)
metadata["ss_network_args"] = str({'algo': 'loha'})
super().save_weights(file, dtype, metadata, extra_state_dict)
def prepare_optimizer_params(self, text_encoder_lr=None, unet_lr=None, default_lr=1e-4):
all_params = []
def enumerate_params(modules, lr):
params = []
for mod in modules:
for param in mod.parameters():
if param.requires_grad:
params.append(param)
return {"params": params, "lr": lr}
if self.train_unet:
unet_mods = [m for m in self.loha_modules if "lora_unet" in m.lora_name]
all_params.append(enumerate_params(unet_mods, unet_lr or default_lr))
if self.train_text_encoder:
te_mods = [m for m in self.loha_modules if "lora_te" in m.lora_name]
all_params.append(enumerate_params(te_mods, text_encoder_lr or default_lr))
return all_params
@torch.no_grad()
def _update_torch_multiplier(self):
multiplier = self._multiplier
try:
first_module = self.get_all_modules()[0]
except IndexError:
return
if hasattr(first_module, 'hada_w1_a'):
device = first_module.hada_w1_a.device
dtype = first_module.hada_w1_a.dtype
if hasattr(first_module.hada_w1_a, '_memory_management_device'):
device = first_module.hada_w1_a._memory_management_device
else:
raise ValueError(f"Unknown module type: {type(first_module)}")
with torch.no_grad():
tensor_multiplier = None
if isinstance(multiplier, int) or isinstance(multiplier, float):
# Create scalar tensor
tensor_multiplier = torch.tensor(multiplier).to(device, dtype=dtype)
elif isinstance(multiplier, list):
tensor_multiplier = torch.tensor(multiplier).to(device, dtype=dtype)
elif isinstance(multiplier, torch.Tensor):
tensor_multiplier = multiplier.clone().detach().to(device, dtype=dtype)
self.torch_multiplier = tensor_multiplier.clone().detach()
+331
View File
@@ -0,0 +1,331 @@
# based heavily on https://github.com/KohakuBlueleaf/LyCORIS/blob/eb460098187f752a5d66406d3affade6f0a07ece/lycoris/modules/lokr.py
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from toolkit.network_mixins import ToolkitModuleMixin
from typing import TYPE_CHECKING, Union, List
from optimum.quanto import QBytesTensor, QTensor
if TYPE_CHECKING:
from toolkit.lora_special import LoRASpecialNetwork
def factorization(dimension: int, factor: int = -1) -> tuple[int, int]:
'''
return a tuple of two value of input dimension decomposed by the number closest to factor
second value is higher or equal than first value.
In LoRA with Kroneckor Product, first value is a value for weight scale.
secon value is a value for weight.
Becuase of non-commutative property, A⊗B ≠ B⊗A. Meaning of two matrices is slightly different.
examples)
factor
-1 2 4 8 16 ...
127 -> 127, 1 127 -> 127, 1 127 -> 127, 1 127 -> 127, 1 127 -> 127, 1
128 -> 16, 8 128 -> 64, 2 128 -> 32, 4 128 -> 16, 8 128 -> 16, 8
250 -> 125, 2 250 -> 125, 2 250 -> 125, 2 250 -> 125, 2 250 -> 125, 2
360 -> 45, 8 360 -> 180, 2 360 -> 90, 4 360 -> 45, 8 360 -> 45, 8
512 -> 32, 16 512 -> 256, 2 512 -> 128, 4 512 -> 64, 8 512 -> 32, 16
1024 -> 32, 32 1024 -> 512, 2 1024 -> 256, 4 1024 -> 128, 8 1024 -> 64, 16
'''
if factor > 0 and (dimension % factor) == 0:
m = factor
n = dimension // factor
return m, n
if factor == -1:
factor = dimension
m, n = 1, dimension
length = m + n
while m < n:
new_m = m + 1
while dimension % new_m != 0:
new_m += 1
new_n = dimension // new_m
if new_m + new_n > length or new_m > factor:
break
else:
m, n = new_m, new_n
if m > n:
n, m = m, n
return m, n
def make_weight_cp(t, wa, wb):
rebuild2 = torch.einsum('i j k l, i p, j r -> p r k l',
t, wa, wb) # [c, d, k1, k2]
return rebuild2
def make_kron(w1, w2, scale):
if len(w2.shape) == 4:
w1 = w1.unsqueeze(2).unsqueeze(2)
w2 = w2.contiguous()
rebuild = torch.kron(w1, w2)
return rebuild*scale
class LokrModule(ToolkitModuleMixin, nn.Module):
def __init__(
self,
lora_name,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.,
rank_dropout=0.,
module_dropout=0.,
use_cp=False,
decompose_both=False,
network: 'LoRASpecialNetwork' = None,
factor: int = -1, # factorization factor
**kwargs,
):
""" if alpha == 0 or None, alpha is rank (no scaling). """
ToolkitModuleMixin.__init__(self, network=network)
torch.nn.Module.__init__(self)
factor = int(factor)
self.lora_name = lora_name
self.lora_dim = lora_dim
self.cp = False
self.use_w1 = False
self.use_w2 = False
self.can_merge_in = True
self.shape = org_module.weight.shape
if org_module.__class__.__name__ == 'Conv2d':
in_dim = org_module.in_channels
k_size = org_module.kernel_size
out_dim = org_module.out_channels
in_m, in_n = factorization(in_dim, factor)
out_l, out_k = factorization(out_dim, factor)
# ((a, b), (c, d), *k_size)
shape = ((out_l, out_k), (in_m, in_n), *k_size)
self.cp = use_cp and k_size != (1, 1)
if decompose_both and lora_dim < max(shape[0][0], shape[1][0])/2:
self.lokr_w1_a = nn.Parameter(
torch.empty(shape[0][0], lora_dim))
self.lokr_w1_b = nn.Parameter(
torch.empty(lora_dim, shape[1][0]))
else:
self.use_w1 = True
self.lokr_w1 = nn.Parameter(torch.empty(
shape[0][0], shape[1][0])) # a*c, 1-mode
if lora_dim >= max(shape[0][1], shape[1][1])/2:
self.use_w2 = True
self.lokr_w2 = nn.Parameter(torch.empty(
shape[0][1], shape[1][1], *k_size))
elif self.cp:
self.lokr_t2 = nn.Parameter(torch.empty(
lora_dim, lora_dim, shape[2], shape[3]))
self.lokr_w2_a = nn.Parameter(
torch.empty(lora_dim, shape[0][1])) # b, 1-mode
self.lokr_w2_b = nn.Parameter(
torch.empty(lora_dim, shape[1][1])) # d, 2-mode
else: # Conv2d not cp
# bigger part. weight and LoRA. [b, dim] x [dim, d*k1*k2]
self.lokr_w2_a = nn.Parameter(
torch.empty(shape[0][1], lora_dim))
self.lokr_w2_b = nn.Parameter(torch.empty(
lora_dim, shape[1][1]*shape[2]*shape[3]))
# w1 ⊗ (w2_a x w2_b) = (a, b)⊗((c, dim)x(dim, d*k1*k2)) = (a, b)⊗(c, d*k1*k2) = (ac, bd*k1*k2)
self.op = F.conv2d
self.extra_args = {
"stride": org_module.stride,
"padding": org_module.padding,
"dilation": org_module.dilation,
"groups": org_module.groups
}
else: # Linear
in_dim = org_module.in_features
out_dim = org_module.out_features
in_m, in_n = factorization(in_dim, factor)
out_l, out_k = factorization(out_dim, factor)
# ((a, b), (c, d)), out_dim = a*c, in_dim = b*d
shape = ((out_l, out_k), (in_m, in_n))
# smaller part. weight scale
if decompose_both and lora_dim < max(shape[0][0], shape[1][0])/2:
self.lokr_w1_a = nn.Parameter(
torch.empty(shape[0][0], lora_dim))
self.lokr_w1_b = nn.Parameter(
torch.empty(lora_dim, shape[1][0]))
else:
self.use_w1 = True
self.lokr_w1 = nn.Parameter(torch.empty(
shape[0][0], shape[1][0])) # a*c, 1-mode
if lora_dim < max(shape[0][1], shape[1][1])/2:
# bigger part. weight and LoRA. [b, dim] x [dim, d]
self.lokr_w2_a = nn.Parameter(
torch.empty(shape[0][1], lora_dim))
self.lokr_w2_b = nn.Parameter(
torch.empty(lora_dim, shape[1][1]))
# w1 ⊗ (w2_a x w2_b) = (a, b)⊗((c, dim)x(dim, d)) = (a, b)⊗(c, d) = (ac, bd)
else:
self.use_w2 = True
self.lokr_w2 = nn.Parameter(
torch.empty(shape[0][1], shape[1][1]))
self.op = F.linear
self.extra_args = {}
self.dropout = dropout
if dropout:
print("[WARN]LoKr haven't implemented normal dropout yet.")
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
if isinstance(alpha, torch.Tensor):
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = lora_dim if alpha is None or alpha == 0 else alpha
if self.use_w2 and self.use_w1:
# use scale = 1
alpha = lora_dim
self.scale = alpha / self.lora_dim
self.register_buffer('alpha', torch.tensor(alpha)) # treat as constant
if self.use_w2:
torch.nn.init.constant_(self.lokr_w2, 0)
else:
if self.cp:
torch.nn.init.kaiming_uniform_(self.lokr_t2, a=math.sqrt(5))
torch.nn.init.kaiming_uniform_(self.lokr_w2_a, a=math.sqrt(5))
torch.nn.init.constant_(self.lokr_w2_b, 0)
if self.use_w1:
torch.nn.init.kaiming_uniform_(self.lokr_w1, a=math.sqrt(5))
else:
torch.nn.init.kaiming_uniform_(self.lokr_w1_a, a=math.sqrt(5))
torch.nn.init.kaiming_uniform_(self.lokr_w1_b, a=math.sqrt(5))
self.multiplier = multiplier
self.org_module = [org_module]
weight = make_kron(
self.lokr_w1 if self.use_w1 else self.lokr_w1_a@self.lokr_w1_b,
(self.lokr_w2 if self.use_w2
else make_weight_cp(self.lokr_t2, self.lokr_w2_a, self.lokr_w2_b) if self.cp
else self.lokr_w2_a@self.lokr_w2_b),
torch.tensor(self.multiplier * self.scale)
)
assert torch.sum(torch.isnan(weight)) == 0, "weight is nan"
# Same as locon.py
def apply_to(self):
self.org_forward = self.org_module[0].forward
self.org_module[0].forward = self.forward
def get_weight(self, orig_weight=None):
weight = make_kron(
self.lokr_w1 if self.use_w1 else self.lokr_w1_a@self.lokr_w1_b,
(self.lokr_w2 if self.use_w2
else make_weight_cp(self.lokr_t2, self.lokr_w2_a, self.lokr_w2_b) if self.cp
else self.lokr_w2_a@self.lokr_w2_b),
torch.tensor(self.scale)
)
if orig_weight is not None:
weight = weight.reshape(orig_weight.shape)
if self.training and self.rank_dropout:
drop = torch.rand(weight.size(0)) < self.rank_dropout
weight *= drop.view(-1, [1] *
len(weight.shape[1:])).to(weight.device)
return weight
@torch.no_grad()
def merge_in(self, merge_weight=1.0):
if not self.can_merge_in:
return
# extract weight from org_module
org_sd = self.org_module[0].state_dict()
# todo find a way to merge in weights when doing quantized model
if 'weight._data' in org_sd:
# quantized weight
return
weight_key = "weight"
if 'weight._data' in org_sd:
# quantized weight
weight_key = "weight._data"
orig_dtype = org_sd[weight_key].dtype
weight = org_sd[weight_key].float()
scale = self.scale
# handle trainable scaler method locon does
if hasattr(self, 'scalar'):
scale = scale * self.scalar
lokr_weight = self.get_weight(weight)
merged_weight = (
weight
+ (lokr_weight * merge_weight).to(weight.device, dtype=weight.dtype)
)
# set weight to org_module
org_sd[weight_key] = merged_weight.to(orig_dtype)
self.org_module[0].load_state_dict(org_sd)
def get_orig_weight(self):
weight = self.org_module[0].weight
if isinstance(weight, QTensor) or isinstance(weight, QBytesTensor):
return weight.dequantize().data.detach()
else:
return weight.data.detach()
def get_orig_bias(self):
if hasattr(self.org_module[0], 'bias') and self.org_module[0].bias is not None:
if isinstance(self.org_module[0].bias, QTensor) or isinstance(self.org_module[0].bias, QBytesTensor):
return self.org_module[0].bias.dequantize().data.detach()
else:
return self.org_module[0].bias.data.detach()
return None
def _call_forward(self, x):
if isinstance(x, QTensor) or isinstance(x, QBytesTensor):
x = x.dequantize()
orig_dtype = x.dtype
orig_weight = self.get_orig_weight()
lokr_weight = self.get_weight(orig_weight).to(dtype=orig_weight.dtype)
multiplier = self.network_ref().torch_multiplier
if x.dtype != orig_weight.dtype:
x = x.to(dtype=orig_weight.dtype)
# we do not currently support split batch multipliers for lokr. Just do a mean
multiplier = torch.mean(multiplier)
weight = (
orig_weight
+ lokr_weight * multiplier
)
bias = self.get_orig_bias()
if bias is not None:
bias = bias.to(weight.device, dtype=weight.dtype)
output = self.op(
x,
weight.view(self.shape),
bias,
**self.extra_args
)
return output.to(orig_dtype)
+617
View File
@@ -0,0 +1,617 @@
import copy
import json
import math
import weakref
import os
import re
import sys
from typing import List, Optional, Dict, Type, Union
import torch
from diffusers import UNet2DConditionModel, PixArtTransformer2DModel, AuraFlowTransformer2DModel, WanTransformer3DModel
from transformers import CLIPTextModel
from toolkit.models.lokr import LokrModule
from .config_modules import NetworkConfig
from .lorm import count_parameters
from .network_mixins import ToolkitNetworkMixin, ToolkitModuleMixin, ExtractableModuleMixin
from toolkit.kohya_lora import LoRANetwork
from toolkit.models.DoRA import DoRAModule
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from toolkit.stable_diffusion_model import StableDiffusion
RE_UPDOWN = re.compile(r"(up|down)_blocks_(\d+)_(resnets|upsamplers|downsamplers|attentions)_(\d+)_")
# diffusers specific stuff
LINEAR_MODULES = [
'Linear',
'LoRACompatibleLinear',
'QLinear',
# 'GroupNorm',
]
CONV_MODULES = [
'Conv2d',
'LoRACompatibleConv',
'QConv2d',
]
class IdentityModule(torch.nn.Module):
def forward(self, x):
return x
class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
"""
replaces forward method of the original Linear, instead of replacing the original Linear module.
"""
def __init__(
self,
lora_name,
org_module: torch.nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=None,
rank_dropout=None,
module_dropout=None,
network: 'LoRASpecialNetwork' = None,
use_bias: bool = False,
is_ara: bool = False,
**kwargs
):
self.can_merge_in = True
"""if alpha == 0 or None, alpha is rank (no scaling)."""
ToolkitModuleMixin.__init__(self, network=network)
torch.nn.Module.__init__(self)
self.lora_name = lora_name
self.orig_module_ref = weakref.ref(org_module)
self.scalar = torch.tensor(1.0, device=org_module.weight.device)
# if is ara lora module, mark it on the layer so memory manager can handle it
if is_ara:
org_module.ara_lora_ref = weakref.ref(self)
# check if parent has bias. if not force use_bias to False
if org_module.bias is None:
use_bias = False
if org_module.__class__.__name__ in CONV_MODULES:
in_dim = org_module.in_channels
out_dim = org_module.out_channels
else:
in_dim = org_module.in_features
out_dim = org_module.out_features
# if limit_rank:
# self.lora_dim = min(lora_dim, in_dim, out_dim)
# if self.lora_dim != lora_dim:
# print(f"{lora_name} dim (rank) is changed to: {self.lora_dim}")
# else:
self.lora_dim = lora_dim
self.full_rank = network.network_type.lower() == "fullrank"
if org_module.__class__.__name__ in CONV_MODULES:
kernel_size = org_module.kernel_size
stride = org_module.stride
padding = org_module.padding
if self.full_rank:
self.lora_down = torch.nn.Conv2d(in_dim, out_dim, kernel_size, stride, padding, bias=False)
self.lora_up = IdentityModule()
else:
self.lora_down = torch.nn.Conv2d(in_dim, self.lora_dim, kernel_size, stride, padding, bias=False)
self.lora_up = torch.nn.Conv2d(self.lora_dim, out_dim, (1, 1), (1, 1), bias=use_bias)
else:
if self.full_rank:
self.lora_down = torch.nn.Linear(in_dim, out_dim, bias=False)
self.lora_up = IdentityModule()
else:
self.lora_down = torch.nn.Linear(in_dim, self.lora_dim, bias=False)
self.lora_up = torch.nn.Linear(self.lora_dim, out_dim, bias=use_bias)
if type(alpha) == torch.Tensor:
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = self.lora_dim if alpha is None or alpha == 0 else alpha
self.scale = alpha / self.lora_dim
self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える
# same as microsoft's
torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5))
if not self.full_rank:
torch.nn.init.zeros_(self.lora_up.weight)
self.multiplier: Union[float, List[float]] = multiplier
# wrap the original module so it doesn't get weights updated
self.org_module = [org_module]
self.dropout = dropout
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
self.is_checkpointing = False
def apply_to(self):
self.org_forward = self.org_module[0].forward
self.org_module[0].forward = self.forward
# del self.org_module
class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
NUM_OF_BLOCKS = 12 # フルモデル相当でのup,downの層の数
# UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel"]
# UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel", "ResnetBlock2D"]
UNET_TARGET_REPLACE_MODULE = ["UNet2DConditionModel"]
# UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["ResnetBlock2D", "Downsample2D", "Upsample2D"]
UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["UNet2DConditionModel"]
TEXT_ENCODER_TARGET_REPLACE_MODULE = [
"CLIPAttention",
"CLIPMLP",
# LLM / Qwen-family text encoders (e.g., z-image-turbo)
"Qwen3ForCausalLM",
"Qwen2ForCausalLM",
"Qwen2VLForConditionalGeneration",
# Some HF model wrapper classes commonly used for decoder-only LMs
"Qwen3Model",
"Qwen2Model",
"LlamaForCausalLM",
"LlamaModel",
"MistralForCausalLM",
"MistralModel",
"GemmaForCausalLM",
"GemmaModel",
"Phi3ForCausalLM",
"Phi3Model",
]
LORA_PREFIX_UNET = "lora_unet"
PEFT_PREFIX_UNET = "unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
# SDXL: must starts with LORA_PREFIX_TEXT_ENCODER
LORA_PREFIX_TEXT_ENCODER1 = "lora_te1"
LORA_PREFIX_TEXT_ENCODER2 = "lora_te2"
def __init__(
self,
text_encoder: Union[List[CLIPTextModel], CLIPTextModel],
unet,
multiplier: float = 1.0,
lora_dim: int = 4,
alpha: float = 1,
dropout: Optional[float] = None,
rank_dropout: Optional[float] = None,
module_dropout: Optional[float] = None,
conv_lora_dim: Optional[int] = None,
conv_alpha: Optional[float] = None,
block_dims: Optional[List[int]] = None,
block_alphas: Optional[List[float]] = None,
conv_block_dims: Optional[List[int]] = None,
conv_block_alphas: Optional[List[float]] = None,
modules_dim: Optional[Dict[str, int]] = None,
modules_alpha: Optional[Dict[str, int]] = None,
module_class: Type[object] = LoRAModule,
varbose: Optional[bool] = False,
train_text_encoder: Optional[bool] = True,
use_text_encoder_1: bool = True,
use_text_encoder_2: bool = True,
train_unet: Optional[bool] = True,
is_sdxl=False,
is_v2=False,
is_v3=False,
is_pixart: bool = False,
is_auraflow: bool = False,
is_flux: bool = False,
is_lumina2: bool = False,
use_bias: bool = False,
is_lorm: bool = False,
ignore_if_contains = None,
only_if_contains = None,
parameter_threshold: float = 0.0,
attn_only: bool = False,
target_lin_modules=LoRANetwork.UNET_TARGET_REPLACE_MODULE,
target_conv_modules=LoRANetwork.UNET_TARGET_REPLACE_MODULE_CONV2D_3X3,
network_type: str = "lora",
full_train_in_out: bool = False,
transformer_only: bool = False,
peft_format: bool = False,
is_assistant_adapter: bool = False,
is_transformer: bool = False,
base_model: 'StableDiffusion' = None,
is_ara: bool = False,
**kwargs
) -> None:
"""
LoRA network: すごく引数が多いが、パターンは以下の通り
1. lora_dimとalphaを指定
2. lora_dim、alpha、conv_lora_dim、conv_alphaを指定
3. block_dimsとblock_alphasを指定 : Conv2d3x3には適用しない
4. block_dims、block_alphas、conv_block_dims、conv_block_alphasを指定 : Conv2d3x3にも適用する
5. modules_dimとmodules_alphaを指定 (推論用)
"""
# call the parent of the parent we are replacing (LoRANetwork) init
torch.nn.Module.__init__(self)
ToolkitNetworkMixin.__init__(
self,
train_text_encoder=train_text_encoder,
train_unet=train_unet,
is_sdxl=is_sdxl,
is_v2=is_v2,
is_lorm=is_lorm,
**kwargs
)
if ignore_if_contains is None:
ignore_if_contains = []
self.ignore_if_contains = ignore_if_contains
self.transformer_only = transformer_only
self.base_model_ref = None
if base_model is not None:
self.base_model_ref = weakref.ref(base_model)
self.only_if_contains: Union[List, None] = only_if_contains
self.lora_dim = lora_dim
self.alpha = alpha
self.conv_lora_dim = conv_lora_dim
self.conv_alpha = conv_alpha
self.dropout = dropout
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
self.is_checkpointing = False
self._multiplier: float = 1.0
self.is_active: bool = False
self.torch_multiplier = None
# triggers the state updates
self.multiplier = multiplier
self.is_sdxl = is_sdxl
self.is_v2 = is_v2
self.is_v3 = is_v3
self.is_pixart = is_pixart
self.is_auraflow = is_auraflow
self.is_flux = is_flux
self.is_lumina2 = is_lumina2
self.network_type = network_type
self.is_assistant_adapter = is_assistant_adapter
self.full_rank = network_type.lower() == "fullrank"
self.is_ara = is_ara
if self.network_type.lower() == "dora":
self.module_class = DoRAModule
module_class = DoRAModule
elif self.network_type.lower() == "lokr":
self.module_class = LokrModule
module_class = LokrModule
self.network_config: NetworkConfig = kwargs.get("network_config", None)
self.peft_format = peft_format
self.is_transformer = is_transformer
# always do peft for flux only for now
if self.is_flux or self.is_v3 or self.is_lumina2 or is_transformer:
# don't do peft format for lokr
if self.network_type.lower() != "lokr":
self.peft_format = True
if self.peft_format:
# no alpha for peft
self.alpha = self.lora_dim
alpha = self.alpha
self.conv_alpha = self.conv_lora_dim
conv_alpha = self.conv_alpha
self.full_train_in_out = full_train_in_out
if modules_dim is not None:
print(f"create LoRA network from weights")
elif block_dims is not None:
print(f"create LoRA network from block_dims")
print(
f"neuron dropout: p={self.dropout}, rank dropout: p={self.rank_dropout}, module dropout: p={self.module_dropout}")
print(f"block_dims: {block_dims}")
print(f"block_alphas: {block_alphas}")
if conv_block_dims is not None:
print(f"conv_block_dims: {conv_block_dims}")
print(f"conv_block_alphas: {conv_block_alphas}")
else:
print(f"create LoRA network. base dim (rank): {lora_dim}, alpha: {alpha}")
print(
f"neuron dropout: p={self.dropout}, rank dropout: p={self.rank_dropout}, module dropout: p={self.module_dropout}")
if self.conv_lora_dim is not None:
print(
f"apply LoRA to Conv2d with kernel size (3,3). dim (rank): {self.conv_lora_dim}, alpha: {self.conv_alpha}")
# create module instances
def create_modules(
is_unet: bool,
text_encoder_idx: Optional[int], # None, 1, 2
root_module: torch.nn.Module,
target_replace_modules: List[torch.nn.Module],
) -> List[LoRAModule]:
unet_prefix = self.LORA_PREFIX_UNET
if self.peft_format:
unet_prefix = self.PEFT_PREFIX_UNET
if is_pixart or is_v3 or is_auraflow or is_flux or is_lumina2 or self.is_transformer:
unet_prefix = f"lora_transformer"
if self.peft_format:
unet_prefix = "transformer"
prefix = (
unet_prefix
if is_unet
else (
self.LORA_PREFIX_TEXT_ENCODER
if text_encoder_idx is None
else (self.LORA_PREFIX_TEXT_ENCODER1 if text_encoder_idx == 1 else self.LORA_PREFIX_TEXT_ENCODER2)
)
)
loras = []
skipped = []
attached_module_ids = set()
seen_lora_names = set()
lora_shape_dict = {}
for name, module in root_module.named_modules():
if module.__class__.__name__ in target_replace_modules:
for child_name, child_module in module.named_modules():
is_linear = child_module.__class__.__name__ in LINEAR_MODULES
is_conv2d = child_module.__class__.__name__ in CONV_MODULES
is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
lora_name = [prefix, name, child_name]
# filter out blank
lora_name = [x for x in lora_name if x and x != ""]
lora_name = ".".join(lora_name)
# if it doesnt have a name, it wil have two dots
lora_name = lora_name.replace("..", ".")
clean_name = lora_name
# Deduplicate: the same child_module can appear multiple times
# when multiple parent modules match target_replace_modules (common for LLMs).
if id(child_module) in attached_module_ids:
continue
# Deduplicate by name as well (safety)
if lora_name in seen_lora_names:
continue
if self.peft_format:
# we replace this on saving
lora_name = lora_name.replace(".", "$$")
else:
lora_name = lora_name.replace(".", "_")
skip = False
if any([word in clean_name for word in self.ignore_if_contains]):
skip = True
# see if it is over threshold
if count_parameters(child_module) < parameter_threshold:
skip = True
if self.transformer_only and is_unet:
transformer_block_names = None
if base_model is not None:
transformer_block_names = base_model.get_transformer_block_names()
if transformer_block_names is not None:
if not any([name in lora_name for name in transformer_block_names]):
skip = True
else:
if self.is_pixart:
if "transformer_blocks" not in lora_name:
skip = True
if self.is_flux:
if "transformer_blocks" not in lora_name:
skip = True
if self.is_lumina2:
if "layers$$" not in lora_name and "noise_refiner$$" not in lora_name and "context_refiner$$" not in lora_name:
skip = True
if self.is_v3:
if "transformer_blocks" not in lora_name:
skip = True
# handle custom models
if hasattr(root_module, 'transformer_blocks'):
if "transformer_blocks" not in lora_name:
skip = True
if hasattr(root_module, 'blocks'):
if "blocks" not in lora_name:
skip = True
if hasattr(root_module, 'single_blocks'):
if "single_blocks" not in lora_name and "double_blocks" not in lora_name:
skip = True
if (is_linear or is_conv2d) and not skip:
if self.only_if_contains is not None:
if not any([word in clean_name for word in self.only_if_contains]) and not any([word in lora_name for word in self.only_if_contains]):
continue
dim = None
alpha = None
if modules_dim is not None:
# モジュール指定あり
if lora_name in modules_dim:
dim = modules_dim[lora_name]
alpha = modules_alpha[lora_name]
else:
# 通常、すべて対象とする
if is_linear or is_conv2d_1x1:
dim = self.lora_dim
alpha = self.alpha
elif self.conv_lora_dim is not None:
dim = self.conv_lora_dim
alpha = self.conv_alpha
if dim is None or dim == 0:
# skipした情報を出力
if is_linear or is_conv2d_1x1 or (
self.conv_lora_dim is not None or conv_block_dims is not None):
skipped.append(lora_name)
continue
module_kwargs = {}
if self.network_type.lower() == "lokr":
module_kwargs["factor"] = self.network_config.lokr_factor
if self.is_ara:
module_kwargs["is_ara"] = True
lora = module_class(
lora_name,
child_module,
self.multiplier,
dim,
alpha,
dropout=dropout,
rank_dropout=rank_dropout,
module_dropout=module_dropout,
network=self,
parent=module,
use_bias=use_bias,
**module_kwargs
)
loras.append(lora)
attached_module_ids.add(id(child_module))
seen_lora_names.add(clean_name)
if self.network_type.lower() == "lokr":
try:
lora_shape_dict[lora_name] = [list(lora.lokr_w1.weight.shape), list(lora.lokr_w2.weight.shape)]
except:
pass
else:
if self.full_rank:
lora_shape_dict[lora_name] = [list(lora.lora_down.weight.shape)]
else:
lora_shape_dict[lora_name] = [list(lora.lora_down.weight.shape), list(lora.lora_up.weight.shape)]
return loras, skipped
text_encoders = text_encoder if type(text_encoder) == list else [text_encoder]
# create LoRA for text encoder
# 毎回すべてのモジュールを作るのは無駄なので要検討
self.text_encoder_loras = []
skipped_te = []
if train_text_encoder:
for i, text_encoder in enumerate(text_encoders):
if not use_text_encoder_1 and i == 0:
continue
if not use_text_encoder_2 and i == 1:
continue
if len(text_encoders) > 1:
index = i + 1
print(f"create LoRA for Text Encoder {index}:")
else:
index = None
print(f"create LoRA for Text Encoder:")
replace_modules = self.TEXT_ENCODER_TARGET_REPLACE_MODULE
if self.is_pixart:
replace_modules = ["T5EncoderModel"]
text_encoder_loras, skipped = create_modules(False, index, text_encoder, replace_modules)
self.text_encoder_loras.extend(text_encoder_loras)
skipped_te += skipped
print(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.")
# extend U-Net target modules if conv2d 3x3 is enabled, or load from weights
target_modules = target_lin_modules
if modules_dim is not None or self.conv_lora_dim is not None or conv_block_dims is not None:
target_modules += target_conv_modules
if is_v3:
target_modules = ["SD3Transformer2DModel"]
if is_pixart:
target_modules = ["PixArtTransformer2DModel"]
if is_auraflow:
target_modules = ["AuraFlowTransformer2DModel"]
if is_flux:
target_modules = ["FluxTransformer2DModel"]
if is_lumina2:
target_modules = ["Lumina2Transformer2DModel"]
if train_unet:
self.unet_loras, skipped_un = create_modules(True, None, unet, target_modules)
else:
self.unet_loras = []
skipped_un = []
print(f"create LoRA for U-Net: {len(self.unet_loras)} modules.")
skipped = skipped_te + skipped_un
if varbose and len(skipped) > 0:
print(
f"because block_lr_weight is 0 or dim (rank) is 0, {len(skipped)} LoRA modules are skipped / block_lr_weightまたはdim (rank)が0の為、次の{len(skipped)}個のLoRAモジュールはスキップされます:"
)
for name in skipped:
print(f"\t{name}")
self.up_lr_weight: List[float] = None
self.down_lr_weight: List[float] = None
self.mid_lr_weight: float = None
self.block_lr = False
# assertion
names = set()
for lora in self.text_encoder_loras + self.unet_loras:
assert lora.lora_name not in names, f"duplicated lora name: {lora.lora_name}"
names.add(lora.lora_name)
if self.full_train_in_out:
print("full train in out")
# we are going to retrain the main in out layers for VAE change usually
if self.is_pixart:
transformer: PixArtTransformer2DModel = unet
self.transformer_pos_embed = copy.deepcopy(transformer.pos_embed)
self.transformer_proj_out = copy.deepcopy(transformer.proj_out)
transformer.pos_embed = self.transformer_pos_embed
transformer.proj_out = self.transformer_proj_out
elif self.is_auraflow:
transformer: AuraFlowTransformer2DModel = unet
self.transformer_pos_embed = copy.deepcopy(transformer.pos_embed)
self.transformer_proj_out = copy.deepcopy(transformer.proj_out)
transformer.pos_embed = self.transformer_pos_embed
transformer.proj_out = self.transformer_proj_out
elif base_model is not None and base_model.arch == "wan21":
transformer: WanTransformer3DModel = unet
self.transformer_pos_embed = copy.deepcopy(transformer.patch_embedding)
self.transformer_proj_out = copy.deepcopy(transformer.proj_out)
transformer.patch_embedding = self.transformer_pos_embed
transformer.proj_out = self.transformer_proj_out
else:
unet: UNet2DConditionModel = unet
unet_conv_in: torch.nn.Conv2d = unet.conv_in
unet_conv_out: torch.nn.Conv2d = unet.conv_out
# clone these and replace their forwards with ours
self.unet_conv_in = copy.deepcopy(unet_conv_in)
self.unet_conv_out = copy.deepcopy(unet_conv_out)
unet.conv_in = self.unet_conv_in
unet.conv_out = self.unet_conv_out
def prepare_optimizer_params(self, text_encoder_lr, unet_lr, default_lr):
# call Lora prepare_optimizer_params
all_params = super().prepare_optimizer_params(text_encoder_lr, unet_lr, default_lr)
if self.full_train_in_out:
base_model = self.base_model_ref() if self.base_model_ref is not None else None
if self.is_pixart or self.is_auraflow or self.is_flux or (base_model is not None and base_model.arch == "wan21"):
all_params.append({"lr": unet_lr, "params": list(self.transformer_pos_embed.parameters())})
all_params.append({"lr": unet_lr, "params": list(self.transformer_proj_out.parameters())})
else:
all_params.append({"lr": unet_lr, "params": list(self.unet_conv_in.parameters())})
all_params.append({"lr": unet_lr, "params": list(self.unet_conv_out.parameters())})
return all_params
+461
View File
@@ -0,0 +1,461 @@
from typing import Union, Tuple, Literal, Optional
import torch
import torch.nn as nn
from diffusers import UNet2DConditionModel
from torch import Tensor
from tqdm import tqdm
from toolkit.config_modules import LoRMConfig
conv = nn.Conv2d
lin = nn.Linear
_size_2_t = Union[int, Tuple[int, int]]
ExtractMode = Union[
'fixed',
'threshold',
'ratio',
'quantile',
'percentage'
]
LINEAR_MODULES = [
'Linear',
'LoRACompatibleLinear'
]
CONV_MODULES = [
# 'Conv2d',
# 'LoRACompatibleConv'
]
UNET_TARGET_REPLACE_MODULE = [
"Transformer2DModel",
# "ResnetBlock2D",
"Downsample2D",
"Upsample2D",
]
LORM_TARGET_REPLACE_MODULE = UNET_TARGET_REPLACE_MODULE
UNET_TARGET_REPLACE_NAME = [
"conv_in",
"conv_out",
"time_embedding.linear_1",
"time_embedding.linear_2",
]
UNET_MODULES_TO_AVOID = [
]
# Low Rank Convolution
class LoRMCon2d(nn.Module):
def __init__(
self,
in_channels: int,
lorm_channels: int,
out_channels: int,
kernel_size: _size_2_t,
stride: _size_2_t = 1,
padding: Union[str, _size_2_t] = 'same',
dilation: _size_2_t = 1,
groups: int = 1,
bias: bool = True,
padding_mode: str = 'zeros',
device=None,
dtype=None
) -> None:
super().__init__()
self.in_channels = in_channels
self.lorm_channels = lorm_channels
self.out_channels = out_channels
self.kernel_size = kernel_size
self.stride = stride
self.padding = padding
self.dilation = dilation
self.groups = groups
self.padding_mode = padding_mode
self.down = nn.Conv2d(
in_channels=in_channels,
out_channels=lorm_channels,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
groups=groups,
bias=False,
padding_mode=padding_mode,
device=device,
dtype=dtype
)
# Kernel size on the up is always 1x1.
# I don't think you could calculate a dual 3x3, or I can't at least
self.up = nn.Conv2d(
in_channels=lorm_channels,
out_channels=out_channels,
kernel_size=(1, 1),
stride=1,
padding='same',
dilation=1,
groups=1,
bias=bias,
padding_mode='zeros',
device=device,
dtype=dtype
)
def forward(self, input: Tensor, *args, **kwargs) -> Tensor:
x = input
x = self.down(x)
x = self.up(x)
return x
class LoRMLinear(nn.Module):
def __init__(
self,
in_features: int,
lorm_features: int,
out_features: int,
bias: bool = True,
device=None,
dtype=None
) -> None:
super().__init__()
self.in_features = in_features
self.lorm_features = lorm_features
self.out_features = out_features
self.down = nn.Linear(
in_features=in_features,
out_features=lorm_features,
bias=False,
device=device,
dtype=dtype
)
self.up = nn.Linear(
in_features=lorm_features,
out_features=out_features,
bias=bias,
# bias=True,
device=device,
dtype=dtype
)
def forward(self, input: Tensor, *args, **kwargs) -> Tensor:
x = input
x = self.down(x)
x = self.up(x)
return x
def extract_conv(
weight: Union[torch.Tensor, nn.Parameter],
mode='fixed',
mode_param=0,
device='cpu'
) -> Tuple[Tensor, Tensor, int, Tensor]:
weight = weight.to(device)
out_ch, in_ch, kernel_size, _ = weight.shape
U, S, Vh = torch.linalg.svd(weight.reshape(out_ch, -1))
if mode == 'percentage':
assert 0 <= mode_param <= 1 # Ensure it's a valid percentage.
original_params = out_ch * in_ch * kernel_size * kernel_size
desired_params = mode_param * original_params
# Solve for lora_rank from the equation
lora_rank = int(desired_params / (in_ch * kernel_size * kernel_size + out_ch))
elif mode == 'fixed':
lora_rank = mode_param
elif mode == 'threshold':
assert mode_param >= 0
lora_rank = torch.sum(S > mode_param).item()
elif mode == 'ratio':
assert 1 >= mode_param >= 0
min_s = torch.max(S) * mode_param
lora_rank = torch.sum(S > min_s).item()
elif mode == 'quantile' or mode == 'percentile':
assert 1 >= mode_param >= 0
s_cum = torch.cumsum(S, dim=0)
min_cum_sum = mode_param * torch.sum(S)
lora_rank = torch.sum(s_cum < min_cum_sum).item()
else:
raise NotImplementedError('Extract mode should be "fixed", "threshold", "ratio" or "quantile"')
lora_rank = max(1, lora_rank)
lora_rank = min(out_ch, in_ch, lora_rank)
if lora_rank >= out_ch / 2:
lora_rank = int(out_ch / 2)
print(f"rank is higher than it should be")
# print(f"Skipping layer as determined rank is too high")
# return None, None, None, None
# return weight, 'full'
U = U[:, :lora_rank]
S = S[:lora_rank]
U = U @ torch.diag(S)
Vh = Vh[:lora_rank, :]
diff = (weight - (U @ Vh).reshape(out_ch, in_ch, kernel_size, kernel_size)).detach()
extract_weight_A = Vh.reshape(lora_rank, in_ch, kernel_size, kernel_size).detach()
extract_weight_B = U.reshape(out_ch, lora_rank, 1, 1).detach()
del U, S, Vh, weight
return extract_weight_A, extract_weight_B, lora_rank, diff
def extract_linear(
weight: Union[torch.Tensor, nn.Parameter],
mode='fixed',
mode_param=0,
device='cpu',
) -> Tuple[Tensor, Tensor, int, Tensor]:
weight = weight.to(device)
out_ch, in_ch = weight.shape
U, S, Vh = torch.linalg.svd(weight)
if mode == 'percentage':
assert 0 <= mode_param <= 1 # Ensure it's a valid percentage.
desired_params = mode_param * out_ch * in_ch
# Solve for lora_rank from the equation
lora_rank = int(desired_params / (in_ch + out_ch))
elif mode == 'fixed':
lora_rank = mode_param
elif mode == 'threshold':
assert mode_param >= 0
lora_rank = torch.sum(S > mode_param).item()
elif mode == 'ratio':
assert 1 >= mode_param >= 0
min_s = torch.max(S) * mode_param
lora_rank = torch.sum(S > min_s).item()
elif mode == 'quantile':
assert 1 >= mode_param >= 0
s_cum = torch.cumsum(S, dim=0)
min_cum_sum = mode_param * torch.sum(S)
lora_rank = torch.sum(s_cum < min_cum_sum).item()
else:
raise NotImplementedError('Extract mode should be "fixed", "threshold", "ratio" or "quantile"')
lora_rank = max(1, lora_rank)
lora_rank = min(out_ch, in_ch, lora_rank)
if lora_rank >= out_ch / 2:
# print(f"rank is higher than it should be")
lora_rank = int(out_ch / 2)
# return weight, 'full'
# print(f"Skipping layer as determined rank is too high")
# return None, None, None, None
U = U[:, :lora_rank]
S = S[:lora_rank]
U = U @ torch.diag(S)
Vh = Vh[:lora_rank, :]
diff = (weight - U @ Vh).detach()
extract_weight_A = Vh.reshape(lora_rank, in_ch).detach()
extract_weight_B = U.reshape(out_ch, lora_rank).detach()
del U, S, Vh, weight
return extract_weight_A, extract_weight_B, lora_rank, diff
def replace_module_by_path(network, name, module):
"""Replace a module in a network by its name."""
name_parts = name.split('.')
current_module = network
for part in name_parts[:-1]:
current_module = getattr(current_module, part)
try:
setattr(current_module, name_parts[-1], module)
except Exception as e:
print(e)
def count_parameters(module):
return sum(p.numel() for p in module.parameters())
def compute_optimal_bias(original_module, linear_down, linear_up, X):
Y_original = original_module(X)
Y_approx = linear_up(linear_down(X))
E = Y_original - Y_approx
optimal_bias = E.mean(dim=0)
return optimal_bias
def format_with_commas(n):
return f"{n:,}"
def print_lorm_extract_details(
start_num_params: int,
end_num_params: int,
num_replaced: int,
):
start_formatted = format_with_commas(start_num_params)
end_formatted = format_with_commas(end_num_params)
num_replaced_formatted = format_with_commas(num_replaced)
width = max(len(start_formatted), len(end_formatted), len(num_replaced_formatted))
print(f"Convert UNet result:")
print(f" - converted: {num_replaced:>{width},} modules")
print(f" - start: {start_num_params:>{width},} params")
print(f" - end: {end_num_params:>{width},} params")
lorm_ignore_if_contains = [
'proj_out', 'proj_in',
]
lorm_parameter_threshold = 1000000
@torch.no_grad()
def convert_diffusers_unet_to_lorm(
unet: UNet2DConditionModel,
config: LoRMConfig,
):
print('Converting UNet to LoRM UNet')
start_num_params = count_parameters(unet)
named_modules = list(unet.named_modules())
num_replaced = 0
pbar = tqdm(total=len(named_modules), desc="UNet -> LoRM UNet")
layer_names_replaced = []
converted_modules = []
ignore_if_contains = [
'proj_out', 'proj_in',
]
for name, module in named_modules:
module_name = module.__class__.__name__
if module_name in UNET_TARGET_REPLACE_MODULE:
for child_name, child_module in module.named_modules():
new_module: Union[LoRMCon2d, LoRMLinear, None] = None
# if child name includes attn, skip it
combined_name = combined_name = f"{name}.{child_name}"
# if child_module.__class__.__name__ in LINEAR_MODULES and child_module.bias is None:
# pass
lorm_config = config.get_config_for_module(combined_name)
extract_mode = lorm_config.extract_mode
extract_mode_param = lorm_config.extract_mode_param
parameter_threshold = lorm_config.parameter_threshold
if any([word in child_name for word in ignore_if_contains]):
pass
elif child_module.__class__.__name__ in LINEAR_MODULES:
if count_parameters(child_module) > parameter_threshold:
# dtype = child_module.weight.dtype
dtype = torch.float32
# extract and convert
down_weight, up_weight, lora_dim, diff = extract_linear(
weight=child_module.weight.clone().detach().float(),
mode=extract_mode,
mode_param=extract_mode_param,
device=child_module.weight.device,
)
if down_weight is None:
continue
down_weight = down_weight.to(dtype=dtype)
up_weight = up_weight.to(dtype=dtype)
bias_weight = None
if child_module.bias is not None:
bias_weight = child_module.bias.data.clone().detach().to(dtype=dtype)
# linear layer weights = (out_features, in_features)
new_module = LoRMLinear(
in_features=down_weight.shape[1],
lorm_features=lora_dim,
out_features=up_weight.shape[0],
bias=bias_weight is not None,
device=down_weight.device,
dtype=down_weight.dtype
)
# replace the weights
new_module.down.weight.data = down_weight
new_module.up.weight.data = up_weight
if bias_weight is not None:
new_module.up.bias.data = bias_weight
# else:
# new_module.up.bias.data = torch.zeros_like(new_module.up.bias.data)
# bias_correction = compute_optimal_bias(
# child_module,
# new_module.down,
# new_module.up,
# torch.randn((1000, down_weight.shape[1])).to(device=down_weight.device, dtype=dtype)
# )
# new_module.up.bias.data += bias_correction
elif child_module.__class__.__name__ in CONV_MODULES:
if count_parameters(child_module) > parameter_threshold:
dtype = child_module.weight.dtype
down_weight, up_weight, lora_dim, diff = extract_conv(
weight=child_module.weight.clone().detach().float(),
mode=extract_mode,
mode_param=extract_mode_param,
device=child_module.weight.device,
)
if down_weight is None:
continue
down_weight = down_weight.to(dtype=dtype)
up_weight = up_weight.to(dtype=dtype)
bias_weight = None
if child_module.bias is not None:
bias_weight = child_module.bias.data.clone().detach().to(dtype=dtype)
new_module = LoRMCon2d(
in_channels=down_weight.shape[1],
lorm_channels=lora_dim,
out_channels=up_weight.shape[0],
kernel_size=child_module.kernel_size,
dilation=child_module.dilation,
padding=child_module.padding,
padding_mode=child_module.padding_mode,
stride=child_module.stride,
bias=bias_weight is not None,
device=down_weight.device,
dtype=down_weight.dtype
)
# replace the weights
new_module.down.weight.data = down_weight
new_module.up.weight.data = up_weight
if bias_weight is not None:
new_module.up.bias.data = bias_weight
if new_module:
combined_name = f"{name}.{child_name}"
replace_module_by_path(unet, combined_name, new_module)
converted_modules.append(new_module)
num_replaced += 1
layer_names_replaced.append(
f"{combined_name} - {format_with_commas(count_parameters(child_module))}")
pbar.update(1)
pbar.close()
end_num_params = count_parameters(unet)
def sorting_key(s):
# Extract the number part, remove commas, and convert to integer
return int(s.split("-")[1].strip().replace(",", ""))
sorted_layer_names_replaced = sorted(layer_names_replaced, key=sorting_key, reverse=True)
for layer_name in sorted_layer_names_replaced:
print(layer_name)
print_lorm_extract_details(
start_num_params=start_num_params,
end_num_params=end_num_params,
num_replaced=num_replaced,
)
return converted_modules
+373
View File
@@ -0,0 +1,373 @@
import math
import os
from typing import Optional, Union, List, Type
import torch
from lycoris.kohya import LycorisNetwork, LoConModule
from lycoris.modules.glora import GLoRAModule
from torch import nn
from transformers import CLIPTextModel
from torch.nn import functional as F
from toolkit.network_mixins import ToolkitNetworkMixin, ToolkitModuleMixin, ExtractableModuleMixin
# diffusers specific stuff
LINEAR_MODULES = [
'Linear',
'LoRACompatibleLinear'
]
CONV_MODULES = [
'Conv2d',
'LoRACompatibleConv'
]
class LoConSpecialModule(ToolkitModuleMixin, LoConModule, ExtractableModuleMixin):
def __init__(
self,
lora_name, org_module: nn.Module,
multiplier=1.0,
lora_dim=4, alpha=1,
dropout=0., rank_dropout=0., module_dropout=0.,
use_cp=False,
network: 'LycorisSpecialNetwork' = None,
use_bias=False,
**kwargs,
):
""" if alpha == 0 or None, alpha is rank (no scaling). """
# call super of super
ToolkitModuleMixin.__init__(self, network=network)
torch.nn.Module.__init__(self)
self.lora_name = lora_name
self.lora_dim = lora_dim
self.cp = False
# check if parent has bias. if not force use_bias to False
if org_module.bias is None:
use_bias = False
self.scalar = nn.Parameter(torch.tensor(0.0))
orig_module_name = org_module.__class__.__name__
if orig_module_name in CONV_MODULES:
self.isconv = True
# For general LoCon
in_dim = org_module.in_channels
k_size = org_module.kernel_size
stride = org_module.stride
padding = org_module.padding
out_dim = org_module.out_channels
self.down_op = F.conv2d
self.up_op = F.conv2d
if use_cp and k_size != (1, 1):
self.lora_down = nn.Conv2d(in_dim, lora_dim, (1, 1), bias=False)
self.lora_mid = nn.Conv2d(lora_dim, lora_dim, k_size, stride, padding, bias=False)
self.cp = True
else:
self.lora_down = nn.Conv2d(in_dim, lora_dim, k_size, stride, padding, bias=False)
self.lora_up = nn.Conv2d(lora_dim, out_dim, (1, 1), bias=use_bias)
elif orig_module_name in LINEAR_MODULES:
self.isconv = False
self.down_op = F.linear
self.up_op = F.linear
if orig_module_name == 'GroupNorm':
# RuntimeError: mat1 and mat2 shapes cannot be multiplied (56320x120 and 320x32)
in_dim = org_module.num_channels
out_dim = org_module.num_channels
else:
in_dim = org_module.in_features
out_dim = org_module.out_features
self.lora_down = nn.Linear(in_dim, lora_dim, bias=False)
self.lora_up = nn.Linear(lora_dim, out_dim, bias=use_bias)
else:
raise NotImplementedError
self.shape = org_module.weight.shape
if dropout:
self.dropout = nn.Dropout(dropout)
else:
self.dropout = nn.Identity()
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
if type(alpha) == torch.Tensor:
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = lora_dim if alpha is None or alpha == 0 else alpha
self.scale = alpha / self.lora_dim
self.register_buffer('alpha', torch.tensor(alpha)) # 定数として扱える
# same as microsoft's
torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5))
torch.nn.init.kaiming_uniform_(self.lora_up.weight)
if self.cp:
torch.nn.init.kaiming_uniform_(self.lora_mid.weight, a=math.sqrt(5))
self.multiplier = multiplier
self.org_module = [org_module]
self.register_load_state_dict_post_hook(self.load_weight_hook)
def load_weight_hook(self, *args, **kwargs):
self.scalar = nn.Parameter(torch.ones_like(self.scalar))
class LycorisSpecialNetwork(ToolkitNetworkMixin, LycorisNetwork):
UNET_TARGET_REPLACE_MODULE = [
"Transformer2DModel",
"ResnetBlock2D",
"Downsample2D",
"Upsample2D",
# 'UNet2DConditionModel',
# 'Conv2d',
# 'Timesteps',
# 'TimestepEmbedding',
# 'Linear',
# 'SiLU',
# 'ModuleList',
# 'DownBlock2D',
# 'ResnetBlock2D', # need
# 'GroupNorm',
# 'LoRACompatibleConv',
# 'LoRACompatibleLinear',
# 'Dropout',
# 'CrossAttnDownBlock2D', # needed
# 'Transformer2DModel', # maybe not, has duplicates
# 'BasicTransformerBlock', # duplicates
# 'LayerNorm',
# 'Attention',
# 'FeedForward',
# 'GEGLU',
# 'UpBlock2D',
# 'UNetMidBlock2DCrossAttn'
]
UNET_TARGET_REPLACE_NAME = [
"conv_in",
"conv_out",
"time_embedding.linear_1",
"time_embedding.linear_2",
]
def __init__(
self,
text_encoder: Union[List[CLIPTextModel], CLIPTextModel],
unet,
multiplier: float = 1.0,
lora_dim: int = 4,
alpha: float = 1,
dropout: Optional[float] = None,
rank_dropout: Optional[float] = None,
module_dropout: Optional[float] = None,
conv_lora_dim: Optional[int] = None,
conv_alpha: Optional[float] = None,
use_cp: Optional[bool] = False,
network_module: Type[object] = LoConSpecialModule,
train_unet: bool = True,
train_text_encoder: bool = True,
use_text_encoder_1: bool = True,
use_text_encoder_2: bool = True,
use_bias: bool = False,
is_lorm: bool = False,
**kwargs,
) -> None:
# call ToolkitNetworkMixin super
ToolkitNetworkMixin.__init__(
self,
train_text_encoder=train_text_encoder,
train_unet=train_unet,
is_lorm=is_lorm,
**kwargs
)
# call the parent of the parent LycorisNetwork
torch.nn.Module.__init__(self)
# LyCORIS unique stuff
if dropout is None:
dropout = 0
if rank_dropout is None:
rank_dropout = 0
if module_dropout is None:
module_dropout = 0
self.train_unet = train_unet
self.train_text_encoder = train_text_encoder
self.torch_multiplier = None
# triggers a tensor update
self.multiplier = multiplier
self.lora_dim = lora_dim
if not self.ENABLE_CONV or conv_lora_dim is None:
conv_lora_dim = 0
conv_alpha = 0
self.conv_lora_dim = int(conv_lora_dim)
if self.conv_lora_dim and self.conv_lora_dim != self.lora_dim:
print('Apply different lora dim for conv layer')
print(f'Conv Dim: {conv_lora_dim}, Linear Dim: {lora_dim}')
elif self.conv_lora_dim == 0:
print('Disable conv layer')
self.alpha = alpha
self.conv_alpha = float(conv_alpha)
if self.conv_lora_dim and self.alpha != self.conv_alpha:
print('Apply different alpha value for conv layer')
print(f'Conv alpha: {conv_alpha}, Linear alpha: {alpha}')
if 1 >= dropout >= 0:
print(f'Use Dropout value: {dropout}')
self.dropout = dropout
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
# create module instances
def create_modules(
prefix,
root_module: torch.nn.Module,
target_replace_modules,
target_replace_names=[]
) -> List[network_module]:
print('Create LyCORIS Module')
loras = []
# remove this
named_modules = root_module.named_modules()
# add a few to tthe generator
for name, module in named_modules:
module_name = module.__class__.__name__
if module_name in target_replace_modules:
if module_name in self.MODULE_ALGO_MAP:
algo = self.MODULE_ALGO_MAP[module_name]
else:
algo = network_module
for child_name, child_module in module.named_modules():
lora_name = prefix + '.' + name + '.' + child_name
lora_name = lora_name.replace('.', '_')
if lora_name.startswith('lora_unet_input_blocks_1_0_emb_layers_1'):
print(f"{lora_name}")
if child_module.__class__.__name__ in LINEAR_MODULES and lora_dim > 0:
lora = algo(
lora_name, child_module, self.multiplier,
self.lora_dim, self.alpha,
self.dropout, self.rank_dropout, self.module_dropout,
use_cp,
network=self,
parent=module,
use_bias=use_bias,
**kwargs
)
elif child_module.__class__.__name__ in CONV_MODULES:
k_size, *_ = child_module.kernel_size
if k_size == 1 and lora_dim > 0:
lora = algo(
lora_name, child_module, self.multiplier,
self.lora_dim, self.alpha,
self.dropout, self.rank_dropout, self.module_dropout,
use_cp,
network=self,
parent=module,
use_bias=use_bias,
**kwargs
)
elif conv_lora_dim > 0:
lora = algo(
lora_name, child_module, self.multiplier,
self.conv_lora_dim, self.conv_alpha,
self.dropout, self.rank_dropout, self.module_dropout,
use_cp,
network=self,
parent=module,
use_bias=use_bias,
**kwargs
)
else:
continue
else:
continue
loras.append(lora)
elif name in target_replace_names:
if name in self.NAME_ALGO_MAP:
algo = self.NAME_ALGO_MAP[name]
else:
algo = network_module
lora_name = prefix + '.' + name
lora_name = lora_name.replace('.', '_')
if module.__class__.__name__ == 'Linear' and lora_dim > 0:
lora = algo(
lora_name, module, self.multiplier,
self.lora_dim, self.alpha,
self.dropout, self.rank_dropout, self.module_dropout,
use_cp,
parent=module,
network=self,
use_bias=use_bias,
**kwargs
)
elif module.__class__.__name__ == 'Conv2d':
k_size, *_ = module.kernel_size
if k_size == 1 and lora_dim > 0:
lora = algo(
lora_name, module, self.multiplier,
self.lora_dim, self.alpha,
self.dropout, self.rank_dropout, self.module_dropout,
use_cp,
network=self,
parent=module,
use_bias=use_bias,
**kwargs
)
elif conv_lora_dim > 0:
lora = algo(
lora_name, module, self.multiplier,
self.conv_lora_dim, self.conv_alpha,
self.dropout, self.rank_dropout, self.module_dropout,
use_cp,
network=self,
parent=module,
use_bias=use_bias,
**kwargs
)
else:
continue
else:
continue
loras.append(lora)
return loras
if network_module == GLoRAModule:
print('GLoRA enabled, only train transformer')
# only train transformer (for GLoRA)
LycorisSpecialNetwork.UNET_TARGET_REPLACE_MODULE = [
"Transformer2DModel",
"Attention",
]
LycorisSpecialNetwork.UNET_TARGET_REPLACE_NAME = []
if isinstance(text_encoder, list):
text_encoders = text_encoder
use_index = True
else:
text_encoders = [text_encoder]
use_index = False
self.text_encoder_loras = []
if self.train_text_encoder:
for i, te in enumerate(text_encoders):
if not use_text_encoder_1 and i == 0:
continue
if not use_text_encoder_2 and i == 1:
continue
self.text_encoder_loras.extend(create_modules(
LycorisSpecialNetwork.LORA_PREFIX_TEXT_ENCODER + (f'{i + 1}' if use_index else ''),
te,
LycorisSpecialNetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE
))
print(f"create LyCORIS for Text Encoder: {len(self.text_encoder_loras)} modules.")
if self.train_unet:
self.unet_loras = create_modules(LycorisSpecialNetwork.LORA_PREFIX_UNET, unet,
LycorisSpecialNetwork.UNET_TARGET_REPLACE_MODULE)
else:
self.unet_loras = []
print(f"create LyCORIS for U-Net: {len(self.unet_loras)} modules.")
self.weights_sd = None
# assertion
names = set()
for lora in self.text_encoder_loras + self.unet_loras:
assert lora.lora_name not in names, f"duplicated lora name: {lora.lora_name}"
names.add(lora.lora_name)
+536
View File
@@ -0,0 +1,536 @@
# heavily based on https://github.com/KohakuBlueleaf/LyCORIS/blob/main/lycoris/utils.py
from typing import *
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.linalg as linalg
from tqdm import tqdm
from collections import OrderedDict
def make_sparse(t: torch.Tensor, sparsity=0.95):
abs_t = torch.abs(t)
np_array = abs_t.detach().cpu().numpy()
quan = float(np.quantile(np_array, sparsity))
sparse_t = t.masked_fill(abs_t < quan, 0)
return sparse_t
def extract_conv(
weight: Union[torch.Tensor, nn.Parameter],
mode='fixed',
mode_param=0,
device='cpu',
is_cp=False,
) -> Tuple[nn.Parameter, nn.Parameter]:
weight = weight.to(device)
out_ch, in_ch, kernel_size, _ = weight.shape
U, S, Vh = linalg.svd(weight.reshape(out_ch, -1))
if mode == 'fixed':
lora_rank = mode_param
elif mode == 'threshold':
assert mode_param >= 0
lora_rank = torch.sum(S > mode_param)
elif mode == 'ratio':
assert 1 >= mode_param >= 0
min_s = torch.max(S) * mode_param
lora_rank = torch.sum(S > min_s)
elif mode == 'quantile' or mode == 'percentile':
assert 1 >= mode_param >= 0
s_cum = torch.cumsum(S, dim=0)
min_cum_sum = mode_param * torch.sum(S)
lora_rank = torch.sum(s_cum < min_cum_sum)
else:
raise NotImplementedError('Extract mode should be "fixed", "threshold", "ratio" or "quantile"')
lora_rank = max(1, lora_rank)
lora_rank = min(out_ch, in_ch, lora_rank)
if lora_rank >= out_ch / 2 and not is_cp:
return weight, 'full'
U = U[:, :lora_rank]
S = S[:lora_rank]
U = U @ torch.diag(S)
Vh = Vh[:lora_rank, :]
diff = (weight - (U @ Vh).reshape(out_ch, in_ch, kernel_size, kernel_size)).detach()
extract_weight_A = Vh.reshape(lora_rank, in_ch, kernel_size, kernel_size).detach()
extract_weight_B = U.reshape(out_ch, lora_rank, 1, 1).detach()
del U, S, Vh, weight
return (extract_weight_A, extract_weight_B, diff), 'low rank'
def extract_linear(
weight: Union[torch.Tensor, nn.Parameter],
mode='fixed',
mode_param=0,
device='cpu',
) -> Tuple[nn.Parameter, nn.Parameter]:
weight = weight.to(device)
out_ch, in_ch = weight.shape
U, S, Vh = linalg.svd(weight)
if mode == 'fixed':
lora_rank = mode_param
elif mode == 'threshold':
assert mode_param >= 0
lora_rank = torch.sum(S > mode_param)
elif mode == 'ratio':
assert 1 >= mode_param >= 0
min_s = torch.max(S) * mode_param
lora_rank = torch.sum(S > min_s)
elif mode == 'quantile' or mode == 'percentile':
assert 1 >= mode_param >= 0
s_cum = torch.cumsum(S, dim=0)
min_cum_sum = mode_param * torch.sum(S)
lora_rank = torch.sum(s_cum < min_cum_sum)
else:
raise NotImplementedError('Extract mode should be "fixed", "threshold", "ratio" or "quantile"')
lora_rank = max(1, lora_rank)
lora_rank = min(out_ch, in_ch, lora_rank)
if lora_rank >= out_ch / 2:
return weight, 'full'
U = U[:, :lora_rank]
S = S[:lora_rank]
U = U @ torch.diag(S)
Vh = Vh[:lora_rank, :]
diff = (weight - U @ Vh).detach()
extract_weight_A = Vh.reshape(lora_rank, in_ch).detach()
extract_weight_B = U.reshape(out_ch, lora_rank).detach()
del U, S, Vh, weight
return (extract_weight_A, extract_weight_B, diff), 'low rank'
def extract_diff(
base_model,
db_model,
mode='fixed',
linear_mode_param=0,
conv_mode_param=0,
extract_device='cpu',
use_bias=False,
sparsity=0.98,
small_conv=True,
linear_only=False,
extract_unet=True,
extract_text_encoder=True,
):
meta = OrderedDict()
UNET_TARGET_REPLACE_MODULE = [
"Transformer2DModel",
"Attention",
"ResnetBlock2D",
"Downsample2D",
"Upsample2D"
]
UNET_TARGET_REPLACE_NAME = [
"conv_in",
"conv_out",
"time_embedding.linear_1",
"time_embedding.linear_2",
]
if linear_only:
UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel", "Attention"]
UNET_TARGET_REPLACE_NAME = [
"conv_in",
"conv_out",
]
if not extract_unet:
UNET_TARGET_REPLACE_MODULE = []
UNET_TARGET_REPLACE_NAME = []
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPMLP"]
if not extract_text_encoder:
TEXT_ENCODER_TARGET_REPLACE_MODULE = []
LORA_PREFIX_UNET = 'lora_unet'
LORA_PREFIX_TEXT_ENCODER = 'lora_te'
def make_state_dict(
prefix,
root_module: torch.nn.Module,
target_module: torch.nn.Module,
target_replace_modules,
target_replace_names=[]
):
loras = {}
temp = {}
temp_name = {}
for name, module in root_module.named_modules():
if module.__class__.__name__ in target_replace_modules:
temp[name] = {}
for child_name, child_module in module.named_modules():
if child_module.__class__.__name__ not in {'Linear', 'LoRACompatibleLinear', 'Conv2d', 'LoRACompatibleConv'}:
continue
temp[name][child_name] = child_module.weight
elif name in target_replace_names:
temp_name[name] = module.weight
for name, module in tqdm(list(target_module.named_modules())):
if name in temp:
weights = temp[name]
for child_name, child_module in module.named_modules():
lora_name = prefix + '.' + name + '.' + child_name
lora_name = lora_name.replace('.', '_')
layer = child_module.__class__.__name__
if layer in {'Linear', 'LoRACompatibleLinear', 'Conv2d', 'LoRACompatibleConv'}:
root_weight = child_module.weight
if torch.allclose(root_weight, weights[child_name]):
continue
if layer == 'Linear' or layer == 'LoRACompatibleLinear':
weight, decompose_mode = extract_linear(
(child_module.weight - weights[child_name]),
mode,
linear_mode_param,
device=extract_device,
)
if decompose_mode == 'low rank':
extract_a, extract_b, diff = weight
elif layer == 'Conv2d' or layer == 'LoRACompatibleConv':
is_linear = (child_module.weight.shape[2] == 1
and child_module.weight.shape[3] == 1)
if not is_linear and linear_only:
continue
weight, decompose_mode = extract_conv(
(child_module.weight - weights[child_name]),
mode,
linear_mode_param if is_linear else conv_mode_param,
device=extract_device,
)
if decompose_mode == 'low rank':
extract_a, extract_b, diff = weight
if small_conv and not is_linear and decompose_mode == 'low rank':
dim = extract_a.size(0)
(extract_c, extract_a, _), _ = extract_conv(
extract_a.transpose(0, 1),
'fixed', dim,
extract_device, True
)
extract_a = extract_a.transpose(0, 1)
extract_c = extract_c.transpose(0, 1)
loras[f'{lora_name}.lora_mid.weight'] = extract_c.detach().cpu().contiguous().half()
diff = child_module.weight - torch.einsum(
'i j k l, j r, p i -> p r k l',
extract_c, extract_a.flatten(1, -1), extract_b.flatten(1, -1)
).detach().cpu().contiguous()
del extract_c
else:
continue
if decompose_mode == 'low rank':
loras[f'{lora_name}.lora_down.weight'] = extract_a.detach().cpu().contiguous().half()
loras[f'{lora_name}.lora_up.weight'] = extract_b.detach().cpu().contiguous().half()
loras[f'{lora_name}.alpha'] = torch.Tensor([extract_a.shape[0]]).half()
if use_bias:
diff = diff.detach().cpu().reshape(extract_b.size(0), -1)
sparse_diff = make_sparse(diff, sparsity).to_sparse().coalesce()
indices = sparse_diff.indices().to(torch.int16)
values = sparse_diff.values().half()
loras[f'{lora_name}.bias_indices'] = indices
loras[f'{lora_name}.bias_values'] = values
loras[f'{lora_name}.bias_size'] = torch.tensor(diff.shape).to(torch.int16)
del extract_a, extract_b, diff
elif decompose_mode == 'full':
loras[f'{lora_name}.diff'] = weight.detach().cpu().contiguous().half()
else:
raise NotImplementedError
elif name in temp_name:
weights = temp_name[name]
lora_name = prefix + '.' + name
lora_name = lora_name.replace('.', '_')
layer = module.__class__.__name__
if layer in {'Linear', 'LoRACompatibleLinear', 'Conv2d', 'LoRACompatibleConv'}:
root_weight = module.weight
if torch.allclose(root_weight, weights):
continue
if layer == 'Linear' or layer == 'LoRACompatibleLinear':
weight, decompose_mode = extract_linear(
(root_weight - weights),
mode,
linear_mode_param,
device=extract_device,
)
if decompose_mode == 'low rank':
extract_a, extract_b, diff = weight
elif layer == 'Conv2d' or layer == 'LoRACompatibleConv':
is_linear = (
root_weight.shape[2] == 1
and root_weight.shape[3] == 1
)
if not is_linear and linear_only:
continue
weight, decompose_mode = extract_conv(
(root_weight - weights),
mode,
linear_mode_param if is_linear else conv_mode_param,
device=extract_device,
)
if decompose_mode == 'low rank':
extract_a, extract_b, diff = weight
if small_conv and not is_linear and decompose_mode == 'low rank':
dim = extract_a.size(0)
(extract_c, extract_a, _), _ = extract_conv(
extract_a.transpose(0, 1),
'fixed', dim,
extract_device, True
)
extract_a = extract_a.transpose(0, 1)
extract_c = extract_c.transpose(0, 1)
loras[f'{lora_name}.lora_mid.weight'] = extract_c.detach().cpu().contiguous().half()
diff = root_weight - torch.einsum(
'i j k l, j r, p i -> p r k l',
extract_c, extract_a.flatten(1, -1), extract_b.flatten(1, -1)
).detach().cpu().contiguous()
del extract_c
else:
continue
if decompose_mode == 'low rank':
loras[f'{lora_name}.lora_down.weight'] = extract_a.detach().cpu().contiguous().half()
loras[f'{lora_name}.lora_up.weight'] = extract_b.detach().cpu().contiguous().half()
loras[f'{lora_name}.alpha'] = torch.Tensor([extract_a.shape[0]]).half()
if use_bias:
diff = diff.detach().cpu().reshape(extract_b.size(0), -1)
sparse_diff = make_sparse(diff, sparsity).to_sparse().coalesce()
indices = sparse_diff.indices().to(torch.int16)
values = sparse_diff.values().half()
loras[f'{lora_name}.bias_indices'] = indices
loras[f'{lora_name}.bias_values'] = values
loras[f'{lora_name}.bias_size'] = torch.tensor(diff.shape).to(torch.int16)
del extract_a, extract_b, diff
elif decompose_mode == 'full':
loras[f'{lora_name}.diff'] = weight.detach().cpu().contiguous().half()
else:
raise NotImplementedError
return loras
text_encoder_loras = make_state_dict(
LORA_PREFIX_TEXT_ENCODER,
base_model[0], db_model[0],
TEXT_ENCODER_TARGET_REPLACE_MODULE
)
unet_loras = make_state_dict(
LORA_PREFIX_UNET,
base_model[2], db_model[2],
UNET_TARGET_REPLACE_MODULE,
UNET_TARGET_REPLACE_NAME
)
print(len(text_encoder_loras), len(unet_loras))
# the | will
return (text_encoder_loras | unet_loras), meta
def get_module(
lyco_state_dict: Dict,
lora_name
):
if f'{lora_name}.lora_up.weight' in lyco_state_dict:
up = lyco_state_dict[f'{lora_name}.lora_up.weight']
down = lyco_state_dict[f'{lora_name}.lora_down.weight']
mid = lyco_state_dict.get(f'{lora_name}.lora_mid.weight', None)
alpha = lyco_state_dict.get(f'{lora_name}.alpha', None)
return 'locon', (up, down, mid, alpha)
elif f'{lora_name}.hada_w1_a' in lyco_state_dict:
w1a = lyco_state_dict[f'{lora_name}.hada_w1_a']
w1b = lyco_state_dict[f'{lora_name}.hada_w1_b']
w2a = lyco_state_dict[f'{lora_name}.hada_w2_a']
w2b = lyco_state_dict[f'{lora_name}.hada_w2_b']
t1 = lyco_state_dict.get(f'{lora_name}.hada_t1', None)
t2 = lyco_state_dict.get(f'{lora_name}.hada_t2', None)
alpha = lyco_state_dict.get(f'{lora_name}.alpha', None)
return 'hada', (w1a, w1b, w2a, w2b, t1, t2, alpha)
elif f'{lora_name}.weight' in lyco_state_dict:
weight = lyco_state_dict[f'{lora_name}.weight']
on_input = lyco_state_dict.get(f'{lora_name}.on_input', False)
return 'ia3', (weight, on_input)
elif (f'{lora_name}.lokr_w1' in lyco_state_dict
or f'{lora_name}.lokr_w1_a' in lyco_state_dict):
w1 = lyco_state_dict.get(f'{lora_name}.lokr_w1', None)
w1a = lyco_state_dict.get(f'{lora_name}.lokr_w1_a', None)
w1b = lyco_state_dict.get(f'{lora_name}.lokr_w1_b', None)
w2 = lyco_state_dict.get(f'{lora_name}.lokr_w2', None)
w2a = lyco_state_dict.get(f'{lora_name}.lokr_w2_a', None)
w2b = lyco_state_dict.get(f'{lora_name}.lokr_w2_b', None)
t1 = lyco_state_dict.get(f'{lora_name}.lokr_t1', None)
t2 = lyco_state_dict.get(f'{lora_name}.lokr_t2', None)
alpha = lyco_state_dict.get(f'{lora_name}.alpha', None)
return 'kron', (w1, w1a, w1b, w2, w2a, w2b, t1, t2, alpha)
elif f'{lora_name}.diff' in lyco_state_dict:
return 'full', lyco_state_dict[f'{lora_name}.diff']
else:
return 'None', ()
def cp_weight_from_conv(
up, down, mid
):
up = up.reshape(up.size(0), up.size(1))
down = down.reshape(down.size(0), down.size(1))
return torch.einsum('m n w h, i m, n j -> i j w h', mid, up, down)
def cp_weight(
wa, wb, t
):
temp = torch.einsum('i j k l, j r -> i r k l', t, wb)
return torch.einsum('i j k l, i r -> r j k l', temp, wa)
@torch.no_grad()
def rebuild_weight(module_type, params, orig_weight, scale=1):
if orig_weight is None:
return orig_weight
merged = orig_weight
if module_type == 'locon':
up, down, mid, alpha = params
if alpha is not None:
scale *= alpha / up.size(1)
if mid is not None:
rebuild = cp_weight_from_conv(up, down, mid)
else:
rebuild = up.reshape(up.size(0), -1) @ down.reshape(down.size(0), -1)
merged = orig_weight + rebuild.reshape(orig_weight.shape) * scale
del up, down, mid, alpha, params, rebuild
elif module_type == 'hada':
w1a, w1b, w2a, w2b, t1, t2, alpha = params
if alpha is not None:
scale *= alpha / w1b.size(0)
if t1 is not None:
rebuild1 = cp_weight(w1a, w1b, t1)
else:
rebuild1 = w1a @ w1b
if t2 is not None:
rebuild2 = cp_weight(w2a, w2b, t2)
else:
rebuild2 = w2a @ w2b
rebuild = (rebuild1 * rebuild2).reshape(orig_weight.shape)
merged = orig_weight + rebuild * scale
del w1a, w1b, w2a, w2b, t1, t2, alpha, params, rebuild, rebuild1, rebuild2
elif module_type == 'ia3':
weight, on_input = params
if not on_input:
weight = weight.reshape(-1, 1)
merged = orig_weight + weight * orig_weight * scale
del weight, on_input, params
elif module_type == 'kron':
w1, w1a, w1b, w2, w2a, w2b, t1, t2, alpha = params
if alpha is not None and (w1b is not None or w2b is not None):
scale *= alpha / (w1b.size(0) if w1b else w2b.size(0))
if w1a is not None and w1b is not None:
if t1:
w1 = cp_weight(w1a, w1b, t1)
else:
w1 = w1a @ w1b
if w2a is not None and w2b is not None:
if t2:
w2 = cp_weight(w2a, w2b, t2)
else:
w2 = w2a @ w2b
rebuild = torch.kron(w1, w2).reshape(orig_weight.shape)
merged = orig_weight + rebuild * scale
del w1, w1a, w1b, w2, w2a, w2b, t1, t2, alpha, params, rebuild
elif module_type == 'full':
rebuild = params.reshape(orig_weight.shape)
merged = orig_weight + rebuild * scale
del params, rebuild
return merged
def merge(
base_model,
lyco_state_dict,
scale: float = 1.0,
device='cpu'
):
UNET_TARGET_REPLACE_MODULE = [
"Transformer2DModel",
"Attention",
"ResnetBlock2D",
"Downsample2D",
"Upsample2D"
]
UNET_TARGET_REPLACE_NAME = [
"conv_in",
"conv_out",
"time_embedding.linear_1",
"time_embedding.linear_2",
]
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPMLP"]
LORA_PREFIX_UNET = 'lora_unet'
LORA_PREFIX_TEXT_ENCODER = 'lora_te'
merged = 0
def merge_state_dict(
prefix,
root_module: torch.nn.Module,
lyco_state_dict: Dict[str, torch.Tensor],
target_replace_modules,
target_replace_names=[]
):
nonlocal merged
for name, module in tqdm(list(root_module.named_modules()), desc=f'Merging {prefix}'):
if module.__class__.__name__ in target_replace_modules:
for child_name, child_module in module.named_modules():
if child_module.__class__.__name__ not in {'Linear', 'LoRACompatibleLinear', 'Conv2d',
'LoRACompatibleConv'}:
continue
lora_name = prefix + '.' + name + '.' + child_name
lora_name = lora_name.replace('.', '_')
result = rebuild_weight(*get_module(
lyco_state_dict, lora_name
), getattr(child_module, 'weight'), scale)
if result is not None:
merged += 1
child_module.requires_grad_(False)
child_module.weight.copy_(result)
elif name in target_replace_names:
lora_name = prefix + '.' + name
lora_name = lora_name.replace('.', '_')
result = rebuild_weight(*get_module(
lyco_state_dict, lora_name
), getattr(module, 'weight'), scale)
if result is not None:
merged += 1
module.requires_grad_(False)
module.weight.copy_(result)
if device == 'cpu':
for k, v in tqdm(list(lyco_state_dict.items()), desc='Converting Dtype'):
lyco_state_dict[k] = v.float()
merge_state_dict(
LORA_PREFIX_TEXT_ENCODER,
base_model[0],
lyco_state_dict,
TEXT_ENCODER_TARGET_REPLACE_MODULE,
UNET_TARGET_REPLACE_NAME
)
merge_state_dict(
LORA_PREFIX_UNET,
base_model[2],
lyco_state_dict,
UNET_TARGET_REPLACE_MODULE,
UNET_TARGET_REPLACE_NAME
)
print(f'{merged} Modules been merged')
+88
View File
@@ -0,0 +1,88 @@
import json
from collections import OrderedDict
from io import BytesIO
import safetensors
from safetensors import safe_open
from info import software_meta
from toolkit.train_tools import addnet_hash_legacy
from toolkit.train_tools import addnet_hash_safetensors
def get_meta_for_safetensors(meta: OrderedDict, name=None, add_software_info=True) -> OrderedDict:
# stringify the meta and reparse OrderedDict to replace [name] with name
meta_string = json.dumps(meta)
if name is not None:
meta_string = meta_string.replace("[name]", name)
save_meta = json.loads(meta_string, object_pairs_hook=OrderedDict)
if add_software_info:
save_meta["software"] = software_meta
# safetensors can only be one level deep
for key, value in save_meta.items():
# if not float, int, bool, or str, convert to json string
if not isinstance(value, str):
save_meta[key] = json.dumps(value)
# add the pt format
save_meta["format"] = "pt"
return save_meta
def add_model_hash_to_meta(state_dict, meta: OrderedDict) -> OrderedDict:
"""Precalculate the model hashes needed by sd-webui-additional-networks to
save time on indexing the model later."""
# Because writing user metadata to the file can change the result of
# sd_models.model_hash(), only retain the training metadata for purposes of
# calculating the hash, as they are meant to be immutable
metadata = {k: v for k, v in meta.items() if k.startswith("ss_")}
bytes = safetensors.torch.save(state_dict, metadata)
b = BytesIO(bytes)
model_hash = addnet_hash_safetensors(b)
legacy_hash = addnet_hash_legacy(b)
meta["sshs_model_hash"] = model_hash
meta["sshs_legacy_hash"] = legacy_hash
return meta
def add_base_model_info_to_meta(
meta: OrderedDict,
base_model: str = None,
is_v1: bool = False,
is_v2: bool = False,
is_xl: bool = False,
) -> OrderedDict:
if base_model is not None:
meta['ss_base_model'] = base_model
elif is_v2:
meta['ss_v2'] = True
meta['ss_base_model_version'] = 'sd_2.1'
elif is_xl:
meta['ss_base_model_version'] = 'sdxl_1.0'
else:
# default to v1.5
meta['ss_base_model_version'] = 'sd_1.5'
return meta
def parse_metadata_from_safetensors(meta: OrderedDict) -> OrderedDict:
parsed_meta = OrderedDict()
for key, value in meta.items():
try:
parsed_meta[key] = json.loads(value)
except json.decoder.JSONDecodeError:
parsed_meta[key] = value
return parsed_meta
def load_metadata_from_safetensors(file_path: str) -> OrderedDict:
try:
with safe_open(file_path, framework="pt") as f:
metadata = f.metadata()
return parse_metadata_from_safetensors(metadata)
except Exception as e:
print(f"Error loading metadata from {file_path}: {e}")
return OrderedDict()
+146
View File
@@ -0,0 +1,146 @@
#based off https://github.com/catid/dora/blob/main/dora.py
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import TYPE_CHECKING, Union, List
from optimum.quanto import QBytesTensor, QTensor
from toolkit.network_mixins import ToolkitModuleMixin, ExtractableModuleMixin
if TYPE_CHECKING:
from toolkit.lora_special import LoRASpecialNetwork
# diffusers specific stuff
LINEAR_MODULES = [
'Linear',
'LoRACompatibleLinear'
# 'GroupNorm',
]
CONV_MODULES = [
'Conv2d',
'LoRACompatibleConv'
]
def transpose(weight, fan_in_fan_out):
if not fan_in_fan_out:
return weight
if isinstance(weight, torch.nn.Parameter):
return torch.nn.Parameter(weight.T)
return weight.T
class DoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
# def __init__(self, d_in, d_out, rank=4, weight=None, bias=None):
def __init__(
self,
lora_name,
org_module: torch.nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=None,
rank_dropout=None,
module_dropout=None,
network: 'LoRASpecialNetwork' = None,
use_bias: bool = False,
**kwargs
):
self.can_merge_in = False
"""if alpha == 0 or None, alpha is rank (no scaling)."""
ToolkitModuleMixin.__init__(self, network=network)
torch.nn.Module.__init__(self)
self.lora_name = lora_name
self.scalar = torch.tensor(1.0)
self.lora_dim = lora_dim
if org_module.__class__.__name__ in CONV_MODULES:
raise NotImplementedError("Convolutional layers are not supported yet")
if type(alpha) == torch.Tensor:
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = self.lora_dim if alpha is None or alpha == 0 else alpha
self.scale = alpha / self.lora_dim
# self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える eng: treat as constant
self.multiplier: Union[float, List[float]] = multiplier
# wrap the original module so it doesn't get weights updated
self.org_module = [org_module]
self.dropout = dropout
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
self.is_checkpointing = False
d_out = org_module.out_features
d_in = org_module.in_features
std_dev = 1 / torch.sqrt(torch.tensor(self.lora_dim).float())
# self.lora_up = nn.Parameter(torch.randn(d_out, self.lora_dim) * std_dev) # lora_A
# self.lora_down = nn.Parameter(torch.zeros(self.lora_dim, d_in)) # lora_B
self.lora_up = nn.Linear(self.lora_dim, d_out, bias=False) # lora_B
# self.lora_up.weight.data = torch.randn_like(self.lora_up.weight.data) * std_dev
self.lora_up.weight.data = torch.zeros_like(self.lora_up.weight.data)
# self.lora_A[adapter_name] = nn.Linear(self.in_features, r, bias=False)
# self.lora_B[adapter_name] = nn.Linear(r, self.out_features, bias=False)
self.lora_down = nn.Linear(d_in, self.lora_dim, bias=False) # lora_A
# self.lora_down.weight.data = torch.zeros_like(self.lora_down.weight.data)
self.lora_down.weight.data = torch.randn_like(self.lora_down.weight.data) * std_dev
# m = Magnitude column-wise across output dimension
weight = self.get_orig_weight()
weight = weight.to(self.lora_up.weight.device, dtype=self.lora_up.weight.dtype)
lora_weight = self.lora_up.weight @ self.lora_down.weight
weight_norm = self._get_weight_norm(weight, lora_weight)
self.magnitude = nn.Parameter(weight_norm.detach().clone(), requires_grad=True)
def apply_to(self):
self.org_forward = self.org_module[0].forward
self.org_module[0].forward = self.forward
# del self.org_module
def get_orig_weight(self):
weight = self.org_module[0].weight
if isinstance(weight, QTensor) or isinstance(weight, QBytesTensor):
return weight.dequantize().data.detach()
else:
return weight.data.detach()
def get_orig_bias(self):
if hasattr(self.org_module[0], 'bias') and self.org_module[0].bias is not None:
return self.org_module[0].bias.data.detach()
return None
# def dora_forward(self, x, *args, **kwargs):
# lora = torch.matmul(self.lora_A, self.lora_B)
# adapted = self.get_orig_weight() + lora
# column_norm = adapted.norm(p=2, dim=0, keepdim=True)
# norm_adapted = adapted / column_norm
# calc_weights = self.magnitude * norm_adapted
# return F.linear(x, calc_weights, self.get_orig_bias())
def _get_weight_norm(self, weight, scaled_lora_weight) -> torch.Tensor:
# calculate L2 norm of weight matrix, column-wise
weight = weight + scaled_lora_weight.to(weight.device)
weight_norm = torch.linalg.norm(weight, dim=1)
return weight_norm
def apply_dora(self, x, scaled_lora_weight):
# ref https://github.com/huggingface/peft/blob/1e6d1d73a0850223b0916052fd8d2382a90eae5a/src/peft/tuners/lora/layer.py#L192
# lora weight is already scaled
# magnitude = self.lora_magnitude_vector[active_adapter]
weight = self.get_orig_weight()
weight = weight.to(scaled_lora_weight.device, dtype=scaled_lora_weight.dtype)
weight_norm = self._get_weight_norm(weight, scaled_lora_weight)
# see section 4.3 of DoRA (https://arxiv.org/abs/2402.09353)
# "[...] we suggest treating ||V +∆V ||_c in
# Eq. (5) as a constant, thereby detaching it from the gradient
# graph. This means that while ||V + ∆V ||_c dynamically
# reflects the updates of ∆V , it won’t receive any gradient
# during backpropagation"
weight_norm = weight_norm.detach()
dora_weight = transpose(weight + scaled_lora_weight, False)
return (self.magnitude / weight_norm - 1).view(1, -1) * F.linear(x.to(dora_weight.dtype), dora_weight)
+147
View File
@@ -0,0 +1,147 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from toolkit.network_mixins import ToolkitModuleMixin
class LoHaModule(ToolkitModuleMixin, nn.Module):
def __init__(
self,
lora_name,
network,
org_module: nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=0.0,
rank_dropout=0.0,
module_dropout=0.0,
use_cp=False,
**kwargs
):
nn.Module.__init__(self)
super().__init__(network=network)
self.lora_name = lora_name
self.org_module = [org_module]
# Capture original forward to avoid recursion
self.org_forward = org_module.forward
self.dropout = dropout
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
self._multiplier = multiplier
self.lora_dim = lora_dim
self.alpha = alpha
if isinstance(org_module, nn.Conv2d):
self.is_conv = True
in_dim = org_module.in_channels
out_dim = org_module.out_channels
k_size = org_module.kernel_size
stride = org_module.stride
padding = org_module.padding
self.down_shape = (lora_dim, in_dim, k_size[0], k_size[1])
self.up_shape = (out_dim, lora_dim, 1, 1)
else:
self.is_conv = False
in_dim = org_module.in_features
out_dim = org_module.out_features
self.down_shape = (lora_dim, in_dim)
self.up_shape = (out_dim, lora_dim)
self.hada_w1_a = nn.Parameter(torch.empty(self.down_shape))
self.hada_w1_b = nn.Parameter(torch.empty(self.up_shape))
self.hada_w2_a = nn.Parameter(torch.empty(self.down_shape))
self.hada_w2_b = nn.Parameter(torch.empty(self.up_shape))
self.scale = alpha / lora_dim
self.reset_parameters()
def reset_parameters(self):
nn.init.normal_(self.hada_w1_a, std=0.1)
nn.init.normal_(self.hada_w1_b, std=0.1)
nn.init.normal_(self.hada_w2_a, std=0.1)
nn.init.constant_(self.hada_w2_b, 0)
def get_diff_weight(self):
if self.is_conv:
w1 = (self.hada_w1_b.flatten(start_dim=1) @ self.hada_w1_a.flatten(start_dim=1)).view(
self.hada_w1_b.shape[0], self.hada_w1_a.shape[1], self.hada_w1_a.shape[2], self.hada_w1_a.shape[3]
)
w2 = (self.hada_w2_b.flatten(start_dim=1) @ self.hada_w2_a.flatten(start_dim=1)).view(
self.hada_w2_b.shape[0], self.hada_w2_a.shape[1], self.hada_w2_a.shape[2], self.hada_w2_a.shape[3]
)
else:
w1 = self.hada_w1_b @ self.hada_w1_a
w2 = self.hada_w2_b @ self.hada_w2_a
return (w1 * w2) * self.scale
def forward(self, x, *args, **kwargs):
network = self.network_ref()
if not network.is_active or network.is_merged_in:
return self.org_forward(x, *args, **kwargs)
org_out = self.org_forward(x, *args, **kwargs)
diff_weight = self.get_diff_weight()
# 1. Sync diff_weight dtype with input (Fixes BFloat16 mismatches)
if diff_weight.dtype != x.dtype:
diff_weight = diff_weight.to(dtype=x.dtype)
# 2. Robust Multiplier Handling
multiplier = network.multiplier if network.multiplier is not None else self._multiplier
# Handle List (Unwrap if possible)
if isinstance(multiplier, list):
if len(multiplier) == 1:
multiplier = multiplier[0]
# If len > 1, it's a vector, keep as list for now, will become Tensor below
# Handle Tensor (Convert scalars to Python float)
if isinstance(multiplier, torch.Tensor):
if multiplier.numel() == 1:
multiplier = multiplier.item()
elif multiplier.dtype != diff_weight.dtype:
multiplier = multiplier.to(dtype=diff_weight.dtype, device=diff_weight.device)
# 3. Apply Multiplier to WEIGHTS
# Using pure float multiplication avoids PyTorch "Integer Tensor" confusion
diff_weight = diff_weight * multiplier
if self.is_conv:
out_diff = F.conv2d(
x,
diff_weight,
bias=None,
stride=self.org_module[0].stride,
padding=self.org_module[0].padding,
dilation=self.org_module[0].dilation,
groups=self.org_module[0].groups
)
else:
out_diff = F.linear(x, diff_weight)
return org_out + out_diff
def merge_in(self, merge_weight=1.0):
if self.network_ref().is_merged_in:
return
with torch.no_grad():
weight = self.org_module[0].weight
diff = self.get_diff_weight() * merge_weight
weight.add_(diff.to(weight.device))
def merge_out(self, merge_weight=1.0):
if not self.network_ref().is_merged_in:
return
with torch.no_grad():
weight = self.org_module[0].weight
diff = self.get_diff_weight() * merge_weight
weight.sub_(diff.to(weight.device))
def parameters(self, recurse: bool = True):
return [self.hada_w1_a, self.hada_w1_b, self.hada_w2_a, self.hada_w2_b]
+845
View File
@@ -0,0 +1,845 @@
import json
import os
from collections import OrderedDict
from typing import Optional, Union, List, Type, TYPE_CHECKING, Dict, Any, Literal
import torch
from optimum.quanto import QTensor
from torch import nn
import weakref
from tqdm import tqdm
from toolkit.config_modules import NetworkConfig
from toolkit.lorm import extract_conv, extract_linear, count_parameters
from toolkit.metadata import add_model_hash_to_meta
from toolkit.paths import KEYMAPS_ROOT
from toolkit.saving import get_lora_keymap_from_model_keymap
from optimum.quanto import QBytesTensor
if TYPE_CHECKING:
from toolkit.lycoris_special import LycorisSpecialNetwork, LoConSpecialModule
from toolkit.lora_special import LoRASpecialNetwork, LoRAModule
from toolkit.stable_diffusion_model import StableDiffusion
from toolkit.models.DoRA import DoRAModule
Network = Union['LycorisSpecialNetwork', 'LoRASpecialNetwork']
Module = Union['LoConSpecialModule', 'LoRAModule', 'DoRAModule']
LINEAR_MODULES = [
'Linear',
'LoRACompatibleLinear',
'QLinear'
# 'GroupNorm',
]
CONV_MODULES = [
'Conv2d',
'LoRACompatibleConv'
]
ExtractMode = Union[
'existing'
'fixed',
'threshold',
'ratio',
'quantile',
'percentage'
]
printed_messages = []
def print_once(msg):
global printed_messages
if msg not in printed_messages:
print(msg)
printed_messages.append(msg)
def broadcast_and_multiply(tensor, multiplier):
# Determine the number of dimensions required
num_extra_dims = tensor.dim() - multiplier.dim()
# Unsqueezing the tensor to match the dimensionality
for _ in range(num_extra_dims):
multiplier = multiplier.unsqueeze(-1)
try:
# Multiplying the broadcasted tensor with the output tensor
result = tensor * multiplier
except RuntimeError as e:
print(e)
print(tensor.size())
print(multiplier.size())
raise e
return result
def add_bias(tensor, bias):
if bias is None:
return tensor
# add batch dim
bias = bias.unsqueeze(0)
bias = torch.cat([bias] * tensor.size(0), dim=0)
# Determine the number of dimensions required
num_extra_dims = tensor.dim() - bias.dim()
# Unsqueezing the tensor to match the dimensionality
for _ in range(num_extra_dims):
bias = bias.unsqueeze(-1)
# we may need to swap -1 for -2
if bias.size(1) != tensor.size(1):
if len(bias.size()) == 3:
bias = bias.permute(0, 2, 1)
elif len(bias.size()) == 4:
bias = bias.permute(0, 3, 1, 2)
# Multiplying the broadcasted tensor with the output tensor
try:
result = tensor + bias
except RuntimeError as e:
print(e)
print(tensor.size())
print(bias.size())
raise e
return result
class ExtractableModuleMixin:
def extract_weight(
self: Module,
extract_mode: ExtractMode = "existing",
extract_mode_param: Union[int, float] = None,
):
device = self.lora_down.weight.device
weight_to_extract = self.org_module[0].weight
if extract_mode == "existing":
extract_mode = 'fixed'
extract_mode_param = self.lora_dim
if isinstance(weight_to_extract, QBytesTensor):
weight_to_extract = weight_to_extract.dequantize()
weight_to_extract = weight_to_extract.clone().detach().float()
if self.org_module[0].__class__.__name__ in CONV_MODULES:
# do conv extraction
down_weight, up_weight, new_dim, diff = extract_conv(
weight=weight_to_extract,
mode=extract_mode,
mode_param=extract_mode_param,
device=device
)
elif self.org_module[0].__class__.__name__ in LINEAR_MODULES:
# do linear extraction
down_weight, up_weight, new_dim, diff = extract_linear(
weight=weight_to_extract,
mode=extract_mode,
mode_param=extract_mode_param,
device=device,
)
else:
raise ValueError(f"Unknown module type: {self.org_module[0].__class__.__name__}")
self.lora_dim = new_dim
# inject weights into the param
self.lora_down.weight.data = down_weight.to(self.lora_down.weight.dtype).clone().detach()
self.lora_up.weight.data = up_weight.to(self.lora_up.weight.dtype).clone().detach()
# copy bias if we have one and are using them
if self.org_module[0].bias is not None and self.lora_up.bias is not None:
self.lora_up.bias.data = self.org_module[0].bias.data.clone().detach()
# set up alphas
self.alpha = (self.alpha * 0) + down_weight.shape[0]
self.scale = self.alpha / self.lora_dim
# assign them
# handle trainable scaler method locon does
if hasattr(self, 'scalar'):
# scaler is a parameter update the value with 1.0
self.scalar.data = torch.tensor(1.0).to(self.scalar.device, self.scalar.dtype)
class ToolkitModuleMixin:
def __init__(
self: Module,
*args,
network: Network,
**kwargs
):
self.network_ref: weakref.ref = weakref.ref(network)
self.is_checkpointing = False
self._multiplier: Union[float, list, torch.Tensor] = None
def _call_forward(self: Module, x):
# module dropout
if self.module_dropout is not None and self.training:
if torch.rand(1) < self.module_dropout:
return 0.0 # added to original forward
if hasattr(self, 'lora_mid') and self.lora_mid is not None:
lx = self.lora_mid(self.lora_down(x))
else:
try:
lx = self.lora_down(x)
except RuntimeError as e:
print(f"Error in {self.__class__.__name__} lora_down")
raise e
if isinstance(self.dropout, nn.Dropout) or isinstance(self.dropout, nn.Identity):
lx = self.dropout(lx)
# normal dropout
elif self.dropout is not None and self.training:
lx = torch.nn.functional.dropout(lx, p=self.dropout)
# rank dropout
if self.rank_dropout is not None and self.rank_dropout > 0 and self.training:
mask = torch.rand((lx.size(0), self.lora_dim), device=lx.device) > self.rank_dropout
if len(lx.size()) == 3:
mask = mask.unsqueeze(1) # for Text Encoder
elif len(lx.size()) == 4:
mask = mask.unsqueeze(-1).unsqueeze(-1) # for Conv2d
lx = lx * mask
# scaling for rank dropout: treat as if the rank is changed
# maskから計算することも考えられるが、augmentation的な効果を期待してrank_dropoutを用いる
scale = self.scale * (1.0 / (1.0 - self.rank_dropout)) # redundant for readability
else:
scale = self.scale
lx = self.lora_up(lx)
# handle trainable scaler method locon does
if hasattr(self, 'scalar'):
scale = scale * self.scalar
return lx * scale
def lorm_forward(self: Network, x, *args, **kwargs):
network: Network = self.network_ref()
if not network.is_active:
return self.org_forward(x, *args, **kwargs)
orig_dtype = x.dtype
if x.dtype != self.lora_down.weight.dtype:
x = x.to(self.lora_down.weight.dtype)
if network.lorm_train_mode == 'local':
# we are going to predict input with both and do a loss on them
inputs = x.detach()
with torch.no_grad():
# get the local prediction
target_pred = self.org_forward(inputs, *args, **kwargs).detach()
with torch.set_grad_enabled(True):
# make a prediction with the lorm
lorm_pred = self.lora_up(self.lora_down(inputs.requires_grad_(True)))
local_loss = torch.nn.functional.mse_loss(target_pred.float(), lorm_pred.float())
# backpropr
local_loss.backward()
network.module_losses.append(local_loss.detach())
# return the original as we dont want our trainer to affect ones down the line
return target_pred
else:
x = self.lora_up(self.lora_down(x))
if x.dtype != orig_dtype:
x = x.to(orig_dtype)
def forward(self: Module, x, *args, **kwargs):
skip = False
network: Network = self.network_ref()
if network.is_lorm:
# we are doing lorm
return self.lorm_forward(x, *args, **kwargs)
# skip if not active
if not network.is_active:
skip = True
# skip if is merged in
if network.is_merged_in:
skip = True
# skip if multiplier is 0
if network._multiplier == 0:
skip = True
if skip:
# network is not active, avoid doing anything
return self.org_forward(x, *args, **kwargs)
# if self.__class__.__name__ == "DoRAModule":
# # return dora forward
# return self.dora_forward(x, *args, **kwargs)
if self.__class__.__name__ == "LokrModule":
return self._call_forward(x)
org_forwarded = self.org_forward(x, *args, **kwargs)
if isinstance(x, QTensor):
x = x.dequantize()
# always cast to float32
lora_input = x.to(self.lora_down.weight.dtype)
lora_output = self._call_forward(lora_input)
multiplier = self.network_ref().torch_multiplier
lora_output_batch_size = lora_output.size(0)
multiplier_batch_size = multiplier.size(0)
if lora_output_batch_size != multiplier_batch_size:
num_interleaves = lora_output_batch_size // multiplier_batch_size
# todo check if this is correct, do we just concat when doing cfg?
multiplier = multiplier.repeat_interleave(num_interleaves)
scaled_lora_output = broadcast_and_multiply(lora_output, multiplier)
scaled_lora_output = scaled_lora_output.to(org_forwarded.dtype)
if self.__class__.__name__ == "DoRAModule":
# ref https://github.com/huggingface/peft/blob/1e6d1d73a0850223b0916052fd8d2382a90eae5a/src/peft/tuners/lora/layer.py#L417
# x = dropout(x)
# todo this wont match the dropout applied to the lora
if isinstance(self.dropout, nn.Dropout) or isinstance(self.dropout, nn.Identity):
lx = self.dropout(x)
# normal dropout
elif self.dropout is not None and self.training:
lx = torch.nn.functional.dropout(x, p=self.dropout)
else:
lx = x
lora_weight = self.lora_up.weight @ self.lora_down.weight
# scale it here
# todo handle our batch split scalers for slider training. For now take the mean of them
scale = multiplier.mean()
scaled_lora_weight = lora_weight * scale
scaled_lora_output = scaled_lora_output + self.apply_dora(lx, scaled_lora_weight).to(org_forwarded.dtype)
try:
x = org_forwarded + scaled_lora_output
except RuntimeError as e:
print(e)
print(org_forwarded.size())
print(scaled_lora_output.size())
raise e
return x
def enable_gradient_checkpointing(self: Module):
self.is_checkpointing = True
def disable_gradient_checkpointing(self: Module):
self.is_checkpointing = False
@torch.no_grad()
def merge_out(self: Module, merge_out_weight=1.0):
# make sure it is positive
merge_out_weight = abs(merge_out_weight)
# merging out is just merging in the negative of the weight
self.merge_in(merge_weight=-merge_out_weight)
@torch.no_grad()
def merge_in(self: Module, merge_weight=1.0):
if not self.can_merge_in:
return
# get up/down weight
if self.full_rank:
up_weight = None
else:
up_weight = self.lora_up.weight.clone().float()
down_weight = self.lora_down.weight.clone().float()
# extract weight from org_module
org_sd = self.org_module[0].state_dict()
# todo find a way to merge in weights when doing quantized model
if 'weight._data' in org_sd:
# quantized weight
return
weight_key = "weight"
if 'weight._data' in org_sd:
# quantized weight
weight_key = "weight._data"
orig_dtype = org_sd[weight_key].dtype
weight = org_sd[weight_key].float()
multiplier = merge_weight
scale = self.scale
# handle trainable scaler method locon does
if hasattr(self, 'scalar'):
scale = scale * self.scalar
weight_device = weight.device
if weight.device != down_weight.device:
weight = weight.to(down_weight.device)
if scale.device != down_weight.device:
scale = scale.to(down_weight.device)
# merge weight
if self.full_rank:
weight = weight + multiplier * down_weight * scale
elif len(weight.size()) == 2:
# linear
weight = weight + multiplier * (up_weight @ down_weight) * scale
elif down_weight.size()[2:4] == (1, 1):
# conv2d 1x1
weight = (
weight
+ multiplier
* (up_weight.squeeze(3).squeeze(2) @ down_weight.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3)
* scale
)
else:
# conv2d 3x3
conved = torch.nn.functional.conv2d(down_weight.permute(1, 0, 2, 3), up_weight).permute(1, 0, 2, 3)
# print(conved.size(), weight.size(), module.stride, module.padding)
weight = weight + multiplier * conved * scale
# set weight to org_module
org_sd[weight_key] = weight.to(weight_device, orig_dtype)
self.org_module[0].load_state_dict(org_sd)
def setup_lorm(self: Module, state_dict: Optional[Dict[str, Any]] = None):
# LoRM (Low Rank Middle) is a method reduce the number of parameters in a module while keeping the inputs and
# outputs the same. It is basically a LoRA but with the original module removed
# if a state dict is passed, use those weights instead of extracting
# todo load from state dict
network: Network = self.network_ref()
lorm_config = network.network_config.lorm_config.get_config_for_module(self.lora_name)
extract_mode = lorm_config.extract_mode
extract_mode_param = lorm_config.extract_mode_param
parameter_threshold = lorm_config.parameter_threshold
self.extract_weight(
extract_mode=extract_mode,
extract_mode_param=extract_mode_param
)
class ToolkitNetworkMixin:
def __init__(
self: Network,
*args,
train_text_encoder: Optional[bool] = True,
train_unet: Optional[bool] = True,
is_sdxl=False,
is_v2=False,
is_ssd=False,
is_vega=False,
network_config: Optional[NetworkConfig] = None,
is_lorm=False,
**kwargs
):
self.train_text_encoder = train_text_encoder
self.train_unet = train_unet
self.is_checkpointing = False
self._multiplier: float = 1.0
self.is_active: bool = False
self.is_sdxl = is_sdxl
self.is_ssd = is_ssd
self.is_vega = is_vega
self.is_v2 = is_v2
self.is_v1 = not is_v2 and not is_sdxl and not is_ssd and not is_vega
self.is_merged_in = False
self.is_lorm = is_lorm
self.network_config: NetworkConfig = network_config
self.module_losses: List[torch.Tensor] = []
self.lorm_train_mode: Literal['local', None] = None
self.can_merge_in = not is_lorm
# will prevent optimizer from loading as it will have double states
self.did_change_weights = False
def get_keymap(self: Network, force_weight_mapping=False):
use_weight_mapping = False
if self.is_ssd:
keymap_tail = 'ssd'
use_weight_mapping = True
elif self.is_vega:
keymap_tail = 'vega'
use_weight_mapping = True
elif self.is_sdxl:
keymap_tail = 'sdxl'
elif self.is_v2:
keymap_tail = 'sd2'
else:
keymap_tail = 'sd1'
# todo double check this
# use_weight_mapping = True
if force_weight_mapping:
use_weight_mapping = True
# load keymap
keymap_name = f"stable_diffusion_locon_{keymap_tail}.json"
if use_weight_mapping:
keymap_name = f"stable_diffusion_{keymap_tail}.json"
keymap_path = os.path.join(KEYMAPS_ROOT, keymap_name)
keymap = None
# check if file exists
if os.path.exists(keymap_path):
with open(keymap_path, 'r') as f:
keymap = json.load(f)['ldm_diffusers_keymap']
if use_weight_mapping and keymap is not None:
# get keymap from weights
keymap = get_lora_keymap_from_model_keymap(keymap)
# upgrade keymaps for DoRA
if self.network_type.lower() == 'dora':
if keymap is not None:
new_keymap = {}
for ldm_key, diffusers_key in keymap.items():
ldm_key = ldm_key.replace('.alpha', '.magnitude')
# ldm_key = ldm_key.replace('.lora_down.weight', '.lora_down')
# ldm_key = ldm_key.replace('.lora_up.weight', '.lora_up')
diffusers_key = diffusers_key.replace('.alpha', '.magnitude')
# diffusers_key = diffusers_key.replace('.lora_down.weight', '.lora_down')
# diffusers_key = diffusers_key.replace('.lora_up.weight', '.lora_up')
new_keymap[ldm_key] = diffusers_key
keymap = new_keymap
return keymap
def get_state_dict(self: Network, extra_state_dict=None, dtype=torch.float16):
keymap = self.get_keymap()
save_keymap = {}
if keymap is not None:
for ldm_key, diffusers_key in keymap.items():
# invert them
save_keymap[diffusers_key] = ldm_key
state_dict = self.state_dict()
save_dict = OrderedDict()
for key in list(state_dict.keys()):
v = state_dict[key]
v = v.detach().clone().to("cpu").to(dtype)
save_key = save_keymap[key] if key in save_keymap else key
save_dict[save_key] = v
del state_dict[key]
if extra_state_dict is not None:
# add extra items to state dict
for key in list(extra_state_dict.keys()):
v = extra_state_dict[key]
v = v.detach().clone().to("cpu").to(dtype)
save_dict[key] = v
if self.peft_format:
# lora_down = lora_A
# lora_up = lora_B
# no alpha
new_save_dict = {}
for key, value in save_dict.items():
if key.endswith('.alpha'):
continue
new_key = key
new_key = new_key.replace('lora_down', 'lora_A')
new_key = new_key.replace('lora_up', 'lora_B')
# replace all $$ with .
new_key = new_key.replace('$$', '.')
new_save_dict[new_key] = value
save_dict = new_save_dict
if self.network_type.lower() == "lokr":
new_save_dict = {}
for key, value in save_dict.items():
# lora_transformer_transformer_blocks_7_attn_to_v.lokr_w1 to lycoris_transformer_blocks_7_attn_to_v.lokr_w1
new_key = key
new_key = new_key.replace('lora_transformer_', 'lycoris_')
new_save_dict[new_key] = value
save_dict = new_save_dict
if self.base_model_ref is not None:
save_dict = self.base_model_ref().convert_lora_weights_before_save(save_dict)
return save_dict
def save_weights(
self: Network,
file, dtype=torch.float16,
metadata=None,
extra_state_dict: Optional[OrderedDict] = None
):
save_dict = self.get_state_dict(extra_state_dict=extra_state_dict, dtype=dtype)
if metadata is not None and len(metadata) == 0:
metadata = None
if metadata is None:
metadata = OrderedDict()
metadata = add_model_hash_to_meta(save_dict, metadata)
# let the model handle the saving
if self.base_model_ref is not None and hasattr(self.base_model_ref(), 'save_lora'):
# call the base model save lora method
self.base_model_ref().save_lora(save_dict, file, metadata)
return
if os.path.splitext(file)[1] == ".safetensors":
from safetensors.torch import save_file
save_file(save_dict, file, metadata)
else:
torch.save(save_dict, file)
def load_weights(self: Network, file, force_weight_mapping=False):
# allows us to save and load to and from ldm weights
keymap = self.get_keymap(force_weight_mapping)
keymap = {} if keymap is None else keymap
if isinstance(file, str):
if self.base_model_ref is not None and hasattr(self.base_model_ref(), 'load_lora'):
# call the base model load lora method
weights_sd = self.base_model_ref().load_lora(file)
else:
if os.path.splitext(file)[1] == ".safetensors":
from safetensors.torch import load_file
weights_sd = load_file(file)
else:
weights_sd = torch.load(file, map_location="cpu")
else:
# probably a state dict
weights_sd = file
if self.base_model_ref is not None:
weights_sd = self.base_model_ref().convert_lora_weights_before_load(weights_sd)
load_sd = OrderedDict()
for key, value in weights_sd.items():
load_key = keymap[key] if key in keymap else key
# replace old double __ with single _
if self.is_pixart:
load_key = load_key.replace('__', '_')
if self.peft_format:
# lora_down = lora_A
# lora_up = lora_B
# no alpha
if load_key.endswith('.alpha'):
continue
load_key = load_key.replace('lora_A', 'lora_down')
load_key = load_key.replace('lora_B', 'lora_up')
# replace all . with $$
load_key = load_key.replace('.', '$$')
load_key = load_key.replace('$$lora_down$$', '.lora_down.')
load_key = load_key.replace('$$lora_up$$', '.lora_up.')
if self.network_type.lower() == "lokr":
# lora_transformer_transformer_blocks_7_attn_to_v.lokr_w1 to lycoris_transformer_blocks_7_attn_to_v.lokr_w1
load_key = load_key.replace('lycoris_', 'lora_transformer_')
load_sd[load_key] = value
# extract extra items from state dict
current_state_dict = self.state_dict()
extra_dict = OrderedDict()
to_delete = []
for key in list(load_sd.keys()):
if key not in current_state_dict:
extra_dict[key] = load_sd[key]
to_delete.append(key)
elif "lora_down" in key or "lora_up" in key:
# handle expanding/shrinking LoRA (linear only)
if len(load_sd[key].shape) == 2:
load_value = load_sd[key] # from checkpoint
blank_val = current_state_dict[key] # shape we need in the target model
tgt_h, tgt_w = blank_val.shape
src_h, src_w = load_value.shape
if (src_h, src_w) == (tgt_h, tgt_w):
# shapes already match: keep original
pass
elif "lora_down" in key and src_h < tgt_h:
print_once(f"Expanding {key} from {load_value.shape} to {blank_val.shape}")
new_val = torch.zeros((tgt_h, tgt_w), device=load_value.device, dtype=load_value.dtype)
new_val[:src_h, :src_w] = load_value # src_w should already match
load_sd[key] = new_val
self.did_change_weights = True
elif "lora_up" in key and src_w < tgt_w:
print_once(f"Expanding {key} from {load_value.shape} to {blank_val.shape}")
new_val = torch.zeros((tgt_h, tgt_w), device=load_value.device, dtype=load_value.dtype)
new_val[:src_h, :src_w] = load_value # src_h should already match
load_sd[key] = new_val
self.did_change_weights = True
elif "lora_down" in key and src_h > tgt_h:
print_once(f"Shrinking {key} from {load_value.shape} to {blank_val.shape}")
load_sd[key] = load_value[:tgt_h, :tgt_w]
self.did_change_weights = True
elif "lora_up" in key and src_w > tgt_w:
print_once(f"Shrinking {key} from {load_value.shape} to {blank_val.shape}")
load_sd[key] = load_value[:tgt_h, :tgt_w]
self.did_change_weights = True
else:
# unexpected mismatch (e.g., both dims differ in a way that doesn't match lora_up/down semantics)
raise ValueError(f"Unhandled LoRA shape change for {key}: src={load_value.shape}, tgt={blank_val.shape}")
for key in to_delete:
del load_sd[key]
print(f"Missing keys: {to_delete}")
if len(to_delete) > 0 and self.is_v1 and not force_weight_mapping and not (
len(to_delete) == 1 and 'emb_params' in to_delete):
print(" Attempting to load with forced keymap")
return self.load_weights(file, force_weight_mapping=True)
info = self.load_state_dict(load_sd, False)
if len(extra_dict.keys()) == 0:
extra_dict = None
return extra_dict
@torch.no_grad()
def _update_torch_multiplier(self: Network):
# builds a tensor for fast usage in the forward pass of the network modules
# without having to set it in every single module every time it changes
multiplier = self._multiplier
# get first module
try:
first_module = self.get_all_modules()[0]
except IndexError:
raise ValueError("There are not any lora modules in this network. Check your config and try again")
if hasattr(first_module, 'lora_down'):
device = first_module.lora_down.weight.device
dtype = first_module.lora_down.weight.dtype
if hasattr(first_module.lora_down, '_memory_management_device'):
device = first_module.lora_down._memory_management_device
elif hasattr(first_module, 'lokr_w1'):
device = first_module.lokr_w1.device
dtype = first_module.lokr_w1.dtype
if hasattr(first_module.lokr_w1, '_memory_management_device'):
device = first_module.lokr_w1._memory_management_device
elif hasattr(first_module, 'lokr_w1_a'):
device = first_module.lokr_w1_a.device
dtype = first_module.lokr_w1_a.dtype
if hasattr(first_module.lokr_w1_a, '_memory_management_device'):
device = first_module.lokr_w1_a._memory_management_device
else:
raise ValueError("Unknown module type")
with torch.no_grad():
tensor_multiplier = None
if isinstance(multiplier, int) or isinstance(multiplier, float):
tensor_multiplier = torch.tensor((multiplier,)).to(device, dtype=dtype)
elif isinstance(multiplier, list):
tensor_multiplier = torch.tensor(multiplier).to(device, dtype=dtype)
elif isinstance(multiplier, torch.Tensor):
tensor_multiplier = multiplier.clone().detach().to(device, dtype=dtype)
self.torch_multiplier = tensor_multiplier.clone().detach()
@property
def multiplier(self) -> Union[float, List[float], List[List[float]]]:
return self._multiplier
@multiplier.setter
def multiplier(self, value: Union[float, List[float], List[List[float]]]):
# it takes time to update all the multipliers, so we only do it if the value has changed
if self._multiplier == value:
return
# if we are setting a single value but have a list, keep the list if every item is the same as value
self._multiplier = value
self._update_torch_multiplier()
# called when the context manager is entered
# ie: with network:
def __enter__(self: Network):
self.is_active = True
def __exit__(self: Network, exc_type, exc_value, tb):
self.is_active = False
def force_to(self: Network, device, dtype):
self.to(device, dtype)
loras = []
if hasattr(self, 'unet_loras'):
loras += self.unet_loras
if hasattr(self, 'text_encoder_loras'):
loras += self.text_encoder_loras
for lora in loras:
lora.to(device, dtype)
def get_all_modules(self: Network) -> List[Module]:
loras = []
if hasattr(self, 'unet_loras'):
loras += self.unet_loras
if hasattr(self, 'text_encoder_loras'):
loras += self.text_encoder_loras
return loras
def _update_checkpointing(self: Network):
for module in self.get_all_modules():
if self.is_checkpointing:
module.enable_gradient_checkpointing()
else:
module.disable_gradient_checkpointing()
def enable_gradient_checkpointing(self: Network):
# not supported
self.is_checkpointing = True
self._update_checkpointing()
def disable_gradient_checkpointing(self: Network):
# not supported
self.is_checkpointing = False
self._update_checkpointing()
def merge_in(self, merge_weight=1.0):
if self.network_type.lower() == 'dora':
return
self.is_merged_in = True
for module in self.get_all_modules():
module.merge_in(merge_weight)
def merge_out(self: Network, merge_weight=1.0):
if not self.is_merged_in:
return
self.is_merged_in = False
for module in self.get_all_modules():
module.merge_out(merge_weight)
def extract_weight(
self: Network,
extract_mode: ExtractMode = "existing",
extract_mode_param: Union[int, float] = None,
):
if extract_mode_param is None:
raise ValueError("extract_mode_param must be set")
for module in tqdm(self.get_all_modules(), desc="Extracting weights"):
module.extract_weight(
extract_mode=extract_mode,
extract_mode_param=extract_mode_param
)
def setup_lorm(self: Network, state_dict: Optional[Dict[str, Any]] = None):
for module in tqdm(self.get_all_modules(), desc="Extracting LoRM"):
module.setup_lorm(state_dict=state_dict)
def calculate_lorem_parameter_reduction(self):
params_reduced = 0
for module in self.get_all_modules():
num_orig_module_params = count_parameters(module.org_module[0])
num_lorem_params = count_parameters(module.lora_down) + count_parameters(module.lora_up)
params_reduced += (num_orig_module_params - num_lorem_params)
return params_reduced
+24
View File
@@ -0,0 +1,24 @@
import os
TOOLKIT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
CONFIG_ROOT = os.path.join(TOOLKIT_ROOT, 'config')
KEYMAPS_ROOT = os.path.join(TOOLKIT_ROOT, "toolkit", "keymaps")
ORIG_CONFIGS_ROOT = os.path.join(TOOLKIT_ROOT, "toolkit", "orig_configs")
DIFFUSERS_CONFIGS_ROOT = os.path.join(TOOLKIT_ROOT, "toolkit", "diffusers_configs")
COMFY_PATH = os.getenv("COMFY_PATH", None)
COMFY_MODELS_PATH = None
if COMFY_PATH:
COMFY_MODELS_PATH = os.path.join(COMFY_PATH, "models")
# check if ENV variable is set
if 'MODELS_PATH' in os.environ:
MODELS_PATH = os.environ['MODELS_PATH']
else:
MODELS_PATH = os.path.join(TOOLKIT_ROOT, "models")
def get_path(path):
# we allow absolute paths, but if it is not absolute, we assume it is relative to the toolkit root
if not os.path.isabs(path):
path = os.path.join(TOOLKIT_ROOT, path)
return path
+738
View File
@@ -0,0 +1,738 @@
import os
from typing import Optional, TYPE_CHECKING, List, Union, Tuple
import torch
from safetensors.torch import load_file, save_file
from tqdm import tqdm
import random
from toolkit.train_tools import get_torch_dtype
import itertools
if TYPE_CHECKING:
from toolkit.config_modules import SliderTargetConfig
class ACTION_TYPES_SLIDER:
ERASE_NEGATIVE = 0
ENHANCE_NEGATIVE = 1
class PromptEmbeds:
# text_embeds: torch.Tensor
# pooled_embeds: Union[torch.Tensor, None]
# attention_mask: Union[torch.Tensor, List[torch.Tensor], None]
def __init__(self, args: Union[Tuple[torch.Tensor], List[torch.Tensor], torch.Tensor], attention_mask=None) -> None:
if isinstance(args, list) or isinstance(args, tuple):
# xl
self.text_embeds = args[0]
self.pooled_embeds = args[1]
else:
# sdv1.x, sdv2.x
self.text_embeds = args
self.pooled_embeds = None
self.attention_mask = attention_mask
def to(self, *args, **kwargs):
if isinstance(self.text_embeds, list) or isinstance(self.text_embeds, tuple):
self.text_embeds = [t.to(*args, **kwargs) for t in self.text_embeds]
else:
self.text_embeds = self.text_embeds.to(*args, **kwargs)
if self.pooled_embeds is not None:
self.pooled_embeds = self.pooled_embeds.to(*args, **kwargs)
if self.attention_mask is not None:
if isinstance(self.attention_mask, list) or isinstance(self.attention_mask, tuple):
self.attention_mask = [t.to(*args, **kwargs) for t in self.attention_mask]
else:
self.attention_mask = self.attention_mask.to(*args, **kwargs)
return self
def detach(self):
new_embeds = self.clone()
if isinstance(new_embeds.text_embeds, list) or isinstance(new_embeds.text_embeds, tuple):
new_embeds.text_embeds = [t.detach() for t in new_embeds.text_embeds]
else:
new_embeds.text_embeds = new_embeds.text_embeds.detach()
if new_embeds.pooled_embeds is not None:
new_embeds.pooled_embeds = new_embeds.pooled_embeds.detach()
if new_embeds.attention_mask is not None:
if isinstance(new_embeds.attention_mask, list) or isinstance(new_embeds.attention_mask, tuple):
new_embeds.attention_mask = [t.detach() for t in new_embeds.attention_mask]
else:
new_embeds.attention_mask = new_embeds.attention_mask.detach()
return new_embeds
def clone(self):
if isinstance(self.text_embeds, list) or isinstance(self.text_embeds, tuple):
cloned_text_embeds = [t.clone() for t in self.text_embeds]
else:
cloned_text_embeds = self.text_embeds.clone()
if self.pooled_embeds is not None:
prompt_embeds = PromptEmbeds([cloned_text_embeds, self.pooled_embeds.clone()])
else:
if isinstance(cloned_text_embeds, list) or isinstance(cloned_text_embeds, tuple):
prompt_embeds = PromptEmbeds([cloned_text_embeds, None])
else:
prompt_embeds = PromptEmbeds(cloned_text_embeds)
if self.attention_mask is not None:
if isinstance(self.attention_mask, list) or isinstance(self.attention_mask, tuple):
prompt_embeds.attention_mask = [t.clone() for t in self.attention_mask]
else:
prompt_embeds.attention_mask = self.attention_mask.clone()
return prompt_embeds
def expand_to_batch(self, batch_size):
pe = self.clone()
if isinstance(pe.text_embeds, list) or isinstance(pe.text_embeds, tuple):
if len(pe.text_embeds[0].shape) == 2:
current_batch_size = len(pe.text_embeds)
else:
current_batch_size = pe.text_embeds[0].shape[0]
else:
current_batch_size = pe.text_embeds.shape[0]
if current_batch_size == batch_size:
return pe
if current_batch_size != 1:
raise Exception("Can only expand batch size for batch size 1")
if isinstance(pe.text_embeds, list) or isinstance(pe.text_embeds, tuple):
if len(pe.text_embeds[0].shape) == 2:
# batch is a list of tensors
pe.text_embeds = pe.text_embeds * batch_size
else:
pe.text_embeds = [t.expand(batch_size, -1) for t in pe.text_embeds]
else:
pe.text_embeds = pe.text_embeds.expand(batch_size, -1)
if pe.pooled_embeds is not None:
pe.pooled_embeds = pe.pooled_embeds.expand(batch_size, -1)
if pe.attention_mask is not None:
if isinstance(pe.attention_mask, list) or isinstance(pe.attention_mask, tuple):
pe.attention_mask = [t.expand(batch_size, -1) for t in pe.attention_mask]
else:
pe.attention_mask = pe.attention_mask.expand(batch_size, -1)
return pe
def save(self, path: str):
"""
Save the prompt embeds to a file.
:param path: The path to save the prompt embeds.
"""
pe = self.clone()
state_dict = {}
if isinstance(pe.text_embeds, list) or isinstance(pe.text_embeds, tuple):
for i, text_embed in enumerate(pe.text_embeds):
state_dict[f"text_embed_{i}"] = text_embed.cpu()
else:
state_dict["text_embed"] = pe.text_embeds.cpu()
if pe.pooled_embeds is not None:
state_dict["pooled_embed"] = pe.pooled_embeds.cpu()
if pe.attention_mask is not None:
if isinstance(pe.attention_mask, list) or isinstance(pe.attention_mask, tuple):
for i, attn in enumerate(pe.attention_mask):
state_dict[f"attention_mask_{i}"] = attn.cpu()
else:
state_dict["attention_mask"] = pe.attention_mask.cpu()
os.makedirs(os.path.dirname(path), exist_ok=True)
save_file(state_dict, path)
@classmethod
def load(cls, path: str) -> 'PromptEmbeds':
"""
Load the prompt embeds from a file.
:param path: The path to load the prompt embeds from.
:return: An instance of PromptEmbeds.
"""
state_dict = load_file(path, device='cpu')
text_embeds = []
pooled_embeds = None
attention_mask = []
is_list = False
for key in sorted(state_dict.keys()):
if key.startswith("text_embed_"):
is_list = True
text_embeds.append(state_dict[key])
elif key == "text_embed":
text_embeds.append(state_dict[key])
elif key == "pooled_embed":
pooled_embeds = state_dict[key]
elif key.startswith("attention_mask_"):
attention_mask.append(state_dict[key])
elif key == "attention_mask":
attention_mask.append(state_dict[key])
pe = cls(None)
pe.text_embeds = text_embeds
if len(text_embeds) == 1 and not is_list:
pe.text_embeds = text_embeds[0]
if pooled_embeds is not None:
pe.pooled_embeds = pooled_embeds
if len(attention_mask) > 0:
if len(attention_mask) == 1:
pe.attention_mask = attention_mask[0]
else:
pe.attention_mask = attention_mask
return pe
class EncodedPromptPair:
def __init__(
self,
target_class,
target_class_with_neutral,
positive_target,
positive_target_with_neutral,
negative_target,
negative_target_with_neutral,
neutral,
empty_prompt,
both_targets,
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
action_list=None,
multiplier=1.0,
multiplier_list=None,
weight=1.0,
target: 'SliderTargetConfig' = None,
):
self.target_class: PromptEmbeds = target_class
self.target_class_with_neutral: PromptEmbeds = target_class_with_neutral
self.positive_target: PromptEmbeds = positive_target
self.positive_target_with_neutral: PromptEmbeds = positive_target_with_neutral
self.negative_target: PromptEmbeds = negative_target
self.negative_target_with_neutral: PromptEmbeds = negative_target_with_neutral
self.neutral: PromptEmbeds = neutral
self.empty_prompt: PromptEmbeds = empty_prompt
self.both_targets: PromptEmbeds = both_targets
self.multiplier: float = multiplier
self.target: 'SliderTargetConfig' = target
if multiplier_list is not None:
self.multiplier_list: list[float] = multiplier_list
else:
self.multiplier_list: list[float] = [multiplier]
self.action: int = action
if action_list is not None:
self.action_list: list[int] = action_list
else:
self.action_list: list[int] = [action]
self.weight: float = weight
# simulate torch to for tensors
def to(self, *args, **kwargs):
self.target_class = self.target_class.to(*args, **kwargs)
self.target_class_with_neutral = self.target_class_with_neutral.to(*args, **kwargs)
self.positive_target = self.positive_target.to(*args, **kwargs)
self.positive_target_with_neutral = self.positive_target_with_neutral.to(*args, **kwargs)
self.negative_target = self.negative_target.to(*args, **kwargs)
self.negative_target_with_neutral = self.negative_target_with_neutral.to(*args, **kwargs)
self.neutral = self.neutral.to(*args, **kwargs)
self.empty_prompt = self.empty_prompt.to(*args, **kwargs)
self.both_targets = self.both_targets.to(*args, **kwargs)
return self
def detach(self):
self.target_class = self.target_class.detach()
self.target_class_with_neutral = self.target_class_with_neutral.detach()
self.positive_target = self.positive_target.detach()
self.positive_target_with_neutral = self.positive_target_with_neutral.detach()
self.negative_target = self.negative_target.detach()
self.negative_target_with_neutral = self.negative_target_with_neutral.detach()
self.neutral = self.neutral.detach()
self.empty_prompt = self.empty_prompt.detach()
self.both_targets = self.both_targets.detach()
return self
def concat_prompt_embeds(prompt_embeds: list["PromptEmbeds"]):
# --- pad text_embeds ---
if isinstance(prompt_embeds[0].text_embeds, (list, tuple)):
embed_list = []
for i in range(len(prompt_embeds[0].text_embeds)):
max_len = max(p.text_embeds[i].shape[1] for p in prompt_embeds)
padded = []
for p in prompt_embeds:
t = p.text_embeds[i]
if t.shape[1] < max_len:
pad = torch.zeros(
(t.shape[0], max_len - t.shape[1], *t.shape[2:]),
dtype=t.dtype,
device=t.device,
)
t = torch.cat([t, pad], dim=1)
padded.append(t)
embed_list.append(torch.cat(padded, dim=0))
text_embeds = embed_list
else:
max_len = max(p.text_embeds.shape[1] for p in prompt_embeds)
padded = []
for p in prompt_embeds:
t = p.text_embeds
if t.shape[1] < max_len:
pad = torch.zeros(
(t.shape[0], max_len - t.shape[1], *t.shape[2:]),
dtype=t.dtype,
device=t.device,
)
t = torch.cat([t, pad], dim=1)
padded.append(t)
text_embeds = torch.cat(padded, dim=0)
# --- pooled embeds ---
pooled_embeds = None
if prompt_embeds[0].pooled_embeds is not None:
pooled_embeds = torch.cat([p.pooled_embeds for p in prompt_embeds], dim=0)
# --- attention mask ---
attention_mask = None
if prompt_embeds[0].attention_mask is not None:
max_len = max(p.attention_mask.shape[1] for p in prompt_embeds)
padded = []
for p in prompt_embeds:
m = p.attention_mask
if m.shape[1] < max_len:
pad = torch.zeros(
(m.shape[0], max_len - m.shape[1]),
dtype=m.dtype,
device=m.device,
)
m = torch.cat([m, pad], dim=1)
padded.append(m)
attention_mask = torch.cat(padded, dim=0)
# wrap back into PromptEmbeds
pe = PromptEmbeds([text_embeds, pooled_embeds])
pe.attention_mask = attention_mask
return pe
def concat_prompt_pairs(prompt_pairs: list[EncodedPromptPair]):
weight = prompt_pairs[0].weight
target_class = concat_prompt_embeds([p.target_class for p in prompt_pairs])
target_class_with_neutral = concat_prompt_embeds([p.target_class_with_neutral for p in prompt_pairs])
positive_target = concat_prompt_embeds([p.positive_target for p in prompt_pairs])
positive_target_with_neutral = concat_prompt_embeds([p.positive_target_with_neutral for p in prompt_pairs])
negative_target = concat_prompt_embeds([p.negative_target for p in prompt_pairs])
negative_target_with_neutral = concat_prompt_embeds([p.negative_target_with_neutral for p in prompt_pairs])
neutral = concat_prompt_embeds([p.neutral for p in prompt_pairs])
empty_prompt = concat_prompt_embeds([p.empty_prompt for p in prompt_pairs])
both_targets = concat_prompt_embeds([p.both_targets for p in prompt_pairs])
# combine all the lists
action_list = []
multiplier_list = []
weight_list = []
for p in prompt_pairs:
action_list += p.action_list
multiplier_list += p.multiplier_list
return EncodedPromptPair(
target_class=target_class,
target_class_with_neutral=target_class_with_neutral,
positive_target=positive_target,
positive_target_with_neutral=positive_target_with_neutral,
negative_target=negative_target,
negative_target_with_neutral=negative_target_with_neutral,
neutral=neutral,
empty_prompt=empty_prompt,
both_targets=both_targets,
action_list=action_list,
multiplier_list=multiplier_list,
weight=weight,
target=prompt_pairs[0].target
)
def split_prompt_embeds(concatenated: PromptEmbeds, num_parts=None) -> List[PromptEmbeds]:
if num_parts is None:
# use batch size
num_parts = concatenated.text_embeds.shape[0]
if isinstance(concatenated.text_embeds, list) or isinstance(concatenated.text_embeds, tuple):
# split each part
text_embeds_splits = [
torch.chunk(text, num_parts, dim=0)
for text in concatenated.text_embeds
]
text_embeds_splits = list(zip(*text_embeds_splits))
else:
text_embeds_splits = torch.chunk(concatenated.text_embeds, num_parts, dim=0)
if concatenated.pooled_embeds is not None:
pooled_embeds_splits = torch.chunk(concatenated.pooled_embeds, num_parts, dim=0)
else:
pooled_embeds_splits = [None] * num_parts
prompt_embeds_list = [
PromptEmbeds([text, pooled])
for text, pooled in zip(text_embeds_splits, pooled_embeds_splits)
]
return prompt_embeds_list
def split_prompt_pairs(concatenated: EncodedPromptPair, num_embeds=None) -> List[EncodedPromptPair]:
target_class_splits = split_prompt_embeds(concatenated.target_class, num_embeds)
target_class_with_neutral_splits = split_prompt_embeds(concatenated.target_class_with_neutral, num_embeds)
positive_target_splits = split_prompt_embeds(concatenated.positive_target, num_embeds)
positive_target_with_neutral_splits = split_prompt_embeds(concatenated.positive_target_with_neutral, num_embeds)
negative_target_splits = split_prompt_embeds(concatenated.negative_target, num_embeds)
negative_target_with_neutral_splits = split_prompt_embeds(concatenated.negative_target_with_neutral, num_embeds)
neutral_splits = split_prompt_embeds(concatenated.neutral, num_embeds)
empty_prompt_splits = split_prompt_embeds(concatenated.empty_prompt, num_embeds)
both_targets_splits = split_prompt_embeds(concatenated.both_targets, num_embeds)
prompt_pairs = []
for i in range(len(target_class_splits)):
action_list_split = concatenated.action_list[i::len(target_class_splits)]
multiplier_list_split = concatenated.multiplier_list[i::len(target_class_splits)]
prompt_pair = EncodedPromptPair(
target_class=target_class_splits[i],
target_class_with_neutral=target_class_with_neutral_splits[i],
positive_target=positive_target_splits[i],
positive_target_with_neutral=positive_target_with_neutral_splits[i],
negative_target=negative_target_splits[i],
negative_target_with_neutral=negative_target_with_neutral_splits[i],
neutral=neutral_splits[i],
empty_prompt=empty_prompt_splits[i],
both_targets=both_targets_splits[i],
action_list=action_list_split,
multiplier_list=multiplier_list_split,
weight=concatenated.weight,
target=concatenated.target
)
prompt_pairs.append(prompt_pair)
return prompt_pairs
class PromptEmbedsCache:
prompts: dict[str, PromptEmbeds] = {}
def __setitem__(self, __name: str, __value: PromptEmbeds) -> None:
self.prompts[__name] = __value
def __getitem__(self, __name: str) -> Optional[PromptEmbeds]:
if __name in self.prompts:
return self.prompts[__name]
else:
return None
class EncodedAnchor:
def __init__(
self,
prompt,
neg_prompt,
multiplier=1.0,
multiplier_list=None
):
self.prompt = prompt
self.neg_prompt = neg_prompt
self.multiplier = multiplier
if multiplier_list is not None:
self.multiplier_list: list[float] = multiplier_list
else:
self.multiplier_list: list[float] = [multiplier]
def to(self, *args, **kwargs):
self.prompt = self.prompt.to(*args, **kwargs)
self.neg_prompt = self.neg_prompt.to(*args, **kwargs)
return self
def concat_anchors(anchors: list[EncodedAnchor]):
prompt = concat_prompt_embeds([a.prompt for a in anchors])
neg_prompt = concat_prompt_embeds([a.neg_prompt for a in anchors])
return EncodedAnchor(
prompt=prompt,
neg_prompt=neg_prompt,
multiplier_list=[a.multiplier for a in anchors]
)
def split_anchors(concatenated: EncodedAnchor, num_anchors: int = 4) -> List[EncodedAnchor]:
prompt_splits = split_prompt_embeds(concatenated.prompt, num_anchors)
neg_prompt_splits = split_prompt_embeds(concatenated.neg_prompt, num_anchors)
multiplier_list_splits = torch.chunk(torch.tensor(concatenated.multiplier_list), num_anchors)
anchors = []
for prompt, neg_prompt, multiplier in zip(prompt_splits, neg_prompt_splits, multiplier_list_splits):
anchor = EncodedAnchor(
prompt=prompt,
neg_prompt=neg_prompt,
multiplier=multiplier.tolist()
)
anchors.append(anchor)
return anchors
def get_permutations(s, max_permutations=8):
# Split the string by comma
phrases = [phrase.strip() for phrase in s.split(',')]
# remove empty strings
phrases = [phrase for phrase in phrases if len(phrase) > 0]
# shuffle the list
random.shuffle(phrases)
# Get all permutations
permutations = list([p for p in itertools.islice(itertools.permutations(phrases), max_permutations)])
# Convert the tuples back to comma separated strings
return [', '.join(permutation) for permutation in permutations]
def get_slider_target_permutations(target: 'SliderTargetConfig', max_permutations=8) -> List['SliderTargetConfig']:
from toolkit.config_modules import SliderTargetConfig
pos_permutations = get_permutations(target.positive, max_permutations=max_permutations)
neg_permutations = get_permutations(target.negative, max_permutations=max_permutations)
permutations = []
for pos, neg in itertools.product(pos_permutations, neg_permutations):
permutations.append(
SliderTargetConfig(
target_class=target.target_class,
positive=pos,
negative=neg,
multiplier=target.multiplier,
weight=target.weight
)
)
# shuffle the list
random.shuffle(permutations)
if len(permutations) > max_permutations:
permutations = permutations[:max_permutations]
return permutations
if TYPE_CHECKING:
from toolkit.stable_diffusion_model import StableDiffusion
@torch.no_grad()
def encode_prompts_to_cache(
prompt_list: list[str],
sd: "StableDiffusion",
cache: Optional[PromptEmbedsCache] = None,
prompt_tensor_file: Optional[str] = None,
) -> PromptEmbedsCache:
# TODO: add support for larger prompts
if cache is None:
cache = PromptEmbedsCache()
if prompt_tensor_file is not None:
# check to see if it exists
if os.path.exists(prompt_tensor_file):
# load it.
print(f"Loading prompt tensors from {prompt_tensor_file}")
prompt_tensors = load_file(prompt_tensor_file, device='cpu')
# add them to the cache
for prompt_txt, prompt_tensor in tqdm(prompt_tensors.items(), desc="Loading prompts", leave=False):
if prompt_txt.startswith("te:"):
prompt = prompt_txt[3:]
# text_embeds
text_embeds = prompt_tensor
pooled_embeds = None
# find pool embeds
if f"pe:{prompt}" in prompt_tensors:
pooled_embeds = prompt_tensors[f"pe:{prompt}"]
# make it
prompt_embeds = PromptEmbeds([text_embeds, pooled_embeds])
cache[prompt] = prompt_embeds.to(device='cpu', dtype=torch.float32)
if len(cache.prompts) == 0:
print("Prompt tensors not found. Encoding prompts..")
empty_prompt = ""
# encode empty_prompt
cache[empty_prompt] = sd.encode_prompt(empty_prompt)
for p in tqdm(prompt_list, desc="Encoding prompts", leave=False):
# build the cache
if cache[p] is None:
cache[p] = sd.encode_prompt(p).to(device="cpu", dtype=torch.float16)
# should we shard? It can get large
if prompt_tensor_file:
print(f"Saving prompt tensors to {prompt_tensor_file}")
state_dict = {}
for prompt_txt, prompt_embeds in cache.prompts.items():
state_dict[f"te:{prompt_txt}"] = prompt_embeds.text_embeds.to(
"cpu", dtype=get_torch_dtype('fp16')
)
if prompt_embeds.pooled_embeds is not None:
state_dict[f"pe:{prompt_txt}"] = prompt_embeds.pooled_embeds.to(
"cpu",
dtype=get_torch_dtype('fp16')
)
save_file(state_dict, prompt_tensor_file)
return cache
@torch.no_grad()
def build_prompt_pair_batch_from_cache(
cache: PromptEmbedsCache,
target: 'SliderTargetConfig',
neutral: Optional[str] = '',
) -> list[EncodedPromptPair]:
erase_negative = len(target.positive.strip()) == 0
enhance_positive = len(target.negative.strip()) == 0
both = not erase_negative and not enhance_positive
prompt_pair_batch = []
if both or erase_negative:
# print("Encoding erase negative")
prompt_pair_batch += [
# erase standard
EncodedPromptPair(
target_class=cache[target.target_class],
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
positive_target=cache[f"{target.positive}"],
positive_target_with_neutral=cache[f"{target.positive} {neutral}"],
negative_target=cache[f"{target.negative}"],
negative_target_with_neutral=cache[f"{target.negative} {neutral}"],
neutral=cache[neutral],
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
multiplier=target.multiplier,
both_targets=cache[f"{target.positive} {target.negative}"],
empty_prompt=cache[""],
weight=target.weight,
target=target
),
]
if both or enhance_positive:
# print("Encoding enhance positive")
prompt_pair_batch += [
# enhance standard, swap pos neg
EncodedPromptPair(
target_class=cache[target.target_class],
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
positive_target=cache[f"{target.negative}"],
positive_target_with_neutral=cache[f"{target.negative} {neutral}"],
negative_target=cache[f"{target.positive}"],
negative_target_with_neutral=cache[f"{target.positive} {neutral}"],
neutral=cache[neutral],
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
multiplier=target.multiplier,
both_targets=cache[f"{target.positive} {target.negative}"],
empty_prompt=cache[""],
weight=target.weight,
target=target
),
]
if both or enhance_positive:
# print("Encoding erase positive (inverse)")
prompt_pair_batch += [
# erase inverted
EncodedPromptPair(
target_class=cache[target.target_class],
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
positive_target=cache[f"{target.negative}"],
positive_target_with_neutral=cache[f"{target.negative} {neutral}"],
negative_target=cache[f"{target.positive}"],
negative_target_with_neutral=cache[f"{target.positive} {neutral}"],
neutral=cache[neutral],
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
both_targets=cache[f"{target.positive} {target.negative}"],
empty_prompt=cache[""],
multiplier=target.multiplier * -1.0,
weight=target.weight,
target=target
),
]
if both or erase_negative:
# print("Encoding enhance negative (inverse)")
prompt_pair_batch += [
# enhance inverted
EncodedPromptPair(
target_class=cache[target.target_class],
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
positive_target=cache[f"{target.positive}"],
positive_target_with_neutral=cache[f"{target.positive} {neutral}"],
negative_target=cache[f"{target.negative}"],
negative_target_with_neutral=cache[f"{target.negative} {neutral}"],
both_targets=cache[f"{target.positive} {target.negative}"],
neutral=cache[neutral],
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
empty_prompt=cache[""],
multiplier=target.multiplier * -1.0,
weight=target.weight,
target=target
),
]
return prompt_pair_batch
def build_latent_image_batch_for_prompt_pair(
pos_latent,
neg_latent,
prompt_pair: EncodedPromptPair,
prompt_chunk_size
):
erase_negative = len(prompt_pair.target.positive.strip()) == 0
enhance_positive = len(prompt_pair.target.negative.strip()) == 0
both = not erase_negative and not enhance_positive
prompt_pair_chunks = split_prompt_pairs(prompt_pair, prompt_chunk_size)
if both and len(prompt_pair_chunks) != 4:
raise Exception("Invalid prompt pair chunks")
if (erase_negative or enhance_positive) and len(prompt_pair_chunks) != 2:
raise Exception("Invalid prompt pair chunks")
latent_list = []
if both or erase_negative:
latent_list.append(pos_latent)
if both or enhance_positive:
latent_list.append(pos_latent)
if both or enhance_positive:
latent_list.append(neg_latent)
if both or erase_negative:
latent_list.append(neg_latent)
return torch.cat(latent_list, dim=0)
def inject_trigger_into_prompt(prompt, trigger=None, to_replace_list=None, add_if_not_present=True):
if trigger is None:
# process as empty string to remove any [trigger] tokens
trigger = ''
output_prompt = prompt
default_replacements = ["[name]", "[trigger]"]
replace_with = trigger
if to_replace_list is None:
to_replace_list = default_replacements
else:
to_replace_list += default_replacements
# remove duplicates
to_replace_list = list(set(to_replace_list))
# replace them all
for to_replace in to_replace_list:
# replace it
output_prompt = output_prompt.replace(to_replace, replace_with)
if trigger.strip() != "":
# see how many times replace_with is in the prompt
num_instances = output_prompt.count(replace_with)
if num_instances == 0 and add_if_not_present:
# add it to the beginning of the prompt
output_prompt = replace_with + " " + output_prompt
# if num_instances > 1:
# print(
# f"Warning: {trigger} token appears {num_instances} times in prompt {output_prompt}. This may cause issues.")
return output_prompt
+7
View File
@@ -0,0 +1,7 @@
lycoris-lora
optimum-quanto
safetensors
diffusers
transformers
accelerate
huggingface-hub
+330
View File
@@ -0,0 +1,330 @@
import json
import os
from collections import OrderedDict
from typing import TYPE_CHECKING, Literal, Optional, Union
import torch
from safetensors.torch import load_file, save_file
from toolkit.train_tools import get_torch_dtype
from toolkit.paths import KEYMAPS_ROOT
if TYPE_CHECKING:
from toolkit.stable_diffusion_model import StableDiffusion
def get_slices_from_string(s: str) -> tuple:
slice_strings = s.split(',')
slices = [eval(f"slice({component.strip()})") for component in slice_strings]
return tuple(slices)
def convert_state_dict_to_ldm_with_mapping(
diffusers_state_dict: 'OrderedDict',
mapping_path: str,
base_path: Union[str, None] = None,
device: str = 'cpu',
dtype: torch.dtype = torch.float32
) -> 'OrderedDict':
converted_state_dict = OrderedDict()
# load mapping
with open(mapping_path, 'r') as f:
mapping = json.load(f, object_pairs_hook=OrderedDict)
# keep track of keys not matched
ldm_matched_keys = []
diffusers_matched_keys = []
ldm_diffusers_keymap = mapping['ldm_diffusers_keymap']
ldm_diffusers_shape_map = mapping['ldm_diffusers_shape_map']
ldm_diffusers_operator_map = mapping['ldm_diffusers_operator_map']
# load base if it exists
# the base just has come keys like timing ids and stuff diffusers doesn't have or they don't match
if base_path is not None:
converted_state_dict = load_file(base_path, device)
# convert to the right dtype
for key in converted_state_dict:
converted_state_dict[key] = converted_state_dict[key].to(device, dtype=dtype)
# process operators first
for ldm_key in ldm_diffusers_operator_map:
# if the key cat is in the ldm key, we need to process it
if 'cat' in ldm_diffusers_operator_map[ldm_key]:
cat_list = []
for diffusers_key in ldm_diffusers_operator_map[ldm_key]['cat']:
cat_list.append(diffusers_state_dict[diffusers_key].detach())
converted_state_dict[ldm_key] = torch.cat(cat_list, dim=0).to(device, dtype=dtype)
diffusers_matched_keys.extend(ldm_diffusers_operator_map[ldm_key]['cat'])
ldm_matched_keys.append(ldm_key)
if 'slice' in ldm_diffusers_operator_map[ldm_key]:
tensor_to_slice = diffusers_state_dict[ldm_diffusers_operator_map[ldm_key]['slice'][0]]
slice_text = diffusers_state_dict[ldm_diffusers_operator_map[ldm_key]['slice'][1]]
converted_state_dict[ldm_key] = tensor_to_slice[get_slices_from_string(slice_text)].detach().to(device,
dtype=dtype)
diffusers_matched_keys.extend(ldm_diffusers_operator_map[ldm_key]['slice'])
ldm_matched_keys.append(ldm_key)
# process the rest of the keys
for ldm_key in ldm_diffusers_keymap:
# if the key is in the ldm key, we need to process it
if ldm_diffusers_keymap[ldm_key] in diffusers_state_dict:
tensor = diffusers_state_dict[ldm_diffusers_keymap[ldm_key]].detach().to(device, dtype=dtype)
# see if we need to reshape
if ldm_key in ldm_diffusers_shape_map:
tensor = tensor.view(ldm_diffusers_shape_map[ldm_key][0])
converted_state_dict[ldm_key] = tensor
diffusers_matched_keys.append(ldm_diffusers_keymap[ldm_key])
ldm_matched_keys.append(ldm_key)
# see if any are missing from know mapping
mapped_diffusers_keys = list(ldm_diffusers_keymap.values())
mapped_ldm_keys = list(ldm_diffusers_keymap.keys())
missing_diffusers_keys = [x for x in mapped_diffusers_keys if x not in diffusers_matched_keys]
missing_ldm_keys = [x for x in mapped_ldm_keys if x not in ldm_matched_keys]
if len(missing_diffusers_keys) > 0:
print(f"WARNING!!!! Missing {len(missing_diffusers_keys)} diffusers keys")
print(missing_diffusers_keys)
if len(missing_ldm_keys) > 0:
print(f"WARNING!!!! Missing {len(missing_ldm_keys)} ldm keys")
print(missing_ldm_keys)
return converted_state_dict
def get_ldm_state_dict_from_diffusers(
state_dict: 'OrderedDict',
sd_version: Literal['1', '2', 'sdxl', 'ssd', 'vega', 'sdxl_refiner'] = '2',
device='cpu',
dtype=get_torch_dtype('fp32'),
):
if sd_version == '1':
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_sd1_ldm_base.safetensors')
mapping_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_sd1.json')
elif sd_version == '2':
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_sd2_ldm_base.safetensors')
mapping_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_sd2.json')
elif sd_version == 'sdxl':
# load our base
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_sdxl_ldm_base.safetensors')
mapping_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_sdxl.json')
elif sd_version == 'ssd':
# load our base
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_ssd_ldm_base.safetensors')
mapping_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_ssd.json')
elif sd_version == 'vega':
# load our base
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_vega_ldm_base.safetensors')
mapping_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_vega.json')
elif sd_version == 'sdxl_refiner':
# load our base
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_refiner_ldm_base.safetensors')
mapping_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_refiner.json')
else:
raise ValueError(f"Invalid sd_version {sd_version}")
# convert the state dict
return convert_state_dict_to_ldm_with_mapping(
state_dict,
mapping_path,
base_path,
device=device,
dtype=dtype
)
def save_ldm_model_from_diffusers(
sd: 'StableDiffusion',
output_file: str,
meta: 'OrderedDict',
save_dtype=get_torch_dtype('fp16'),
sd_version: Literal['1', '2', 'sdxl', 'ssd', 'vega'] = '2'
):
converted_state_dict = get_ldm_state_dict_from_diffusers(
sd.state_dict(),
sd_version,
device='cpu',
dtype=save_dtype
)
# make sure parent folder exists
os.makedirs(os.path.dirname(output_file), exist_ok=True)
save_file(converted_state_dict, output_file, metadata=meta)
def save_lora_from_diffusers(
lora_state_dict: 'OrderedDict',
output_file: str,
meta: 'OrderedDict',
save_dtype=get_torch_dtype('fp16'),
sd_version: Literal['1', '2', 'sdxl', 'ssd', 'vega'] = '2'
):
converted_state_dict = OrderedDict()
# only handle sxdxl for now
if sd_version != 'sdxl' and sd_version != 'ssd' and sd_version != 'vega':
raise ValueError(f"Invalid sd_version {sd_version}")
for key, value in lora_state_dict.items():
# todo verify if this works with ssd
# test encoders share keys for some reason
if key.begins_with('lora_te'):
converted_state_dict[key] = value.detach().to('cpu', dtype=save_dtype)
else:
converted_key = key
# make sure parent folder exists
os.makedirs(os.path.dirname(output_file), exist_ok=True)
save_file(converted_state_dict, output_file, metadata=meta)
def save_t2i_from_diffusers(
t2i_state_dict: 'OrderedDict',
output_file: str,
meta: 'OrderedDict',
dtype=get_torch_dtype('fp16'),
):
# todo: test compatibility with non diffusers
converted_state_dict = OrderedDict()
for key, value in t2i_state_dict.items():
converted_state_dict[key] = value.detach().to('cpu', dtype=dtype)
# make sure parent folder exists
os.makedirs(os.path.dirname(output_file), exist_ok=True)
save_file(converted_state_dict, output_file, metadata=meta)
def load_t2i_model(
path_to_file,
device: Union[str] = 'cpu',
dtype: torch.dtype = torch.float32
):
raw_state_dict = load_file(path_to_file, device)
converted_state_dict = OrderedDict()
for key, value in raw_state_dict.items():
# todo see if we need to convert dict
converted_state_dict[key] = value.detach().to(device, dtype=dtype)
return converted_state_dict
def save_ip_adapter_from_diffusers(
combined_state_dict: 'OrderedDict',
output_file: str,
meta: 'OrderedDict',
dtype=get_torch_dtype('fp16'),
direct_save: bool = False
):
# todo: test compatibility with non diffusers
converted_state_dict = OrderedDict()
for module_name, state_dict in combined_state_dict.items():
if direct_save:
converted_state_dict[module_name] = state_dict.detach().to('cpu', dtype=dtype)
else:
for key, value in state_dict.items():
converted_state_dict[f"{module_name}.{key}"] = value.detach().to('cpu', dtype=dtype)
# make sure parent folder exists
os.makedirs(os.path.dirname(output_file), exist_ok=True)
save_file(converted_state_dict, output_file, metadata=meta)
def load_ip_adapter_model(
path_to_file,
device: Union[str] = 'cpu',
dtype: torch.dtype = torch.float32,
direct_load: bool = False
):
# check if it is safetensors or checkpoint
if path_to_file.endswith('.safetensors'):
raw_state_dict = load_file(path_to_file, device)
combined_state_dict = OrderedDict()
if direct_load:
return raw_state_dict
for combo_key, value in raw_state_dict.items():
key_split = combo_key.split('.')
module_name = key_split.pop(0)
if module_name not in combined_state_dict:
combined_state_dict[module_name] = OrderedDict()
combined_state_dict[module_name]['.'.join(key_split)] = value.detach().to(device, dtype=dtype)
return combined_state_dict
else:
return torch.load(path_to_file, map_location=device)
def load_custom_adapter_model(
path_to_file,
device: Union[str] = 'cpu',
dtype: torch.dtype = torch.float32
):
# check if it is safetensors or checkpoint
if path_to_file.endswith('.safetensors'):
raw_state_dict = load_file(path_to_file, device)
combined_state_dict = OrderedDict()
device = device if isinstance(device, torch.device) else torch.device(device)
dtype = dtype if isinstance(dtype, torch.dtype) else get_torch_dtype(dtype)
for combo_key, value in raw_state_dict.items():
key_split = combo_key.split('.')
module_name = key_split.pop(0)
if module_name not in combined_state_dict:
combined_state_dict[module_name] = OrderedDict()
combined_state_dict[module_name]['.'.join(key_split)] = value.detach().to(device, dtype=dtype)
return combined_state_dict
else:
return torch.load(path_to_file, map_location=device)
def get_lora_keymap_from_model_keymap(model_keymap: 'OrderedDict') -> 'OrderedDict':
lora_keymap = OrderedDict()
# see if we have dual text encoders " a key that starts with conditioner.embedders.1
has_dual_text_encoders = False
for key in model_keymap:
if key.startswith('conditioner.embedders.1'):
has_dual_text_encoders = True
break
# map through the keys and values
for key, value in model_keymap.items():
# ignore bias weights
if key.endswith('bias'):
continue
if key.endswith('.weight'):
# remove the .weight
key = key[:-7]
if value.endswith(".weight"):
# remove the .weight
value = value[:-7]
# unet for all
key = key.replace('model.diffusion_model', 'lora_unet')
if value.startswith('unet'):
value = f"lora_{value}"
# text encoder
if has_dual_text_encoders:
key = key.replace('conditioner.embedders.0', 'lora_te1')
key = key.replace('conditioner.embedders.1', 'lora_te2')
if value.startswith('te0') or value.startswith('te1'):
value = f"lora_{value}"
value.replace('lora_te1', 'lora_te2')
value.replace('lora_te0', 'lora_te1')
key = key.replace('cond_stage_model.transformer', 'lora_te')
if value.startswith('te_'):
value = f"lora_{value}"
# replace periods with underscores
key = key.replace('.', '_')
value = value.replace('.', '_')
# add all the weights
lora_keymap[f"{key}.lora_down.weight"] = f"{value}.lora_down.weight"
lora_keymap[f"{key}.lora_down.bias"] = f"{value}.lora_down.bias"
lora_keymap[f"{key}.lora_up.weight"] = f"{value}.lora_up.weight"
lora_keymap[f"{key}.lora_up.bias"] = f"{value}.lora_up.bias"
lora_keymap[f"{key}.alpha"] = f"{value}.alpha"
return lora_keymap
+765
View File
@@ -0,0 +1,765 @@
import argparse
import hashlib
import json
import os
import time
from typing import TYPE_CHECKING, Union, List
import sys
from diffusers import (
DDPMScheduler,
EulerAncestralDiscreteScheduler,
DPMSolverMultistepScheduler,
DPMSolverSinglestepScheduler,
LMSDiscreteScheduler,
PNDMScheduler,
DDIMScheduler,
EulerDiscreteScheduler,
HeunDiscreteScheduler,
KDPM2DiscreteScheduler,
KDPM2AncestralDiscreteScheduler
)
import torch
import re
from transformers import T5Tokenizer, T5EncoderModel, UMT5EncoderModel
SCHEDULER_LINEAR_START = 0.00085
SCHEDULER_LINEAR_END = 0.0120
SCHEDULER_TIMESTEPS = 1000
SCHEDLER_SCHEDULE = "scaled_linear"
UNET_ATTENTION_TIME_EMBED_DIM = 256 # XL
TEXT_ENCODER_2_PROJECTION_DIM = 1280
UNET_PROJECTION_CLASS_EMBEDDING_INPUT_DIM = 2816
def get_torch_dtype(dtype_str):
# if it is a torch dtype, return it
if isinstance(dtype_str, torch.dtype):
return dtype_str
if dtype_str == "float" or dtype_str == "fp32" or dtype_str == "single" or dtype_str == "float32":
return torch.float
if dtype_str == "fp16" or dtype_str == "half" or dtype_str == "float16":
return torch.float16
if dtype_str == "bf16" or dtype_str == "bfloat16":
return torch.bfloat16
if dtype_str == "8bit" or dtype_str == "e4m3fn" or dtype_str == "float8":
return torch.float8_e4m3fn
return dtype_str
def replace_filewords_prompt(prompt, args: argparse.Namespace):
# if name_replace attr in args (may not be)
if hasattr(args, "name_replace") and args.name_replace is not None:
# replace [name] to args.name_replace
prompt = prompt.replace("[name]", args.name_replace)
if hasattr(args, "prepend") and args.prepend is not None:
# prepend to every item in prompt file
prompt = args.prepend + ' ' + prompt
if hasattr(args, "append") and args.append is not None:
# append to every item in prompt file
prompt = prompt + ' ' + args.append
return prompt
def replace_filewords_in_dataset_group(dataset_group, args: argparse.Namespace):
# if name_replace attr in args (may not be)
if hasattr(args, "name_replace") and args.name_replace is not None:
if not len(dataset_group.image_data) > 0:
# throw error
raise ValueError("dataset_group.image_data is empty")
for key in dataset_group.image_data:
dataset_group.image_data[key].caption = dataset_group.image_data[key].caption.replace(
"[name]", args.name_replace)
return dataset_group
def get_seeds_from_latents(latents):
# latents shape = (batch_size, 4, height, width)
# for speed we only use 8x8 slice of the first channel
seeds = []
# split batch up
for i in range(latents.shape[0]):
# use only first channel, multiply by 255 and convert to int
tensor = latents[i, 0, :, :] * 255.0 # shape = (height, width)
# slice 8x8
tensor = tensor[:8, :8]
# clip to 0-255
tensor = torch.clamp(tensor, 0, 255)
# convert to 8bit int
tensor = tensor.to(torch.uint8)
# convert to bytes
tensor_bytes = tensor.cpu().numpy().tobytes()
# hash
hash_object = hashlib.sha256(tensor_bytes)
# get hex
hex_dig = hash_object.hexdigest()
# convert to int
seed = int(hex_dig, 16) % (2 ** 32)
# append
seeds.append(seed)
return seeds
def get_noise_from_latents(latents):
seed_list = get_seeds_from_latents(latents)
noise = []
for seed in seed_list:
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
noise.append(torch.randn_like(latents[0]))
return torch.stack(noise)
# mix 0 is completely noise mean, mix 1 is completely target mean
def match_noise_to_target_mean_offset(noise, target, mix=0.5, dim=None):
dim = dim or (1, 2, 3)
# reduce mean of noise on dim 2, 3, keeping 0 and 1 intact
noise_mean = noise.mean(dim=dim, keepdim=True)
target_mean = target.mean(dim=dim, keepdim=True)
new_noise_mean = mix * target_mean + (1 - mix) * noise_mean
noise = noise - noise_mean + new_noise_mean
return noise
# https://www.crosslabs.org//blog/diffusion-with-offset-noise
def apply_noise_offset(noise, noise_offset):
if noise_offset is None or (noise_offset < 0.000001 and noise_offset > -0.000001):
return noise
if len(noise.shape) > 4:
raise ValueError("Applying noise offset not supported for video models at this time.")
noise = noise + noise_offset * torch.randn((noise.shape[0], noise.shape[1], 1, 1), device=noise.device)
return noise
if TYPE_CHECKING:
from toolkit.stable_diffusion_model import PromptEmbeds
def concat_prompt_embeddings(
unconditional: 'PromptEmbeds',
conditional: 'PromptEmbeds',
n_imgs: int=0,
):
from toolkit.stable_diffusion_model import PromptEmbeds
text_embeds = torch.cat(
[unconditional.text_embeds, conditional.text_embeds]
).repeat_interleave(n_imgs, dim=0)
pooled_embeds = None
if unconditional.pooled_embeds is not None and conditional.pooled_embeds is not None:
pooled_embeds = torch.cat(
[unconditional.pooled_embeds, conditional.pooled_embeds]
).repeat_interleave(n_imgs, dim=0)
return PromptEmbeds([text_embeds, pooled_embeds])
def addnet_hash_safetensors(b):
"""New model hash used by sd-webui-additional-networks for .safetensors format files"""
hash_sha256 = hashlib.sha256()
blksize = 1024 * 1024
b.seek(0)
header = b.read(8)
n = int.from_bytes(header, "little")
offset = n + 8
b.seek(offset)
for chunk in iter(lambda: b.read(blksize), b""):
hash_sha256.update(chunk)
return hash_sha256.hexdigest()
def addnet_hash_legacy(b):
"""Old model hash used by sd-webui-additional-networks for .safetensors format files"""
m = hashlib.sha256()
b.seek(0x100000)
m.update(b.read(0x10000))
return m.hexdigest()[0:8]
if TYPE_CHECKING:
from transformers import CLIPTextModel, CLIPTokenizer, CLIPTextModelWithProjection
def text_tokenize(
tokenizer: 'CLIPTokenizer',
prompts: list[str],
truncate: bool = True,
max_length: int = None,
max_length_multiplier: int = 4,
):
# allow fo up to 4x the max length for long prompts
if max_length is None:
if truncate:
max_length = tokenizer.model_max_length
else:
# allow up to 4x the max length for long prompts
max_length = tokenizer.model_max_length * max_length_multiplier
input_ids = tokenizer(
prompts,
padding='max_length',
max_length=max_length,
truncation=True,
return_tensors="pt",
).input_ids
if truncate or max_length == tokenizer.model_max_length:
return input_ids
else:
# remove additional padding
num_chunks = input_ids.shape[1] // tokenizer.model_max_length
chunks = torch.chunk(input_ids, chunks=num_chunks, dim=1)
# New list to store non-redundant chunks
non_redundant_chunks = []
for chunk in chunks:
if not chunk.eq(chunk[0, 0]).all(): # Check if all elements in the chunk are the same as the first element
non_redundant_chunks.append(chunk)
input_ids = torch.cat(non_redundant_chunks, dim=1)
return input_ids
# https://github.com/huggingface/diffusers/blob/78922ed7c7e66c20aa95159c7b7a6057ba7d590d/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py#L334-L348
def text_encode_xl(
text_encoder: Union['CLIPTextModel', 'CLIPTextModelWithProjection'],
tokens: torch.FloatTensor,
num_images_per_prompt: int = 1,
max_length: int = 77, # not sure what default to put here, always pass one?
truncate: bool = True,
):
if truncate:
# normal short prompt 77 tokens max
prompt_embeds = text_encoder(
tokens.to(text_encoder.device), output_hidden_states=True
)
pooled_prompt_embeds = prompt_embeds[0]
prompt_embeds = prompt_embeds.hidden_states[-2] # always penultimate layer
else:
# handle long prompts
prompt_embeds_list = []
tokens = tokens.to(text_encoder.device)
pooled_prompt_embeds = None
for i in range(0, tokens.shape[-1], max_length):
# todo run it through the in a single batch
section_tokens = tokens[:, i: i + max_length]
embeds = text_encoder(section_tokens, output_hidden_states=True)
pooled_prompt_embed = embeds[0]
if pooled_prompt_embeds is None:
# we only want the first ( I think??)
pooled_prompt_embeds = pooled_prompt_embed
prompt_embed = embeds.hidden_states[-2] # always penultimate layer
prompt_embeds_list.append(prompt_embed)
prompt_embeds = torch.cat(prompt_embeds_list, dim=1)
bs_embed, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1)
return prompt_embeds, pooled_prompt_embeds
def encode_prompts_xl(
tokenizers: list['CLIPTokenizer'],
text_encoders: list[Union['CLIPTextModel', 'CLIPTextModelWithProjection']],
prompts: list[str],
prompts2: Union[list[str], None],
num_images_per_prompt: int = 1,
use_text_encoder_1: bool = True, # sdxl
use_text_encoder_2: bool = True, # sdxl
truncate: bool = True,
max_length=None,
dropout_prob=0.0,
) -> tuple[torch.FloatTensor, torch.FloatTensor]:
# text_encoder and text_encoder_2's penuultimate layer's output
text_embeds_list = []
pooled_text_embeds = None # always text_encoder_2's pool
if prompts2 is None:
prompts2 = prompts
for idx, (tokenizer, text_encoder) in enumerate(zip(tokenizers, text_encoders)):
# todo, we are using a blank string to ignore that encoder for now.
# find a better way to do this (zeroing?, removing it from the unet?)
prompt_list_to_use = prompts if idx == 0 else prompts2
if idx == 0 and not use_text_encoder_1:
prompt_list_to_use = ["" for _ in prompts]
if idx == 1 and not use_text_encoder_2:
prompt_list_to_use = ["" for _ in prompts]
if dropout_prob > 0.0:
# randomly drop out prompts
prompt_list_to_use = [
prompt if torch.rand(1).item() > dropout_prob else "" for prompt in prompt_list_to_use
]
text_tokens_input_ids = text_tokenize(tokenizer, prompt_list_to_use, truncate=truncate, max_length=max_length)
# set the max length for the next one
if idx == 0:
max_length = text_tokens_input_ids.shape[-1]
text_embeds, pooled_text_embeds = text_encode_xl(
text_encoder, text_tokens_input_ids, num_images_per_prompt, max_length=tokenizer.model_max_length,
truncate=truncate
)
text_embeds_list.append(text_embeds)
bs_embed = pooled_text_embeds.shape[0]
pooled_text_embeds = pooled_text_embeds.repeat(1, num_images_per_prompt).view(
bs_embed * num_images_per_prompt, -1
)
return torch.concat(text_embeds_list, dim=-1), pooled_text_embeds
def encode_prompts_sd3(
tokenizers: list['CLIPTokenizer'],
text_encoders: list[Union['CLIPTextModel', 'CLIPTextModelWithProjection', T5EncoderModel]],
prompts: list[str],
num_images_per_prompt: int = 1,
truncate: bool = True,
max_length=None,
dropout_prob=0.0,
pipeline = None,
):
text_embeds_list = []
pooled_text_embeds = None # always text_encoder_2's pool
prompt_2 = prompts
prompt_2 = [prompt_2] if isinstance(prompt_2, str) else prompt_2
prompt_3 = prompts
prompt_3 = [prompt_3] if isinstance(prompt_3, str) else prompt_3
device = text_encoders[0].device
prompt_embed, pooled_prompt_embed = pipeline._get_clip_prompt_embeds(
prompt=prompts,
device=device,
num_images_per_prompt=num_images_per_prompt,
clip_skip=None,
clip_model_index=0,
)
prompt_2_embed, pooled_prompt_2_embed = pipeline._get_clip_prompt_embeds(
prompt=prompt_2,
device=device,
num_images_per_prompt=num_images_per_prompt,
clip_skip=None,
clip_model_index=1,
)
clip_prompt_embeds = torch.cat([prompt_embed, prompt_2_embed], dim=-1)
t5_prompt_embed = pipeline._get_t5_prompt_embeds(
prompt=prompt_3,
num_images_per_prompt=num_images_per_prompt,
device=device
)
clip_prompt_embeds = torch.nn.functional.pad(
clip_prompt_embeds, (0, t5_prompt_embed.shape[-1] - clip_prompt_embeds.shape[-1])
)
prompt_embeds = torch.cat([clip_prompt_embeds, t5_prompt_embed], dim=-2)
pooled_prompt_embeds = torch.cat([pooled_prompt_embed, pooled_prompt_2_embed], dim=-1)
return prompt_embeds, pooled_prompt_embeds
# ref for long prompts https://github.com/huggingface/diffusers/issues/2136
def text_encode(text_encoder: 'CLIPTextModel', tokens, truncate: bool = True, max_length=None):
if max_length is None and not truncate:
raise ValueError("max_length must be set if truncate is True")
try:
tokens = tokens.to(text_encoder.device)
except Exception as e:
print(e)
print("tokens.device", tokens.device)
print("text_encoder.device", text_encoder.device)
raise e
if truncate:
return text_encoder(tokens)[0]
else:
# handle long prompts
prompt_embeds_list = []
for i in range(0, tokens.shape[-1], max_length):
prompt_embeds = text_encoder(tokens[:, i: i + max_length])[0]
prompt_embeds_list.append(prompt_embeds)
return torch.cat(prompt_embeds_list, dim=1)
def encode_prompts(
tokenizer: 'CLIPTokenizer',
text_encoder: 'CLIPTextModel',
prompts: list[str],
truncate: bool = True,
max_length=None,
dropout_prob=0.0,
):
if max_length is None:
max_length = tokenizer.model_max_length
if dropout_prob > 0.0:
# randomly drop out prompts
prompts = [
prompt if torch.rand(1).item() > dropout_prob else "" for prompt in prompts
]
text_tokens = text_tokenize(tokenizer, prompts, truncate=truncate, max_length=max_length)
text_embeddings = text_encode(text_encoder, text_tokens, truncate=truncate, max_length=max_length)
return text_embeddings
def encode_prompts_pixart(
tokenizer: 'T5Tokenizer',
text_encoder: 'T5EncoderModel',
prompts: list[str],
truncate: bool = True,
max_length=None,
dropout_prob=0.0,
):
if max_length is None:
# See Section 3.1. of the paper.
max_length = 120
if dropout_prob > 0.0:
# randomly drop out prompts
prompts = [
prompt if torch.rand(1).item() > dropout_prob else "" for prompt in prompts
]
text_inputs = tokenizer(
prompts,
padding="max_length",
max_length=max_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
untruncated_ids = tokenizer(prompts, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
text_input_ids, untruncated_ids
):
removed_text = tokenizer.batch_decode(untruncated_ids[:, max_length - 1: -1])
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.to(text_encoder.device)
text_input_ids = text_input_ids.to(text_encoder.device)
prompt_embeds = text_encoder(text_input_ids, attention_mask=prompt_attention_mask)
return prompt_embeds.last_hidden_state, prompt_attention_mask
def encode_prompts_auraflow(
tokenizer: 'T5Tokenizer',
text_encoder: 'UMT5EncoderModel',
prompts: list[str],
truncate: bool = True,
max_length=None,
dropout_prob=0.0,
):
if max_length is None:
max_length = 256
if dropout_prob > 0.0:
# randomly drop out prompts
prompts = [
prompt if torch.rand(1).item() > dropout_prob else "" for prompt in prompts
]
device = text_encoder.device
text_inputs = tokenizer(
prompts,
truncation=True,
max_length=max_length,
padding="max_length",
return_tensors="pt",
)
text_input_ids = text_inputs["input_ids"]
untruncated_ids = tokenizer(prompts, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
text_input_ids, untruncated_ids
):
removed_text = tokenizer.batch_decode(untruncated_ids[:, max_length - 1: -1])
text_inputs = {k: v.to(device) for k, v in text_inputs.items()}
prompt_embeds = text_encoder(**text_inputs)[0]
prompt_attention_mask = text_inputs["attention_mask"].unsqueeze(-1).expand(prompt_embeds.shape)
prompt_embeds = prompt_embeds * prompt_attention_mask
return prompt_embeds, prompt_attention_mask
def encode_prompts_flux(
tokenizer: List[Union['CLIPTokenizer','T5Tokenizer']],
text_encoder: List[Union['CLIPTextModel', 'T5EncoderModel']],
prompts: list[str],
truncate: bool = True,
max_length=None,
dropout_prob=0.0,
attn_mask: bool = False,
):
if max_length is None:
max_length = 512
if dropout_prob > 0.0:
# randomly drop out prompts
prompts = [
prompt if torch.rand(1).item() > dropout_prob else "" for prompt in prompts
]
device = text_encoder[0].device
dtype = text_encoder[0].dtype
batch_size = len(prompts)
# clip
text_inputs = tokenizer[0](
prompts,
padding="max_length",
max_length=tokenizer[0].model_max_length,
truncation=True,
return_overflowing_tokens=False,
return_length=False,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_embeds = text_encoder[0](text_input_ids.to(device), output_hidden_states=False)
# Use pooled output of CLIPTextModel
pooled_prompt_embeds = prompt_embeds.pooler_output
pooled_prompt_embeds = pooled_prompt_embeds.to(dtype=dtype, device=device)
# T5
text_inputs = tokenizer[1](
prompts,
padding="max_length",
max_length=max_length,
truncation=True,
return_length=False,
return_overflowing_tokens=False,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_embeds = text_encoder[1](text_input_ids.to(device), output_hidden_states=False)[0]
dtype = text_encoder[1].dtype
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
if attn_mask:
prompt_attention_mask = text_inputs["attention_mask"].unsqueeze(-1).expand(prompt_embeds.shape)
prompt_embeds = prompt_embeds * prompt_attention_mask.to(dtype=prompt_embeds.dtype, device=prompt_embeds.device)
return prompt_embeds, pooled_prompt_embeds
# for XL
def get_add_time_ids(
height: int,
width: int,
dynamic_crops: bool = False,
dtype: torch.dtype = torch.float32,
):
if dynamic_crops:
# random float scale between 1 and 3
random_scale = torch.rand(1).item() * 2 + 1
original_size = (int(height * random_scale), int(width * random_scale))
# random position
crops_coords_top_left = (
torch.randint(0, original_size[0] - height, (1,)).item(),
torch.randint(0, original_size[1] - width, (1,)).item(),
)
target_size = (height, width)
else:
original_size = (height, width)
crops_coords_top_left = (0, 0)
target_size = (height, width)
# this is expected as 6
add_time_ids = list(original_size + crops_coords_top_left + target_size)
# this is expected as 2816
passed_add_embed_dim = (
UNET_ATTENTION_TIME_EMBED_DIM * len(add_time_ids) # 256 * 6
+ TEXT_ENCODER_2_PROJECTION_DIM # + 1280
)
if passed_add_embed_dim != UNET_PROJECTION_CLASS_EMBEDDING_INPUT_DIM:
raise ValueError(
f"Model expects an added time embedding vector of length {UNET_PROJECTION_CLASS_EMBEDDING_INPUT_DIM}, but a vector of {passed_add_embed_dim} was created. The model has an incorrect config. Please check `unet.config.time_embedding_type` and `text_encoder_2.config.projection_dim`."
)
add_time_ids = torch.tensor([add_time_ids], dtype=dtype)
return add_time_ids
def concat_embeddings(
unconditional: torch.FloatTensor,
conditional: torch.FloatTensor,
n_imgs: int,
):
return torch.cat([unconditional, conditional]).repeat_interleave(n_imgs, dim=0)
def add_all_snr_to_noise_scheduler(noise_scheduler, device):
try:
if hasattr(noise_scheduler, "all_snr"):
return
# compute it
with torch.no_grad():
alphas_cumprod = noise_scheduler.alphas_cumprod
sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
alpha = sqrt_alphas_cumprod
sigma = sqrt_one_minus_alphas_cumprod
all_snr = (alpha / sigma) ** 2
all_snr.requires_grad = False
noise_scheduler.all_snr = all_snr.to(device)
except Exception as e:
# just move on
pass
def get_all_snr(noise_scheduler, device):
if hasattr(noise_scheduler, "all_snr"):
return noise_scheduler.all_snr.to(device)
# compute it
with torch.no_grad():
alphas_cumprod = noise_scheduler.alphas_cumprod
sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
alpha = sqrt_alphas_cumprod
sigma = sqrt_one_minus_alphas_cumprod
all_snr = (alpha / sigma) ** 2
all_snr.requires_grad = False
return all_snr.to(device)
class LearnableSNRGamma:
"""
This is a trainer for learnable snr gamma
It will adapt to the dataset and attempt to adjust the snr multiplier to balance the loss over the timesteps
"""
def __init__(self, noise_scheduler: Union['DDPMScheduler'], device='cuda'):
self.device = device
self.noise_scheduler: Union['DDPMScheduler'] = noise_scheduler
self.offset_1 = torch.nn.Parameter(torch.tensor(0.0, dtype=torch.float32, device=device))
self.offset_2 = torch.nn.Parameter(torch.tensor(0.777, dtype=torch.float32, device=device))
self.scale = torch.nn.Parameter(torch.tensor(4.14, dtype=torch.float32, device=device))
self.gamma = torch.nn.Parameter(torch.tensor(2.03, dtype=torch.float32, device=device))
self.optimizer = torch.optim.AdamW([self.offset_1, self.offset_2, self.gamma, self.scale], lr=0.01)
self.buffer = []
self.max_buffer_size = 20
def forward(self, loss, timesteps):
# do a our train loop for lsnr here and return our values detached
loss = loss.detach()
with torch.no_grad():
loss_chunks = torch.chunk(loss, loss.shape[0], dim=0)
for loss_chunk in loss_chunks:
self.buffer.append(loss_chunk.mean().detach())
if len(self.buffer) > self.max_buffer_size:
self.buffer.pop(0)
all_snr = get_all_snr(self.noise_scheduler, loss.device)
snr: torch.Tensor = torch.stack([all_snr[t] for t in timesteps]).detach().float().to(loss.device)
base_snrs = snr.clone().detach()
snr.requires_grad = True
snr = (snr + self.offset_1) * self.scale + self.offset_2
gamma_over_snr = torch.div(torch.ones_like(snr) * self.gamma, snr)
snr_weight = torch.abs(gamma_over_snr).float().to(loss.device) # directly using gamma over snr
snr_adjusted_loss = loss * snr_weight
with torch.no_grad():
target = torch.mean(torch.stack(self.buffer)).detach()
# local_loss = torch.mean(torch.abs(snr_adjusted_loss - target))
squared_differences = (snr_adjusted_loss - target) ** 2
local_loss = torch.mean(squared_differences)
local_loss.backward()
self.optimizer.step()
self.optimizer.zero_grad()
return base_snrs, self.gamma.detach(), self.offset_1.detach(), self.offset_2.detach(), self.scale.detach()
def apply_learnable_snr_gos(
loss,
timesteps,
learnable_snr_trainer: LearnableSNRGamma
):
snr, gamma, offset_1, offset_2, scale = learnable_snr_trainer.forward(loss, timesteps)
snr = (snr + offset_1) * scale + offset_2
gamma_over_snr = torch.div(torch.ones_like(snr) * gamma, snr)
snr_weight = torch.abs(gamma_over_snr).float().to(loss.device) # directly using gamma over snr
snr_adjusted_loss = loss * snr_weight
return snr_adjusted_loss
def apply_snr_weight(
loss,
timesteps,
noise_scheduler: Union['DDPMScheduler'],
gamma,
fixed=False,
):
# will get it from noise scheduler if exist or will calculate it if not
all_snr = get_all_snr(noise_scheduler, loss.device)
# step_indices = []
# for t in timesteps:
# for i, st in enumerate(noise_scheduler.timesteps):
# if st == t:
# step_indices.append(i)
# break
# this breaks on some schedulers
# step_indices = [(noise_scheduler.timesteps == t).nonzero().item() for t in timesteps]
offset = 0
if noise_scheduler.timesteps[0] == 1000:
offset = 1
snr = torch.stack([all_snr[(t - offset).int()] for t in timesteps])
gamma_over_snr = torch.div(torch.ones_like(snr) * gamma, snr)
if fixed:
snr_weight = gamma_over_snr.float().to(loss.device) # directly using gamma over snr
else:
snr_weight = torch.minimum(gamma_over_snr, torch.ones_like(gamma_over_snr)).float().to(loss.device)
snr_adjusted_loss = loss * snr_weight
return snr_adjusted_loss
def precondition_model_outputs_flow_match(model_output, model_input, timestep_tensor, noise_scheduler):
mo_chunks = torch.chunk(model_output, model_output.shape[0], dim=0)
mi_chunks = torch.chunk(model_input, model_input.shape[0], dim=0)
timestep_chunks = torch.chunk(timestep_tensor, timestep_tensor.shape[0], dim=0)
out_chunks = []
# unsqueeze if timestep is zero dim
for idx in range(model_output.shape[0]):
sigmas = noise_scheduler.get_sigmas(timestep_chunks[idx], n_dim=model_output.ndim,
dtype=model_output.dtype, device=model_output.device)
# Follow: Section 5 of https://arxiv.org/abs/2206.00364.
# Preconditioning of the model outputs.
out = mo_chunks[idx] * (-sigmas) + mi_chunks[idx]
out_chunks.append(out)
return torch.cat(out_chunks, dim=0)
+559
View File
@@ -0,0 +1,559 @@
import os
import types
from typing import List, Optional
import huggingface_hub
import torch
import torch.utils.checkpoint
import yaml
from toolkit.config_modules import GenerateImageConfig, ModelConfig, NetworkConfig
from toolkit.lora_special import LoRASpecialNetwork
from toolkit.models.base_model import BaseModel
from toolkit.basic import flush
from toolkit.prompt_utils import PromptEmbeds
from toolkit.samplers.custom_flowmatch_sampler import (
CustomFlowMatchEulerDiscreteScheduler,
)
from toolkit.accelerator import unwrap_model
from optimum.quanto import freeze
from toolkit.util.quantize import quantize, get_qtype, quantize_model
from toolkit.memory_management import MemoryManager
from safetensors.torch import load_file
from transformers import AutoTokenizer, Qwen3ForCausalLM
from diffusers import AutoencoderKL
try:
from diffusers import ZImagePipeline
from diffusers.models.transformers import ZImageTransformer2DModel
except ImportError:
raise ImportError(
"Diffusers is out of date. Update diffusers to the latest version."
)
scheduler_config = {
"num_train_timesteps": 1000,
"use_dynamic_shifting": False,
"shift": 3.0,
}
# --- SURGICAL TOOL: MONKEY PATCH ---
# Prevents AI Toolkit from crashing when it tries to set requires_grad=True
# on quantized (integer) weights.
def safe_requires_grad_(self, requires_grad=True):
for param in self.parameters():
if param.dtype.is_floating_point:
param.requires_grad = requires_grad
return self
class ZImageModel(BaseModel):
arch = "zimage"
def __init__(
self,
device,
model_config: ModelConfig,
dtype="bf16",
custom_pipeline=None,
noise_scheduler=None,
**kwargs,
):
super().__init__(
device, model_config, dtype, custom_pipeline, noise_scheduler, **kwargs
)
self.is_flow_matching = True
self.is_transformer = True
self.target_lora_modules = ["ZImageTransformer2DModel", "Qwen3ForCausalLM", "Qwen2ForCausalLM"]
@staticmethod
def get_train_scheduler():
return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
def get_bucket_divisibility(self):
return 16 * 2
def load_training_adapter(self, transformer: ZImageTransformer2DModel):
self.print_and_status_update("Loading assistant LoRA")
lora_path = self.model_config.assistant_lora_path
if not os.path.exists(lora_path):
lora_splits = lora_path.split("/")
if len(lora_splits) != 3:
raise ValueError(f"Invalid LoRA path: {lora_path}")
repo_id = "/".join(lora_splits[:2])
filename = lora_splits[2]
try:
lora_path = huggingface_hub.hf_hub_download(repo_id=repo_id, filename=filename)
self.model_config.assistant_lora_path = lora_path
except Exception as e:
raise ValueError(f"Failed to download assistant LoRA: {e}")
lora_state_dict = load_file(lora_path)
dim = int(lora_state_dict["diffusion_model.layers.0.attention.to_k.lora_A.weight"].shape[0])
new_sd = {}
for key, value in lora_state_dict.items():
new_key = key.replace("diffusion_model.", "transformer.")
new_sd[new_key] = value
lora_state_dict = new_sd
network_config = NetworkConfig(type="lora", linear=dim, linear_alpha=dim, transformer_only=True)
LoRASpecialNetwork.LORA_PREFIX_UNET = "lora_transformer"
network = LoRASpecialNetwork(
text_encoder=None,
unet=transformer,
lora_dim=network_config.linear,
multiplier=1.0,
alpha=network_config.linear_alpha,
train_unet=True,
train_text_encoder=False,
network_config=network_config,
network_type=network_config.type,
transformer_only=network_config.transformer_only,
is_transformer=True,
target_lin_modules=self.target_lora_modules,
is_assistant_adapter=True,
is_ara=True,
)
network.apply_to(None, transformer, apply_text_encoder=False, apply_unet=True)
self.print_and_status_update("Merging in assistant LoRA")
network.force_to(transformer.device, dtype=self.torch_dtype)
network._update_torch_multiplier()
network.load_weights(lora_state_dict)
network.merge_in(merge_weight=1.0)
network.is_merged_in = False
self.assistant_lora = network
self.assistant_lora.multiplier = -1.0
self.assistant_lora.is_active = False
self.invert_assistant_lora = True
def load_model(self):
dtype = self.torch_dtype
self.print_and_status_update("Loading ZImage model")
model_path = self.model_config.name_or_path
base_model_path = self.model_config.extras_name_or_path
self.print_and_status_update("Loading transformer")
transformer_path = model_path
transformer_subfolder = "transformer"
if os.path.exists(transformer_path):
transformer_subfolder = None
transformer_path = os.path.join(transformer_path, "transformer")
te_folder_path = os.path.join(model_path, "text_encoder")
if os.path.exists(te_folder_path):
base_model_path = model_path
# 1. Load Transformer
transformer = ZImageTransformer2DModel.from_pretrained(
transformer_path,
subfolder=transformer_subfolder,
torch_dtype=dtype,
low_cpu_mem_usage=False,
ignore_mismatched_sizes=True
)
if self.model_config.assistant_lora_path is not None:
self.load_training_adapter(transformer)
# --- SURGICAL MODIFICATION: QUANTIZATION (QUANTO) ---
should_quantize_transformer = self.model_config.quantize or (
self.model_config.low_vram and not self.model_config.train_unet
)
if should_quantize_transformer:
if self.model_config.qtype == "qfloat8":
self.model_config.qtype = "float8"
self.print_and_status_update(f"Surgical Plan: Quantizing Transformer (Quanto)")
transformer.requires_grad_(False)
# Monkey Patch BEFORE quantize
transformer.requires_grad_ = types.MethodType(safe_requires_grad_, transformer)
quantize_model(self, transformer)
flush()
# --- FIX: TROJAN HORSE PARAMETER ---
transformer.dummy_param = torch.nn.Parameter(torch.zeros(1, dtype=dtype, device=self.device_torch))
transformer.dummy_param.requires_grad = True
# Enable input grads (backup mechanism)
if hasattr(transformer, "enable_input_require_grads"):
transformer.enable_input_require_grads()
else:
def make_inputs_require_grad(module, input, output):
output.requires_grad_(True)
if hasattr(transformer, "patch_embed"):
transformer.patch_embed.register_forward_hook(make_inputs_require_grad)
elif hasattr(transformer, "pos_embed"):
transformer.pos_embed.register_forward_hook(make_inputs_require_grad)
if (self.model_config.layer_offloading and self.model_config.layer_offloading_transformer_percent > 0):
MemoryManager.attach(
transformer,
self.device_torch,
offload_percent=self.model_config.layer_offloading_transformer_percent,
ignore_modules=[transformer.x_pad_token, transformer.cap_pad_token]
)
# --- SURGICAL FIX: RESIDENCY ---
train_te = getattr(self.model_config, 'train_text_encoder', False)
if self.model_config.low_vram and not train_te:
self.print_and_status_update("Moving transformer to CPU")
transformer.to("cpu")
else:
self.print_and_status_update("Surgical Plan: Keeping Transformer on GPU for Gradient Flow")
transformer.to(self.device_torch)
flush()
self.print_and_status_update("Text Encoder")
tokenizer = AutoTokenizer.from_pretrained(
base_model_path, subfolder="tokenizer", torch_dtype=dtype
)
text_encoder = Qwen3ForCausalLM.from_pretrained(
base_model_path, subfolder="text_encoder", torch_dtype=dtype
)
if (self.model_config.layer_offloading and self.model_config.layer_offloading_text_encoder_percent > 0):
MemoryManager.attach(
text_encoder,
self.device_torch,
offload_percent=self.model_config.layer_offloading_text_encoder_percent,
)
# --- SURGICAL MODIFICATION: TEXT ENCODER HANDLING ---
if train_te:
self.print_and_status_update("Surgical Plan: Text Encoder Training Active - Preserving BF16")
text_encoder.to(self.device_torch, dtype=dtype)
text_encoder.requires_grad_(False)
if hasattr(text_encoder, "config"):
text_encoder.config.use_cache = False
self.print_and_status_update("Enabling Gradient Checkpointing for Text Encoder")
text_encoder.gradient_checkpointing_enable()
if hasattr(text_encoder, "enable_input_require_grads"):
text_encoder.enable_input_require_grads()
text_encoder.train()
else:
text_encoder.to(self.device_torch, dtype=dtype)
if self.model_config.quantize_te:
self.print_and_status_update("Quantizing Text Encoder")
quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te))
freeze(text_encoder)
flush()
self.print_and_status_update("Loading VAE")
vae = AutoencoderKL.from_pretrained(
base_model_path, subfolder="vae", torch_dtype=dtype
)
self.noise_scheduler = ZImageModel.get_train_scheduler()
self.print_and_status_update("Making pipe")
kwargs = {} # Fixed kwargs error
pipe: ZImagePipeline = ZImagePipeline(
scheduler=self.noise_scheduler,
text_encoder=None,
tokenizer=tokenizer,
vae=vae,
transformer=None,
**kwargs,
)
pipe.text_encoder = text_encoder
pipe.transformer = transformer
self.print_and_status_update("Preparing Model")
text_encoder = [pipe.text_encoder]
tokenizer = [pipe.tokenizer]
if not self.low_vram:
pipe.transformer = pipe.transformer.to(self.device_torch)
flush()
text_encoder[0].to(self.device_torch)
if not train_te:
text_encoder[0].requires_grad_(False)
text_encoder[0].eval()
flush()
self.vae = vae
self.text_encoder = text_encoder
self.tokenizer = tokenizer
self.model = pipe.transformer
# --- FIX: Alias UNet for BaseModel compatibility ---
self.unet = self.model
self.pipeline = pipe
self.print_and_status_update("Model Loaded")
# --- SURGICAL FIX: Custom Device State Handler ---
def set_device_state(self, state):
# Helper to get attributes safe for dict or object
def get_state_attr(obj, name, default=None):
if isinstance(obj, dict):
return obj.get(name, default)
return getattr(obj, name, default)
target_device = get_state_attr(state, 'device')
if self.text_encoder is not None:
if isinstance(self.text_encoder, list):
for te in self.text_encoder:
te.to(target_device)
else:
self.text_encoder.to(target_device)
if self.vae is not None:
self.vae.to(target_device)
if self.transformer is not None:
self.transformer.to(target_device)
should_train_unet = get_state_attr(state, 'train_unet', False)
require_grads = get_state_attr(state, 'require_grads', False)
if should_train_unet:
self.transformer.train()
else:
self.transformer.eval()
# Force train mode if using checkpointing for TE training
if getattr(self.model_config, 'train_text_encoder', False):
self.transformer.train()
# Apply grads SAFELY using monkey patched logic logic if needed,
# or manual check here
target_grad = should_train_unet or require_grads
for param in self.transformer.parameters():
if param.dtype.is_floating_point:
param.requires_grad_(target_grad)
else:
param.requires_grad_(False)
def get_generation_pipeline(self):
scheduler = ZImageModel.get_train_scheduler()
pipeline: ZImagePipeline = ZImagePipeline(
scheduler=scheduler,
text_encoder=unwrap_model(self.text_encoder[0]),
tokenizer=self.tokenizer[0],
vae=unwrap_model(self.vae),
transformer=unwrap_model(self.transformer),
)
pipeline = pipeline.to(self.device_torch)
return pipeline
def generate_single_image(
self,
pipeline: ZImagePipeline,
gen_config: GenerateImageConfig,
conditional_embeds: PromptEmbeds,
unconditional_embeds: PromptEmbeds,
generator: torch.Generator,
extra: dict,
):
self.model.to(self.device_torch, dtype=self.torch_dtype)
self.model.to(self.device_torch)
sc = self.get_bucket_divisibility()
gen_config.width = int(gen_config.width // sc * sc)
gen_config.height = int(gen_config.height // sc * sc)
img = pipeline(
prompt_embeds=conditional_embeds.text_embeds,
negative_prompt_embeds=unconditional_embeds.text_embeds,
height=gen_config.height,
width=gen_config.width,
num_inference_steps=gen_config.num_inference_steps,
guidance_scale=gen_config.guidance_scale,
latents=gen_config.latents,
generator=generator,
**extra,
).images[0]
return img
def get_noise_prediction(
self,
latent_model_input: torch.Tensor,
timestep: torch.Tensor,
text_embeddings: PromptEmbeds,
**kwargs,
):
self.model.to(self.device_torch)
self.transformer.train() # Force train for checkpointing
# --- SURGICAL FIX: Jumper Cable (External Checkpointing) ---
# 1. Prepare Inputs (Must require grad)
latent_model_input.requires_grad_(True)
timestep_model_input = (1000 - timestep) / 1000
encoder_hidden_states = text_embeddings.text_embeds
if isinstance(encoder_hidden_states, list):
encoder_hidden_states = encoder_hidden_states[0]
text_embeddings.text_embeds = encoder_hidden_states
# FIX: Ensure 3D (Batch, Seq, Dim)
if torch.is_tensor(encoder_hidden_states):
if encoder_hidden_states.ndim == 2:
encoder_hidden_states = encoder_hidden_states.unsqueeze(0)
encoder_hidden_states.requires_grad_(True)
# 2. Define Wrapper (Handles unpacking for Z-Image)
def _forward_wrapper(latents_batch, t_batch, encoder_hidden_batch):
# Unpack Latents: (B, C, H, W) -> List[(C, 1, H, W)]
latents_list = list(latents_batch.unsqueeze(2).unbind(0))
# Unpack Embeds: (B, Seq, Dim) -> List[(Seq, Dim)]
# Fixes "1 and 2" dimensions error by strictly providing 2D tensors
encoder_list = list(encoder_hidden_batch.unbind(0))
# Run Transformer (Black Box)
output = self.transformer(latents_list, t_batch, encoder_list)
# Repack: List[Tensor] -> Tensor
return torch.stack([t.float() for t in output[0]], dim=0)
# 3. Execute via Jumper Cable
noise_pred = torch.utils.checkpoint.checkpoint(
_forward_wrapper,
latent_model_input,
timestep_model_input,
encoder_hidden_states,
use_reentrant=False
)
noise_pred = noise_pred.squeeze(2)
noise_pred = -noise_pred
# --- SURGICAL FIX: TROJAN HORSE CONNECTION ---
if hasattr(self.transformer, "dummy_param"):
if self.transformer.dummy_param.device != noise_pred.device:
self.transformer.dummy_param.data = self.transformer.dummy_param.data.to(noise_pred.device)
loss_proxy = self.transformer.dummy_param.sum() * 0
noise_pred = noise_pred + loss_proxy
return noise_pred
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
# SURGICAL FIX: Manual Pipeline Bypass & Direct Input Injection
train_te = getattr(self.model_config, 'train_text_encoder', False)
if not train_te:
if self.pipeline.text_encoder.device != self.device_torch:
self.pipeline.text_encoder.to(self.device_torch)
prompt_embeds, _ = self.pipeline.encode_prompt(
prompt,
do_classifier_free_guidance=False,
device=self.device_torch,
)
return PromptEmbeds([prompt_embeds, None])
tokenizer = self.tokenizer[0] if isinstance(self.tokenizer, list) else self.tokenizer
text_encoder = self.text_encoder[0] if isinstance(self.text_encoder, list) else self.text_encoder
text_encoder.to(self.device_torch)
if isinstance(prompt, str):
prompt = [prompt]
max_len = getattr(tokenizer, 'model_max_length', 512)
if max_len > 1024: max_len = 512
text_inputs = tokenizer(
prompt,
padding="max_length",
max_length=max_len,
truncation=True,
return_tensors="pt",
).to(self.device_torch)
# Direct Injection Logic
with torch.set_grad_enabled(True):
input_embed_layer = text_encoder.get_input_embeddings()
inputs_embeds = input_embed_layer(text_inputs.input_ids)
inputs_embeds.requires_grad_(True)
outputs = text_encoder(
inputs_embeds=inputs_embeds,
attention_mask=text_inputs.attention_mask,
output_hidden_states=True
)
if hasattr(outputs, "hidden_states"):
prompt_embeds = outputs.hidden_states[-1]
else:
prompt_embeds = outputs[0]
# FORCE 3D SHAPE (Batch, Seq, Dim)
if prompt_embeds.ndim == 2:
prompt_embeds = prompt_embeds.unsqueeze(0)
return PromptEmbeds([prompt_embeds, None])
def get_model_has_grad(self):
if self.model is None:
return False
return any(p.requires_grad for p in self.model.parameters())
def get_te_has_grad(self):
if self.text_encoder is None:
return False
te0 = self.text_encoder[0] if isinstance(self.text_encoder, list) else self.text_encoder
return any(p.requires_grad for p in te0.parameters())
def save_model(self, output_path, meta, save_dtype):
transformer: ZImageTransformer2DModel = unwrap_model(self.model)
transformer.save_pretrained(
save_directory=os.path.join(output_path, "transformer"),
safe_serialization=True,
)
if self.get_te_has_grad():
te0 = self.text_encoder[0] if isinstance(self.text_encoder, list) else self.text_encoder
te0 = unwrap_model(te0)
te0.save_pretrained(
save_directory=os.path.join(output_path, "text_encoder"),
safe_serialization=True,
)
tok0 = self.tokenizer[0] if isinstance(self.tokenizer, list) else self.tokenizer
tok0.save_pretrained(os.path.join(output_path, "tokenizer"))
meta_path = os.path.join(output_path, "aitk_meta.yaml")
with open(meta_path, "w") as f:
yaml.dump(meta, f)
def get_loss_target(self, *args, **kwargs):
noise = kwargs.get("noise")
batch = kwargs.get("batch")
return (noise - batch.latents).detach()
def get_base_model_version(self):
return "zimage"
def get_transformer_block_names(self) -> Optional[List[str]]:
return ["layers"]
def convert_lora_weights_before_save(self, state_dict):
new_sd = {}
for key, value in state_dict.items():
new_key = key.replace("transformer.", "diffusion_model.")
new_sd[new_key] = value
return new_sd
def convert_lora_weights_before_load(self, state_dict):
new_sd = {}
for key, value in state_dict.items():
new_key = key.replace("diffusion_model.", "transformer.")
new_sd[new_key] = value
return new_sd
+114
View File
@@ -0,0 +1,114 @@
import torch
import copy
from comfy.model_patcher import ModelPatcher
class ZImageVectorMerge:
"""
Implements Task Arithmetic Merging:
Result = Base + Strength * (Turbo - Base)
This treats the difference between Turbo and Base as a 'Task Vector'
and injects that vector into the Base model.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_base": ("MODEL",),
"model_turbo": ("MODEL",),
"strength": ("FLOAT", {
"default": 0.3,
"min": -2.0,
"max": 2.0,
"step": 0.01,
"display": "number"
}),
}
}
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("merged_model",)
FUNCTION = "apply_vector_merge"
CATEGORY = "Experimental"
def apply_vector_merge(self, model_base, model_turbo, strength):
print(f"Applying Vector Merge with strength: {strength}")
# Clone the base model structure so we don't corrupt the loaded checkpoint
# We use ModelPatcher.clone() if available, otherwise manual copy
if isinstance(model_base, ModelPatcher):
new_model_patcher = model_base.clone()
base_model_obj = new_model_patcher.model
else:
# Fallback for raw model objects
base_model_obj = model_base
new_model_patcher = copy.deepcopy(model_base)
# Get the underlying state dicts
# Note: We access the diffusion_model directly to avoid VAE/TextEncoder noise
base_sd = base_model_obj.diffusion_model.state_dict()
# specific handling for getting the turbo state dict
if isinstance(model_turbo, ModelPatcher):
turbo_sd = model_turbo.model.diffusion_model.state_dict()
else:
turbo_sd = model_turbo.diffusion_model.state_dict()
# Prepare the new state dict
merged_sd = {}
keys_processed = 0
for key in base_sd.keys():
if key in turbo_sd:
# Get weights
w_base = base_sd[key]
w_turbo = turbo_sd[key]
# Check for shape mismatch (safety)
if w_base.shape != w_turbo.shape:
print(f"Warning: Shape mismatch for key {key}. Skipping. Base: {w_base.shape}, Turbo: {w_turbo.shape}")
merged_sd[key] = w_base
continue
# METHOD 1 MATH:
# Vector = (Turbo - Base)
# New = Base + Strength * Vector
# This simplifies to: New = Base + Strength * Turbo - Strength * Base
# Or: New = (1 - Strength) * Base + Strength * Turbo
# We perform operation on correct device to save VRAM/Time
# Using float32 for precision during merge is recommended
w_base_f = w_base.to(dtype=torch.float32)
w_turbo_f = w_turbo.to(dtype=torch.float32)
# Calculate the vector difference
task_vector = w_turbo_f - w_base_f
# Apply vector
merged_weight = w_base_f + (strength * task_vector)
# Cast back to original dtype (usually float16 or bfloat16)
merged_sd[key] = merged_weight.to(w_base.dtype)
keys_processed += 1
else:
# If key missing in Turbo, keep Base
merged_sd[key] = base_sd[key]
print(f"Merge complete. Processed {keys_processed} keys.")
# Load the new weights into our cloned model
# We use strict=False just in case, but keys should match based on your check
base_model_obj.diffusion_model.load_state_dict(merged_sd, strict=False)
return (new_model_patcher,)
# Node Mapping for ComfyUI
NODE_CLASS_MAPPINGS = {
"ZImageVectorMerge": ZImageVectorMerge
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZImageVectorMerge": "Z-Image Vector Merge (Method 1)"
}