Add files via upload
This commit is contained in:
+69
@@ -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
@@ -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']
|
||||
+153
@@ -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)"
|
||||
}
|
||||
@@ -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)"
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
software_meta = {
|
||||
"name": "AI Toolkit",
|
||||
"version": "1.0.0",
|
||||
"url": "https://github.com/ostris/ai-toolkit"
|
||||
}
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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)]
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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.
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
@@ -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
@@ -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
|
||||
@@ -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)
|
||||
@@ -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')
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -0,0 +1,7 @@
|
||||
lycoris-lora
|
||||
optimum-quanto
|
||||
safetensors
|
||||
diffusers
|
||||
transformers
|
||||
accelerate
|
||||
huggingface-hub
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)"
|
||||
}
|
||||
Reference in New Issue
Block a user