diff --git a/CondMul.py b/CondMul.py new file mode 100644 index 0000000..055884d --- /dev/null +++ b/CondMul.py @@ -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" +} \ No newline at end of file diff --git a/GGUF_RAW.py b/GGUF_RAW.py new file mode 100644 index 0000000..0b0bfd8 --- /dev/null +++ b/GGUF_RAW.py @@ -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'] \ No newline at end of file diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..c7f86d6 --- /dev/null +++ b/LICENSE @@ -0,0 +1,4 @@ +MIT License + +Copyright (c) 2026 TripleHeadedMonkey +... \ No newline at end of file diff --git a/Qwen_Lora.py b/Qwen_Lora.py new file mode 100644 index 0000000..a283bba --- /dev/null +++ b/Qwen_Lora.py @@ -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. + e.g. lora_te.model.layers.0.self_attn.q_proj + + actual weight key: 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)" +} \ No newline at end of file diff --git a/TIES.py b/TIES.py new file mode 100644 index 0000000..397580f --- /dev/null +++ b/TIES.py @@ -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)" +} \ No newline at end of file diff --git a/Universal_LoRA.py b/Universal_LoRA.py new file mode 100644 index 0000000..9c7ce68 --- /dev/null +++ b/Universal_LoRA.py @@ -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" +} \ No newline at end of file diff --git a/ZCondAdv.py b/ZCondAdv.py new file mode 100644 index 0000000..793f66e --- /dev/null +++ b/ZCondAdv.py @@ -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" +} \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..3779e6e --- /dev/null +++ b/__init__.py @@ -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"] \ No newline at end of file diff --git a/dequant.py b/dequant.py new file mode 100644 index 0000000..c0c2231 --- /dev/null +++ b/dequant.py @@ -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, +} diff --git a/diffz.py b/diffz.py new file mode 100644 index 0000000..4b88582 --- /dev/null +++ b/diffz.py @@ -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" +} \ No newline at end of file diff --git a/getimagesizeplus.py b/getimagesizeplus.py new file mode 100644 index 0000000..31ebc0e --- /dev/null +++ b/getimagesizeplus.py @@ -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" +} \ No newline at end of file diff --git a/info.py b/info.py new file mode 100644 index 0000000..f794720 --- /dev/null +++ b/info.py @@ -0,0 +1,5 @@ +software_meta = { + "name": "AI Toolkit", + "version": "1.0.0", + "url": "https://github.com/ostris/ai-toolkit" +} \ No newline at end of file diff --git a/keyword_match_gate.py b/keyword_match_gate.py new file mode 100644 index 0000000..31b7846 --- /dev/null +++ b/keyword_match_gate.py @@ -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", +} diff --git a/loader.py b/loader.py new file mode 100644 index 0000000..413ce3f --- /dev/null +++ b/loader.py @@ -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 diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..b896302 --- /dev/null +++ b/nodes.py @@ -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" +} \ No newline at end of file diff --git a/nodes2.py b/nodes2.py new file mode 100644 index 0000000..ff5aaf0 --- /dev/null +++ b/nodes2.py @@ -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, +} diff --git a/ops.py b/ops.py new file mode 100644 index 0000000..555c754 --- /dev/null +++ b/ops.py @@ -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)] \ No newline at end of file diff --git a/primitive_widget_to_string.py b/primitive_widget_to_string.py new file mode 100644 index 0000000..0104705 --- /dev/null +++ b/primitive_widget_to_string.py @@ -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", +} diff --git a/readme.md b/readme.md new file mode 100644 index 0000000..ccc831a --- /dev/null +++ b/readme.md @@ -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. diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..ad0b050 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,15 @@ +torch +numpy +safetensors>=0.4.0 +gguf +packaging +transformers +diffusers +huggingface_hub +accelerate +optimum +sentencepiece +tqdm +pyyaml +torchaudio +lycoris \ No newline at end of file diff --git a/toolkit/BaseSDTrainProcess.py b/toolkit/BaseSDTrainProcess.py new file mode 100644 index 0000000..a007ab2 --- /dev/null +++ b/toolkit/BaseSDTrainProcess.py @@ -0,0 +1,2532 @@ +import copy +import glob +import inspect +import json +import random +import shutil +from collections import OrderedDict +import os +import re +import traceback +from typing import Union, List, Optional + +import numpy as np +import yaml +from diffusers import T2IAdapter, ControlNetModel +from diffusers.training_utils import compute_density_for_timestep_sampling +from safetensors.torch import save_file, load_file +# from lycoris.config import PRESET +from torch.utils.data import DataLoader +import torch +import torch.backends.cuda +from huggingface_hub import HfApi, Repository, interpreter_login +from huggingface_hub.utils import HfFolder +from toolkit.memory_management import MemoryManager + +from toolkit.basic import value_map +from toolkit.clip_vision_adapter import ClipVisionAdapter +from toolkit.custom_adapter import CustomAdapter +from toolkit.data_loader import get_dataloader_from_datasets, trigger_dataloader_setup_epoch +from toolkit.data_transfer_object.data_loader import FileItemDTO, DataLoaderBatchDTO +from toolkit.ema import ExponentialMovingAverage +from toolkit.embedding import Embedding +from toolkit.image_utils import show_tensors, show_latents, reduce_contrast +from toolkit.ip_adapter import IPAdapter +from toolkit.lora_special import LoRASpecialNetwork +from toolkit.lorm import convert_diffusers_unet_to_lorm, count_parameters, print_lorm_extract_details, \ + lorm_ignore_if_contains, lorm_parameter_threshold, LORM_TARGET_REPLACE_MODULE +from toolkit.lycoris_special import LycorisSpecialNetwork +from toolkit.models.decorator import Decorator +from toolkit.network_mixins import Network +from toolkit.optimizer import get_optimizer +from toolkit.paths import CONFIG_ROOT +from toolkit.progress_bar import ToolkitProgressBar +from toolkit.reference_adapter import ReferenceAdapter +from toolkit.sampler import get_sampler +from toolkit.saving import save_t2i_from_diffusers, load_t2i_model, save_ip_adapter_from_diffusers, \ + load_ip_adapter_model, load_custom_adapter_model + +from toolkit.scheduler import get_lr_scheduler +from toolkit.sd_device_states_presets import get_train_sd_device_state_preset +from toolkit.stable_diffusion_model import StableDiffusion + +from jobs.process import BaseTrainProcess +from toolkit.metadata import get_meta_for_safetensors, load_metadata_from_safetensors, add_base_model_info_to_meta, \ + parse_metadata_from_safetensors +from toolkit.train_tools import get_torch_dtype, LearnableSNRGamma, apply_learnable_snr_gos, apply_snr_weight +import gc + +from tqdm import tqdm + +from toolkit.config_modules import SaveConfig, LoggingConfig, SampleConfig, NetworkConfig, TrainConfig, ModelConfig, \ + GenerateImageConfig, EmbeddingConfig, DatasetConfig, preprocess_dataset_raw_config, AdapterConfig, GuidanceConfig, validate_configs, \ + DecoratorConfig +from toolkit.logging_aitk import create_logger +from diffusers import FluxTransformer2DModel +from toolkit.accelerator import get_accelerator, unwrap_model +from toolkit.print import print_acc +from accelerate import Accelerator +import transformers +import diffusers +import hashlib + +from toolkit.util.blended_blur_noise import get_blended_blur_noise +from toolkit.util.get_model import get_model_class + +def flush(): + torch.cuda.empty_cache() + gc.collect() + + +class BaseSDTrainProcess(BaseTrainProcess): + + def __init__(self, process_id: int, job, config: OrderedDict, custom_pipeline=None): + super().__init__(process_id, job, config) + self.accelerator: Accelerator = get_accelerator() + if self.accelerator.is_local_main_process: + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_error() + else: + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + self.sd: StableDiffusion + self.embedding: Union[Embedding, None] = None + + self.custom_pipeline = custom_pipeline + self.step_num = 0 + self.start_step = 0 + self.epoch_num = 0 + self.last_save_step = 0 + # start at 1 so we can do a sample at the start + self.grad_accumulation_step = 1 + # if true, then we do not do an optimizer step. We are accumulating gradients + self.is_grad_accumulation_step = False + self.device = str(self.accelerator.device) + self.device_torch = self.accelerator.device + network_config = self.get_conf('network', None) + if network_config is not None: + self.network_config = NetworkConfig(**network_config) + else: + self.network_config = None + # Allow overriding optimizer LR separately for TE / UNet without requiring TrainConfig schema changes. + _train_raw = dict(self.get_conf('train', {}) or {}) + # These are consumed by BaseSDTrainProcess when building optimizer param groups. + # They are intentionally popped so TrainConfig doesn't need to know about them. + self.text_encoder_lr_override = _train_raw.pop('text_encoder_lr', None) + self.unet_lr_override = _train_raw.pop('unet_lr', None) + + self.train_config = TrainConfig(**_train_raw) + model_config = self.get_conf('model', {}) + self.modules_being_trained: List[torch.nn.Module] = [] + + # update modelconfig dtype to match train + model_config['dtype'] = self.train_config.dtype + self.model_config = ModelConfig(**model_config) + + self.save_config = SaveConfig(**self.get_conf('save', {})) + self.sample_config = SampleConfig(**self.get_conf('sample', {})) + first_sample_config = self.get_conf('first_sample', None) + if first_sample_config is not None: + self.has_first_sample_requested = True + self.first_sample_config = SampleConfig(**first_sample_config) + else: + self.has_first_sample_requested = False + self.first_sample_config = self.sample_config + self.logging_config = LoggingConfig(**self.get_conf('logging', {})) + self.logger = create_logger(self.logging_config, config) + self.optimizer: torch.optim.Optimizer = None + self.lr_scheduler = None + self.data_loader: Union[DataLoader, None] = None + self.data_loader_reg: Union[DataLoader, None] = None + self.trigger_word = self.get_conf('trigger_word', None) + + self.guidance_config: Union[GuidanceConfig, None] = None + guidance_config_raw = self.get_conf('guidance', None) + if guidance_config_raw is not None: + self.guidance_config = GuidanceConfig(**guidance_config_raw) + + # store is all are cached. Allows us to not load vae if we don't need to + self.is_latents_cached = True + raw_datasets = self.get_conf('datasets', None) + if raw_datasets is not None and len(raw_datasets) > 0: + raw_datasets = preprocess_dataset_raw_config(raw_datasets) + self.datasets = None + self.datasets_reg = None + self.dataset_configs: List[DatasetConfig] = [] + self.params = [] + + # add dataset text embedding cache to their config + if self.train_config.cache_text_embeddings: + for raw_dataset in raw_datasets: + raw_dataset['cache_text_embeddings'] = True + + if raw_datasets is not None and len(raw_datasets) > 0: + for raw_dataset in raw_datasets: + dataset = DatasetConfig(**raw_dataset) + # handle trigger word per dataset + if dataset.trigger_word is None and self.trigger_word is not None: + dataset.trigger_word = self.trigger_word + is_caching = dataset.cache_latents or dataset.cache_latents_to_disk + if not is_caching: + self.is_latents_cached = False + if dataset.is_reg: + if self.datasets_reg is None: + self.datasets_reg = [] + self.datasets_reg.append(dataset) + else: + if self.datasets is None: + self.datasets = [] + self.datasets.append(dataset) + self.dataset_configs.append(dataset) + + self.is_caching_text_embeddings = any( + dataset.cache_text_embeddings for dataset in self.dataset_configs + ) + + self.embed_config = None + embedding_raw = self.get_conf('embedding', None) + if embedding_raw is not None: + self.embed_config = EmbeddingConfig(**embedding_raw) + + self.decorator_config: DecoratorConfig = None + decorator_raw = self.get_conf('decorator', None) + if decorator_raw is not None: + if not self.model_config.is_flux: + raise ValueError("Decorators are only supported for Flux models currently") + self.decorator_config = DecoratorConfig(**decorator_raw) + + # t2i adapter + self.adapter_config = None + adapter_raw = self.get_conf('adapter', None) + if adapter_raw is not None: + self.adapter_config = AdapterConfig(**adapter_raw) + # sdxl adapters end in _xl. Only full_adapter_xl for now + if self.model_config.is_xl and not self.adapter_config.adapter_type.endswith('_xl'): + self.adapter_config.adapter_type += '_xl' + + # to hold network if there is one + self.network: Union[Network, None] = None + self.adapter: Union[T2IAdapter, IPAdapter, ClipVisionAdapter, ReferenceAdapter, CustomAdapter, ControlNetModel, None] = None + self.embedding: Union[Embedding, None] = None + self.decorator: Union[Decorator, None] = None + + is_training_adapter = self.adapter_config is not None and self.adapter_config.train + + self.do_lorm = self.get_conf('do_lorm', False) + self.lorm_extract_mode = self.get_conf('lorm_extract_mode', 'ratio') + self.lorm_extract_mode_param = self.get_conf('lorm_extract_mode_param', 0.25) + # 'ratio', 0.25) + + # get the device state preset based on what we are training + self.train_device_state_preset = get_train_sd_device_state_preset( + device=self.device_torch, + train_unet=self.train_config.train_unet, + train_text_encoder=self.train_config.train_text_encoder, + cached_latents=self.is_latents_cached, + train_lora=self.network_config is not None, + train_adapter=is_training_adapter, + train_embedding=self.embed_config is not None, + train_decorator=self.decorator_config is not None, + train_refiner=self.train_config.train_refiner, + unload_text_encoder=self.train_config.unload_text_encoder or self.is_caching_text_embeddings, + require_grads=False # we ensure them later + ) + + self.get_params_device_state_preset = get_train_sd_device_state_preset( + device=self.device_torch, + train_unet=self.train_config.train_unet, + train_text_encoder=self.train_config.train_text_encoder, + cached_latents=self.is_latents_cached, + train_lora=self.network_config is not None, + train_adapter=is_training_adapter, + train_embedding=self.embed_config is not None, + train_decorator=self.decorator_config is not None, + train_refiner=self.train_config.train_refiner, + unload_text_encoder=self.train_config.unload_text_encoder or self.is_caching_text_embeddings, + require_grads=True # We check for grads when getting params + ) + + # fine_tuning here is for training actual SD network, not LoRA, embeddings, etc. it is (Dreambooth, etc) + self.is_fine_tuning = True + if self.network_config is not None or is_training_adapter or self.embed_config is not None or self.decorator_config is not None: + self.is_fine_tuning = False + + self.named_lora = False + if self.embed_config is not None or is_training_adapter: + self.named_lora = True + self.snr_gos: Union[LearnableSNRGamma, None] = None + self.ema: ExponentialMovingAverage = None + + validate_configs(self.train_config, self.model_config, self.save_config, self.dataset_configs) + + do_profiler = self.get_conf('torch_profiler', False) + self.torch_profiler = None if not do_profiler else torch.profiler.profile( + activities=[ + torch.profiler.ProfilerActivity.CPU, + torch.profiler.ProfilerActivity.CUDA, + ], + ) + + self.current_boundary_index = 0 + self.steps_this_boundary = 0 + self.num_consecutive_oom = 0 + + def post_process_generate_image_config_list(self, generate_image_config_list: List[GenerateImageConfig]): + # override in subclass + return generate_image_config_list + + def sample(self, step=None, is_first=False): + if not self.accelerator.is_main_process: + return + flush() + sample_folder = os.path.join(self.save_root, 'samples') + gen_img_config_list = [] + + sample_config = self.first_sample_config if is_first else self.sample_config + start_seed = sample_config.seed + current_seed = start_seed + + test_image_paths = [] + if self.adapter_config is not None and self.adapter_config.test_img_path is not None: + test_image_path_list = self.adapter_config.test_img_path + # divide up images so they are evenly distributed across prompts + for i in range(len(sample_config.prompts)): + test_image_paths.append(test_image_path_list[i % len(test_image_path_list)]) + + for i in range(len(sample_config.prompts)): + if sample_config.walk_seed: + current_seed = start_seed + i + + step_num = '' + if step is not None: + # zero-pad 9 digits + step_num = f"_{str(step).zfill(9)}" + + filename = f"[time]_{step_num}_[count].{self.sample_config.ext}" + + output_path = os.path.join(sample_folder, filename) + + prompt = sample_config.prompts[i] + + # add embedding if there is one + # note: diffusers will automatically expand the trigger to the number of added tokens + # ie test123 will become test123 test123_1 test123_2 etc. Do not add this yourself here + if self.embedding is not None: + prompt = self.embedding.inject_embedding_to_prompt( + prompt, expand_token=True, add_if_not_present=False + ) + if self.adapter is not None and isinstance(self.adapter, ClipVisionAdapter): + prompt = self.adapter.inject_trigger_into_prompt( + prompt, expand_token=True, add_if_not_present=False + ) + if self.trigger_word is not None: + prompt = self.sd.inject_trigger_into_prompt( + prompt, self.trigger_word, add_if_not_present=False + ) + + extra_args = {} + if self.adapter_config is not None and self.adapter_config.test_img_path is not None: + extra_args['adapter_image_path'] = test_image_paths[i] + + sample_item = sample_config.samples[i] + if sample_item.seed is not None: + current_seed = sample_item.seed + + gen_img_config_list.append(GenerateImageConfig( + prompt=prompt, # it will autoparse the prompt + width=sample_item.width, + height=sample_item.height, + negative_prompt=sample_item.neg, + seed=current_seed, + guidance_scale=sample_item.guidance_scale, + guidance_rescale=sample_config.guidance_rescale, + num_inference_steps=sample_item.sample_steps, + network_multiplier=sample_item.network_multiplier, + output_path=output_path, + output_ext=sample_config.ext, + adapter_conditioning_scale=sample_config.adapter_conditioning_scale, + refiner_start_at=sample_config.refiner_start_at, + extra_values=sample_config.extra_values, + logger=self.logger, + num_frames=sample_item.num_frames, + fps=sample_item.fps, + ctrl_img=sample_item.ctrl_img, + ctrl_idx=sample_item.ctrl_idx, + ctrl_img_1=sample_item.ctrl_img_1, + ctrl_img_2=sample_item.ctrl_img_2, + ctrl_img_3=sample_item.ctrl_img_3, + do_cfg_norm=sample_config.do_cfg_norm, + **extra_args + )) + + # post process + gen_img_config_list = self.post_process_generate_image_config_list(gen_img_config_list) + + # if we have an ema, set it to validation mode + if self.ema is not None: + self.ema.eval() + + # let adapter know we are sampling + if self.adapter is not None and isinstance(self.adapter, CustomAdapter): + self.adapter.is_sampling = True + + # send to be generated + self.sd.generate_images(gen_img_config_list, sampler=sample_config.sampler) + + + if self.adapter is not None and isinstance(self.adapter, CustomAdapter): + self.adapter.is_sampling = False + + if self.ema is not None: + self.ema.train() + + def update_training_metadata(self): + o_dict = OrderedDict({ + "training_info": self.get_training_info() + }) + o_dict['ss_base_model_version'] = self.sd.get_base_model_version() + + # o_dict = add_base_model_info_to_meta( + # o_dict, + # is_v2=self.model_config.is_v2, + # is_xl=self.model_config.is_xl, + # ) + o_dict['ss_output_name'] = self.job.name + + if self.trigger_word is not None: + # just so auto1111 will pick it up + o_dict['ss_tag_frequency'] = { + f"1_{self.trigger_word}": { + f"{self.trigger_word}": 1 + } + } + + self.add_meta(o_dict) + + def get_training_info(self): + info = OrderedDict({ + 'step': self.step_num, + 'epoch': self.epoch_num, + }) + return info + + def clean_up_saves(self): + if not self.accelerator.is_main_process: + return + # remove old saves + # get latest saved step + latest_item = None + if os.path.exists(self.save_root): + # pattern is {job_name}_{zero_filled_step} for both files and directories + pattern = f"{self.job.name}_*" + items = glob.glob(os.path.join(self.save_root, pattern)) + # Separate files and directories + safetensors_files = [f for f in items if f.endswith('.safetensors')] + pt_files = [f for f in items if f.endswith('.pt')] + directories = [d for d in items if os.path.isdir(d) and not d.endswith('.safetensors')] + embed_files = [] + # do embedding files + if self.embed_config is not None: + embed_pattern = f"{self.embed_config.trigger}_*" + embed_items = glob.glob(os.path.join(self.save_root, embed_pattern)) + # will end in safetensors or pt + embed_files = [f for f in embed_items if f.endswith('.safetensors') or f.endswith('.pt')] + + # check for critic files + critic_pattern = f"CRITIC_{self.job.name}_*" + critic_items = glob.glob(os.path.join(self.save_root, critic_pattern)) + + # Sort the lists by creation time if they are not empty + if safetensors_files: + safetensors_files.sort(key=os.path.getctime) + if pt_files: + pt_files.sort(key=os.path.getctime) + if directories: + directories.sort(key=os.path.getctime) + if embed_files: + embed_files.sort(key=os.path.getctime) + if critic_items: + critic_items.sort(key=os.path.getctime) + + # Combine and sort the lists + combined_items = safetensors_files + directories + pt_files + combined_items.sort(key=os.path.getctime) + + num_saves_to_keep = self.save_config.max_step_saves_to_keep + + if hasattr(self.sd, 'max_step_saves_to_keep_multiplier'): + num_saves_to_keep *= self.sd.max_step_saves_to_keep_multiplier + + # Use slicing with a check to avoid 'NoneType' error + safetensors_to_remove = safetensors_files[ + :-num_saves_to_keep] if safetensors_files else [] + pt_files_to_remove = pt_files[:-num_saves_to_keep] if pt_files else [] + directories_to_remove = directories[:-num_saves_to_keep] if directories else [] + embeddings_to_remove = embed_files[:-num_saves_to_keep] if embed_files else [] + critic_to_remove = critic_items[:-num_saves_to_keep] if critic_items else [] + + items_to_remove = safetensors_to_remove + pt_files_to_remove + directories_to_remove + embeddings_to_remove + critic_to_remove + + # remove all but the latest max_step_saves_to_keep + # items_to_remove = combined_items[:-num_saves_to_keep] + + # remove duplicates + items_to_remove = list(dict.fromkeys(items_to_remove)) + + for item in items_to_remove: + print_acc(f"Removing old save: {item}") + if os.path.isdir(item): + shutil.rmtree(item) + else: + os.remove(item) + # see if a yaml file with same name exists + yaml_file = os.path.splitext(item)[0] + ".yaml" + if os.path.exists(yaml_file): + os.remove(yaml_file) + if combined_items: + latest_item = combined_items[-1] + return latest_item + + def post_save_hook(self, save_path): + # override in subclass + pass + + def done_hook(self): + pass + + def end_step_hook(self): + pass + + def save(self, step=None): + if not self.accelerator.is_main_process: + return + flush() + if self.ema is not None: + # always save params as ema + self.ema.eval() + + if not os.path.exists(self.save_root): + os.makedirs(self.save_root, exist_ok=True) + + step_num = '' + if step is not None: + self.last_save_step = step + # zeropad 9 digits + step_num = f"_{str(step).zfill(9)}" + + self.update_training_metadata() + filename = f'{self.job.name}{step_num}.safetensors' + file_path = os.path.join(self.save_root, filename) + + save_meta = copy.deepcopy(self.meta) + # get extra meta + if self.adapter is not None and isinstance(self.adapter, CustomAdapter): + additional_save_meta = self.adapter.get_additional_save_metadata() + if additional_save_meta is not None: + for key, value in additional_save_meta.items(): + save_meta[key] = value + + # prepare meta + save_meta = get_meta_for_safetensors(save_meta, self.job.name) + if not self.is_fine_tuning: + if self.network is not None: + lora_name = self.job.name + if self.named_lora: + # add _lora to name + lora_name += '_LoRA' + + filename = f'{lora_name}{step_num}.safetensors' + file_path = os.path.join(self.save_root, filename) + prev_multiplier = self.network.multiplier + self.network.multiplier = 1.0 + + # if we are doing embedding training as well, add that + embedding_dict = self.embedding.state_dict() if self.embedding else None + self.network.save_weights( + file_path, + dtype=get_torch_dtype(self.save_config.dtype), + metadata=save_meta, + extra_state_dict=embedding_dict + ) + self.network.multiplier = prev_multiplier + # if we have an embedding as well, pair it with the network + + # even if added to lora, still save the trigger version + if self.embedding is not None: + emb_filename = f'{self.embed_config.trigger}{step_num}.safetensors' + emb_file_path = os.path.join(self.save_root, emb_filename) + # for combo, above will get it + # set current step + self.embedding.step = self.step_num + # change filename to pt if that is set + if self.embed_config.save_format == "pt": + # replace extension + emb_file_path = os.path.splitext(emb_file_path)[0] + ".pt" + self.embedding.save(emb_file_path) + + if self.decorator is not None: + dec_filename = f'{self.job.name}{step_num}.safetensors' + dec_file_path = os.path.join(self.save_root, dec_filename) + decorator_state_dict = self.decorator.state_dict() + for key, value in decorator_state_dict.items(): + if isinstance(value, torch.Tensor): + decorator_state_dict[key] = value.clone().to('cpu', dtype=get_torch_dtype(self.save_config.dtype)) + save_file( + decorator_state_dict, + dec_file_path, + metadata=save_meta, + ) + + if self.adapter is not None and self.adapter_config.train: + adapter_name = self.job.name + if self.network_config is not None or self.embedding is not None: + # add _lora to name + if self.adapter_config.type == 't2i': + adapter_name += '_t2i' + elif self.adapter_config.type == 'control_net': + adapter_name += '_cn' + elif self.adapter_config.type == 'clip': + adapter_name += '_clip' + elif self.adapter_config.type.startswith('ip'): + adapter_name += '_ip' + else: + adapter_name += '_adapter' + + filename = f'{adapter_name}{step_num}.safetensors' + file_path = os.path.join(self.save_root, filename) + # save adapter + state_dict = self.adapter.state_dict() + if self.adapter_config.type == 't2i': + save_t2i_from_diffusers( + state_dict, + output_file=file_path, + meta=save_meta, + dtype=get_torch_dtype(self.save_config.dtype) + ) + elif self.adapter_config.type == 'control_net': + # save in diffusers format + name_or_path = file_path.replace('.safetensors', '') + # move it to the new dtype and cpu + orig_device = self.adapter.device + orig_dtype = self.adapter.dtype + self.adapter = self.adapter.to(torch.device('cpu'), dtype=get_torch_dtype(self.save_config.dtype)) + self.adapter.save_pretrained( + name_or_path, + dtype=get_torch_dtype(self.save_config.dtype), + safe_serialization=True + ) + meta_path = os.path.join(name_or_path, 'aitk_meta.yaml') + with open(meta_path, 'w') as f: + yaml.dump(self.meta, f) + # move it back + self.adapter = self.adapter.to(orig_device, dtype=orig_dtype) + else: + direct_save = False + if self.adapter_config.train_only_image_encoder: + direct_save = True + elif isinstance(self.adapter, CustomAdapter): + direct_save = self.adapter.do_direct_save + save_ip_adapter_from_diffusers( + state_dict, + output_file=file_path, + meta=save_meta, + dtype=get_torch_dtype(self.save_config.dtype), + direct_save=direct_save + ) + else: + if self.save_config.save_format == "diffusers": + # saving as a folder path + file_path = file_path.replace('.safetensors', '') + # convert it back to normal object + save_meta = parse_metadata_from_safetensors(save_meta) + + if self.sd.refiner_unet and self.train_config.train_refiner: + # save refiner + refiner_name = self.job.name + '_refiner' + filename = f'{refiner_name}{step_num}.safetensors' + file_path = os.path.join(self.save_root, filename) + self.sd.save_refiner( + file_path, + save_meta, + get_torch_dtype(self.save_config.dtype) + ) + if self.train_config.train_unet or self.train_config.train_text_encoder: + self.sd.save( + file_path, + save_meta, + get_torch_dtype(self.save_config.dtype) + ) + + # save learnable params as json if we have thim + if self.snr_gos: + json_data = { + 'offset_1': self.snr_gos.offset_1.item(), + 'offset_2': self.snr_gos.offset_2.item(), + 'scale': self.snr_gos.scale.item(), + 'gamma': self.snr_gos.gamma.item(), + } + path_to_save = file_path = os.path.join(self.save_root, 'learnable_snr.json') + with open(path_to_save, 'w') as f: + json.dump(json_data, f, indent=4) + + print_acc(f"Saved checkpoint to {file_path}") + + # save optimizer + if self.optimizer is not None: + try: + filename = f'optimizer.pt' + file_path = os.path.join(self.save_root, filename) + try: + state_dict = unwrap_model(self.optimizer).state_dict() + except Exception as e: + state_dict = self.optimizer.state_dict() + torch.save(state_dict, file_path) + print_acc(f"Saved optimizer to {file_path}") + except Exception as e: + print_acc(e) + print_acc("Could not save optimizer") + + self.clean_up_saves() + self.post_save_hook(file_path) + + if self.ema is not None: + self.ema.train() + flush() + + # Called before the model is loaded + def hook_before_model_load(self): + # override in subclass + pass + + def hook_after_model_load(self): + # override in subclass + pass + + def hook_add_extra_train_params(self, params): + # override in subclass + return params + + def hook_before_train_loop(self): + if self.accelerator.is_main_process: + self.logger.start() + self.prepare_accelerator() + + def sample_step_hook(self, img_num, total_imgs): + pass + + def prepare_accelerator(self): + # set some config + self.accelerator.even_batches=False + + # # prepare all the models stuff for accelerator (hopefully we dont miss any) + self.sd.vae = self.accelerator.prepare(self.sd.vae) + if self.sd.unet is not None: + self.sd.unet = self.accelerator.prepare(self.sd.unet) + # todo always tdo it? + self.modules_being_trained.append(self.sd.unet) + if self.sd.text_encoder is not None and self.train_config.train_text_encoder: + if isinstance(self.sd.text_encoder, list): + self.sd.text_encoder = [self.accelerator.prepare(model) for model in self.sd.text_encoder] + self.modules_being_trained.extend(self.sd.text_encoder) + else: + self.sd.text_encoder = self.accelerator.prepare(self.sd.text_encoder) + self.modules_being_trained.append(self.sd.text_encoder) + if self.sd.refiner_unet is not None and self.train_config.train_refiner: + self.sd.refiner_unet = self.accelerator.prepare(self.sd.refiner_unet) + self.modules_being_trained.append(self.sd.refiner_unet) + # todo, do we need to do the network or will "unet" get it? + if self.sd.network is not None: + self.sd.network = self.accelerator.prepare(self.sd.network) + self.modules_being_trained.append(self.sd.network) + if self.adapter is not None and self.adapter_config.train: + # todo adapters may not be a module. need to check + self.adapter = self.accelerator.prepare(self.adapter) + self.modules_being_trained.append(self.adapter) + + # prepare other things + self.optimizer = self.accelerator.prepare(self.optimizer) + if self.lr_scheduler is not None: + self.lr_scheduler = self.accelerator.prepare(self.lr_scheduler) + # self.data_loader = self.accelerator.prepare(self.data_loader) + # if self.data_loader_reg is not None: + # self.data_loader_reg = self.accelerator.prepare(self.data_loader_reg) + + + def ensure_params_requires_grad(self, force=False): + if self.train_config.do_paramiter_swapping and not force: + # the optimizer will handle this if we are not forcing + return + for group in self.params: + for param in group['params']: + if isinstance(param, torch.nn.Parameter): # Ensure it's a proper parameter + param.requires_grad_(True) + + def setup_ema(self): + if self.train_config.ema_config.use_ema: + # our params are in groups. We need them as a single iterable + params = [] + for group in self.optimizer.param_groups: + for param in group['params']: + params.append(param) + self.ema = ExponentialMovingAverage( + params, + decay=self.train_config.ema_config.ema_decay, + use_feedback=self.train_config.ema_config.use_feedback, + param_multiplier=self.train_config.ema_config.param_multiplier, + ) + + def before_dataset_load(self): + pass + + def get_params(self): + # you can extend this in subclass to get params + # otherwise params will be gathered through normal means + return None + + def hook_train_loop(self, batch): + # return loss + return 0.0 + + def hook_after_sd_init_before_load(self): + pass + + def get_latest_save_path(self, name=None, post=''): + if name == None: + name = self.job.name + # get latest saved step + latest_path = None + if os.path.exists(self.save_root): + # Define patterns for both files and directories + patterns = [ + f"{name}*{post}.safetensors", + f"{name}*{post}.pt", + f"{name}*{post}" + ] + # Search for both files and directories + paths = [] + for pattern in patterns: + paths.extend(glob.glob(os.path.join(self.save_root, pattern))) + + # Filter out non-existent paths and sort by creation time + if paths: + paths = [p for p in paths if os.path.exists(p)] + # remove false positives + if '_LoRA' not in name: + paths = [p for p in paths if '_LoRA' not in p] + if '_refiner' not in name: + paths = [p for p in paths if '_refiner' not in p] + if '_t2i' not in name: + paths = [p for p in paths if '_t2i' not in p] + if '_cn' not in name: + paths = [p for p in paths if '_cn' not in p] + + if len(paths) > 0: + latest_path = max(paths, key=os.path.getctime) + + return latest_path + + def load_training_state_from_metadata(self, path): + if not self.accelerator.is_main_process: + return + meta = None + # if path is folder, then it is diffusers + if os.path.isdir(path): + meta_path = os.path.join(path, 'aitk_meta.yaml') + # load it + if os.path.exists(meta_path): + with open(meta_path, 'r') as f: + meta = yaml.load(f, Loader=yaml.FullLoader) + else: + meta = load_metadata_from_safetensors(path) + # if 'training_info' in Orderdict keys + if meta is not None and 'training_info' in meta and 'step' in meta['training_info'] and self.train_config.start_step is None: + self.step_num = meta['training_info']['step'] + if 'epoch' in meta['training_info']: + self.epoch_num = meta['training_info']['epoch'] + self.start_step = self.step_num + print_acc(f"Found step {self.step_num} in metadata, starting from there") + + def load_weights(self, path): + if self.network is not None: + extra_weights = self.network.load_weights(path) + self.load_training_state_from_metadata(path) + return extra_weights + else: + print_acc("load_weights not implemented for non-network models") + return None + + def apply_snr(self, seperated_loss, timesteps): + if self.train_config.learnable_snr_gos: + # add snr_gamma + seperated_loss = apply_learnable_snr_gos(seperated_loss, timesteps, self.snr_gos) + elif self.train_config.snr_gamma is not None and self.train_config.snr_gamma > 0.000001: + # add snr_gamma + seperated_loss = apply_snr_weight(seperated_loss, timesteps, self.sd.noise_scheduler, self.train_config.snr_gamma, fixed=True) + elif self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001: + # add min_snr_gamma + seperated_loss = apply_snr_weight(seperated_loss, timesteps, self.sd.noise_scheduler, self.train_config.min_snr_gamma) + + return seperated_loss + + def load_lorm(self): + latest_save_path = self.get_latest_save_path() + if latest_save_path is not None: + # hacky way to reload weights for now + # todo, do this + state_dict = load_file(latest_save_path, device=self.device) + self.sd.unet.load_state_dict(state_dict) + + meta = load_metadata_from_safetensors(latest_save_path) + # if 'training_info' in Orderdict keys + if 'training_info' in meta and 'step' in meta['training_info']: + self.step_num = meta['training_info']['step'] + if 'epoch' in meta['training_info']: + self.epoch_num = meta['training_info']['epoch'] + self.start_step = self.step_num + print_acc(f"Found step {self.step_num} in metadata, starting from there") + + # def get_sigmas(self, timesteps, n_dim=4, dtype=torch.float32): + # self.sd.noise_scheduler.set_timesteps(1000, device=self.device_torch) + # sigmas = self.sd.noise_scheduler.sigmas.to(device=self.device_torch, dtype=dtype) + # schedule_timesteps = self.sd.noise_scheduler.timesteps.to(self.device_torch, ) + # timesteps = timesteps.to(self.device_torch, ) + # + # # step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + # step_indices = [t for t in timesteps] + # + # sigma = sigmas[step_indices].flatten() + # while len(sigma.shape) < n_dim: + # sigma = sigma.unsqueeze(-1) + # return sigma + + def load_additional_training_modules(self, params): + # override in subclass + return params + + def get_sigmas(self, timesteps, n_dim=4, dtype=torch.float32): + sigmas = self.sd.noise_scheduler.sigmas.to(device=self.device, dtype=dtype) + schedule_timesteps = self.sd.noise_scheduler.timesteps.to(self.device) + timesteps = timesteps.to(self.device) + + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + def get_optimal_noise(self, latents, dtype=torch.float32): + batch_num = latents.shape[0] + chunks = torch.chunk(latents, batch_num, dim=0) + noise_chunks = [] + for chunk in chunks: + noise_samples = [torch.randn_like(chunk, device=chunk.device, dtype=dtype) for _ in range(self.train_config.optimal_noise_pairing_samples)] + # find the one most similar to the chunk + lowest_loss = 999999999999 + best_noise = None + for noise in noise_samples: + loss = torch.nn.functional.mse_loss(chunk, noise) + if loss < lowest_loss: + lowest_loss = loss + best_noise = noise + noise_chunks.append(best_noise) + noise = torch.cat(noise_chunks, dim=0) + return noise + + def get_consistent_noise(self, latents, batch: 'DataLoaderBatchDTO', dtype=torch.float32): + batch_num = latents.shape[0] + chunks = torch.chunk(latents, batch_num, dim=0) + noise_chunks = [] + for idx, chunk in enumerate(chunks): + # get seed from path + file_item = batch.file_items[idx] + img_path = file_item.path + # add augmentors + if file_item.flip_x: + img_path += '_fx' + if file_item.flip_y: + img_path += '_fy' + seed = int(hashlib.md5(img_path.encode()).hexdigest(), 16) & 0xffffffff + generator = torch.Generator("cpu").manual_seed(seed) + noise_chunk = torch.randn(chunk.shape, generator=generator).to(chunk.device, dtype=dtype) + noise_chunks.append(noise_chunk) + noise = torch.cat(noise_chunks, dim=0).to(dtype=dtype) + return noise + + + def get_noise( + self, + latents, + batch_size, + dtype=torch.float32, + batch: 'DataLoaderBatchDTO' = None, + timestep=None, + ): + if self.train_config.optimal_noise_pairing_samples > 1: + noise = self.get_optimal_noise(latents, dtype=dtype) + elif self.train_config.force_consistent_noise: + if batch is None: + raise ValueError("Batch must be provided for consistent noise") + noise = self.get_consistent_noise(latents, batch, dtype=dtype) + else: + if hasattr(self.sd, 'get_latent_noise_from_latents'): + noise = self.sd.get_latent_noise_from_latents( + latents, + noise_offset=self.train_config.noise_offset + ).to(self.device_torch, dtype=dtype) + else: + # get noise + noise = self.sd.get_latent_noise( + height=latents.shape[2], + width=latents.shape[3], + num_channels=latents.shape[1], + batch_size=batch_size, + noise_offset=self.train_config.noise_offset, + ).to(self.device_torch, dtype=dtype) + + if self.train_config.blended_blur_noise: + noise = get_blended_blur_noise( + latents, noise, timestep + ) + + return noise + + def process_general_training_batch(self, batch: 'DataLoaderBatchDTO'): + with torch.no_grad(): + with self.timer('prepare_prompt'): + prompts = batch.get_caption_list() + is_reg_list = batch.get_is_reg_list() + + is_any_reg = any([is_reg for is_reg in is_reg_list]) + + do_double = self.train_config.short_and_long_captions and not is_any_reg + + if self.train_config.short_and_long_captions and do_double: + # dont do this with regs. No point + + # double batch and add short captions to the end + prompts = prompts + batch.get_caption_short_list() + is_reg_list = is_reg_list + is_reg_list + if self.model_config.refiner_name_or_path is not None and self.train_config.train_unet: + prompts = prompts + prompts + is_reg_list = is_reg_list + is_reg_list + + conditioned_prompts = [] + + for prompt, is_reg in zip(prompts, is_reg_list): + + # make sure the embedding is in the prompts + if self.embedding is not None: + prompt = self.embedding.inject_embedding_to_prompt( + prompt, + expand_token=True, + add_if_not_present=not is_reg, + ) + + if self.adapter and isinstance(self.adapter, ClipVisionAdapter): + prompt = self.adapter.inject_trigger_into_prompt( + prompt, + expand_token=True, + add_if_not_present=not is_reg, + ) + + # make sure trigger is in the prompts if not a regularization run + if self.trigger_word is not None: + prompt = self.sd.inject_trigger_into_prompt( + prompt, + trigger=self.trigger_word, + add_if_not_present=not is_reg, + ) + + if not is_reg and self.train_config.prompt_saturation_chance > 0.0: + # do random prompt saturation by expanding the prompt to hit at least 77 tokens + if random.random() < self.train_config.prompt_saturation_chance: + est_num_tokens = len(prompt.split(' ')) + if est_num_tokens < 77: + num_repeats = int(77 / est_num_tokens) + 1 + prompt = ', '.join([prompt] * num_repeats) + + + conditioned_prompts.append(prompt) + + with self.timer('prepare_latents'): + dtype = get_torch_dtype(self.train_config.dtype) + imgs = None + is_reg = any(batch.get_is_reg_list()) + if batch.tensor is not None: + imgs = batch.tensor + imgs = imgs.to(self.device_torch, dtype=dtype) + # dont adjust for regs. + if self.train_config.img_multiplier is not None and not is_reg: + # do it ad contrast + imgs = reduce_contrast(imgs, self.train_config.img_multiplier) + if batch.latents is not None: + latents = batch.latents.to(self.device_torch, dtype=dtype) + batch.latents = latents + else: + # normalize to + if self.train_config.standardize_images: + if self.sd.is_xl or self.sd.is_vega or self.sd.is_ssd: + target_mean_list = [0.0002, -0.1034, -0.1879] + target_std_list = [0.5436, 0.5116, 0.5033] + else: + target_mean_list = [-0.0739, -0.1597, -0.2380] + target_std_list = [0.5623, 0.5295, 0.5347] + # Mean: tensor([-0.0739, -0.1597, -0.2380]) + # Standard Deviation: tensor([0.5623, 0.5295, 0.5347]) + imgs_channel_mean = imgs.mean(dim=(2, 3), keepdim=True) + imgs_channel_std = imgs.std(dim=(2, 3), keepdim=True) + imgs = (imgs - imgs_channel_mean) / imgs_channel_std + target_mean = torch.tensor(target_mean_list, device=self.device_torch, dtype=dtype) + target_std = torch.tensor(target_std_list, device=self.device_torch, dtype=dtype) + # expand them to match dim + target_mean = target_mean.unsqueeze(0).unsqueeze(2).unsqueeze(3) + target_std = target_std.unsqueeze(0).unsqueeze(2).unsqueeze(3) + + imgs = imgs * target_std + target_mean + batch.tensor = imgs + + # show_tensors(imgs, 'imgs') + + latents = self.sd.encode_images(imgs) + batch.latents = latents + + if self.train_config.standardize_latents: + if self.sd.is_xl or self.sd.is_vega or self.sd.is_ssd: + target_mean_list = [-0.1075, 0.0231, -0.0135, 0.2164] + target_std_list = [0.8979, 0.7505, 0.9150, 0.7451] + else: + target_mean_list = [0.2949, -0.3188, 0.0807, 0.1929] + target_std_list = [0.8560, 0.9629, 0.7778, 0.6719] + + latents_channel_mean = latents.mean(dim=(2, 3), keepdim=True) + latents_channel_std = latents.std(dim=(2, 3), keepdim=True) + latents = (latents - latents_channel_mean) / latents_channel_std + target_mean = torch.tensor(target_mean_list, device=self.device_torch, dtype=dtype) + target_std = torch.tensor(target_std_list, device=self.device_torch, dtype=dtype) + # expand them to match dim + target_mean = target_mean.unsqueeze(0).unsqueeze(2).unsqueeze(3) + target_std = target_std.unsqueeze(0).unsqueeze(2).unsqueeze(3) + + latents = latents * target_std + target_mean + batch.latents = latents + + # show_latents(latents, self.sd.vae, 'latents') + + + if batch.unconditional_tensor is not None and batch.unconditional_latents is None: + unconditional_imgs = batch.unconditional_tensor + unconditional_imgs = unconditional_imgs.to(self.device_torch, dtype=dtype) + unconditional_latents = self.sd.encode_images(unconditional_imgs) + batch.unconditional_latents = unconditional_latents * self.train_config.latent_multiplier + + unaugmented_latents = None + if self.train_config.loss_target == 'differential_noise': + # we determine noise from the differential of the latents + unaugmented_latents = self.sd.encode_images(batch.unaugmented_tensor) + + with self.timer('prepare_scheduler'): + + batch_size = len(batch.file_items) + min_noise_steps = self.train_config.min_denoising_steps + max_noise_steps = self.train_config.max_denoising_steps + if self.model_config.refiner_name_or_path is not None: + # if we are not training the unet, then we are only doing refiner and do not need to double up + if self.train_config.train_unet: + max_noise_steps = round(self.train_config.max_denoising_steps * self.model_config.refiner_start_at) + do_double = True + else: + min_noise_steps = round(self.train_config.max_denoising_steps * self.model_config.refiner_start_at) + do_double = False + + num_train_timesteps = self.train_config.num_train_timesteps + + if self.train_config.noise_scheduler in ['custom_lcm']: + # we store this value on our custom one + self.sd.noise_scheduler.set_timesteps( + self.sd.noise_scheduler.train_timesteps, device=self.device_torch + ) + elif self.train_config.noise_scheduler in ['lcm']: + self.sd.noise_scheduler.set_timesteps( + num_train_timesteps, device=self.device_torch, original_inference_steps=num_train_timesteps + ) + elif self.train_config.noise_scheduler == 'flowmatch': + linear_timesteps = any([ + self.train_config.linear_timesteps, + self.train_config.linear_timesteps2, + self.train_config.timestep_type == 'linear', + self.train_config.timestep_type == 'one_step', + ]) + + timestep_type = 'linear' if linear_timesteps else None + if timestep_type is None: + timestep_type = self.train_config.timestep_type + + if self.train_config.timestep_type == 'next_sample': + # simulate a sample + num_train_timesteps = self.train_config.next_sample_timesteps + timestep_type = 'shift' + + patch_size = 1 + if self.sd.is_flux or 'flex' in self.sd.arch: + # flux is a patch size of 1, but latents are divided by 2, so we need to double it + patch_size = 2 + elif hasattr(self.sd.unet.config, 'patch_size'): + patch_size = self.sd.unet.config.patch_size + + self.sd.noise_scheduler.set_train_timesteps( + num_train_timesteps, + device=self.device_torch, + timestep_type=timestep_type, + latents=latents, + patch_size=patch_size, + ) + else: + self.sd.noise_scheduler.set_timesteps( + num_train_timesteps, device=self.device_torch + ) + if self.sd.is_multistage: + with self.timer('adjust_multistage_timesteps'): + # get our current sample range + boundaries = [1] + self.sd.multistage_boundaries + boundary_max, boundary_min = boundaries[self.current_boundary_index], boundaries[self.current_boundary_index + 1] + asc_timesteps = torch.flip(self.sd.noise_scheduler.timesteps, dims=[0]) + lo = len(asc_timesteps) - torch.searchsorted(asc_timesteps, torch.tensor(boundary_max * 1000, device=asc_timesteps.device), right=False) + hi = len(asc_timesteps) - torch.searchsorted(asc_timesteps, torch.tensor(boundary_min * 1000, device=asc_timesteps.device), right=True) + first_idx = (lo - 1).item() if hi > lo else 0 + last_idx = (hi - 1).item() if hi > lo else 999 + min_noise_steps = first_idx + max_noise_steps = last_idx + + # clip min max indicies + min_noise_steps = max(min_noise_steps, 0) + max_noise_steps = min(max_noise_steps, num_train_timesteps - 1) + + + with self.timer('prepare_timesteps_indices'): + + content_or_style = self.train_config.content_or_style + if is_reg: + content_or_style = self.train_config.content_or_style_reg + + # if self.train_config.timestep_sampling == 'style' or self.train_config.timestep_sampling == 'content': + if self.train_config.timestep_type == 'next_sample': + timestep_indices = torch.randint( + 0, + num_train_timesteps - 2, # -1 for 0 idx, -1 so we can step + (batch_size,), + device=self.device_torch + ) + timestep_indices = timestep_indices.long() + elif self.train_config.timestep_type == 'one_step': + timestep_indices = torch.zeros((batch_size,), device=self.device_torch, dtype=torch.long) + elif content_or_style in ['style', 'content']: + # this is from diffusers training code + # Cubic sampling for favoring later or earlier timesteps + # For more details about why cubic sampling is used for content / structure, + # refer to section 3.4 of https://arxiv.org/abs/2302.08453 + + # for content / structure, it is best to favor earlier timesteps + # for style, it is best to favor later timesteps + + orig_timesteps = torch.rand((batch_size,), device=latents.device) + + if content_or_style == 'content': + timestep_indices = orig_timesteps ** 3 * self.train_config.num_train_timesteps + elif content_or_style == 'style': + timestep_indices = (1 - orig_timesteps ** 3) * self.train_config.num_train_timesteps + + timestep_indices = value_map( + timestep_indices, + 0, + self.train_config.num_train_timesteps - 1, + min_noise_steps, + max_noise_steps + ) + timestep_indices = timestep_indices.long().clamp( + min_noise_steps, + max_noise_steps + ) + + elif content_or_style == 'balanced': + if min_noise_steps == max_noise_steps: + timestep_indices = torch.ones((batch_size,), device=self.device_torch) * min_noise_steps + else: + # todo, some schedulers use indices, otheres use timesteps. Not sure what to do here + min_idx = min_noise_steps + 1 + max_idx = max_noise_steps - 1 + if self.train_config.noise_scheduler == 'flowmatch': + # flowmatch uses indices, so we need to use indices + min_idx = min_noise_steps + max_idx = max_noise_steps + timestep_indices = torch.randint( + min_idx, + max_idx, + (batch_size,), + device=self.device_torch + ) + timestep_indices = timestep_indices.long() + else: + raise ValueError(f"Unknown content_or_style {content_or_style}") + with self.timer('convert_timestep_indices_to_timesteps'): + # convert the timestep_indices to a timestep + timesteps = self.sd.noise_scheduler.timesteps[timestep_indices.long()] + + with self.timer('prepare_noise'): + # get noise + noise = self.get_noise(latents, batch_size, dtype=dtype, batch=batch, timestep=timesteps) + + # add dynamic noise offset. Dynamic noise is offsetting the noise to the same channelwise mean as the latents + # this will negate any noise offsets + if self.train_config.dynamic_noise_offset and not is_reg: + latents_channel_mean = latents.mean(dim=(2, 3), keepdim=True) / 2 + # subtract channel mean to that we compensate for the mean of the latents on the noise offset per channel + noise = noise + latents_channel_mean + + if self.train_config.loss_target == 'differential_noise': + differential = latents - unaugmented_latents + # add noise to differential + # noise = noise + differential + noise = noise + (differential * 0.5) + # noise = value_map(differential, 0, torch.abs(differential).max(), 0, torch.abs(noise).max()) + latents = unaugmented_latents + + noise_multiplier = self.train_config.noise_multiplier + + s = (noise.shape[0], noise.shape[1], 1, 1) + if len(noise.shape) == 5: + # if we have a 5d tensor, then we need to do it on a per batch item, per channel basis, per frame + s = (noise.shape[0], noise.shape[1], noise.shape[2], 1, 1) + + if self.train_config.random_noise_multiplier > 0.0: + + # do it on a per batch item, per channel basis + noise_multiplier = 1 + torch.randn( + s, + device=noise.device, + dtype=noise.dtype + ) * self.train_config.random_noise_multiplier + + with self.timer('make_noisy_latents'): + + noise = noise * noise_multiplier + + if self.train_config.random_noise_shift > 0.0: + # get random noise -1 to 1 + noise_shift = torch.randn( + s, + device=noise.device, + dtype=noise.dtype + ) * self.train_config.random_noise_shift + # add to noise + noise += noise_shift + + latent_multiplier = self.train_config.latent_multiplier + + # handle adaptive scaling mased on std + if self.train_config.adaptive_scaling_factor: + std = latents.std(dim=(2, 3), keepdim=True) + normalizer = 1 / (std + 1e-6) + latent_multiplier = normalizer + + latents = latents * latent_multiplier + batch.latents = latents + + # normalize latents to a mean of 0 and an std of 1 + # mean_zero_latents = latents - latents.mean() + # latents = mean_zero_latents / mean_zero_latents.std() + + if batch.unconditional_latents is not None: + batch.unconditional_latents = batch.unconditional_latents * self.train_config.latent_multiplier + + + noisy_latents = self.sd.add_noise(latents, noise, timesteps) + + # determine scaled noise + # todo do we need to scale this or does it always predict full intensity + # noise = noisy_latents - latents + + # https://github.com/huggingface/diffusers/blob/324d18fba23f6c9d7475b0ff7c777685f7128d40/examples/t2i_adapter/train_t2i_adapter_sdxl.py#L1170C17-L1171C77 + if self.train_config.loss_target == 'source' or self.train_config.loss_target == 'unaugmented': + sigmas = self.get_sigmas(timesteps, len(noisy_latents.shape), noisy_latents.dtype) + # add it to the batch + batch.sigmas = sigmas + # todo is this for sdxl? find out where this came from originally + # noisy_latents = noisy_latents / ((sigmas ** 2 + 1) ** 0.5) + + def double_up_tensor(tensor: torch.Tensor): + if tensor is None: + return None + return torch.cat([tensor, tensor], dim=0) + + if do_double: + if self.model_config.refiner_name_or_path: + # apply refiner double up + refiner_timesteps = torch.randint( + max_noise_steps, + self.train_config.max_denoising_steps, + (batch_size,), + device=self.device_torch + ) + refiner_timesteps = refiner_timesteps.long() + # add our new timesteps on to end + timesteps = torch.cat([timesteps, refiner_timesteps], dim=0) + + refiner_noisy_latents = self.sd.noise_scheduler.add_noise(latents, noise, refiner_timesteps) + noisy_latents = torch.cat([noisy_latents, refiner_noisy_latents], dim=0) + + else: + # just double it + noisy_latents = double_up_tensor(noisy_latents) + timesteps = double_up_tensor(timesteps) + + noise = double_up_tensor(noise) + # prompts are already updated above + imgs = double_up_tensor(imgs) + batch.mask_tensor = double_up_tensor(batch.mask_tensor) + batch.control_tensor = double_up_tensor(batch.control_tensor) + + noisy_latent_multiplier = self.train_config.noisy_latent_multiplier + + if noisy_latent_multiplier != 1.0: + noisy_latents = noisy_latents * noisy_latent_multiplier + + # remove grads for these + noisy_latents.requires_grad = False + noisy_latents = noisy_latents.detach() + noise.requires_grad = False + noise = noise.detach() + + return noisy_latents, noise, timesteps, conditioned_prompts, imgs + + def setup_adapter(self): + # t2i adapter + is_t2i = self.adapter_config.type == 't2i' + is_control_net = self.adapter_config.type == 'control_net' + if self.adapter_config.type == 't2i': + suffix = 't2i' + elif self.adapter_config.type == 'control_net': + suffix = 'cn' + elif self.adapter_config.type == 'clip': + suffix = 'clip' + elif self.adapter_config.type == 'reference': + suffix = 'ref' + elif self.adapter_config.type.startswith('ip'): + suffix = 'ip' + else: + suffix = 'adapter' + adapter_name = self.name + if self.network_config is not None: + adapter_name = f"{adapter_name}_{suffix}" + latest_save_path = self.get_latest_save_path(adapter_name) + + if latest_save_path is not None and not self.adapter_config.train: + # the save path is for something else since we are not training + latest_save_path = self.adapter_config.name_or_path + + dtype = get_torch_dtype(self.train_config.dtype) + if is_t2i: + # if we do not have a last save path and we have a name_or_path, + # load from that + if latest_save_path is None and self.adapter_config.name_or_path is not None: + self.adapter = T2IAdapter.from_pretrained( + self.adapter_config.name_or_path, + torch_dtype=get_torch_dtype(self.train_config.dtype), + varient="fp16", + # use_safetensors=True, + ) + else: + self.adapter = T2IAdapter( + in_channels=self.adapter_config.in_channels, + channels=self.adapter_config.channels, + num_res_blocks=self.adapter_config.num_res_blocks, + downscale_factor=self.adapter_config.downscale_factor, + adapter_type=self.adapter_config.adapter_type, + ) + elif is_control_net: + if self.adapter_config.name_or_path is None: + raise ValueError("ControlNet requires a name_or_path to load from currently") + load_from_path = self.adapter_config.name_or_path + if latest_save_path is not None: + load_from_path = latest_save_path + self.adapter = ControlNetModel.from_pretrained( + load_from_path, + torch_dtype=get_torch_dtype(self.train_config.dtype), + ) + elif self.adapter_config.type == 'clip': + self.adapter = ClipVisionAdapter( + sd=self.sd, + adapter_config=self.adapter_config, + ) + elif self.adapter_config.type == 'reference': + self.adapter = ReferenceAdapter( + sd=self.sd, + adapter_config=self.adapter_config, + ) + elif self.adapter_config.type.startswith('ip'): + self.adapter = IPAdapter( + sd=self.sd, + adapter_config=self.adapter_config, + ) + if self.train_config.gradient_checkpointing: + self.adapter.enable_gradient_checkpointing() + else: + self.adapter = CustomAdapter( + sd=self.sd, + adapter_config=self.adapter_config, + train_config=self.train_config, + ) + self.adapter.to(self.device_torch, dtype=dtype) + if latest_save_path is not None and not is_control_net: + # load adapter from path + print_acc(f"Loading adapter from {latest_save_path}") + if is_t2i: + loaded_state_dict = load_t2i_model( + latest_save_path, + self.device, + dtype=dtype + ) + self.adapter.load_state_dict(loaded_state_dict) + elif self.adapter_config.type.startswith('ip'): + # ip adapter + loaded_state_dict = load_ip_adapter_model( + latest_save_path, + self.device, + dtype=dtype, + direct_load=self.adapter_config.train_only_image_encoder + ) + self.adapter.load_state_dict(loaded_state_dict) + else: + # custom adapter + loaded_state_dict = load_custom_adapter_model( + latest_save_path, + self.device, + dtype=dtype + ) + self.adapter.load_state_dict(loaded_state_dict) + if latest_save_path is not None and self.adapter_config.train: + self.load_training_state_from_metadata(latest_save_path) + # set trainable params + self.sd.adapter = self.adapter + + def run(self): + # torch.autograd.set_detect_anomaly(True) + # run base process run + BaseTrainProcess.run(self) + params = [] + + ### HOOK ### + self.hook_before_model_load() + model_config_to_load = copy.deepcopy(self.model_config) + + if self.is_fine_tuning: + # get the latest checkpoint + # check to see if we have a latest save + latest_save_path = self.get_latest_save_path() + + if latest_save_path is not None: + print_acc(f"#### IMPORTANT RESUMING FROM {latest_save_path} ####") + model_config_to_load.name_or_path = latest_save_path + self.load_training_state_from_metadata(latest_save_path) + + ModelClass = get_model_class(self.model_config) + # if the model class has get_train_scheduler static method + if hasattr(ModelClass, 'get_train_scheduler'): + sampler = ModelClass.get_train_scheduler() + else: + # get the noise scheduler + arch = 'sd' + if self.model_config.is_pixart: + arch = 'pixart' + if self.model_config.is_flux: + arch = 'flux' + if self.model_config.is_lumina2: + arch = 'lumina2' + sampler = get_sampler( + self.train_config.noise_scheduler, + { + "prediction_type": "v_prediction" if self.model_config.is_v_pred else "epsilon", + }, + arch=arch, + ) + + if self.train_config.train_refiner and self.model_config.refiner_name_or_path is not None and self.network_config is None: + previous_refiner_save = self.get_latest_save_path(self.job.name + '_refiner') + if previous_refiner_save is not None: + model_config_to_load.refiner_name_or_path = previous_refiner_save + self.load_training_state_from_metadata(previous_refiner_save) + + self.sd = ModelClass( + # todo handle single gpu and multi gpu here + # device=self.device, + device=self.accelerator.device, + model_config=model_config_to_load, + dtype=self.train_config.dtype, + custom_pipeline=self.custom_pipeline, + noise_scheduler=sampler, + ) + + self.hook_after_sd_init_before_load() + # run base sd process run + self.sd.load_model() + + # compile the model if needed + if self.model_config.compile: + try: + torch.compile(self.sd.unet, dynamic=True, fullgraph=True, mode='max-autotune') + except Exception as e: + print_acc(f"Failed to compile model: {e}") + print_acc("Continuing without compilation") + + self.sd.add_after_sample_image_hook(self.sample_step_hook) + + dtype = get_torch_dtype(self.train_config.dtype) + + # model is loaded from BaseSDProcess + unet = self.sd.unet + vae = self.sd.vae + tokenizer = self.sd.tokenizer + text_encoder = self.sd.text_encoder + noise_scheduler = self.sd.noise_scheduler + + if self.train_config.xformers: + vae.enable_xformers_memory_efficient_attention() + unet.enable_xformers_memory_efficient_attention() + if isinstance(text_encoder, list): + for te in text_encoder: + # if it has it + if hasattr(te, 'enable_xformers_memory_efficient_attention'): + te.enable_xformers_memory_efficient_attention() + + if self.train_config.attention_backend != 'native': + if hasattr(vae, 'set_attention_backend'): + vae.set_attention_backend(self.train_config.attention_backend) + if hasattr(unet, 'set_attention_backend'): + unet.set_attention_backend(self.train_config.attention_backend) + if isinstance(text_encoder, list): + for te in text_encoder: + if hasattr(te, 'set_attention_backend'): + te.set_attention_backend(self.train_config.attention_backend) + else: + if hasattr(text_encoder, 'set_attention_backend'): + text_encoder.set_attention_backend(self.train_config.attention_backend) + if self.train_config.sdp: + torch.backends.cuda.enable_math_sdp(True) + torch.backends.cuda.enable_flash_sdp(True) + torch.backends.cuda.enable_mem_efficient_sdp(True) + + # # check if we have sage and is flux + # if self.sd.is_flux: + # # try_to_activate_sage_attn() + # try: + # from sageattention import sageattn + # from toolkit.models.flux_sage_attn import FluxSageAttnProcessor2_0 + # model: FluxTransformer2DModel = self.sd.unet + # # enable sage attention on each block + # for block in model.transformer_blocks: + # processor = FluxSageAttnProcessor2_0() + # block.attn.set_processor(processor) + # for block in model.single_transformer_blocks: + # processor = FluxSageAttnProcessor2_0() + # block.attn.set_processor(processor) + + # except ImportError: + # print_acc("sage attention is not installed. Using SDP instead") + + if self.train_config.gradient_checkpointing: + # if has method enable_gradient_checkpointing + if hasattr(unet, 'enable_gradient_checkpointing'): + unet.enable_gradient_checkpointing() + elif hasattr(unet, 'gradient_checkpointing'): + unet.gradient_checkpointing = True + else: + print("Gradient checkpointing not supported on this model") + if isinstance(text_encoder, list): + for te in text_encoder: + if hasattr(te, 'enable_gradient_checkpointing'): + te.enable_gradient_checkpointing() + if hasattr(te, "gradient_checkpointing_enable"): + te.gradient_checkpointing_enable() + else: + if hasattr(text_encoder, 'enable_gradient_checkpointing'): + text_encoder.enable_gradient_checkpointing() + if hasattr(text_encoder, "gradient_checkpointing_enable"): + text_encoder.gradient_checkpointing_enable() + + if self.sd.refiner_unet is not None: + self.sd.refiner_unet.to(self.device_torch, dtype=dtype) + self.sd.refiner_unet.requires_grad_(False) + self.sd.refiner_unet.eval() + if self.train_config.xformers: + self.sd.refiner_unet.enable_xformers_memory_efficient_attention() + if self.train_config.gradient_checkpointing: + self.sd.refiner_unet.enable_gradient_checkpointing() + + # Respect training config for text encoder. + # If text embeddings are cached or TE is quantized, TE must remain frozen. + te_should_train = bool(getattr(self.train_config, "train_text_encoder", False)) + if getattr(self.train_config, "cache_text_embeddings", False) or getattr(self, "is_caching_text_embeddings", False): + te_should_train = False + if bool(getattr(self.model_config, "quantize_te", False)): + te_should_train = False + + if isinstance(text_encoder, list): + for te in text_encoder: + te.requires_grad_(te_should_train) + te.train(te_should_train) + if not te_should_train: + te.eval() + else: + text_encoder.requires_grad_(te_should_train) + text_encoder.train(te_should_train) + if not te_should_train: + text_encoder.eval() + unet.to(self.device_torch, dtype=dtype) + unet.requires_grad_(False) + unet.eval() + vae = vae.to(torch.device('cpu'), dtype=dtype) + vae.requires_grad_(False) + vae.eval() + if self.train_config.learnable_snr_gos: + self.snr_gos = LearnableSNRGamma( + self.sd.noise_scheduler, device=self.device_torch + ) + # check to see if previous settings exist + path_to_load = os.path.join(self.save_root, 'learnable_snr.json') + if os.path.exists(path_to_load): + with open(path_to_load, 'r') as f: + json_data = json.load(f) + if 'offset' in json_data: + # legacy + self.snr_gos.offset_2.data = torch.tensor(json_data['offset'], device=self.device_torch) + else: + self.snr_gos.offset_1.data = torch.tensor(json_data['offset_1'], device=self.device_torch) + self.snr_gos.offset_2.data = torch.tensor(json_data['offset_2'], device=self.device_torch) + self.snr_gos.scale.data = torch.tensor(json_data['scale'], device=self.device_torch) + self.snr_gos.gamma.data = torch.tensor(json_data['gamma'], device=self.device_torch) + + self.hook_after_model_load() + flush() + if not self.is_fine_tuning: + if self.network_config is not None: + # TODO should we completely switch to LycorisSpecialNetwork? + network_kwargs = self.network_config.network_kwargs + is_lycoris = False + is_lorm = self.network_config.type.lower() == 'lorm' + # default to LoCON if there are any conv layers or if it is named + NetworkClass = LoRASpecialNetwork + if self.network_config.type.lower() == 'locon' or self.network_config.type.lower() == 'lycoris': + NetworkClass = LycorisSpecialNetwork + is_lycoris = True + + if is_lorm: + network_kwargs['ignore_if_contains'] = lorm_ignore_if_contains + network_kwargs['parameter_threshold'] = lorm_parameter_threshold + network_kwargs['target_lin_modules'] = LORM_TARGET_REPLACE_MODULE + + if self.network_config.type.lower() == 'loha': + from toolkit.loha_network import LoHaNetwork + NetworkClass = LoHaNetwork + # LoHa doesn't use standard lora_dim/alpha names in init sometimes, + # but our LoHaNetwork wrapper standardizes this. + # We ensure kwargs don't conflict + if 'conv_lora_dim' in network_kwargs: del network_kwargs['conv_lora_dim'] + if 'conv_alpha' in network_kwargs: del network_kwargs['conv_alpha'] + + # if is_lycoris: + # preset = PRESET['full'] + # NetworkClass.apply_preset(preset) + + if hasattr(self.sd, 'target_lora_modules'): + network_kwargs['target_lin_modules'] = self.sd.target_lora_modules + + self.network = NetworkClass( + text_encoder=text_encoder, + unet=self.sd.get_model_to_train(), + lora_dim=self.network_config.linear, + multiplier=1.0, + alpha=self.network_config.linear_alpha, + train_unet=self.train_config.train_unet, + train_text_encoder=self.train_config.train_text_encoder, + conv_lora_dim=self.network_config.conv, + conv_alpha=self.network_config.conv_alpha, + is_sdxl=self.model_config.is_xl or self.model_config.is_ssd, + is_v2=self.model_config.is_v2, + is_v3=self.model_config.is_v3, + is_pixart=self.model_config.is_pixart, + is_auraflow=self.model_config.is_auraflow, + is_flux=self.model_config.is_flux, + is_lumina2=self.model_config.is_lumina2, + is_ssd=self.model_config.is_ssd, + is_vega=self.model_config.is_vega, + dropout=self.network_config.dropout, + use_text_encoder_1=self.model_config.use_text_encoder_1, + use_text_encoder_2=self.model_config.use_text_encoder_2, + use_bias=is_lorm, + is_lorm=is_lorm, + network_config=self.network_config, + network_type=self.network_config.type, + transformer_only=self.network_config.transformer_only, + is_transformer=self.sd.is_transformer, + base_model=self.sd, + **network_kwargs + ) + + + # todo switch everything to proper mixed precision like this + self.network.force_to(self.device_torch, dtype=torch.float32) + # give network to sd so it can use it + self.sd.network = self.network + self.network._update_torch_multiplier() + + self.network.apply_to( + text_encoder, + unet, + self.train_config.train_text_encoder, + self.train_config.train_unet + ) + + # we cannot merge in if quantized + if self.model_config.quantize or self.model_config.layer_offloading: + # todo find a way around this + self.network.can_merge_in = False + + if is_lorm: + self.network.is_lorm = True + # make sure it is on the right device + self.sd.unet.to(self.sd.device, dtype=dtype) + original_unet_param_count = count_parameters(self.sd.unet) + self.network.setup_lorm() + new_unet_param_count = original_unet_param_count - self.network.calculate_lorem_parameter_reduction() + + print_lorm_extract_details( + start_num_params=original_unet_param_count, + end_num_params=new_unet_param_count, + num_replaced=len(self.network.get_all_modules()), + ) + + self.network.prepare_grad_etc(text_encoder, unet) + flush() + + # LyCORIS doesnt have default_lr + config = { + # allow separate LR for text encoder when provided in train config + 'text_encoder_lr': (self.text_encoder_lr_override if self.text_encoder_lr_override is not None else self.train_config.lr), + 'unet_lr': (self.unet_lr_override if self.unet_lr_override is not None else self.train_config.lr), + } + sig = inspect.signature(self.network.prepare_optimizer_params) + if 'default_lr' in sig.parameters: + config['default_lr'] = self.train_config.lr + if 'learning_rate' in sig.parameters: + config['learning_rate'] = self.train_config.lr + params_net = self.network.prepare_optimizer_params( + **config + ) + + params += params_net + + if self.train_config.gradient_checkpointing: + self.network.enable_gradient_checkpointing() + + lora_name = self.name + # need to adapt name so they are not mixed up + if self.named_lora: + lora_name = f"{lora_name}_LoRA" + + latest_save_path = self.get_latest_save_path(lora_name) + extra_weights = None + if latest_save_path is not None: + print_acc(f"#### IMPORTANT RESUMING FROM {latest_save_path} ####") + print_acc(f"Loading from {latest_save_path}") + extra_weights = self.load_weights(latest_save_path) + self.network.multiplier = 1.0 + + if self.network_config.layer_offloading: + MemoryManager.attach( + self.network, + self.device_torch + ) + + if self.embed_config is not None: + # we are doing embedding training as well + self.embedding = Embedding( + sd=self.sd, + embed_config=self.embed_config + ) + latest_save_path = self.get_latest_save_path(self.embed_config.trigger) + # load last saved weights + if latest_save_path is not None: + self.embedding.load_embedding_from_file(latest_save_path, self.device_torch) + if self.embedding.step > 1: + self.step_num = self.embedding.step + self.start_step = self.step_num + + # self.step_num = self.embedding.step + # self.start_step = self.step_num + params.append({ + 'params': list(self.embedding.get_trainable_params()), + 'lr': self.train_config.embedding_lr + }) + + flush() + + if self.decorator_config is not None: + self.decorator = Decorator( + num_tokens=self.decorator_config.num_tokens, + token_size=4096 # t5xxl hidden size for flux + ) + latest_save_path = self.get_latest_save_path() + # load last saved weights + if latest_save_path is not None: + state_dict = load_file(latest_save_path) + self.decorator.load_state_dict(state_dict) + self.load_training_state_from_metadata(latest_save_path) + + params.append({ + 'params': list(self.decorator.parameters()), + 'lr': self.train_config.lr + }) + + # give it to the sd network + self.sd.decorator = self.decorator + self.decorator.to(self.device_torch, dtype=torch.float32) + self.decorator.train() + + flush() + + if self.adapter_config is not None: + self.setup_adapter() + if self.adapter_config.train: + + if isinstance(self.adapter, IPAdapter): + # we have custom LR groups for IPAdapter + adapter_param_groups = self.adapter.get_parameter_groups(self.train_config.adapter_lr) + for group in adapter_param_groups: + params.append(group) + else: + # set trainable params + params.append({ + 'params': list(self.adapter.parameters()), + 'lr': self.train_config.adapter_lr + }) + + if self.train_config.gradient_checkpointing: + self.adapter.enable_gradient_checkpointing() + flush() + + params = self.load_additional_training_modules(params) + + else: # no network, embedding or adapter + # set the device state preset before getting params + self.sd.set_device_state(self.get_params_device_state_preset) + + # params = self.get_params() + if len(params) == 0: + # will only return savable weights and ones with grad + params = self.sd.prepare_optimizer_params( + unet=self.train_config.train_unet, + text_encoder=self.train_config.train_text_encoder, + text_encoder_lr=(self.text_encoder_lr_override if self.text_encoder_lr_override is not None else self.train_config.lr), + unet_lr=(self.unet_lr_override if self.unet_lr_override is not None else self.train_config.lr), + default_lr=self.train_config.lr, + refiner=self.train_config.train_refiner and self.sd.refiner_unet is not None, + refiner_lr=self.train_config.refiner_lr, + ) + # we may be using it for prompt injections + if self.adapter_config is not None and self.adapter is None: + self.setup_adapter() + flush() + ### HOOK ### + params = self.hook_add_extra_train_params(params) + self.params = params + # self.params = [] + + # for param in params: + # if isinstance(param, dict): + # self.params += param['params'] + # else: + # self.params.append(param) + + if self.train_config.start_step is not None: + self.step_num = self.train_config.start_step + self.start_step = self.step_num + + optimizer_type = self.train_config.optimizer.lower() + + # esure params require grad + self.ensure_params_requires_grad(force=True) + optimizer = get_optimizer(self.params, optimizer_type, learning_rate=self.train_config.lr, + optimizer_params=self.train_config.optimizer_params) + self.optimizer = optimizer + + # set it to do paramiter swapping + if self.train_config.do_paramiter_swapping: + # only works for adafactor, but it should have thrown an error prior to this otherwise + self.optimizer.enable_paramiter_swapping(self.train_config.paramiter_swapping_factor) + + # check if it exists + optimizer_state_filename = f'optimizer.pt' + optimizer_state_file_path = os.path.join(self.save_root, optimizer_state_filename) + if os.path.exists(optimizer_state_file_path): + # try to load + # previous param groups + # previous_params = copy.deepcopy(optimizer.param_groups) + previous_lrs = [] + for group in optimizer.param_groups: + previous_lrs.append(group['lr']) + + load_optimizer = True + if self.network is not None: + if self.network.did_change_weights: + # do not load optimizer if the network changed, it will result in + # a double state that will oom. + load_optimizer = False + + if load_optimizer: + try: + print_acc(f"Loading optimizer state from {optimizer_state_file_path}") + optimizer_state_dict = torch.load(optimizer_state_file_path, weights_only=True) + optimizer.load_state_dict(optimizer_state_dict) + del optimizer_state_dict + flush() + except Exception as e: + print_acc(f"Failed to load optimizer state from {optimizer_state_file_path}") + print_acc(e) + + # update the optimizer LR from the params + print_acc(f"Updating optimizer LR from params") + if len(previous_lrs) > 0: + for i, group in enumerate(optimizer.param_groups): + group['lr'] = previous_lrs[i] + group['initial_lr'] = previous_lrs[i] + + # Update the learning rates if they changed + # optimizer.param_groups = previous_params + + lr_scheduler_params = self.train_config.lr_scheduler_params + + # make sure it had bare minimum + if 'max_iterations' not in lr_scheduler_params: + lr_scheduler_params['total_iters'] = self.train_config.steps + + lr_scheduler = get_lr_scheduler( + self.train_config.lr_scheduler, + optimizer, + **lr_scheduler_params + ) + self.lr_scheduler = lr_scheduler + + ### HOOk ### + self.before_dataset_load() + # load datasets if passed in the root process + if self.datasets is not None: + self.data_loader = get_dataloader_from_datasets(self.datasets, self.train_config.batch_size, self.sd) + if self.datasets_reg is not None: + self.data_loader_reg = get_dataloader_from_datasets(self.datasets_reg, self.train_config.batch_size, + self.sd) + + flush() + self.last_save_step = self.step_num + ### HOOK ### + self.hook_before_train_loop() + + if self.has_first_sample_requested and self.step_num <= 1 and not self.train_config.disable_sampling: + print_acc("Generating first sample from first sample config") + self.sample(0, is_first=True) + + # sample first + if self.train_config.skip_first_sample or self.train_config.disable_sampling: + print_acc("Skipping first sample due to config setting") + elif self.step_num <= 1 or self.train_config.force_first_sample: + print_acc("Generating baseline samples before training") + self.sample(self.step_num) + + if self.accelerator.is_local_main_process: + self.progress_bar = ToolkitProgressBar( + total=self.train_config.steps, + desc=self.job.name, + leave=True, + initial=self.step_num, + iterable=range(0, self.train_config.steps), + ) + self.progress_bar.pause() + else: + self.progress_bar = None + + if self.data_loader is not None: + dataloader = self.data_loader + dataloader_iterator = iter(dataloader) + else: + dataloader = None + dataloader_iterator = None + + if self.data_loader_reg is not None: + dataloader_reg = self.data_loader_reg + dataloader_iterator_reg = iter(dataloader_reg) + else: + dataloader_reg = None + dataloader_iterator_reg = None + + # zero any gradients + optimizer.zero_grad() + + self.lr_scheduler.step(self.step_num) + + self.sd.set_device_state(self.train_device_state_preset) + flush() + # self.step_num = 0 + + # print_acc(f"Compiling Model") + # torch.compile(self.sd.unet, dynamic=True) + + # make sure all params require grad + self.ensure_params_requires_grad(force=True) + + + ################################################################### + # TRAIN LOOP + ################################################################### + + + start_step_num = self.step_num + did_first_flush = False + flush_next = False + for step in range(start_step_num, self.train_config.steps): + if self.train_config.do_paramiter_swapping: + self.optimizer.optimizer.swap_paramiters() + self.timer.start('train_loop') + if flush_next: + flush() + flush_next = False + if self.train_config.do_random_cfg: + self.train_config.do_cfg = True + self.train_config.cfg_scale = value_map(random.random(), 0, 1, 1.0, self.train_config.max_cfg_scale) + self.step_num = step + # default to true so various things can turn it off + self.is_grad_accumulation_step = True + if self.train_config.free_u: + self.sd.pipeline.enable_freeu(s1=0.9, s2=0.2, b1=1.1, b2=1.2) + if self.progress_bar is not None: + self.progress_bar.unpause() + with torch.no_grad(): + # if is even step and we have a reg dataset, use that + # todo improve this logic to send one of each through if we can buckets and batch size might be an issue + is_reg_step = False + is_save_step = self.save_config.save_every and self.step_num % self.save_config.save_every == 0 + is_sample_step = self.sample_config.sample_every and self.step_num % self.sample_config.sample_every == 0 + if self.train_config.disable_sampling: + is_sample_step = False + + batch_list = [] + + for b in range(self.train_config.gradient_accumulation): + # keep track to alternate on an accumulation step for reg + batch_step = step + # don't do a reg step on sample or save steps as we dont want to normalize on those + if batch_step % 2 == 0 and dataloader_reg is not None and not is_save_step and not is_sample_step: + try: + with self.timer('get_batch:reg'): + batch = next(dataloader_iterator_reg) + except StopIteration: + with self.timer('reset_batch:reg'): + # hit the end of an epoch, reset + if self.progress_bar is not None: + self.progress_bar.pause() + dataloader_iterator_reg = iter(dataloader_reg) + trigger_dataloader_setup_epoch(dataloader_reg) + + with self.timer('get_batch:reg'): + batch = next(dataloader_iterator_reg) + if self.progress_bar is not None: + self.progress_bar.unpause() + is_reg_step = True + elif dataloader is not None: + try: + with self.timer('get_batch'): + batch = next(dataloader_iterator) + except StopIteration: + with self.timer('reset_batch'): + # hit the end of an epoch, reset + if self.progress_bar is not None: + self.progress_bar.pause() + dataloader_iterator = iter(dataloader) + trigger_dataloader_setup_epoch(dataloader) + self.epoch_num += 1 + if self.train_config.gradient_accumulation_steps == -1: + # if we are accumulating for an entire epoch, trigger a step + self.is_grad_accumulation_step = False + self.grad_accumulation_step = 0 + with self.timer('get_batch'): + batch = next(dataloader_iterator) + if self.progress_bar is not None: + self.progress_bar.unpause() + else: + batch = None + batch_list.append(batch) + batch_step += 1 + + # setup accumulation + if self.train_config.gradient_accumulation_steps == -1: + # epoch is handling the accumulation, dont touch it + pass + else: + # determine if we are accumulating or not + # since optimizer step happens in the loop, we trigger it a step early + # since we cannot reprocess it before them + optimizer_step_at = self.train_config.gradient_accumulation_steps + is_optimizer_step = self.grad_accumulation_step >= optimizer_step_at + self.is_grad_accumulation_step = not is_optimizer_step + if is_optimizer_step: + self.grad_accumulation_step = 0 + + # flush() + ### HOOK ### + if self.torch_profiler is not None: + self.torch_profiler.start() + did_oom = False + loss_dict = None + try: + with self.accelerator.accumulate(self.modules_being_trained): + loss_dict = self.hook_train_loop(batch_list) + except torch.cuda.OutOfMemoryError: + did_oom = True + except RuntimeError as e: + if "CUDA out of memory" in str(e): + did_oom = True + else: + raise # not an OOM; surface real errors + if did_oom: + self.num_consecutive_oom += 1 + if self.num_consecutive_oom > 3: + raise RuntimeError("OOM during training step 3 times in a row, aborting training") + optimizer.zero_grad(set_to_none=True) + flush() + torch.cuda.ipc_collect() + # skip this step and keep going + print_acc("") + print_acc("################################################") + print_acc(f"# OOM during training step, skipping batch {self.num_consecutive_oom}/3 #") + print_acc("################################################") + print_acc("") + else: + self.num_consecutive_oom = 0 + if self.torch_profiler is not None: + torch.cuda.synchronize() # Make sure all CUDA ops are done + self.torch_profiler.stop() + + print("\n==== Profile Results ====") + print(self.torch_profiler.key_averages().table(sort_by="cpu_time_total", row_limit=1000)) + self.timer.stop('train_loop') + if not did_first_flush: + flush() + did_first_flush = True + # flush() + # setup the networks to gradient checkpointing and everything works + if self.adapter is not None and isinstance(self.adapter, ReferenceAdapter): + self.adapter.clear_memory() + + with torch.no_grad(): + # torch.cuda.empty_cache() + # if optimizer has get_lrs method, then use it + if not did_oom and loss_dict is not None: + if hasattr(optimizer, 'get_avg_learning_rate'): + learning_rate = optimizer.get_avg_learning_rate() + elif hasattr(optimizer, 'get_learning_rates'): + learning_rate = optimizer.get_learning_rates()[0] + elif self.train_config.optimizer.lower().startswith('dadaptation') or \ + self.train_config.optimizer.lower().startswith('prodigy'): + learning_rate = ( + optimizer.param_groups[0]["d"] * + optimizer.param_groups[0]["lr"] + ) + else: + learning_rate = optimizer.param_groups[0]['lr'] + + prog_bar_string = f"lr: {learning_rate:.1e}" + for key, value in loss_dict.items(): + prog_bar_string += f" {key}: {value:.3e}" + + if self.progress_bar is not None: + self.progress_bar.set_postfix_str(prog_bar_string) + + # if the batch is a DataLoaderBatchDTO, then we need to clean it up + if isinstance(batch, DataLoaderBatchDTO): + with self.timer('batch_cleanup'): + batch.cleanup() + + # don't do on first step + if self.step_num != self.start_step: + if is_sample_step or is_save_step: + self.accelerator.wait_for_everyone() + + if is_save_step: + self.accelerator + # print above the progress bar + if self.progress_bar is not None: + self.progress_bar.pause() + print_acc(f"\nSaving at step {self.step_num}") + self.save(self.step_num) + self.ensure_params_requires_grad() + # clear any grads + optimizer.zero_grad() + flush() + flush_next = True + if self.progress_bar is not None: + self.progress_bar.unpause() + + if is_sample_step: + if self.progress_bar is not None: + self.progress_bar.pause() + flush() + # print above the progress bar + if self.train_config.free_u: + self.sd.pipeline.disable_freeu() + self.sample(self.step_num) + if self.train_config.unload_text_encoder: + # make sure the text encoder is unloaded + self.sd.text_encoder_to('cpu') + flush() + + self.ensure_params_requires_grad() + if self.progress_bar is not None: + self.progress_bar.unpause() + + if self.logging_config.log_every and self.step_num % self.logging_config.log_every == 0: + if self.progress_bar is not None: + self.progress_bar.pause() + with self.timer('log_to_tensorboard'): + # log to tensorboard + if self.accelerator.is_main_process: + if self.writer is not None: + for key, value in loss_dict.items(): + self.writer.add_scalar(f"{key}", value, self.step_num) + self.writer.add_scalar(f"lr", learning_rate, self.step_num) + if self.progress_bar is not None: + self.progress_bar.unpause() + + if self.accelerator.is_main_process: + # log to logger + self.logger.log({ + 'learning_rate': learning_rate, + }) + for key, value in loss_dict.items(): + self.logger.log({ + f'loss/{key}': value, + }) + elif self.logging_config.log_every is None: + if self.accelerator.is_main_process: + # log every step + self.logger.log({ + 'learning_rate': learning_rate, + }) + for key, value in loss_dict.items(): + self.logger.log({ + f'loss/{key}': value, + }) + + + if self.performance_log_every > 0 and self.step_num % self.performance_log_every == 0: + if self.progress_bar is not None: + self.progress_bar.pause() + # print the timers and clear them + self.timer.print() + self.timer.reset() + if self.progress_bar is not None: + self.progress_bar.unpause() + + # commit log + if self.accelerator.is_main_process: + self.logger.commit(step=self.step_num) + + # sets progress bar to match out step + if self.progress_bar is not None: + self.progress_bar.update(step - self.progress_bar.n) + + ############################# + # End of step + ############################# + + # update various steps + self.step_num = step + 1 + self.grad_accumulation_step += 1 + self.end_step_hook() + + + ################################################################### + ## END TRAIN LOOP + ################################################################### + self.accelerator.wait_for_everyone() + if self.progress_bar is not None: + self.progress_bar.close() + if self.train_config.free_u: + self.sd.pipeline.disable_freeu() + if not self.train_config.disable_sampling: + self.sample(self.step_num) + self.logger.commit(step=self.step_num) + print_acc("") + if self.accelerator.is_main_process: + self.save() + self.logger.finish() + self.accelerator.end_training() + + if self.accelerator.is_main_process: + # push to hub + if self.save_config.push_to_hub: + if("HF_TOKEN" not in os.environ): + interpreter_login(new_session=False, write_permission=True) + self.push_to_hub( + repo_id=self.save_config.hf_repo_id, + private=self.save_config.hf_private + ) + del ( + self.sd, + unet, + noise_scheduler, + optimizer, + self.network, + tokenizer, + text_encoder, + ) + + flush() + self.done_hook() + + def push_to_hub( + self, + repo_id: str, + private: bool = False, + ): + if not self.accelerator.is_main_process: + return + readme_content = self._generate_readme(repo_id) + readme_path = os.path.join(self.save_root, "README.md") + with open(readme_path, "w", encoding="utf-8") as f: + f.write(readme_content) + + api = HfApi() + + api.create_repo( + repo_id, + private=private, + exist_ok=True + ) + + api.upload_folder( + repo_id=repo_id, + folder_path=self.save_root, + ignore_patterns=["*.yaml", "*.pt"], + repo_type="model", + ) + + + def _generate_readme(self, repo_id: str) -> str: + """Generates the content of the README.md file.""" + + # Gather model info + base_model = self.model_config.name_or_path + instance_prompt = self.trigger_word if hasattr(self, "trigger_word") else None + if base_model == "black-forest-labs/FLUX.1-schnell": + license = "apache-2.0" + elif base_model == "black-forest-labs/FLUX.1-dev": + license = "other" + license_name = "flux-1-dev-non-commercial-license" + license_link = "https://huggingface.co/black-forest-labs/FLUX.1-dev/blob/main/LICENSE.md" + else: + license = "creativeml-openrail-m" + tags = [ + "text-to-image", + ] + if self.model_config.is_xl: + tags.append("stable-diffusion-xl") + if self.model_config.is_flux: + tags.append("flux") + if self.model_config.is_lumina2: + tags.append("lumina2") + if self.model_config.is_v3: + tags.append("sd3") + if self.network_config: + tags.extend( + [ + "lora", + "diffusers", + "template:sd-lora", + "ai-toolkit", + ] + ) + + # Generate the widget section + widgets = [] + sample_image_paths = [] + samples_dir = os.path.join(self.save_root, "samples") + if os.path.isdir(samples_dir): + for filename in os.listdir(samples_dir): + #The filenames are structured as 1724085406830__00000500_0.jpg + #So here we capture the 2nd part (steps) and 3rd (index the matches the prompt) + match = re.search(r"__(\d+)_(\d+)\.jpg$", filename) + if match: + steps, index = int(match.group(1)), int(match.group(2)) + #Here we only care about uploading the latest samples, the match with the # of steps + if steps == self.train_config.steps: + sample_image_paths.append((index, f"samples/{filename}")) + + # Sort by numeric index + sample_image_paths.sort(key=lambda x: x[0]) + + # Create widgets matching prompt with the index + for i, prompt in enumerate(self.sample_config.prompts): + if i < len(sample_image_paths): + # Associate prompts with sample image paths based on the extracted index + _, image_path = sample_image_paths[i] + widgets.append( + { + "text": prompt, + "output": { + "url": image_path + }, + } + ) + dtype = "torch.bfloat16" if self.model_config.is_flux else "torch.float16" + # Construct the README content + readme_content = f"""--- +tags: +{yaml.dump(tags, indent=4).strip()} +{"widget:" if os.path.isdir(samples_dir) else ""} +{yaml.dump(widgets, indent=4).strip() if widgets else ""} +base_model: {base_model} +{"instance_prompt: " + instance_prompt if instance_prompt else ""} +license: {license} +{'license_name: ' + license_name if license == "other" else ""} +{'license_link: ' + license_link if license == "other" else ""} +--- + +# {self.job.name} +Model trained with [AI Toolkit by Ostris](https://github.com/ostris/ai-toolkit) + + +## Trigger words + +{"You should use `" + instance_prompt + "` to trigger the image generation." if instance_prompt else "No trigger words defined."} + +## Download model and use it with ComfyUI, AUTOMATIC1111, SD.Next, Invoke AI, etc. + +Weights for this model are available in Safetensors format. + +[Download](/{repo_id}/tree/main) them in the Files & versions tab. + +## Use it with the [🧨 diffusers library](https://github.com/huggingface/diffusers) + +```py +from diffusers import AutoPipelineForText2Image +import torch + +pipeline = AutoPipelineForText2Image.from_pretrained('{base_model}', torch_dtype={dtype}).to('cuda') +pipeline.load_lora_weights('{repo_id}', weight_name='{self.job.name}.safetensors') +image = pipeline('{instance_prompt if not widgets else self.sample_config.prompts[0]}').images[0] +image.save("my_image.png") +``` + +For more details, including weighting, merging and fusing LoRAs, check the [documentation on loading LoRAs in diffusers](https://huggingface.co/docs/diffusers/main/en/using-diffusers/loading_adapters) + +""" + return readme_content diff --git a/toolkit/config_modules.py b/toolkit/config_modules.py new file mode 100644 index 0000000..242fa8a --- /dev/null +++ b/toolkit/config_modules.py @@ -0,0 +1,1338 @@ +import os +import time +from typing import List, Optional, Literal, Tuple, Union, TYPE_CHECKING, Dict +import random + +import torch +import torchaudio + +from toolkit.prompt_utils import PromptEmbeds + +ImgExt = Literal['jpg', 'png', 'webp'] + +SaveFormat = Literal['safetensors', 'diffusers'] + +if TYPE_CHECKING: + from toolkit.guidance import GuidanceType + from toolkit.logging_aitk import EmptyLogger +else: + EmptyLogger = None + +class SaveConfig: + def __init__(self, **kwargs): + self.save_every: int = kwargs.get('save_every', 1000) + self.dtype: str = kwargs.get('dtype', 'float16') + self.max_step_saves_to_keep: int = kwargs.get('max_step_saves_to_keep', 5) + self.save_format: SaveFormat = kwargs.get('save_format', 'safetensors') + if self.save_format not in ['safetensors', 'diffusers']: + raise ValueError(f"save_format must be safetensors or diffusers, got {self.save_format}") + self.push_to_hub: bool = kwargs.get("push_to_hub", False) + self.hf_repo_id: Optional[str] = kwargs.get("hf_repo_id", None) + self.hf_private: Optional[str] = kwargs.get("hf_private", False) + +class LoggingConfig: + def __init__(self, **kwargs): + self.log_every: int = kwargs.get('log_every', 100) + self.verbose: bool = kwargs.get('verbose', False) + self.use_wandb: bool = kwargs.get('use_wandb', False) + self.project_name: str = kwargs.get('project_name', 'ai-toolkit') + self.run_name: str = kwargs.get('run_name', None) + +class SampleItem: + def __init__( + self, + sample_config: 'SampleConfig', + **kwargs + ): + # prompt should always be in the kwargs + self.prompt = kwargs.get('prompt', None) + self.width: int = kwargs.get('width', sample_config.width) + self.height: int = kwargs.get('height', sample_config.height) + self.neg: str = kwargs.get('neg', sample_config.neg) + self.seed: Optional[int] = kwargs.get('seed', None) # if none, default to autogen seed + self.guidance_scale: float = kwargs.get('guidance_scale', sample_config.guidance_scale) + self.sample_steps: int = kwargs.get('sample_steps', sample_config.sample_steps) + self.fps: int = kwargs.get('fps', sample_config.fps) + self.num_frames: int = kwargs.get('num_frames', sample_config.num_frames) + self.ctrl_img: Optional[str] = kwargs.get('ctrl_img', None) + self.ctrl_idx: int = kwargs.get('ctrl_idx', 0) + # for multi control image models + self.ctrl_img_1: Optional[str] = kwargs.get('ctrl_img_1', self.ctrl_img) + self.ctrl_img_2: Optional[str] = kwargs.get('ctrl_img_2', None) + self.ctrl_img_3: Optional[str] = kwargs.get('ctrl_img_3', None) + + self.network_multiplier: float = kwargs.get('network_multiplier', sample_config.network_multiplier) + # convert to a number if it is a string + if isinstance(self.network_multiplier, str): + try: + self.network_multiplier = float(self.network_multiplier) + except: + print(f"Invalid network_multiplier {self.network_multiplier}, defaulting to 1.0") + self.network_multiplier = 1.0 + + # only for models that support it, (qwen image edit 2509 for now) + self.do_cfg_norm: bool = kwargs.get('do_cfg_norm', False) + +class SampleConfig: + def __init__(self, **kwargs): + self.sampler: str = kwargs.get('sampler', 'ddpm') + self.sample_every: int = kwargs.get('sample_every', 100) + self.width: int = kwargs.get('width', 512) + self.height: int = kwargs.get('height', 512) + self.neg = kwargs.get('neg', False) + self.seed = kwargs.get('seed', 0) + self.walk_seed = kwargs.get('walk_seed', False) + self.guidance_scale = kwargs.get('guidance_scale', 7) + self.sample_steps = kwargs.get('sample_steps', 20) + self.network_multiplier = kwargs.get('network_multiplier', 1) + self.guidance_rescale = kwargs.get('guidance_rescale', 0.0) + self.ext: ImgExt = kwargs.get('format', 'jpg') + self.adapter_conditioning_scale = kwargs.get('adapter_conditioning_scale', 1.0) + self.refiner_start_at = kwargs.get('refiner_start_at', + 0.5) # step to start using refiner on sample if it exists + self.extra_values = kwargs.get('extra_values', []) + self.num_frames = kwargs.get('num_frames', 1) + self.fps: int = kwargs.get('fps', 16) + if self.num_frames > 1 and self.ext not in ['webp']: + print("Changing sample extention to animated webp") + self.ext = 'webp' + + prompts: list[str] = kwargs.get('prompts', []) + + self.samples: Optional[List[SampleItem]] = None + # use the legacy prompts if it is passed that way to get samples object + default_samples_kwargs = [ + {"prompt": x} for x in prompts + ] + raw_samples = kwargs.get('samples', default_samples_kwargs) + self.samples = [SampleItem(self, **item) for item in raw_samples] + # only for models that support it, (qwen image edit 2509 for now) + self.do_cfg_norm: bool = kwargs.get('do_cfg_norm', False) + + @property + def prompts(self): + # for backwards compatibility as this is checked for length frequently + return [sample.prompt for sample in self.samples if sample.prompt is not None] + + + + +class LormModuleSettingsConfig: + def __init__(self, **kwargs): + self.contains: str = kwargs.get('contains', '4nt$3') + self.extract_mode: str = kwargs.get('extract_mode', 'ratio') + # min num parameters to attach to + self.parameter_threshold: int = kwargs.get('parameter_threshold', 0) + self.extract_mode_param: dict = kwargs.get('extract_mode_param', 0.25) + + +class LoRMConfig: + def __init__(self, **kwargs): + self.extract_mode: str = kwargs.get('extract_mode', 'ratio') + self.do_conv: bool = kwargs.get('do_conv', False) + self.extract_mode_param: dict = kwargs.get('extract_mode_param', 0.25) + self.parameter_threshold: int = kwargs.get('parameter_threshold', 0) + module_settings = kwargs.get('module_settings', []) + default_module_settings = { + 'extract_mode': self.extract_mode, + 'extract_mode_param': self.extract_mode_param, + 'parameter_threshold': self.parameter_threshold, + } + module_settings = [{**default_module_settings, **module_setting, } for module_setting in module_settings] + self.module_settings: List[LormModuleSettingsConfig] = [LormModuleSettingsConfig(**module_setting) for + module_setting in module_settings] + + def get_config_for_module(self, block_name): + for setting in self.module_settings: + contain_pieces = setting.contains.split('|') + if all(contain_piece in block_name for contain_piece in contain_pieces): + return setting + # try replacing the . with _ + contain_pieces = setting.contains.replace('.', '_').split('|') + if all(contain_piece in block_name for contain_piece in contain_pieces): + return setting + # do default + return LormModuleSettingsConfig(**{ + 'extract_mode': self.extract_mode, + 'extract_mode_param': self.extract_mode_param, + 'parameter_threshold': self.parameter_threshold, + }) + + +NetworkType = Literal['lora', 'locon', 'lorm', 'lokr'] + + +class NetworkConfig: + def __init__(self, **kwargs): + self.type: NetworkType = kwargs.get('type', 'lora') + rank = kwargs.get('rank', None) + linear = kwargs.get('linear', None) + if rank is not None: + self.rank: int = rank # rank for backward compatibility + self.linear: int = rank + elif linear is not None: + self.rank: int = linear + self.linear: int = linear + else: + self.rank: int = 4 + self.linear: int = 4 + self.conv: int = kwargs.get('conv', None) + self.alpha: float = kwargs.get('alpha', 1.0) + self.linear_alpha: float = kwargs.get('linear_alpha', self.alpha) + self.conv_alpha: float = kwargs.get('conv_alpha', self.conv) + self.dropout: Union[float, None] = kwargs.get('dropout', None) + self.network_kwargs: dict = kwargs.get('network_kwargs', {}) + + self.lorm_config: Union[LoRMConfig, None] = None + lorm = kwargs.get('lorm', None) + if lorm is not None: + self.lorm_config: LoRMConfig = LoRMConfig(**lorm) + + if self.type == 'lorm': + # set linear to arbitrary values so it makes them + self.linear = 4 + self.rank = 4 + if self.lorm_config.do_conv: + self.conv = 4 + + self.transformer_only = kwargs.get('transformer_only', True) + + self.lokr_full_rank = kwargs.get('lokr_full_rank', False) + if self.lokr_full_rank and self.type.lower() == 'lokr': + self.linear = 9999999999 + self.linear_alpha = 9999999999 + self.conv = 9999999999 + self.conv_alpha = 9999999999 + # -1 automatically finds the largest factor + self.lokr_factor = kwargs.get('lokr_factor', -1) + + # for multi stage models + self.split_multistage_loras = kwargs.get('split_multistage_loras', True) + + # ramtorch, doesn't work yet + self.layer_offloading = kwargs.get('layer_offloading', False) + + +AdapterTypes = Literal['t2i', 'ip', 'ip+', 'clip', 'ilora', 'photo_maker', 'control_net', 'control_lora', 'i2v'] + +CLIPLayer = Literal['penultimate_hidden_states', 'image_embeds', 'last_hidden_state'] + + +class AdapterConfig: + def __init__(self, **kwargs): + self.type: AdapterTypes = kwargs.get('type', 't2i') # t2i, ip, clip, control_net, i2v + self.in_channels: int = kwargs.get('in_channels', 3) + self.channels: List[int] = kwargs.get('channels', [320, 640, 1280, 1280]) + self.num_res_blocks: int = kwargs.get('num_res_blocks', 2) + self.downscale_factor: int = kwargs.get('downscale_factor', 8) + self.adapter_type: str = kwargs.get('adapter_type', 'full_adapter') + self.image_dir: str = kwargs.get('image_dir', None) + self.test_img_path: List[str] = kwargs.get('test_img_path', None) + if self.test_img_path is not None: + if isinstance(self.test_img_path, str): + self.test_img_path = self.test_img_path.split(',') + self.test_img_path = [p.strip() for p in self.test_img_path] + self.test_img_path = [p for p in self.test_img_path if p != ''] + + self.train: str = kwargs.get('train', False) + self.image_encoder_path: str = kwargs.get('image_encoder_path', None) + self.name_or_path = kwargs.get('name_or_path', None) + + num_tokens = kwargs.get('num_tokens', None) + if num_tokens is None and self.type.startswith('ip'): + if self.type == 'ip+': + num_tokens = 16 + num_tokens = 16 + elif self.type == 'ip': + num_tokens = 4 + + self.num_tokens: int = num_tokens + self.train_image_encoder: bool = kwargs.get('train_image_encoder', False) + self.train_only_image_encoder: bool = kwargs.get('train_only_image_encoder', False) + if self.train_only_image_encoder: + self.train_image_encoder = True + self.train_only_image_encoder_positional_embedding: bool = kwargs.get( + 'train_only_image_encoder_positional_embedding', False) + self.image_encoder_arch: str = kwargs.get('image_encoder_arch', 'clip') # clip vit vit_hybrid, safe + self.safe_reducer_channels: int = kwargs.get('safe_reducer_channels', 512) + self.safe_channels: int = kwargs.get('safe_channels', 2048) + self.safe_tokens: int = kwargs.get('safe_tokens', 8) + self.quad_image: bool = kwargs.get('quad_image', False) + + # clip vision + self.trigger = kwargs.get('trigger', 'tri993r') + self.trigger_class_name = kwargs.get('trigger_class_name', None) + + self.class_names = kwargs.get('class_names', []) + + self.clip_layer: CLIPLayer = kwargs.get('clip_layer', None) + if self.clip_layer is None: + if self.type.startswith('ip+'): + self.clip_layer = 'penultimate_hidden_states' + else: + self.clip_layer = 'last_hidden_state' + + # text encoder + self.text_encoder_path: str = kwargs.get('text_encoder_path', None) + self.text_encoder_arch: str = kwargs.get('text_encoder_arch', 'clip') # clip t5 + + self.train_scaler: bool = kwargs.get('train_scaler', False) + self.scaler_lr: Optional[float] = kwargs.get('scaler_lr', None) + + # trains with a scaler to easy channel bias but merges it in on save + self.merge_scaler: bool = kwargs.get('merge_scaler', False) + + # for ilora + self.head_dim: int = kwargs.get('head_dim', 1024) + self.num_heads: int = kwargs.get('num_heads', 1) + self.ilora_down: bool = kwargs.get('ilora_down', True) + self.ilora_mid: bool = kwargs.get('ilora_mid', True) + self.ilora_up: bool = kwargs.get('ilora_up', True) + + self.pixtral_max_image_size: int = kwargs.get('pixtral_max_image_size', 512) + self.pixtral_random_image_size: int = kwargs.get('pixtral_random_image_size', False) + + self.flux_only_double: bool = kwargs.get('flux_only_double', False) + + # train and use a conv layer to pool the embedding + self.conv_pooling: bool = kwargs.get('conv_pooling', False) + self.conv_pooling_stacks: int = kwargs.get('conv_pooling_stacks', 1) + self.sparse_autoencoder_dim: Optional[int] = kwargs.get('sparse_autoencoder_dim', None) + + # for llm adapter + self.num_cloned_blocks: int = kwargs.get('num_cloned_blocks', 0) + self.quantize_llm: bool = kwargs.get('quantize_llm', False) + + # for control lora only + lora_config: dict = kwargs.get('lora_config', None) + if lora_config is not None: + self.lora_config: NetworkConfig = NetworkConfig(**lora_config) + else: + self.lora_config = None + self.num_control_images: int = kwargs.get('num_control_images', 1) + # decimal for how often the control is dropped out and replaced with noise 1.0 is 100% + self.control_image_dropout: float = kwargs.get('control_image_dropout', 0.0) + self.has_inpainting_input: bool = kwargs.get('has_inpainting_input', False) + self.invert_inpaint_mask_chance: float = kwargs.get('invert_inpaint_mask_chance', 0.0) + + # for subpixel adapter + self.subpixel_downscale_factor: int = kwargs.get('subpixel_downscale_factor', 8) + + # for i2v adapter + # append the masked start frame. During pretraining we will only do the vision encoder + self.i2v_do_start_frame: bool = kwargs.get('i2v_do_start_frame', False) + + +class EmbeddingConfig: + def __init__(self, **kwargs): + self.trigger = kwargs.get('trigger', 'custom_embedding') + self.tokens = kwargs.get('tokens', 4) + self.init_words = kwargs.get('init_words', '*') + self.save_format = kwargs.get('save_format', 'safetensors') + self.trigger_class_name = kwargs.get('trigger_class_name', None) # used for inverted masked prior + + +class DecoratorConfig: + def __init__(self, **kwargs): + self.num_tokens: str = kwargs.get('num_tokens', 4) + + +ContentOrStyleType = Literal['balanced', 'style', 'content'] +LossTarget = Literal['noise', 'source', 'unaugmented', 'differential_noise'] + + +class TrainConfig: + def __init__(self, **kwargs): + self.noise_scheduler = kwargs.get('noise_scheduler', 'ddpm') + self.content_or_style: ContentOrStyleType = kwargs.get('content_or_style', 'balanced') + self.content_or_style_reg: ContentOrStyleType = kwargs.get('content_or_style', 'balanced') + self.steps: int = kwargs.get('steps', 1000) + self.lr = kwargs.get('lr', 1e-6) + self.unet_lr = kwargs.get('unet_lr', self.lr) + self.text_encoder_lr = kwargs.get('text_encoder_lr', self.lr) + self.refiner_lr = kwargs.get('refiner_lr', self.lr) + self.embedding_lr = kwargs.get('embedding_lr', self.lr) + self.adapter_lr = kwargs.get('adapter_lr', self.lr) + self.optimizer = kwargs.get('optimizer', 'adamw') + self.optimizer_params = kwargs.get('optimizer_params', {}) + self.lr_scheduler = kwargs.get('lr_scheduler', 'constant') + self.lr_scheduler_params = kwargs.get('lr_scheduler_params', {}) + self.min_denoising_steps: int = kwargs.get('min_denoising_steps', 0) + self.max_denoising_steps: int = kwargs.get('max_denoising_steps', 999) + self.batch_size: int = kwargs.get('batch_size', 1) + self.orig_batch_size: int = self.batch_size + self.dtype: str = kwargs.get('dtype', 'fp32') + self.xformers = kwargs.get('xformers', False) + self.sdp = kwargs.get('sdp', False) + # see https://huggingface.co/docs/diffusers/main/optimization/attention_backends#available-backends for options + self.attention_backend: str = kwargs.get('attention_backend', 'native') # native, flash, _flash_3_hub, _flash_3, + self.train_unet = kwargs.get('train_unet', True) + self.train_text_encoder = kwargs.get('train_text_encoder', False) + self.train_refiner = kwargs.get('train_refiner', True) + self.train_turbo = kwargs.get('train_turbo', False) + self.show_turbo_outputs = kwargs.get('show_turbo_outputs', False) + self.min_snr_gamma = kwargs.get('min_snr_gamma', None) + self.snr_gamma = kwargs.get('snr_gamma', None) + # trains a gamma, offset, and scale to adjust loss to adapt to timestep differentials + # this should balance the learning rate across all timesteps over time + self.learnable_snr_gos = kwargs.get('learnable_snr_gos', False) + self.noise_offset = kwargs.get('noise_offset', 0.0) + self.skip_first_sample = kwargs.get('skip_first_sample', False) + self.force_first_sample = kwargs.get('force_first_sample', False) + self.gradient_checkpointing = kwargs.get('gradient_checkpointing', True) + self.weight_jitter = kwargs.get('weight_jitter', 0.0) + self.merge_network_on_save = kwargs.get('merge_network_on_save', False) + self.max_grad_norm = kwargs.get('max_grad_norm', 1.0) + self.start_step = kwargs.get('start_step', None) + self.free_u = kwargs.get('free_u', False) + self.adapter_assist_name_or_path: Optional[str] = kwargs.get('adapter_assist_name_or_path', None) + self.adapter_assist_type: Optional[str] = kwargs.get('adapter_assist_type', 't2i') # t2i, control_net + self.noise_multiplier = kwargs.get('noise_multiplier', 1.0) + self.target_noise_multiplier = kwargs.get('target_noise_multiplier', 1.0) + self.random_noise_multiplier = kwargs.get('random_noise_multiplier', 0.0) + self.random_noise_shift = kwargs.get('random_noise_shift', 0.0) + self.img_multiplier = kwargs.get('img_multiplier', 1.0) + self.noisy_latent_multiplier = kwargs.get('noisy_latent_multiplier', 1.0) + self.latent_multiplier = kwargs.get('latent_multiplier', 1.0) + self.negative_prompt = kwargs.get('negative_prompt', None) + self.max_negative_prompts = kwargs.get('max_negative_prompts', 1) + # multiplier applied to loos on regularization images + self.reg_weight = kwargs.get('reg_weight', 1.0) + self.num_train_timesteps = kwargs.get('num_train_timesteps', 1000) + # automatically adapte the vae scaling based on the image norm + self.adaptive_scaling_factor = kwargs.get('adaptive_scaling_factor', False) + + # dropout that happens before encoding. It functions independently per text encoder + self.prompt_dropout_prob = kwargs.get('prompt_dropout_prob', 0.0) + + # match the norm of the noise before computing loss. This will help the model maintain its + # current understandin of the brightness of images. + + self.match_noise_norm = kwargs.get('match_noise_norm', False) + + # set to -1 to accumulate gradients for entire epoch + # warning, only do this with a small dataset or you will run out of memory + # This is legacy but left in for backwards compatibility + self.gradient_accumulation_steps = kwargs.get('gradient_accumulation_steps', 1) + + # this will do proper gradient accumulation where you will not see a step until the end of the accumulation + # the method above will show a step every accumulation + self.gradient_accumulation = kwargs.get('gradient_accumulation', 1) + if self.gradient_accumulation > 1: + if self.gradient_accumulation_steps != 1: + raise ValueError("gradient_accumulation and gradient_accumulation_steps are mutually exclusive") + + # short long captions will double your batch size. This only works when a dataset is + # prepared with a json caption file that has both short and long captions in it. It will + # Double up every image and run it through with both short and long captions. The idea + # is that the network will learn how to generate good images with both short and long captions + self.short_and_long_captions = kwargs.get('short_and_long_captions', False) + # if above is NOT true, this will make it so the long caption foes to te2 and the short caption goes to te1 for sdxl only + self.short_and_long_captions_encoder_split = kwargs.get('short_and_long_captions_encoder_split', False) + + # basically gradient accumulation but we run just 1 item through the network + # and accumulate gradients. This can be used as basic gradient accumulation but is very helpful + # for training tricks that increase batch size but need a single gradient step + self.single_item_batching = kwargs.get('single_item_batching', False) + + match_adapter_assist = kwargs.get('match_adapter_assist', False) + self.match_adapter_chance = kwargs.get('match_adapter_chance', 0.0) + self.loss_target: LossTarget = kwargs.get('loss_target', + 'noise') # noise, source, unaugmented, differential_noise + + # When a mask is passed in a dataset, and this is true, + # we will predict noise without a the LoRa network and use the prediction as a target for + # unmasked reign. It is unmasked regularization basically + self.inverted_mask_prior = kwargs.get('inverted_mask_prior', False) + self.inverted_mask_prior_multiplier = kwargs.get('inverted_mask_prior_multiplier', 0.5) + + # DOP will will run the same image and prompt through the network without the trigger word blank and use it as a target + self.diff_output_preservation = kwargs.get('diff_output_preservation', False) + self.diff_output_preservation_multiplier = kwargs.get('diff_output_preservation_multiplier', 1.0) + # If the trigger word is in the prompt, we will use this class name to replace it eg. "sks woman" -> "woman" + self.diff_output_preservation_class = kwargs.get('diff_output_preservation_class', '') + + # blank prompt preservation will preserve the model's knowledge of a blank prompt + self.blank_prompt_preservation = kwargs.get('blank_prompt_preservation', False) + self.blank_prompt_preservation_multiplier = kwargs.get('blank_prompt_preservation_multiplier', 1.0) + + # legacy + if match_adapter_assist and self.match_adapter_chance == 0.0: + self.match_adapter_chance = 1.0 + + # standardize inputs to the meand std of the model knowledge + self.standardize_images = kwargs.get('standardize_images', False) + self.standardize_latents = kwargs.get('standardize_latents', False) + + # if self.train_turbo and not self.noise_scheduler.startswith("euler"): + # raise ValueError(f"train_turbo is only supported with euler and wuler_a noise schedulers") + + self.dynamic_noise_offset = kwargs.get('dynamic_noise_offset', False) + self.do_cfg = kwargs.get('do_cfg', False) + self.do_random_cfg = kwargs.get('do_random_cfg', False) + self.cfg_scale = kwargs.get('cfg_scale', 1.0) + self.max_cfg_scale = kwargs.get('max_cfg_scale', self.cfg_scale) + self.cfg_rescale = kwargs.get('cfg_rescale', None) + if self.cfg_rescale is None: + self.cfg_rescale = self.cfg_scale + + # applies the inverse of the prediction mean and std to the target to correct + # for norm drift + self.correct_pred_norm = kwargs.get('correct_pred_norm', False) + self.correct_pred_norm_multiplier = kwargs.get('correct_pred_norm_multiplier', 1.0) + + self.loss_type = kwargs.get('loss_type', 'mse') # mse, mae, wavelet, pixelspace, mean_flow + + # scale the prediction by this. Increase for more detail, decrease for less + self.pred_scaler = kwargs.get('pred_scaler', 1.0) + + # repeats the prompt a few times to saturate the encoder + self.prompt_saturation_chance = kwargs.get('prompt_saturation_chance', 0.0) + + # applies negative loss on the prior to encourage network to diverge from it + self.do_prior_divergence = kwargs.get('do_prior_divergence', False) + + ema_config: Union[Dict, None] = kwargs.get('ema_config', None) + # if it is set explicitly to false, leave it false. + if ema_config is not None and ema_config.get('use_ema', False): + ema_config['use_ema'] = True + print(f"Using EMA") + else: + ema_config = {'use_ema': False} + + self.ema_config: EMAConfig = EMAConfig(**ema_config) + + # adds an additional loss to the network to encourage it output a normalized standard deviation + self.target_norm_std = kwargs.get('target_norm_std', None) + self.target_norm_std_value = kwargs.get('target_norm_std_value', 1.0) + self.timestep_type = kwargs.get('timestep_type', 'sigmoid') # sigmoid, linear, lognorm_blend, next_sample, weighted, one_step + self.next_sample_timesteps = kwargs.get('next_sample_timesteps', 8) + self.linear_timesteps = kwargs.get('linear_timesteps', False) + self.linear_timesteps2 = kwargs.get('linear_timesteps2', False) + self.disable_sampling = kwargs.get('disable_sampling', False) + + # will cache a blank prompt or the trigger word, and unload the text encoder to cpu + # will make training faster and use less vram + self.unload_text_encoder = kwargs.get('unload_text_encoder', False) + # will toggle all datasets to cache text embeddings + self.cache_text_embeddings: bool = kwargs.get('cache_text_embeddings', False) + # for swapping which parameters are trained during training + self.do_paramiter_swapping = kwargs.get('do_paramiter_swapping', False) + # 0.1 is 10% of the parameters active at a time lower is less vram, higher is more + self.paramiter_swapping_factor = kwargs.get('paramiter_swapping_factor', 0.1) + # bypass the guidance embedding for training. For open flux with guidance embedding + self.bypass_guidance_embedding = kwargs.get('bypass_guidance_embedding', False) + + # diffusion feature extractor + self.latent_feature_extractor_path = kwargs.get('latent_feature_extractor_path', None) + self.latent_feature_loss_weight = kwargs.get('latent_feature_loss_weight', 1.0) + + # we use this in the code, but it really needs to be called latent_feature_extractor as that makes more sense with new architecture + self.diffusion_feature_extractor_path = kwargs.get('diffusion_feature_extractor_path', self.latent_feature_extractor_path) + self.diffusion_feature_extractor_weight = kwargs.get('diffusion_feature_extractor_weight', self.latent_feature_loss_weight) + + # optimal noise pairing + self.optimal_noise_pairing_samples = kwargs.get('optimal_noise_pairing_samples', 1) + + # forces same noise for the same image at a given size. + self.force_consistent_noise = kwargs.get('force_consistent_noise', False) + self.blended_blur_noise = kwargs.get('blended_blur_noise', False) + + # contrastive loss + self.do_guidance_loss = kwargs.get('do_guidance_loss', False) + self.guidance_loss_target: Union[int, List[int, int]] = kwargs.get('guidance_loss_target', 3.0) + self.do_guidance_loss_cfg_zero: bool = kwargs.get('do_guidance_loss_cfg_zero', False) + self.unconditional_prompt: str = kwargs.get('unconditional_prompt', '') + if isinstance(self.guidance_loss_target, tuple): + self.guidance_loss_target = list(self.guidance_loss_target) + + self.do_differential_guidance = kwargs.get('do_differential_guidance', False) + self.differential_guidance_scale = kwargs.get('differential_guidance_scale', 3.0) + + # for multi stage models, how often to switch the boundary + self.switch_boundary_every: int = kwargs.get('switch_boundary_every', 1) + + +ModelArch = Literal['sd1', 'sd2', 'sd3', 'sdxl', 'pixart', 'pixart_sigma', 'auraflow', 'flux', 'flex1', 'flex2', 'lumina2', 'vega', 'ssd', 'wan21'] + + +class ModelConfig: + def __init__(self, **kwargs): + self.name_or_path: str = kwargs.get('name_or_path', None) + # name or path is updated on fine tuning. Keep a copy of the original + self.name_or_path_original: str = self.name_or_path + self.is_v2: bool = kwargs.get('is_v2', False) + self.is_xl: bool = kwargs.get('is_xl', False) + self.is_pixart: bool = kwargs.get('is_pixart', False) + self.is_pixart_sigma: bool = kwargs.get('is_pixart_sigma', False) + self.is_auraflow: bool = kwargs.get('is_auraflow', False) + self.is_v3: bool = kwargs.get('is_v3', False) + self.is_flux: bool = kwargs.get('is_flux', False) + self.is_lumina2: bool = kwargs.get('is_lumina2', False) + if self.is_pixart_sigma: + self.is_pixart = True + self.use_flux_cfg = kwargs.get('use_flux_cfg', False) + self.is_ssd: bool = kwargs.get('is_ssd', False) + self.is_vega: bool = kwargs.get('is_vega', False) + self.is_v_pred: bool = kwargs.get('is_v_pred', False) + self.dtype: str = kwargs.get('dtype', 'float16') + self.vae_path = kwargs.get('vae_path', None) + self.refiner_name_or_path = kwargs.get('refiner_name_or_path', None) + self._original_refiner_name_or_path = self.refiner_name_or_path + self.refiner_start_at = kwargs.get('refiner_start_at', 0.5) + self.lora_path = kwargs.get('lora_path', None) + # mainly for decompression loras for distilled models + self.assistant_lora_path = kwargs.get('assistant_lora_path', None) + self.inference_lora_path = kwargs.get('inference_lora_path', None) + self.latent_space_version = kwargs.get('latent_space_version', None) + + # only for SDXL models for now + self.use_text_encoder_1: bool = kwargs.get('use_text_encoder_1', True) + self.use_text_encoder_2: bool = kwargs.get('use_text_encoder_2', True) + + self.experimental_xl: bool = kwargs.get('experimental_xl', False) + + if self.name_or_path is None: + raise ValueError('name_or_path must be specified') + + if self.is_ssd: + # sed sdxl as true since it is mostly the same architecture + self.is_xl = True + + if self.is_vega: + self.is_xl = True + + # for text encoder quant. Only works with pixart currently + self.text_encoder_bits = kwargs.get('text_encoder_bits', 16) # 16, 8, 4 + self.unet_path = kwargs.get("unet_path", None) + self.unet_sample_size = kwargs.get("unet_sample_size", None) + self.vae_device = kwargs.get("vae_device", None) + self.vae_dtype = kwargs.get("vae_dtype", self.dtype) + self.te_device = kwargs.get("te_device", None) + self.te_dtype = kwargs.get("te_dtype", self.dtype) + + # only for flux for now + self.quantize = kwargs.get("quantize", False) + self.quantize_te = kwargs.get("quantize_te", self.quantize) + self.qtype = kwargs.get("qtype", "qfloat8") + self.qtype_te = kwargs.get("qtype_te", "qfloat8") + self.low_vram = kwargs.get("low_vram", False) + self.attn_masking = kwargs.get("attn_masking", False) + if self.attn_masking and not self.is_flux: + raise ValueError("attn_masking is only supported with flux models currently") + # for targeting a specific layers + self.ignore_if_contains: Optional[List[str]] = kwargs.get("ignore_if_contains", None) + self.only_if_contains: Optional[List[str]] = kwargs.get("only_if_contains", None) + self.quantize_kwargs = kwargs.get("quantize_kwargs", {}) + + # splits the model over the available gpus WIP + self.split_model_over_gpus = kwargs.get("split_model_over_gpus", False) + if self.split_model_over_gpus and not self.is_flux: + raise ValueError("split_model_over_gpus is only supported with flux models currently") + self.split_model_other_module_param_count_scale = kwargs.get("split_model_other_module_param_count_scale", 0.3) + + self.te_name_or_path = kwargs.get("te_name_or_path", None) + + self.arch: ModelArch = kwargs.get("arch", None) + + # auto memory management, only for some models + self.auto_memory = kwargs.get("auto_memory", False) + # auto memory is deprecated, use layer offloading instead + if self.auto_memory: + print("auto_memory is deprecated, use layer_offloading instead") + self.layer_offloading = kwargs.get("layer_offloading", self.auto_memory ) + if self.layer_offloading and self.qtype == "qfloat8": + self.qtype = "float8" + if self.layer_offloading and self.qtype_te == "qfloat8": + self.qtype_te = "float8" + + # 0 is off and 1.0 is 100% of the layers + self.layer_offloading_transformer_percent = kwargs.get("layer_offloading_transformer_percent", 1.0) + self.layer_offloading_text_encoder_percent = kwargs.get("layer_offloading_text_encoder_percent", 1.0) + + # can be used to load the extras like text encoder or vae from here + # only setup for some models but will prevent having to download the te for + # 20 different model variants + self.extras_name_or_path = kwargs.get("extras_name_or_path", self.name_or_path) + + # path to an accuracy recovery adapter, either local or remote + self.accuracy_recovery_adapter = kwargs.get("accuracy_recovery_adapter", None) + + # parse ARA from qtype + if self.qtype is not None and "|" in self.qtype: + self.qtype, self.accuracy_recovery_adapter = self.qtype.split('|') + + # compile the model with torch compile + self.compile = kwargs.get("compile", False) + + # kwargs to pass to the model + self.model_kwargs = kwargs.get("model_kwargs", {}) + + # allow frontend to pass arch with a color like arch:tag + # but remove the tag + if self.arch is not None: + if ':' in self.arch: + self.arch = self.arch.split(':')[0] + + if self.arch == "flex1": + self.arch = "flux" + + + # handle migrating to new model arch + if self.arch is not None: + # reverse the arch to the old style + if self.arch == 'sd2': + self.is_v2 = True + elif self.arch == 'sd3': + self.is_v3 = True + elif self.arch == 'sdxl': + self.is_xl = True + elif self.arch == 'pixart': + self.is_pixart = True + elif self.arch == 'pixart_sigma': + self.is_pixart_sigma = True + elif self.arch == 'auraflow': + self.is_auraflow = True + elif self.arch == 'flux': + self.is_flux = True + elif self.arch == 'lumina2': + self.is_lumina2 = True + elif self.arch == 'vega': + self.is_vega = True + elif self.arch == 'ssd': + self.is_ssd = True + else: + pass + if self.arch is None: + if kwargs.get('is_v2', False): + self.arch = 'sd2' + elif kwargs.get('is_v3', False): + self.arch = 'sd3' + elif kwargs.get('is_xl', False): + self.arch = 'sdxl' + elif kwargs.get('is_pixart', False): + self.arch = 'pixart' + elif kwargs.get('is_pixart_sigma', False): + self.arch = 'pixart_sigma' + elif kwargs.get('is_auraflow', False): + self.arch = 'auraflow' + elif kwargs.get('is_flux', False): + self.arch = 'flux' + elif kwargs.get('is_lumina2', False): + self.arch = 'lumina2' + elif kwargs.get('is_vega', False): + self.arch = 'vega' + elif kwargs.get('is_ssd', False): + self.arch = 'ssd' + else: + self.arch = 'sd1' + + + +class EMAConfig: + def __init__(self, **kwargs): + self.use_ema: bool = kwargs.get('use_ema', False) + self.ema_decay: float = kwargs.get('ema_decay', 0.999) + # feeds back the decay difference into the parameter + self.use_feedback: bool = kwargs.get('use_feedback', False) + + # every update, the params are multiplied by this amount + # only use for things without a bias like lora + # similar to a decay in an optimizer but the opposite + self.param_multiplier: float = kwargs.get('param_multiplier', 1.0) + + +class ReferenceDatasetConfig: + def __init__(self, **kwargs): + # can pass with a side by side pait or a folder with pos and neg folder + self.pair_folder: str = kwargs.get('pair_folder', None) + self.pos_folder: str = kwargs.get('pos_folder', None) + self.neg_folder: str = kwargs.get('neg_folder', None) + + self.network_weight: float = float(kwargs.get('network_weight', 1.0)) + self.pos_weight: float = float(kwargs.get('pos_weight', self.network_weight)) + self.neg_weight: float = float(kwargs.get('neg_weight', self.network_weight)) + # make sure they are all absolute values no negatives + self.pos_weight = abs(self.pos_weight) + self.neg_weight = abs(self.neg_weight) + + self.target_class: str = kwargs.get('target_class', '') + self.size: int = kwargs.get('size', 512) + + +class SliderTargetConfig: + def __init__(self, **kwargs): + self.target_class: str = kwargs.get('target_class', '') + self.positive: str = kwargs.get('positive', '') + self.negative: str = kwargs.get('negative', '') + self.multiplier: float = kwargs.get('multiplier', 1.0) + self.weight: float = kwargs.get('weight', 1.0) + self.shuffle: bool = kwargs.get('shuffle', False) + + +class GuidanceConfig: + def __init__(self, **kwargs): + self.target_class: str = kwargs.get('target_class', '') + self.guidance_scale: float = kwargs.get('guidance_scale', 1.0) + self.positive_prompt: str = kwargs.get('positive_prompt', '') + self.negative_prompt: str = kwargs.get('negative_prompt', '') + + +class SliderConfigAnchors: + def __init__(self, **kwargs): + self.prompt = kwargs.get('prompt', '') + self.neg_prompt = kwargs.get('neg_prompt', '') + self.multiplier = kwargs.get('multiplier', 1.0) + + +class SliderConfig: + def __init__(self, **kwargs): + targets = kwargs.get('targets', []) + anchors = kwargs.get('anchors', []) + anchors = [SliderConfigAnchors(**anchor) for anchor in anchors] + self.anchors: List[SliderConfigAnchors] = anchors + self.resolutions: List[List[int]] = kwargs.get('resolutions', [[512, 512]]) + self.prompt_file: str = kwargs.get('prompt_file', None) + self.prompt_tensors: str = kwargs.get('prompt_tensors', None) + self.batch_full_slide: bool = kwargs.get('batch_full_slide', True) + self.use_adapter: bool = kwargs.get('use_adapter', None) # depth + self.adapter_img_dir = kwargs.get('adapter_img_dir', None) + self.low_ram = kwargs.get('low_ram', False) + + # expand targets if shuffling + from toolkit.prompt_utils import get_slider_target_permutations + self.targets: List[SliderTargetConfig] = [] + targets = [SliderTargetConfig(**target) for target in targets] + # do permutations if shuffle is true + print(f"Building slider targets") + for target in targets: + if target.shuffle: + target_permutations = get_slider_target_permutations(target, max_permutations=8) + self.targets = self.targets + target_permutations + else: + self.targets.append(target) + print(f"Built {len(self.targets)} slider targets (with permutations)") + +ControlTypes = Literal['depth', 'line', 'pose', 'inpaint', 'mask'] + +class DatasetConfig: + """ + Dataset config for sd-datasets + + """ + + def __init__(self, **kwargs): + self.type = kwargs.get('type', 'image') # sd, slider, reference + # will be legacy + self.folder_path: str = kwargs.get('folder_path', None) + # can be json or folder path + self.dataset_path: str = kwargs.get('dataset_path', None) + + self.default_caption: str = kwargs.get('default_caption', None) + # trigger word for just this dataset + self.trigger_word: str = kwargs.get('trigger_word', None) + random_triggers = kwargs.get('random_triggers', []) + # if they are a string, load them from a file + if isinstance(random_triggers, str) and os.path.exists(random_triggers): + with open(random_triggers, 'r') as f: + random_triggers = f.read().splitlines() + # remove empty lines + random_triggers = [line for line in random_triggers if line.strip() != ''] + self.random_triggers: List[str] = random_triggers + self.random_triggers_max: int = kwargs.get('random_triggers_max', 1) + self.caption_ext: str = kwargs.get('caption_ext', '.txt') + # if caption_ext doesnt start with a dot, add it + if self.caption_ext and not self.caption_ext.startswith('.'): + self.caption_ext = '.' + self.caption_ext + self.random_scale: bool = kwargs.get('random_scale', False) + self.random_crop: bool = kwargs.get('random_crop', False) + self.resolution: int = kwargs.get('resolution', 512) + self.scale: float = kwargs.get('scale', 1.0) + self.buckets: bool = kwargs.get('buckets', True) + self.bucket_tolerance: int = kwargs.get('bucket_tolerance', 64) + self.is_reg: bool = kwargs.get('is_reg', False) + self.prior_reg: bool = kwargs.get('prior_reg', False) + self.network_weight: float = float(kwargs.get('network_weight', 1.0)) + self.token_dropout_rate: float = float(kwargs.get('token_dropout_rate', 0.0)) + self.shuffle_tokens: bool = kwargs.get('shuffle_tokens', False) + self.caption_dropout_rate: float = float(kwargs.get('caption_dropout_rate', 0.0)) + self.keep_tokens: int = kwargs.get('keep_tokens', 0) # #of first tokens to always keep unless caption dropped + self.flip_x: bool = kwargs.get('flip_x', False) + self.flip_y: bool = kwargs.get('flip_y', False) + self.augments: List[str] = kwargs.get('augments', []) + self.control_path: Union[str,List[str]] = kwargs.get('control_path', None) # depth maps, etc + if self.control_path == '': + self.control_path = None + + # handle multi control inputs from the ui. It is just easier to handle it here for a cleaner ui experience + control_path_1 = kwargs.get('control_path_1', None) + control_path_2 = kwargs.get('control_path_2', None) + control_path_3 = kwargs.get('control_path_3', None) + + if any([control_path_1, control_path_2, control_path_3]): + control_paths = [] + if control_path_1: + control_paths.append(control_path_1) + if control_path_2: + control_paths.append(control_path_2) + if control_path_3: + control_paths.append(control_path_3) + self.control_path = control_paths + + # color for transparent reigon of control images with transparency + self.control_transparent_color: List[int] = kwargs.get('control_transparent_color', [0, 0, 0]) + # inpaint images should be webp/png images with alpha channel. The alpha 0 (invisible) section will + # be the part conditioned to be inpainted. The alpha 1 (visible) section will be the part that is ignored + self.inpaint_path: Union[str,List[str]] = kwargs.get('inpaint_path', None) + # instead of cropping ot match image, it will serve the full size control image (clip images ie for ip adapters) + self.full_size_control_images: bool = kwargs.get('full_size_control_images', True) + self.alpha_mask: bool = kwargs.get('alpha_mask', False) # if true, will use alpha channel as mask + self.mask_path: str = kwargs.get('mask_path', + None) # focus mask (black and white. White has higher loss than black) + self.unconditional_path: str = kwargs.get('unconditional_path', + None) # path where matching unconditional images are located + self.invert_mask: bool = kwargs.get('invert_mask', False) # invert mask + self.mask_min_value: float = kwargs.get('mask_min_value', 0.0) # min value for . 0 - 1 + self.poi: Union[str, None] = kwargs.get('poi', + None) # if one is set and in json data, will be used as auto crop scale point of interes + self.use_short_captions: bool = kwargs.get('use_short_captions', False) # if true, will use 'caption_short' from json + self.num_repeats: int = kwargs.get('num_repeats', 1) # number of times to repeat dataset + # cache latents will store them in memory + self.cache_latents: bool = kwargs.get('cache_latents', False) + # cache latents to disk will store them on disk. If both are true, it will save to disk, but keep in memory + self.cache_latents_to_disk: bool = kwargs.get('cache_latents_to_disk', False) + self.cache_clip_vision_to_disk: bool = kwargs.get('cache_clip_vision_to_disk', False) + self.cache_text_embeddings: bool = kwargs.get('cache_text_embeddings', False) + + self.standardize_images: bool = kwargs.get('standardize_images', False) + + # https://albumentations.ai/docs/api_reference/augmentations/transforms + # augmentations are returned as a separate image and cannot currently be cached + self.augmentations: List[dict] = kwargs.get('augmentations', None) + self.shuffle_augmentations: bool = kwargs.get('shuffle_augmentations', False) + + has_augmentations = self.augmentations is not None and len(self.augmentations) > 0 + + if (len(self.augments) > 0 or has_augmentations) and (self.cache_latents or self.cache_latents_to_disk): + print(f"WARNING: Augments are not supported with caching latents. Setting cache_latents to False") + self.cache_latents = False + self.cache_latents_to_disk = False + + # legacy compatability + legacy_caption_type = kwargs.get('caption_type', None) + if legacy_caption_type: + self.caption_ext = legacy_caption_type + self.caption_type = self.caption_ext + self.guidance_type: GuidanceType = kwargs.get('guidance_type', 'targeted') + + # ip adapter / reference dataset + self.clip_image_path: str = kwargs.get('clip_image_path', None) # depth maps, etc + # get the clip image randomly from the same folder as the image. Useful for folder grouped pairs. + self.clip_image_from_same_folder: bool = kwargs.get('clip_image_from_same_folder', False) + self.clip_image_augmentations: List[dict] = kwargs.get('clip_image_augmentations', None) + self.clip_image_shuffle_augmentations: bool = kwargs.get('clip_image_shuffle_augmentations', False) + self.replacements: List[str] = kwargs.get('replacements', []) + self.loss_multiplier: float = kwargs.get('loss_multiplier', 1.0) + + self.num_workers: int = kwargs.get('num_workers', 2) + self.prefetch_factor: int = kwargs.get('prefetch_factor', 2) + self.extra_values: List[float] = kwargs.get('extra_values', []) + self.square_crop: bool = kwargs.get('square_crop', False) + # apply same augmentations to control images. Usually want this true unless special case + self.replay_transforms: bool = kwargs.get('replay_transforms', True) + + # for video + # if num_frames is greater than 1, the dataloader will look for video files. + # num_frames will be the number of frames in the training batch. If num_frames is 1, it will look for images + self.num_frames: int = kwargs.get('num_frames', 1) + # if true, will shrink video to our frames. For instance, if we have a video with 100 frames and num_frames is 10, + # we would pull frame 0, 10, 20, 30, 40, 50, 60, 70, 80, 90 so they are evenly spaced + self.shrink_video_to_frames: bool = kwargs.get('shrink_video_to_frames', True) + # fps is only used if shrink_video_to_frames is false. This will attempt to pull the num_frames at the given fps + # it will select a random start frame and pull the frames at the given fps + # this could have various issues with shorter videos and videos with variable fps + # I recommend trimming your videos to the desired length and using shrink_video_to_frames(default) + self.fps: int = kwargs.get('fps', 16) + + # debug the frame count and frame selection. You dont need this. It is for debugging. + self.debug: bool = kwargs.get('debug', False) + + # automatic controls + self.controls: List[ControlTypes] = kwargs.get('controls', []) + if isinstance(self.controls, str): + self.controls = [self.controls] + # remove empty strings + self.controls = [control for control in self.controls if control.strip() != ''] + + # if true, will use a fask method to get image sizes. This can result in errors. Do not use unless you know what you are doing + self.fast_image_size: bool = kwargs.get('fast_image_size', False) + + self.do_i2v: bool = kwargs.get('do_i2v', True) # do image to video on models that are both t2i and i2v capable + + +def preprocess_dataset_raw_config(raw_config: List[dict]) -> List[dict]: + """ + This just splits up the datasets by resolutions so you dont have to do it manually + :param raw_config: + :return: + """ + # split up datasets by resolutions + new_config = [] + for dataset in raw_config: + resolution = dataset.get('resolution', 512) + if isinstance(resolution, list): + resolution_list = resolution + else: + resolution_list = [resolution] + for res in resolution_list: + dataset_copy = dataset.copy() + dataset_copy['resolution'] = res + new_config.append(dataset_copy) + return new_config + + +class GenerateImageConfig: + def __init__( + self, + prompt: str = '', + prompt_2: Optional[str] = None, + width: int = 512, + height: int = 512, + num_inference_steps: int = 50, + guidance_scale: float = 7.5, + negative_prompt: str = '', + negative_prompt_2: Optional[str] = None, + seed: int = -1, + network_multiplier: float = 1.0, + guidance_rescale: float = 0.0, + # the tag [time] will be replaced with milliseconds since epoch + output_path: str = None, # full image path + output_folder: str = None, # folder to save image in if output_path is not specified + output_ext: str = ImgExt, # extension to save image as if output_path is not specified + output_tail: str = '', # tail to add to output filename + add_prompt_file: bool = False, # add a prompt file with generated image + adapter_image_path: str = None, # path to adapter image + adapter_conditioning_scale: float = 1.0, # scale for adapter conditioning + latents: Union[torch.Tensor | None] = None, # input latent to start with, + extra_kwargs: dict = None, # extra data to save with prompt file + refiner_start_at: float = 0.5, # start at this percentage of a step. 0.0 to 1.0 . 1.0 is the end + extra_values: List[float] = None, # extra values to save with prompt file + logger: Optional[EmptyLogger] = None, + ctrl_img: Optional[str] = None, # control image for controlnet + ctrl_img_1: Optional[str] = None, # first control image for multi control model + ctrl_img_2: Optional[str] = None, # second control image for multi control model + ctrl_img_3: Optional[str] = None, # third control image for multi control model + num_frames: int = 1, + fps: int = 15, + ctrl_idx: int = 0, + do_cfg_norm: bool = False, + ): + self.width: int = width + self.height: int = height + self.num_inference_steps: int = num_inference_steps + self.guidance_scale: float = guidance_scale + self.guidance_rescale: float = guidance_rescale + self.prompt: str = prompt + self.prompt_2: str = prompt_2 + self.negative_prompt: str = negative_prompt + self.negative_prompt_2: str = negative_prompt_2 + self.latents: Union[torch.Tensor | None] = latents + + self.output_path: str = output_path + self.seed: int = seed + if self.seed == -1: + # generate random one + self.seed = random.randint(0, 2 ** 32 - 1) + self.network_multiplier: float = network_multiplier + self.output_folder: str = output_folder + self.output_ext: str = output_ext + self.add_prompt_file: bool = add_prompt_file + self.output_tail: str = output_tail + self.gen_time: int = int(time.time() * 1000) + self.adapter_image_path: str = adapter_image_path + self.adapter_conditioning_scale: float = adapter_conditioning_scale + self.extra_kwargs = extra_kwargs if extra_kwargs is not None else {} + self.refiner_start_at = refiner_start_at + self.extra_values = extra_values if extra_values is not None else [] + self.num_frames = num_frames + self.fps = fps + self.ctrl_img = ctrl_img + self.ctrl_idx = ctrl_idx + + if ctrl_img_1 is None and ctrl_img is not None: + ctrl_img_1 = ctrl_img + + self.ctrl_img_1 = ctrl_img_1 + self.ctrl_img_2 = ctrl_img_2 + self.ctrl_img_3 = ctrl_img_3 + + # prompt string will override any settings above + self._process_prompt_string() + + # handle dual text encoder prompts if nothing passed + if negative_prompt_2 is None: + self.negative_prompt_2 = negative_prompt + + if prompt_2 is None: + self.prompt_2 = self.prompt + + # parse prompt paths + if self.output_path is None and self.output_folder is None: + raise ValueError('output_path or output_folder must be specified') + elif self.output_path is not None: + self.output_folder = os.path.dirname(self.output_path) + self.output_ext = os.path.splitext(self.output_path)[1][1:] + self.output_filename_no_ext = os.path.splitext(os.path.basename(self.output_path))[0] + + else: + self.output_filename_no_ext = '[time]_[count]' + if len(self.output_tail) > 0: + self.output_filename_no_ext += '_' + self.output_tail + self.output_path = os.path.join(self.output_folder, self.output_filename_no_ext + '.' + self.output_ext) + + # adjust height + self.height = max(64, self.height - self.height % 8) # round to divisible by 8 + self.width = max(64, self.width - self.width % 8) # round to divisible by 8 + + self.logger = logger + + self.do_cfg_norm: bool = do_cfg_norm + + def set_gen_time(self, gen_time: int = None): + if gen_time is not None: + self.gen_time = gen_time + else: + self.gen_time = int(time.time() * 1000) + + def _get_path_no_ext(self, count: int = 0, max_count=0): + # zero pad count + count_str = str(count).zfill(len(str(max_count))) + # replace [time] with gen time + filename = self.output_filename_no_ext.replace('[time]', str(self.gen_time)) + # replace [count] with count + filename = filename.replace('[count]', count_str) + return filename + + def get_image_path(self, count: int = 0, max_count=0): + filename = self._get_path_no_ext(count, max_count) + ext = self.output_ext + # if it does not start with a dot add one + if ext[0] != '.': + ext = '.' + ext + filename += ext + # join with folder + return os.path.join(self.output_folder, filename) + + def get_prompt_path(self, count: int = 0, max_count=0): + filename = self._get_path_no_ext(count, max_count) + filename += '.txt' + # join with folder + return os.path.join(self.output_folder, filename) + + def save_image(self, image, count: int = 0, max_count=0): + # make parent dirs + os.makedirs(self.output_folder, exist_ok=True) + self.set_gen_time() + if isinstance(image, list): + # video + if self.num_frames == 1: + raise ValueError(f"Expected 1 img but got a list {len(image)}") + if self.num_frames > 1 and self.output_ext not in ['webp']: + self.output_ext = 'webp' + if self.output_ext == 'webp': + # save as animated webp + duration = 1000 // self.fps # Convert fps to milliseconds per frame + image[0].save( + self.get_image_path(count, max_count), + format='WEBP', + append_images=image[1:], + save_all=True, + duration=duration, # Duration per frame in milliseconds + loop=0, # 0 means loop forever + quality=80 # Quality setting (0-100) + ) + else: + raise ValueError(f"Unsupported video format {self.output_ext}") + elif self.output_ext in ['wav', 'mp3']: + # save audio file + torchaudio.save( + self.get_image_path(count, max_count), + image[0].to('cpu'), + sample_rate=48000, + format=None, + backend=None + ) + else: + # TODO save image gen header info for A1111 and us, our seeds probably wont match + image.save(self.get_image_path(count, max_count)) + # do prompt file + if self.add_prompt_file: + self.save_prompt_file(count, max_count) + + def save_prompt_file(self, count: int = 0, max_count=0): + # save prompt file + with open(self.get_prompt_path(count, max_count), 'w') as f: + prompt = self.prompt + if self.prompt_2 is not None: + prompt += ' --p2 ' + self.prompt_2 + if self.negative_prompt is not None: + prompt += ' --n ' + self.negative_prompt + if self.negative_prompt_2 is not None: + prompt += ' --n2 ' + self.negative_prompt_2 + prompt += ' --w ' + str(self.width) + prompt += ' --h ' + str(self.height) + prompt += ' --seed ' + str(self.seed) + prompt += ' --cfg ' + str(self.guidance_scale) + prompt += ' --steps ' + str(self.num_inference_steps) + prompt += ' --m ' + str(self.network_multiplier) + prompt += ' --gr ' + str(self.guidance_rescale) + + # get gen info + try: + f.write(self.prompt) + except Exception as e: + print(f"Error writing prompt file. Prompt contains non-unicode characters. {e}") + + def _process_prompt_string(self): + # we will try to support all sd-scripts where we can + + # FROM SD-SCRIPTS + # --n Treat everything until the next option as a negative prompt. + # --w Specify the width of the generated image. + # --h Specify the height of the generated image. + # --d Specify the seed for the generated image. + # --l Specify the CFG scale for the generated image. + # --s Specify the number of steps during generation. + + # OURS and some QOL additions + # --m Specify the network multiplier for the generated image. + # --p2 Prompt for the second text encoder (SDXL only) + # --n2 Negative prompt for the second text encoder (SDXL only) + # --gr Specify the guidance rescale for the generated image (SDXL only) + + # --seed Specify the seed for the generated image same as --d + # --cfg Specify the CFG scale for the generated image same as --l + # --steps Specify the number of steps during generation same as --s + # --network_multiplier Specify the network multiplier for the generated image same as --m + + # process prompt string and update values if it has some + if self.prompt is not None and len(self.prompt) > 0: + # process prompt string + prompt = self.prompt + prompt = prompt.strip() + p_split = prompt.split('--') + self.prompt = p_split[0].strip() + + if len(p_split) > 1: + for split in p_split[1:]: + # allows multi char flags + flag = split.split(' ')[0].strip() + content = split[len(flag):].strip() + if flag == 'p2': + self.prompt_2 = content + elif flag == 'n': + self.negative_prompt = content + elif flag == 'n2': + self.negative_prompt_2 = content + elif flag == 'w': + self.width = int(content) + elif flag == 'h': + self.height = int(content) + elif flag == 'd': + self.seed = int(content) + elif flag == 'seed': + self.seed = int(content) + elif flag == 'l': + self.guidance_scale = float(content) + elif flag == 'cfg': + self.guidance_scale = float(content) + elif flag == 's': + self.num_inference_steps = int(content) + elif flag == 'steps': + self.num_inference_steps = int(content) + elif flag == 'm': + self.network_multiplier = float(content) + elif flag == 'network_multiplier': + self.network_multiplier = float(content) + elif flag == 'gr': + self.guidance_rescale = float(content) + elif flag == 'a': + self.adapter_conditioning_scale = float(content) + elif flag == 'ref': + self.refiner_start_at = float(content) + elif flag == 'ev': + # split by comma + self.extra_values = [float(val) for val in content.split(',')] + elif flag == 'extra_values': + # split by comma + self.extra_values = [float(val) for val in content.split(',')] + elif flag == 'frames': + self.num_frames = int(content) + elif flag == 'num_frames': + self.num_frames = int(content) + elif flag == 'fps': + self.fps = int(content) + elif flag == 'ctrl_img': + self.ctrl_img = content + elif flag == 'ctrl_idx': + self.ctrl_idx = int(content) + + def post_process_embeddings( + self, + conditional_prompt_embeds: PromptEmbeds, + unconditional_prompt_embeds: Optional[PromptEmbeds] = None, + ): + # this is called after prompt embeds are encoded. We can override them in the future here + pass + + def log_image(self, image, count: int = 0, max_count=0): + if self.logger is None: + return + + self.logger.log_image(image, count, self.prompt) + + +def validate_configs( + train_config: TrainConfig, + model_config: ModelConfig, + save_config: SaveConfig, + dataset_configs: List[DatasetConfig] +): + if model_config.is_flux: + if save_config.save_format != 'diffusers': + # make it diffusers + save_config.save_format = 'diffusers' + if model_config.use_flux_cfg: + # bypass the embedding + train_config.bypass_guidance_embedding = True + if train_config.bypass_guidance_embedding and train_config.do_guidance_loss: + raise ValueError("Cannot bypass guidance embedding and do guidance loss at the same time. " + "Please set bypass_guidance_embedding to False or do_guidance_loss to False.") + + if model_config.accuracy_recovery_adapter is not None: + if model_config.assistant_lora_path is not None: + raise ValueError("Cannot use accuracy recovery adapter and assistant lora at the same time. " + "Please set one of them to None.") + + # see if any datasets are caching text embeddings + is_caching_text_embeddings = any(dataset.cache_text_embeddings for dataset in dataset_configs) + if is_caching_text_embeddings: + + # check if they are doing differential output preservation + if train_config.diff_output_preservation: + raise ValueError("Cannot use differential output preservation with caching text embeddings. Please set diff_output_preservation to False.") + + # make sure they are all cached + for dataset in dataset_configs: + if not dataset.cache_text_embeddings: + raise ValueError("All datasets must have cache_text_embeddings set to True when caching text embeddings is enabled.") + + # qwen image edit cannot cache text embeddings + if model_config.arch == 'qwen_image_edit': + if train_config.unload_text_encoder: + raise ValueError("Cannot cache unload text encoder with qwen_image_edit model. Control images are encoded with text embeddings. You can cache the text embeddings though") + + if train_config.diff_output_preservation and train_config.blank_prompt_preservation: + raise ValueError("Cannot use both differential output preservation and blank prompt preservation at the same time. Please set one of them to False.") + + diff --git a/toolkit/kohya_lora.py b/toolkit/kohya_lora.py new file mode 100644 index 0000000..d53179a --- /dev/null +++ b/toolkit/kohya_lora.py @@ -0,0 +1,1221 @@ +# LoRA network module +# reference: +# https://github.com/microsoft/LoRA/blob/main/loralib/layers.py +# https://github.com/cloneofsimo/lora/blob/master/lora_diffusion/lora.py + +# taken from kohya lora sd scripts + +import math +import os +from typing import Dict, List, Optional, Tuple, Type, Union +from diffusers import AutoencoderKL +from transformers import CLIPTextModel +import numpy as np +import torch +import re + + +RE_UPDOWN = re.compile(r"(up|down)_blocks_(\d+)_(resnets|upsamplers|downsamplers|attentions)_(\d+)_") + + +class LoRAModule(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, + ): + """if alpha == 0 or None, alpha is rank (no scaling).""" + super().__init__() + self.lora_name = lora_name + + if org_module.__class__.__name__ == "Conv2d": + 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 + + if org_module.__class__.__name__ == "Conv2d": + kernel_size = org_module.kernel_size + stride = org_module.stride + padding = org_module.padding + 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=False) + 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=False) + + 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)) + torch.nn.init.zeros_(self.lora_up.weight) + + self.multiplier = multiplier + self.org_module = org_module # remove in applying + self.dropout = dropout + self.rank_dropout = rank_dropout + self.module_dropout = module_dropout + + def apply_to(self): + self.org_forward = self.org_module.forward + self.org_module.forward = self.forward + del self.org_module + + def forward(self, x): + org_forwarded = self.org_forward(x) + + # module dropout + if self.module_dropout is not None and self.training: + if torch.rand(1) < self.module_dropout: + return org_forwarded + + lx = self.lora_down(x) + + # normal dropout + if 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.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) + + return org_forwarded + lx * self.multiplier * scale + + +class LoRAInfModule(LoRAModule): + def __init__( + self, + lora_name, + org_module: torch.nn.Module, + multiplier=1.0, + lora_dim=4, + alpha=1, + **kwargs, + ): + # no dropout for inference + super().__init__(lora_name, org_module, multiplier, lora_dim, alpha) + + self.org_module_ref = [org_module] # 後から参照できるように + self.enabled = True + + # check regional or not by lora_name + self.text_encoder = False + if lora_name.startswith("lora_te_"): + self.regional = False + self.use_sub_prompt = True + self.text_encoder = True + elif "attn2_to_k" in lora_name or "attn2_to_v" in lora_name: + self.regional = False + self.use_sub_prompt = True + elif "time_emb" in lora_name: + self.regional = False + self.use_sub_prompt = False + else: + self.regional = True + self.use_sub_prompt = False + + self.network: LoRANetwork = None + + def set_network(self, network): + self.network = network + + # freezeしてマージする + def merge_to(self, sd, dtype, device): + # get up/down weight + up_weight = sd["lora_up.weight"].to(torch.float).to(device) + down_weight = sd["lora_down.weight"].to(torch.float).to(device) + + # extract weight from org_module + org_sd = self.org_module.state_dict() + weight = org_sd["weight"].to(torch.float) + + # merge weight + if len(weight.size()) == 2: + # linear + weight = weight + self.multiplier * (up_weight @ down_weight) * self.scale + elif down_weight.size()[2:4] == (1, 1): + # conv2d 1x1 + weight = ( + weight + + self.multiplier + * (up_weight.squeeze(3).squeeze(2) @ down_weight.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3) + * self.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 + self.multiplier * conved * self.scale + + # set weight to org_module + org_sd["weight"] = weight.to(dtype) + self.org_module.load_state_dict(org_sd) + + # 復元できるマージのため、このモジュールのweightを返す + def get_weight(self, multiplier=None): + if multiplier is None: + multiplier = self.multiplier + + # get up/down weight from module + up_weight = self.lora_up.weight.to(torch.float) + down_weight = self.lora_down.weight.to(torch.float) + + # pre-calculated weight + if len(down_weight.size()) == 2: + # linear + weight = self.multiplier * (up_weight @ down_weight) * self.scale + elif down_weight.size()[2:4] == (1, 1): + # conv2d 1x1 + weight = ( + self.multiplier + * (up_weight.squeeze(3).squeeze(2) @ down_weight.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3) + * self.scale + ) + else: + # conv2d 3x3 + conved = torch.nn.functional.conv2d(down_weight.permute(1, 0, 2, 3), up_weight).permute(1, 0, 2, 3) + weight = self.multiplier * conved * self.scale + + return weight + + def set_region(self, region): + self.region = region + self.region_mask = None + + def default_forward(self, x): + # print("default_forward", self.lora_name, x.size()) + return self.org_forward(x) + self.lora_up(self.lora_down(x)) * self.multiplier * self.scale + + def forward(self, x): + if not self.enabled: + return self.org_forward(x) + + if self.network is None or self.network.sub_prompt_index is None: + return self.default_forward(x) + if not self.regional and not self.use_sub_prompt: + return self.default_forward(x) + + if self.regional: + return self.regional_forward(x) + else: + return self.sub_prompt_forward(x) + + def get_mask_for_x(self, x): + # calculate size from shape of x + if len(x.size()) == 4: + h, w = x.size()[2:4] + area = h * w + else: + area = x.size()[1] + + mask = self.network.mask_dic[area] + if mask is None: + raise ValueError(f"mask is None for resolution {area}") + if len(x.size()) != 4: + mask = torch.reshape(mask, (1, -1, 1)) + return mask + + def regional_forward(self, x): + if "attn2_to_out" in self.lora_name: + return self.to_out_forward(x) + + if self.network.mask_dic is None: # sub_prompt_index >= 3 + return self.default_forward(x) + + # apply mask for LoRA result + lx = self.lora_up(self.lora_down(x)) * self.multiplier * self.scale + mask = self.get_mask_for_x(lx) + # print("regional", self.lora_name, self.network.sub_prompt_index, lx.size(), mask.size()) + lx = lx * mask + + x = self.org_forward(x) + x = x + lx + + if "attn2_to_q" in self.lora_name and self.network.is_last_network: + x = self.postp_to_q(x) + + return x + + def postp_to_q(self, x): + # repeat x to num_sub_prompts + has_real_uncond = x.size()[0] // self.network.batch_size == 3 + qc = self.network.batch_size # uncond + qc += self.network.batch_size * self.network.num_sub_prompts # cond + if has_real_uncond: + qc += self.network.batch_size # real_uncond + + query = torch.zeros((qc, x.size()[1], x.size()[2]), device=x.device, dtype=x.dtype) + query[: self.network.batch_size] = x[: self.network.batch_size] + + for i in range(self.network.batch_size): + qi = self.network.batch_size + i * self.network.num_sub_prompts + query[qi : qi + self.network.num_sub_prompts] = x[self.network.batch_size + i] + + if has_real_uncond: + query[-self.network.batch_size :] = x[-self.network.batch_size :] + + # print("postp_to_q", self.lora_name, x.size(), query.size(), self.network.num_sub_prompts) + return query + + def sub_prompt_forward(self, x): + if x.size()[0] == self.network.batch_size: # if uncond in text_encoder, do not apply LoRA + return self.org_forward(x) + + emb_idx = self.network.sub_prompt_index + if not self.text_encoder: + emb_idx += self.network.batch_size + + # apply sub prompt of X + lx = x[emb_idx :: self.network.num_sub_prompts] + lx = self.lora_up(self.lora_down(lx)) * self.multiplier * self.scale + + # print("sub_prompt_forward", self.lora_name, x.size(), lx.size(), emb_idx) + + x = self.org_forward(x) + x[emb_idx :: self.network.num_sub_prompts] += lx + + return x + + def to_out_forward(self, x): + # print("to_out_forward", self.lora_name, x.size(), self.network.is_last_network) + + if self.network.is_last_network: + masks = [None] * self.network.num_sub_prompts + self.network.shared[self.lora_name] = (None, masks) + else: + lx, masks = self.network.shared[self.lora_name] + + # call own LoRA + x1 = x[self.network.batch_size + self.network.sub_prompt_index :: self.network.num_sub_prompts] + lx1 = self.lora_up(self.lora_down(x1)) * self.multiplier * self.scale + + if self.network.is_last_network: + lx = torch.zeros( + (self.network.num_sub_prompts * self.network.batch_size, *lx1.size()[1:]), device=lx1.device, dtype=lx1.dtype + ) + self.network.shared[self.lora_name] = (lx, masks) + + # print("to_out_forward", lx.size(), lx1.size(), self.network.sub_prompt_index, self.network.num_sub_prompts) + lx[self.network.sub_prompt_index :: self.network.num_sub_prompts] += lx1 + masks[self.network.sub_prompt_index] = self.get_mask_for_x(lx1) + + # if not last network, return x and masks + x = self.org_forward(x) + if not self.network.is_last_network: + return x + + lx, masks = self.network.shared.pop(self.lora_name) + + # if last network, combine separated x with mask weighted sum + has_real_uncond = x.size()[0] // self.network.batch_size == self.network.num_sub_prompts + 2 + + out = torch.zeros((self.network.batch_size * (3 if has_real_uncond else 2), *x.size()[1:]), device=x.device, dtype=x.dtype) + out[: self.network.batch_size] = x[: self.network.batch_size] # uncond + if has_real_uncond: + out[-self.network.batch_size :] = x[-self.network.batch_size :] # real_uncond + + # print("to_out_forward", self.lora_name, self.network.sub_prompt_index, self.network.num_sub_prompts) + # for i in range(len(masks)): + # if masks[i] is None: + # masks[i] = torch.zeros_like(masks[-1]) + + mask = torch.cat(masks) + mask_sum = torch.sum(mask, dim=0) + 1e-4 + for i in range(self.network.batch_size): + # 1枚の画像ごとに処理する + lx1 = lx[i * self.network.num_sub_prompts : (i + 1) * self.network.num_sub_prompts] + lx1 = lx1 * mask + lx1 = torch.sum(lx1, dim=0) + + xi = self.network.batch_size + i * self.network.num_sub_prompts + x1 = x[xi : xi + self.network.num_sub_prompts] + x1 = x1 * mask + x1 = torch.sum(x1, dim=0) + x1 = x1 / mask_sum + + x1 = x1 + lx1 + out[self.network.batch_size + i] = x1 + + # print("to_out_forward", x.size(), out.size(), has_real_uncond) + return out + + +def parse_block_lr_kwargs(nw_kwargs): + down_lr_weight = nw_kwargs.get("down_lr_weight", None) + mid_lr_weight = nw_kwargs.get("mid_lr_weight", None) + up_lr_weight = nw_kwargs.get("up_lr_weight", None) + + # 以上のいずれにも設定がない場合は無効としてNoneを返す + if down_lr_weight is None and mid_lr_weight is None and up_lr_weight is None: + return None, None, None + + # extract learning rate weight for each block + if down_lr_weight is not None: + # if some parameters are not set, use zero + if "," in down_lr_weight: + down_lr_weight = [(float(s) if s else 0.0) for s in down_lr_weight.split(",")] + + if mid_lr_weight is not None: + mid_lr_weight = float(mid_lr_weight) + + if up_lr_weight is not None: + if "," in up_lr_weight: + up_lr_weight = [(float(s) if s else 0.0) for s in up_lr_weight.split(",")] + + down_lr_weight, mid_lr_weight, up_lr_weight = get_block_lr_weight( + down_lr_weight, mid_lr_weight, up_lr_weight, float(nw_kwargs.get("block_lr_zero_threshold", 0.0)) + ) + + return down_lr_weight, mid_lr_weight, up_lr_weight + + +def create_network( + multiplier: float, + network_dim: Optional[int], + network_alpha: Optional[float], + vae: AutoencoderKL, + text_encoder: Union[CLIPTextModel, List[CLIPTextModel]], + unet, + neuron_dropout: Optional[float] = None, + **kwargs, +): + if network_dim is None: + network_dim = 4 # default + if network_alpha is None: + network_alpha = 1.0 + + # extract dim/alpha for conv2d, and block dim + conv_dim = kwargs.get("conv_dim", None) + conv_alpha = kwargs.get("conv_alpha", None) + if conv_dim is not None: + conv_dim = int(conv_dim) + if conv_alpha is None: + conv_alpha = 1.0 + else: + conv_alpha = float(conv_alpha) + + # block dim/alpha/lr + block_dims = kwargs.get("block_dims", None) + down_lr_weight, mid_lr_weight, up_lr_weight = parse_block_lr_kwargs(kwargs) + + # 以上のいずれかに指定があればblockごとのdim(rank)を有効にする + if block_dims is not None or down_lr_weight is not None or mid_lr_weight is not None or up_lr_weight is not None: + block_alphas = kwargs.get("block_alphas", None) + conv_block_dims = kwargs.get("conv_block_dims", None) + conv_block_alphas = kwargs.get("conv_block_alphas", None) + + block_dims, block_alphas, conv_block_dims, conv_block_alphas = get_block_dims_and_alphas( + block_dims, block_alphas, network_dim, network_alpha, conv_block_dims, conv_block_alphas, conv_dim, conv_alpha + ) + + # remove block dim/alpha without learning rate + block_dims, block_alphas, conv_block_dims, conv_block_alphas = remove_block_dims_and_alphas( + block_dims, block_alphas, conv_block_dims, conv_block_alphas, down_lr_weight, mid_lr_weight, up_lr_weight + ) + + else: + block_alphas = None + conv_block_dims = None + conv_block_alphas = None + + # rank/module dropout + rank_dropout = kwargs.get("rank_dropout", None) + if rank_dropout is not None: + rank_dropout = float(rank_dropout) + module_dropout = kwargs.get("module_dropout", None) + if module_dropout is not None: + module_dropout = float(module_dropout) + + # すごく引数が多いな ( ^ω^)・・・ + network = LoRANetwork( + text_encoder, + unet, + multiplier=multiplier, + lora_dim=network_dim, + alpha=network_alpha, + dropout=neuron_dropout, + rank_dropout=rank_dropout, + module_dropout=module_dropout, + conv_lora_dim=conv_dim, + conv_alpha=conv_alpha, + block_dims=block_dims, + block_alphas=block_alphas, + conv_block_dims=conv_block_dims, + conv_block_alphas=conv_block_alphas, + varbose=True, + ) + + if up_lr_weight is not None or mid_lr_weight is not None or down_lr_weight is not None: + network.set_block_lr_weight(up_lr_weight, mid_lr_weight, down_lr_weight) + + return network + + +# このメソッドは外部から呼び出される可能性を考慮しておく +# network_dim, network_alpha にはデフォルト値が入っている。 +# block_dims, block_alphas は両方ともNoneまたは両方とも値が入っている +# conv_dim, conv_alpha は両方ともNoneまたは両方とも値が入っている +def get_block_dims_and_alphas( + block_dims, block_alphas, network_dim, network_alpha, conv_block_dims, conv_block_alphas, conv_dim, conv_alpha +): + num_total_blocks = LoRANetwork.NUM_OF_BLOCKS * 2 + 1 + + def parse_ints(s): + return [int(i) for i in s.split(",")] + + def parse_floats(s): + return [float(i) for i in s.split(",")] + + # block_dimsとblock_alphasをパースする。必ず値が入る + if block_dims is not None: + block_dims = parse_ints(block_dims) + assert ( + len(block_dims) == num_total_blocks + ), f"block_dims must have {num_total_blocks} elements / block_dimsは{num_total_blocks}個指定してください" + else: + print(f"block_dims is not specified. all dims are set to {network_dim} / block_dimsが指定されていません。すべてのdimは{network_dim}になります") + block_dims = [network_dim] * num_total_blocks + + if block_alphas is not None: + block_alphas = parse_floats(block_alphas) + assert ( + len(block_alphas) == num_total_blocks + ), f"block_alphas must have {num_total_blocks} elements / block_alphasは{num_total_blocks}個指定してください" + else: + print( + f"block_alphas is not specified. all alphas are set to {network_alpha} / block_alphasが指定されていません。すべてのalphaは{network_alpha}になります" + ) + block_alphas = [network_alpha] * num_total_blocks + + # conv_block_dimsとconv_block_alphasを、指定がある場合のみパースする。指定がなければconv_dimとconv_alphaを使う + if conv_block_dims is not None: + conv_block_dims = parse_ints(conv_block_dims) + assert ( + len(conv_block_dims) == num_total_blocks + ), f"conv_block_dims must have {num_total_blocks} elements / conv_block_dimsは{num_total_blocks}個指定してください" + + if conv_block_alphas is not None: + conv_block_alphas = parse_floats(conv_block_alphas) + assert ( + len(conv_block_alphas) == num_total_blocks + ), f"conv_block_alphas must have {num_total_blocks} elements / conv_block_alphasは{num_total_blocks}個指定してください" + else: + if conv_alpha is None: + conv_alpha = 1.0 + print( + f"conv_block_alphas is not specified. all alphas are set to {conv_alpha} / conv_block_alphasが指定されていません。すべてのalphaは{conv_alpha}になります" + ) + conv_block_alphas = [conv_alpha] * num_total_blocks + else: + if conv_dim is not None: + print( + f"conv_dim/alpha for all blocks are set to {conv_dim} and {conv_alpha} / すべてのブロックのconv_dimとalphaは{conv_dim}および{conv_alpha}になります" + ) + conv_block_dims = [conv_dim] * num_total_blocks + conv_block_alphas = [conv_alpha] * num_total_blocks + else: + conv_block_dims = None + conv_block_alphas = None + + return block_dims, block_alphas, conv_block_dims, conv_block_alphas + + +# 層別学習率用に層ごとの学習率に対する倍率を定義する、外部から呼び出される可能性を考慮しておく +def get_block_lr_weight( + down_lr_weight, mid_lr_weight, up_lr_weight, zero_threshold +) -> Tuple[List[float], List[float], List[float]]: + # パラメータ未指定時は何もせず、今までと同じ動作とする + if up_lr_weight is None and mid_lr_weight is None and down_lr_weight is None: + return None, None, None + + max_len = LoRANetwork.NUM_OF_BLOCKS # フルモデル相当でのup,downの層の数 + + def get_list(name_with_suffix) -> List[float]: + import math + + tokens = name_with_suffix.split("+") + name = tokens[0] + base_lr = float(tokens[1]) if len(tokens) > 1 else 0.0 + + if name == "cosine": + return [math.sin(math.pi * (i / (max_len - 1)) / 2) + base_lr for i in reversed(range(max_len))] + elif name == "sine": + return [math.sin(math.pi * (i / (max_len - 1)) / 2) + base_lr for i in range(max_len)] + elif name == "linear": + return [i / (max_len - 1) + base_lr for i in range(max_len)] + elif name == "reverse_linear": + return [i / (max_len - 1) + base_lr for i in reversed(range(max_len))] + elif name == "zeros": + return [0.0 + base_lr] * max_len + else: + print( + "Unknown lr_weight argument %s is used. Valid arguments: / 不明なlr_weightの引数 %s が使われました。有効な引数:\n\tcosine, sine, linear, reverse_linear, zeros" + % (name) + ) + return None + + if type(down_lr_weight) == str: + down_lr_weight = get_list(down_lr_weight) + if type(up_lr_weight) == str: + up_lr_weight = get_list(up_lr_weight) + + if (up_lr_weight != None and len(up_lr_weight) > max_len) or (down_lr_weight != None and len(down_lr_weight) > max_len): + print("down_weight or up_weight is too long. Parameters after %d-th are ignored." % max_len) + print("down_weightもしくはup_weightが長すぎます。%d個目以降のパラメータは無視されます。" % max_len) + up_lr_weight = up_lr_weight[:max_len] + down_lr_weight = down_lr_weight[:max_len] + + if (up_lr_weight != None and len(up_lr_weight) < max_len) or (down_lr_weight != None and len(down_lr_weight) < max_len): + print("down_weight or up_weight is too short. Parameters after %d-th are filled with 1." % max_len) + print("down_weightもしくはup_weightが短すぎます。%d個目までの不足したパラメータは1で補われます。" % max_len) + + if down_lr_weight != None and len(down_lr_weight) < max_len: + down_lr_weight = down_lr_weight + [1.0] * (max_len - len(down_lr_weight)) + if up_lr_weight != None and len(up_lr_weight) < max_len: + up_lr_weight = up_lr_weight + [1.0] * (max_len - len(up_lr_weight)) + + if (up_lr_weight != None) or (mid_lr_weight != None) or (down_lr_weight != None): + print("apply block learning rate / 階層別学習率を適用します。") + if down_lr_weight != None: + down_lr_weight = [w if w > zero_threshold else 0 for w in down_lr_weight] + print("down_lr_weight (shallower -> deeper, 浅い層->深い層):", down_lr_weight) + else: + print("down_lr_weight: all 1.0, すべて1.0") + + if mid_lr_weight != None: + mid_lr_weight = mid_lr_weight if mid_lr_weight > zero_threshold else 0 + print("mid_lr_weight:", mid_lr_weight) + else: + print("mid_lr_weight: 1.0") + + if up_lr_weight != None: + up_lr_weight = [w if w > zero_threshold else 0 for w in up_lr_weight] + print("up_lr_weight (deeper -> shallower, 深い層->浅い層):", up_lr_weight) + else: + print("up_lr_weight: all 1.0, すべて1.0") + + return down_lr_weight, mid_lr_weight, up_lr_weight + + +# lr_weightが0のblockをblock_dimsから除外する、外部から呼び出す可能性を考慮しておく +def remove_block_dims_and_alphas( + block_dims, block_alphas, conv_block_dims, conv_block_alphas, down_lr_weight, mid_lr_weight, up_lr_weight +): + # set 0 to block dim without learning rate to remove the block + if down_lr_weight != None: + for i, lr in enumerate(down_lr_weight): + if lr == 0: + block_dims[i] = 0 + if conv_block_dims is not None: + conv_block_dims[i] = 0 + if mid_lr_weight != None: + if mid_lr_weight == 0: + block_dims[LoRANetwork.NUM_OF_BLOCKS] = 0 + if conv_block_dims is not None: + conv_block_dims[LoRANetwork.NUM_OF_BLOCKS] = 0 + if up_lr_weight != None: + for i, lr in enumerate(up_lr_weight): + if lr == 0: + block_dims[LoRANetwork.NUM_OF_BLOCKS + 1 + i] = 0 + if conv_block_dims is not None: + conv_block_dims[LoRANetwork.NUM_OF_BLOCKS + 1 + i] = 0 + + return block_dims, block_alphas, conv_block_dims, conv_block_alphas + + +# 外部から呼び出す可能性を考慮しておく +def get_block_index(lora_name: str) -> int: + block_idx = -1 # invalid lora name + + m = RE_UPDOWN.search(lora_name) + if m: + g = m.groups() + i = int(g[1]) + j = int(g[3]) + if g[2] == "resnets": + idx = 3 * i + j + elif g[2] == "attentions": + idx = 3 * i + j + elif g[2] == "upsamplers" or g[2] == "downsamplers": + idx = 3 * i + 2 + + if g[0] == "down": + block_idx = 1 + idx # 0に該当するLoRAは存在しない + elif g[0] == "up": + block_idx = LoRANetwork.NUM_OF_BLOCKS + 1 + idx + + elif "mid_block_" in lora_name: + block_idx = LoRANetwork.NUM_OF_BLOCKS # idx=12 + + return block_idx + + +# Create network from weights for inference, weights are not loaded here (because can be merged) +def create_network_from_weights(multiplier, file, vae, text_encoder, unet, weights_sd=None, for_inference=False, **kwargs): + if weights_sd is None: + if os.path.splitext(file)[1] == ".safetensors": + from safetensors.torch import load_file, safe_open + + weights_sd = load_file(file) + else: + weights_sd = torch.load(file, map_location="cpu") + + # get dim/alpha mapping + modules_dim = {} + modules_alpha = {} + for key, value in weights_sd.items(): + if "." not in key: + continue + + lora_name = key.split(".")[0] + if "alpha" in key: + modules_alpha[lora_name] = value + elif "lora_down" in key: + dim = value.size()[0] + modules_dim[lora_name] = dim + # print(lora_name, value.size(), dim) + + # support old LoRA without alpha + for key in modules_dim.keys(): + if key not in modules_alpha: + modules_alpha[key] = modules_dim[key] + + module_class = LoRAInfModule if for_inference else LoRAModule + + network = LoRANetwork( + text_encoder, unet, multiplier=multiplier, modules_dim=modules_dim, modules_alpha=modules_alpha, module_class=module_class + ) + + # block lr + down_lr_weight, mid_lr_weight, up_lr_weight = parse_block_lr_kwargs(kwargs) + if up_lr_weight is not None or mid_lr_weight is not None or down_lr_weight is not None: + network.set_block_lr_weight(up_lr_weight, mid_lr_weight, down_lr_weight) + + return network, weights_sd + + +class LoRANetwork(torch.nn.Module): + NUM_OF_BLOCKS = 12 # フルモデル相当でのup,downの層の数 + + UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel"] + UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["ResnetBlock2D", "Downsample2D", "Upsample2D"] + TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPMLP"] + LORA_PREFIX_UNET = "lora_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, + ) -> 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を指定 (推論用) + """ + super().__init__() + self.multiplier = multiplier + + 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 + + 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]: + prefix = ( + self.LORA_PREFIX_UNET + 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 = [] + 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__ == "Linear" + is_conv2d = child_module.__class__.__name__ == "Conv2d" + is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1) + + if is_linear or is_conv2d: + lora_name = prefix + "." + name + "." + child_name + lora_name = lora_name.replace(".", "_") + + 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] + elif is_unet and block_dims is not None: + # U-Netでblock_dims指定あり + block_idx = get_block_index(lora_name) + if is_linear or is_conv2d_1x1: + dim = block_dims[block_idx] + alpha = block_alphas[block_idx] + elif conv_block_dims is not None: + dim = conv_block_dims[block_idx] + alpha = conv_block_alphas[block_idx] + 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 + + lora = module_class( + lora_name, + child_module, + self.multiplier, + dim, + alpha, + dropout=dropout, + rank_dropout=rank_dropout, + module_dropout=module_dropout, + ) + loras.append(lora) + 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 = [] + for i, text_encoder in enumerate(text_encoders): + 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:") + + text_encoder_loras, skipped = create_modules(False, index, text_encoder, LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE) + 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 = LoRANetwork.UNET_TARGET_REPLACE_MODULE + if modules_dim is not None or self.conv_lora_dim is not None or conv_block_dims is not None: + target_modules += LoRANetwork.UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 + + self.unet_loras, skipped_un = create_modules(True, None, unet, target_modules) + 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) + + def set_multiplier(self, multiplier): + self.multiplier = multiplier + for lora in self.text_encoder_loras + self.unet_loras: + lora.multiplier = self.multiplier + + def load_weights(self, file): + 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") + + info = self.load_state_dict(weights_sd, False) + return info + + def apply_to(self, text_encoder, unet, apply_text_encoder=True, apply_unet=True): + if apply_text_encoder: + print("enable LoRA for text encoder") + else: + self.text_encoder_loras = [] + + if apply_unet: + print("enable LoRA for U-Net") + else: + self.unet_loras = [] + + for lora in self.text_encoder_loras + self.unet_loras: + lora.apply_to() + self.add_module(lora.lora_name, lora) + + # マージできるかどうかを返す + def is_mergeable(self): + return True + + # TODO refactor to common function with apply_to + def merge_to(self, text_encoder, unet, weights_sd, dtype, device): + apply_text_encoder = apply_unet = False + for key in weights_sd.keys(): + if key.startswith(LoRANetwork.LORA_PREFIX_TEXT_ENCODER): + apply_text_encoder = True + elif key.startswith(LoRANetwork.LORA_PREFIX_UNET): + apply_unet = True + + if apply_text_encoder: + print("enable LoRA for text encoder") + else: + self.text_encoder_loras = [] + + if apply_unet: + print("enable LoRA for U-Net") + else: + self.unet_loras = [] + + for lora in self.text_encoder_loras + self.unet_loras: + sd_for_lora = {} + for key in weights_sd.keys(): + if key.startswith(lora.lora_name): + sd_for_lora[key[len(lora.lora_name) + 1 :]] = weights_sd[key] + lora.merge_to(sd_for_lora, dtype, device) + + print(f"weights are merged") + + # 層別学習率用に層ごとの学習率に対する倍率を定義する 引数の順番が逆だがとりあえず気にしない + def set_block_lr_weight( + self, + up_lr_weight: List[float] = None, + mid_lr_weight: float = None, + down_lr_weight: List[float] = None, + ): + self.block_lr = True + self.down_lr_weight = down_lr_weight + self.mid_lr_weight = mid_lr_weight + self.up_lr_weight = up_lr_weight + + def get_lr_weight(self, lora: LoRAModule) -> float: + lr_weight = 1.0 + block_idx = get_block_index(lora.lora_name) + if block_idx < 0: + return lr_weight + + if block_idx < LoRANetwork.NUM_OF_BLOCKS: + if self.down_lr_weight != None: + lr_weight = self.down_lr_weight[block_idx] + elif block_idx == LoRANetwork.NUM_OF_BLOCKS: + if self.mid_lr_weight != None: + lr_weight = self.mid_lr_weight + elif block_idx > LoRANetwork.NUM_OF_BLOCKS: + if self.up_lr_weight != None: + lr_weight = self.up_lr_weight[block_idx - LoRANetwork.NUM_OF_BLOCKS - 1] + + return lr_weight + + # 二つのText Encoderに別々の学習率を設定できるようにするといいかも + def prepare_optimizer_params(self, text_encoder_lr, unet_lr, default_lr): + self.requires_grad_(True) + all_params = [] + + def enumerate_params(loras): + params = [] + for lora in loras: + params.extend(lora.parameters()) + return params + + if self.text_encoder_loras: + param_data = {"params": enumerate_params(self.text_encoder_loras)} + if text_encoder_lr is not None: + param_data["lr"] = text_encoder_lr + all_params.append(param_data) + + if self.unet_loras: + if self.block_lr: + # 学習率のグラフをblockごとにしたいので、blockごとにloraを分類 + block_idx_to_lora = {} + for lora in self.unet_loras: + idx = get_block_index(lora.lora_name) + if idx not in block_idx_to_lora: + block_idx_to_lora[idx] = [] + block_idx_to_lora[idx].append(lora) + + # blockごとにパラメータを設定する + for idx, block_loras in block_idx_to_lora.items(): + param_data = {"params": enumerate_params(block_loras)} + + if unet_lr is not None: + param_data["lr"] = unet_lr * self.get_lr_weight(block_loras[0]) + elif default_lr is not None: + param_data["lr"] = default_lr * self.get_lr_weight(block_loras[0]) + if ("lr" in param_data) and (param_data["lr"] == 0): + continue + all_params.append(param_data) + + else: + param_data = {"params": enumerate_params(self.unet_loras)} + if unet_lr is not None: + param_data["lr"] = unet_lr + all_params.append(param_data) + + return all_params + + def enable_gradient_checkpointing(self): + # not supported + pass + + def prepare_grad_etc(self, text_encoder, unet): + self.requires_grad_(True) + + def on_epoch_start(self, text_encoder, unet): + self.train() + + def get_trainable_params(self): + return self.parameters() + + def save_weights(self, file, dtype, metadata): + if metadata is not None and len(metadata) == 0: + metadata = None + + state_dict = self.state_dict() + + if dtype is not None: + for key in list(state_dict.keys()): + v = state_dict[key] + v = v.detach().clone().to("cpu").to(dtype) + state_dict[key] = v + + if os.path.splitext(file)[1] == ".safetensors": + from safetensors.torch import save_file + + # Precalculate model hashes to save time on indexing + if metadata is None: + metadata = {} + # model_hash, legacy_hash = train_util.precalculate_safetensors_hashes(state_dict, metadata) + # metadata["sshs_model_hash"] = model_hash + # metadata["sshs_legacy_hash"] = legacy_hash + + save_file(state_dict, file, metadata) + else: + torch.save(state_dict, file) + + # mask is a tensor with values from 0 to 1 + def set_region(self, sub_prompt_index, is_last_network, mask): + if mask.max() == 0: + mask = torch.ones_like(mask) + + self.mask = mask + self.sub_prompt_index = sub_prompt_index + self.is_last_network = is_last_network + + for lora in self.text_encoder_loras + self.unet_loras: + lora.set_network(self) + + def set_current_generation(self, batch_size, num_sub_prompts, width, height, shared): + self.batch_size = batch_size + self.num_sub_prompts = num_sub_prompts + self.current_size = (height, width) + self.shared = shared + + # create masks + mask = self.mask + mask_dic = {} + mask = mask.unsqueeze(0).unsqueeze(1) # b(1),c(1),h,w + ref_weight = self.text_encoder_loras[0].lora_down.weight if self.text_encoder_loras else self.unet_loras[0].lora_down.weight + dtype = ref_weight.dtype + device = ref_weight.device + + def resize_add(mh, mw): + # print(mh, mw, mh * mw) + m = torch.nn.functional.interpolate(mask, (mh, mw), mode="bilinear") # doesn't work in bf16 + m = m.to(device, dtype=dtype) + mask_dic[mh * mw] = m + + h = height // 8 + w = width // 8 + for _ in range(4): + resize_add(h, w) + if h % 2 == 1 or w % 2 == 1: # add extra shape if h/w is not divisible by 2 + resize_add(h + h % 2, w + w % 2) + h = (h + 1) // 2 + w = (w + 1) // 2 + + self.mask_dic = mask_dic + + def backup_weights(self): + # 重みのバックアップを行う + loras: List[LoRAInfModule] = self.text_encoder_loras + self.unet_loras + for lora in loras: + org_module = lora.org_module_ref[0] + if not hasattr(org_module, "_lora_org_weight"): + sd = org_module.state_dict() + org_module._lora_org_weight = sd["weight"].detach().clone() + org_module._lora_restored = True + + def restore_weights(self): + # 重みのリストアを行う + loras: List[LoRAInfModule] = self.text_encoder_loras + self.unet_loras + for lora in loras: + org_module = lora.org_module_ref[0] + if not org_module._lora_restored: + sd = org_module.state_dict() + sd["weight"] = org_module._lora_org_weight + org_module.load_state_dict(sd) + org_module._lora_restored = True + + def pre_calculation(self): + # 事前計算を行う + loras: List[LoRAInfModule] = self.text_encoder_loras + self.unet_loras + for lora in loras: + org_module = lora.org_module_ref[0] + sd = org_module.state_dict() + + org_weight = sd["weight"] + lora_weight = lora.get_weight().to(org_weight.device, dtype=org_weight.dtype) + sd["weight"] = org_weight + lora_weight + assert sd["weight"].shape == org_weight.shape + org_module.load_state_dict(sd) + + org_module._lora_restored = False + lora.enabled = False + + def apply_max_norm_regularization(self, max_norm_value, device): + downkeys = [] + upkeys = [] + alphakeys = [] + norms = [] + keys_scaled = 0 + + state_dict = self.state_dict() + for key in state_dict.keys(): + if "lora_down" in key and "weight" in key: + downkeys.append(key) + upkeys.append(key.replace("lora_down", "lora_up")) + alphakeys.append(key.replace("lora_down.weight", "alpha")) + + for i in range(len(downkeys)): + down = state_dict[downkeys[i]].to(device) + up = state_dict[upkeys[i]].to(device) + alpha = state_dict[alphakeys[i]].to(device) + dim = down.shape[0] + scale = alpha / dim + + if up.shape[2:] == (1, 1) and down.shape[2:] == (1, 1): + updown = (up.squeeze(2).squeeze(2) @ down.squeeze(2).squeeze(2)).unsqueeze(2).unsqueeze(3) + elif up.shape[2:] == (3, 3) or down.shape[2:] == (3, 3): + updown = torch.nn.functional.conv2d(down.permute(1, 0, 2, 3), up).permute(1, 0, 2, 3) + else: + updown = up @ down + + updown *= scale + + norm = updown.norm().clamp(min=max_norm_value / 2) + desired = torch.clamp(norm, max=max_norm_value) + ratio = desired.cpu() / norm.cpu() + sqrt_ratio = ratio**0.5 + if ratio != 1: + keys_scaled += 1 + state_dict[upkeys[i]] *= sqrt_ratio + state_dict[downkeys[i]] *= sqrt_ratio + scalednorm = updown.norm() * ratio + norms.append(scalednorm.item()) + + return keys_scaled, sum(norms) / len(norms), max(norms) diff --git a/toolkit/loha_network.py b/toolkit/loha_network.py new file mode 100644 index 0000000..2c08881 --- /dev/null +++ b/toolkit/loha_network.py @@ -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() \ No newline at end of file diff --git a/toolkit/lokr.py b/toolkit/lokr.py new file mode 100644 index 0000000..e953041 --- /dev/null +++ b/toolkit/lokr.py @@ -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) diff --git a/toolkit/lora_special.py b/toolkit/lora_special.py new file mode 100644 index 0000000..aec1678 --- /dev/null +++ b/toolkit/lora_special.py @@ -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 + diff --git a/toolkit/lorm.py b/toolkit/lorm.py new file mode 100644 index 0000000..de8f74b --- /dev/null +++ b/toolkit/lorm.py @@ -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 diff --git a/toolkit/lycoris_special.py b/toolkit/lycoris_special.py new file mode 100644 index 0000000..9c958a7 --- /dev/null +++ b/toolkit/lycoris_special.py @@ -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) diff --git a/toolkit/lycoris_utils.py b/toolkit/lycoris_utils.py new file mode 100644 index 0000000..526a28d --- /dev/null +++ b/toolkit/lycoris_utils.py @@ -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') diff --git a/toolkit/metadata.py b/toolkit/metadata.py new file mode 100644 index 0000000..6119126 --- /dev/null +++ b/toolkit/metadata.py @@ -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() diff --git a/toolkit/models/DoRA.py b/toolkit/models/DoRA.py new file mode 100644 index 0000000..c919a17 --- /dev/null +++ b/toolkit/models/DoRA.py @@ -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) diff --git a/toolkit/models/loha.py b/toolkit/models/loha.py new file mode 100644 index 0000000..3d55bdc --- /dev/null +++ b/toolkit/models/loha.py @@ -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] \ No newline at end of file diff --git a/toolkit/network_mixins.py b/toolkit/network_mixins.py new file mode 100644 index 0000000..224f0e3 --- /dev/null +++ b/toolkit/network_mixins.py @@ -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 diff --git a/toolkit/paths.py b/toolkit/paths.py new file mode 100644 index 0000000..87ef458 --- /dev/null +++ b/toolkit/paths.py @@ -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 diff --git a/toolkit/prompt_utils.py b/toolkit/prompt_utils.py new file mode 100644 index 0000000..8b00e81 --- /dev/null +++ b/toolkit/prompt_utils.py @@ -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 diff --git a/toolkit/requirements.txt b/toolkit/requirements.txt new file mode 100644 index 0000000..d26fea8 --- /dev/null +++ b/toolkit/requirements.txt @@ -0,0 +1,7 @@ +lycoris-lora +optimum-quanto +safetensors +diffusers +transformers +accelerate +huggingface-hub \ No newline at end of file diff --git a/toolkit/saving.py b/toolkit/saving.py new file mode 100644 index 0000000..f18ccff --- /dev/null +++ b/toolkit/saving.py @@ -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 diff --git a/toolkit/train_tools.py b/toolkit/train_tools.py new file mode 100644 index 0000000..e6a0715 --- /dev/null +++ b/toolkit/train_tools.py @@ -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) diff --git a/toolkit/z_image.py b/toolkit/z_image.py new file mode 100644 index 0000000..fd2679a --- /dev/null +++ b/toolkit/z_image.py @@ -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 \ No newline at end of file diff --git a/z_image_vector_merge.py b/z_image_vector_merge.py new file mode 100644 index 0000000..53334ed --- /dev/null +++ b/z_image_vector_merge.py @@ -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)" +} \ No newline at end of file