From 932d928847141727072caa74b4cb5fe91223d786 Mon Sep 17 00:00:00 2001 From: Jordan Thompson Date: Tue, 23 Dec 2025 23:51:14 -0800 Subject: [PATCH] Add Power Lora Merger --- __init__.py | 5 +- modules/__init__.py | 0 modules/cli/__init__.py | 0 modules/cli/lora_spec.py | 148 ++ modules/cli/merge_loras_z.py | 1973 ++++++++++++++++++ nodes/WASEdgeSafeLatentUpscale.py | 397 ++++ nodes/WASLatentContrastLimitedDetailBoost.py | 232 ++ nodes/WASPowerLoraMerger.py | 862 ++++++++ requirements.txt | 4 +- web/was_power_lora_merger.js | 719 +++++++ 10 files changed, 4338 insertions(+), 2 deletions(-) create mode 100644 modules/__init__.py create mode 100644 modules/cli/__init__.py create mode 100644 modules/cli/lora_spec.py create mode 100644 modules/cli/merge_loras_z.py create mode 100644 nodes/WASEdgeSafeLatentUpscale.py create mode 100644 nodes/WASLatentContrastLimitedDetailBoost.py create mode 100644 nodes/WASPowerLoraMerger.py create mode 100644 web/was_power_lora_merger.js diff --git a/__init__.py b/__init__.py index 7508f59..35fb8aa 100644 --- a/__init__.py +++ b/__init__.py @@ -2,6 +2,7 @@ import importlib import pkgutil import time import traceback +import os try: @@ -20,6 +21,8 @@ PREFIX = "[WAS Extras] " NODE_CLASS_MAPPINGS = {} NODE_DISPLAY_NAME_MAPPINGS = {} +WEB_DIRECTORY = os.path.join(os.path.dirname(__file__), "web") + class NodeLoader: def __init__(self, package_name: str, prefix: str = PREFIX): @@ -131,4 +134,4 @@ class NodeLoader: _loader = NodeLoader(package_name=__name__, prefix=PREFIX) _loader.load_all() -__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"] diff --git a/modules/__init__.py b/modules/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/modules/cli/__init__.py b/modules/cli/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/modules/cli/lora_spec.py b/modules/cli/lora_spec.py new file mode 100644 index 0000000..d731c25 --- /dev/null +++ b/modules/cli/lora_spec.py @@ -0,0 +1,148 @@ +import argparse +import os + +from collections import Counter +from safetensors.torch import safe_open + + +def format_bytes(num_bytes: int) -> str: + units = ["B", "KB", "MB", "GB", "TB"] + size = float(num_bytes) + for unit in units: + if size < 1024.0 or unit == units[-1]: + if unit == "B": + return f"{int(size)} {unit}" + return f"{size:.2f} {unit}" + size /= 1024.0 + return f"{num_bytes} B" + + +def estimate_tensor_bytes(tensor) -> int: + return int(tensor.numel()) * int(tensor.element_size()) + + +def common_prefix(a: str, b: str) -> str: + n = min(len(a), len(b)) + i = 0 + while i < n and a[i] == b[i]: + i += 1 + return a[:i] + + +def derive_schema_prefixes(keys, max_depth: int = 4) -> Counter: + separators = [".", "/", ":", "_"] + counts = Counter() + + for key in keys: + parts = [key] + for sep in separators: + new_parts = [] + for p in parts: + new_parts.extend(p.split(sep)) + parts = new_parts + + parts = [p for p in parts if p] + if not parts: + continue + + depth = min(max_depth, len(parts)) + prefix = ".".join(parts[:depth]) + counts[prefix] += 1 + + return counts + + +def inspect_and_write_report( + path: str, + limit: int = 30, + schema_depth: int = 4, + schema_top: int = 25, +) -> None: + if not os.path.isfile(path): + raise FileNotFoundError(path) + + base_name = os.path.splitext(os.path.basename(path))[0] + output_path = os.path.join( + os.path.dirname(path), + f"{base_name}.schematics.txt" + ) + + file_size = os.path.getsize(path) + + with open(output_path, "w", encoding="utf-8") as out: + out.write("Tensor Container Schematic Report\n") + out.write("=" * 80 + "\n\n") + out.write(f"Source file: {path}\n") + out.write(f"File size: {format_bytes(file_size)}\n\n") + + with safe_open(path, framework="pt", device="cpu") as f: + keys = list(f.keys()) + total = len(keys) + + out.write(f"Total tensors: {total}\n") + out.write(f"Showing first {min(limit, total)} entries\n\n") + + for idx, key in enumerate(keys[:limit], start=1): + tensor = f.get_tensor(key) + + shape = tuple(tensor.shape) + ndim = int(tensor.ndim) + numel = int(tensor.numel()) + dtype = str(tensor.dtype).replace("torch.", "") + bytes_est = estimate_tensor_bytes(tensor) + + out.write(f"{idx:02d}. {key}\n") + out.write(f" shape: {shape}\n") + out.write(f" ndim: {ndim}\n") + out.write(f" numel: {numel}\n") + out.write(f" dtype: {dtype}\n") + out.write(f" bytes: {format_bytes(bytes_est)}\n") + out.write("-" * 80 + "\n") + + if total > 0: + out.write("\nLexical Schema Summary (Purely Name-Based)\n") + out.write("=" * 80 + "\n") + out.write( + f"Grouping token depth: {schema_depth} | " + f"Top groups shown: {schema_top}\n\n" + ) + + counts = derive_schema_prefixes(keys, max_depth=schema_depth) + for prefix, count in counts.most_common(schema_top): + out.write(f"{count:6d} {prefix}\n") + + keys_sorted = sorted(keys) + shared = keys_sorted[0] + for k in keys_sorted[1:]: + shared = common_prefix(shared, k) + if not shared: + break + + out.write("\nCommon character prefix across all keys:\n") + out.write(shared if shared else "(none)") + out.write("\n") + + print(f"Report written to:\n {output_path}") + + +def main() -> None: + parser = argparse.ArgumentParser( + description="Generate a schematic report for a .safetensors file." + ) + parser.add_argument("path", type=str, help="Path to .safetensors file") + parser.add_argument("--limit", type=int, default=30, help="Number of keys to list (default: 30)") + parser.add_argument("--schema-depth", type=int, default=4, help="Token depth for schema grouping") + parser.add_argument("--schema-top", type=int, default=25, help="Number of schema groups to show") + + args = parser.parse_args() + + inspect_and_write_report( + path=args.path, + limit=args.limit, + schema_depth=args.schema_depth, + schema_top=args.schema_top, + ) + + +if __name__ == "__main__": + main() diff --git a/modules/cli/merge_loras_z.py b/modules/cli/merge_loras_z.py new file mode 100644 index 0000000..61db016 --- /dev/null +++ b/modules/cli/merge_loras_z.py @@ -0,0 +1,1973 @@ +#!/usr/bin/env python3 +import argparse +import json +import os +from dataclasses import dataclass +from typing import Any, Callable, Dict, List, Tuple, Optional + +import torch +from safetensors import safe_open +from safetensors.torch import save_file +from tqdm import tqdm + + +ProgressCallback = Callable[[str, int, int, Optional[str]], None] + + +def _progress(cb: Optional[ProgressCallback], stage: str, current: int, total: int, message: Optional[str] = None): + if cb is None: + return + try: + cb(stage, int(current), int(total), message) + except Exception: + pass + + +@dataclass +class LoraPair: + down: torch.Tensor + up: torch.Tensor + alpha: float + is_conv: bool + in_hw: Tuple[int, int] + + +def parse_weighted_paths(items: List[str]) -> List[Tuple[str, float]]: + result: List[Tuple[str, float]] = [] + for item in items: + if "@" in item: + path, weight_str = item.rsplit("@", 1) + weight = float(weight_str) + result.append((path, weight)) + else: + result.append((item, 1.0)) + return result + + +def get_device(force_cpu: bool) -> torch.device: + if force_cpu: + return torch.device("cpu") + if torch.cuda.is_available(): + return torch.device("cuda") + return torch.device("cpu") + + +def get_dtype(dtype_name: str) -> torch.dtype: + name = dtype_name.lower() + if name in ("fp16", "float16", "half"): + return torch.float16 + if name in ("fp32", "float32"): + return torch.float32 + if name in ("bf16", "bfloat16"): + return torch.bfloat16 + raise ValueError(f"Unsupported dtype: {dtype_name}") + + +def get_compute_dtype(dtype_name: str) -> Optional[torch.dtype]: + """ + Compute dtype is for internal merge math alignment. If None => auto behavior. + """ + name = dtype_name.lower() + if name in ("auto", "none"): + return None + return get_dtype(name) + + +def is_lora_key(key: str) -> bool: + return ( + key.endswith("lora_down.weight") + or key.endswith("lora_up.weight") + or key.endswith("lora_down") + or key.endswith("lora_up") + or key.endswith("lora_A.weight") + or key.endswith("lora_B.weight") + or key.endswith("lora_A") + or key.endswith("lora_B") + ) + + +def get_module_prefix(key: str) -> str: + if ".lora_down" in key: + return key.split(".lora_down", 1)[0] + if ".lora_up" in key: + return key.split(".lora_up", 1)[0] + if ".lora_A" in key: + return key.split(".lora_A", 1)[0] + if ".lora_B" in key: + return key.split(".lora_B", 1)[0] + parts = key.split(".") + return ".".join(parts[:-2]) + + +def get_alpha_key(prefix: str) -> str: + return f"{prefix}.alpha" + + +def _get_save_key_style(metadata_ref: Dict[str, str]) -> str: + style = (metadata_ref or {}).get("_was_lora_key_style", "") + style_norm = str(style).strip().lower() + if style_norm in ("ab", "a/b", "lora_a_b"): + return "ab" + return "downup" + + +def _down_key(prefix: str, key_style: str) -> str: + if key_style == "ab": + return f"{prefix}.lora_A.weight" + return f"{prefix}.lora_down.weight" + + +def _up_key(prefix: str, key_style: str) -> str: + if key_style == "ab": + return f"{prefix}.lora_B.weight" + return f"{prefix}.lora_up.weight" + + +def module_is_included( + prefix: str, + include_patterns: List[str], + exclude_patterns: List[str], +) -> bool: + if include_patterns: + if not any(pat in prefix for pat in include_patterns): + return False + if exclude_patterns: + if any(pat in prefix for pat in exclude_patterns): + return False + return True + + +def load_lora_pairs( + path: str, + device: torch.device, + progress_cb: Optional[ProgressCallback] = None, +) -> Tuple[Dict[str, LoraPair], Dict[str, str]]: + pairs: Dict[str, LoraPair] = {} + metadata: Dict[str, str] = {} + + with safe_open(path, framework="pt", device=str(device)) as f: + keys = list(f.keys()) + try: + raw_meta = f.metadata() + metadata = dict(raw_meta) if isinstance(raw_meta, dict) else {} + except Exception: + metadata = {} + + detected_key_style: Optional[str] = None + for k in keys: + if k.endswith("lora_A.weight") or k.endswith("lora_B.weight") or k.endswith("lora_A") or k.endswith("lora_B"): + detected_key_style = "ab" + break + if detected_key_style is None: + for k in keys: + if k.endswith("lora_down.weight") or k.endswith("lora_up.weight") or k.endswith("lora_down") or k.endswith("lora_up"): + detected_key_style = "downup" + break + if detected_key_style is not None: + metadata["_was_lora_key_style"] = detected_key_style + + downs: Dict[str, torch.Tensor] = {} + ups: Dict[str, torch.Tensor] = {} + alphas: Dict[str, float] = {} + + total_keys = len(keys) + for i, key in enumerate(keys, start=1): + _progress(progress_cb, "load.keys", i, total_keys, key) + if is_lora_key(key): + tensor = f.get_tensor(key).to(device) + prefix = get_module_prefix(key) + if "lora_down" in key or "lora_A" in key: + downs[prefix] = tensor + elif "lora_up" in key or "lora_B" in key: + ups[prefix] = tensor + elif key.endswith(".alpha"): + tensor = f.get_tensor(key) + value = float(tensor.item()) + try: + prefix = get_module_prefix(key.replace(".alpha", ".lora_up")) + alphas[prefix] = value + except Exception: + try: + prefix = get_module_prefix(key.replace(".alpha", ".lora_down")) + alphas[prefix] = value + except Exception: + pass + + module_keys = set(downs.keys()) | set(ups.keys()) + + module_keys_list = list(module_keys) + total_modules = len(module_keys_list) + for i, prefix in enumerate(module_keys_list, start=1): + _progress(progress_cb, "load.modules", i, total_modules, prefix) + if prefix not in downs or prefix not in ups: + continue + + down = downs[prefix] + up = ups[prefix] + + if down.ndim == 4 or up.ndim == 4: + is_conv = True + k_h = down.shape[2] if down.ndim == 4 else 1 + k_w = down.shape[3] if down.ndim == 4 else 1 + in_hw = (k_h, k_w) + else: + is_conv = False + in_hw = (1, 1) + + rank = down.shape[0] + alpha = float(alphas.get(prefix, rank)) + pairs[prefix] = LoraPair( + down=down, + up=up, + alpha=alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + + return pairs, metadata + + +def _is_float_tensor(t: torch.Tensor) -> bool: + return t.is_floating_point() + + +def _prefer_bf16_when_mixed(dtypes: List[torch.dtype]) -> torch.dtype: + """ + If there is any mixed dtype situation and no explicit compute dtype was given, + we favor bf16. + """ + uniq = list({dt for dt in dtypes}) + if len(uniq) == 1: + return uniq[0] + return torch.bfloat16 + + +def _resolve_working_dtype( + tensors: List[torch.Tensor], + explicit_compute_dtype: Optional[torch.dtype], +) -> torch.dtype: + float_dtypes = [t.dtype for t in tensors if _is_float_tensor(t)] + if not float_dtypes: + return explicit_compute_dtype or torch.bfloat16 + if explicit_compute_dtype is not None: + return explicit_compute_dtype + return _prefer_bf16_when_mixed(float_dtypes) + + +def _cast_pair(lp: LoraPair, dtype: torch.dtype) -> LoraPair: + if lp.down.dtype == dtype and lp.up.dtype == dtype: + return lp + return LoraPair( + down=lp.down.to(dtype=dtype), + up=lp.up.to(dtype=dtype), + alpha=lp.alpha, + is_conv=lp.is_conv, + in_hw=lp.in_hw, + ) + + +def _safe_dot(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + """ + Dot product that guarantees dtype/device alignment and uses fp32 for stability. + """ + if a.device != b.device: + b = b.to(a.device) + a32 = a.to(torch.float32) + b32 = b.to(torch.float32) + return a32 @ b32 + + +def compute_delta(lp: LoraPair, compute_dtype: torch.dtype) -> torch.Tensor: + """ + Compute full weight-space delta produced by this LoRA module. + + - Inputs are aligned to compute_dtype. + - Matmul is done in fp32 for stability. + - Output delta is fp32. + """ + down = lp.down + up = lp.up + + if down.dtype != compute_dtype: + down = down.to(compute_dtype) + if up.dtype != compute_dtype: + up = up.to(compute_dtype) + + alpha = lp.alpha + rank = down.shape[0] + scale = alpha / max(rank, 1) + + if lp.is_conv: + r, c_in, k_h, k_w = down.shape + out_channels = up.shape[0] + down_flat = down.view(r, c_in * k_h * k_w).to(torch.float32) + up_flat = up.view(out_channels, r).to(torch.float32) + delta_flat = (up_flat @ down_flat) * float(scale) + delta = delta_flat.view(out_channels, c_in, k_h, k_w) + return delta + + return (up.to(torch.float32) @ down.to(torch.float32)) * float(scale) + + +def tensor_norm(t: torch.Tensor) -> float: + return float(torch.norm(t).item()) + + +def choose_auto_rank(singular_values: torch.Tensor, threshold: float) -> int: + if singular_values.numel() == 0: + return 0 + s2 = singular_values * singular_values + total = torch.sum(s2).item() + if total <= 0.0: + return int(singular_values.shape[0]) + cumulative = torch.cumsum(s2, dim=0) + target = threshold * total + mask = cumulative >= target + indices = torch.nonzero(mask, as_tuple=False) + if indices.numel() == 0: + return int(singular_values.shape[0]) + r = int(indices[0, 0].item() + 1) + return r + + +def delta_to_svd_factors( + delta: torch.Tensor, + rank: int, + is_conv: bool, + auto_rank_threshold: float, + out_dtype: torch.dtype, +) -> Tuple[torch.Tensor, torch.Tensor]: + if is_conv: + out_channels, c_in, k_h, k_w = delta.shape + mat = delta.view(out_channels, c_in * k_h * k_w) + else: + out_channels, c_in = delta.shape + mat = delta + + mat32 = mat.to(torch.float32) + u, s, vh = torch.linalg.svd(mat32, full_matrices=False) + + if rank == -1: + r = choose_auto_rank(s, auto_rank_threshold) + elif rank <= 0: + r = s.shape[0] + else: + r = min(rank, s.shape[0]) + + if r <= 0: + if is_conv: + return ( + torch.zeros((0, 0, 0, 0), dtype=out_dtype, device=delta.device), + torch.zeros((0, 0), dtype=out_dtype, device=delta.device), + ) + return ( + torch.zeros((0, 0), dtype=out_dtype, device=delta.device), + torch.zeros((0, 0), dtype=out_dtype, device=delta.device), + ) + + u_r = u[:, :r] + s_r = s[:r] + vh_r = vh[:r, :] + + sqrt_s = torch.sqrt(s_r.clamp_min(0.0)) + up = u_r * sqrt_s.unsqueeze(0) + down = sqrt_s.unsqueeze(1) * vh_r + + if is_conv: + down = down.view(r, c_in, k_h, k_w) + up = up.view(out_channels, r, 1, 1) + + down = down.to(out_dtype) + up = up.to(out_dtype) + + return down.contiguous(), up.contiguous() + + +def merge_mode_add( + loras: List[Tuple[Dict[str, LoraPair], float]], + include_patterns: List[str], + exclude_patterns: List[str], + explicit_compute_dtype: Optional[torch.dtype], + progress_cb: Optional[ProgressCallback] = None, +) -> Dict[str, LoraPair]: + all_prefixes = set() + for pairs, _w in loras: + all_prefixes.update(pairs.keys()) + + merged: Dict[str, LoraPair] = {} + + prefixes_sorted = sorted(all_prefixes) + total_modules = len(prefixes_sorted) + for i, prefix in enumerate( + tqdm(prefixes_sorted, desc="Merging (add)", unit="module", disable=progress_cb is not None), + start=1, + ): + _progress(progress_cb, "merge.add", i, total_modules, prefix) + if not module_is_included(prefix, include_patterns, exclude_patterns): + continue + + present: List[Tuple[LoraPair, float]] = [] + for pairs, weight in loras: + if prefix in pairs: + present.append((pairs[prefix], weight)) + + if not present: + continue + + is_conv = present[0][0].is_conv + in_hw = present[0][0].in_hw + + for lp, _w in present[1:]: + if lp.is_conv != is_conv or lp.in_hw != in_hw: + raise ValueError(f"Incompatible shapes for module '{prefix}' between LoRAs") + + compute_dtype = _resolve_working_dtype( + [p[0].down for p in present] + [p[0].up for p in present], + explicit_compute_dtype, + ) + + downs_list: List[torch.Tensor] = [] + ups_list: List[torch.Tensor] = [] + total_rank = 0 + + for lp, w in present: + lp = _cast_pair(lp, compute_dtype) + rank_i = lp.down.shape[0] + scale_i = w * (lp.alpha / max(rank_i, 1)) + + downs_list.append(lp.down) + ups_list.append(lp.up * float(scale_i)) + total_rank += rank_i + + new_down = torch.cat(downs_list, dim=0) + new_up = torch.cat(ups_list, dim=1) + + new_alpha = float(total_rank) + + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + + return merged + + +def merge_mode_block_mix_weighted( + loras: List[Tuple[Dict[str, LoraPair], float]], + method: str, + rank_value: int, + preset: str, + recipe: str, + concept_mix: float, + style_mix: float, + include_patterns: List[str], + exclude_patterns: List[str], + auto_rank_threshold: float, + explicit_compute_dtype: Optional[torch.dtype], + progress_cb: Optional[ProgressCallback] = None, +) -> Dict[str, LoraPair]: + if len(loras) != 2: + raise ValueError("block-mix weighted mode requires exactly two LoRAs (A and B).") + + method_norm = (method or "svd").lower() + if method_norm not in ("stack", "svd"): + raise ValueError("block-mix method must be 'stack' or 'svd'.") + + a_pairs, a_weight = loras[0] + b_pairs, b_weight = loras[1] + + all_prefixes = set(a_pairs.keys()) | set(b_pairs.keys()) + prefixes_sorted = sorted(all_prefixes) + + preset_norm = (preset or "auto").lower().replace("_", "-") + if preset_norm == "auto": + preset_norm = _infer_block_mix_preset(prefixes_sorted) + + c_mix = float(concept_mix) + s_mix = float(style_mix) + c_mix = 0.0 if c_mix < 0.0 else (1.0 if c_mix > 1.0 else c_mix) + s_mix = 0.0 if s_mix < 0.0 else (1.0 if s_mix > 1.0 else s_mix) + + merged: Dict[str, LoraPair] = {} + total_modules = len(prefixes_sorted) + for i, prefix in enumerate( + tqdm(prefixes_sorted, desc=f"Merging (block-mix-weighted:{method_norm})", unit="module", disable=progress_cb is not None), + start=1, + ): + _progress(progress_cb, f"merge.block-mix-weighted.{method_norm}", i, total_modules, prefix) + if not module_is_included(prefix, include_patterns, exclude_patterns): + continue + + role = _block_mix_role(prefix, preset_norm) + mix_a = c_mix if role == "concept" else s_mix + mix_b = 1.0 - mix_a + + present: List[Tuple[LoraPair, float]] = [] + if prefix in a_pairs and mix_a != 0.0: + present.append((a_pairs[prefix], float(a_weight) * float(mix_a))) + if prefix in b_pairs and mix_b != 0.0: + present.append((b_pairs[prefix], float(b_weight) * float(mix_b))) + + if not present: + continue + + ref_lp = present[0][0] + is_conv = ref_lp.is_conv + in_hw = ref_lp.in_hw + for lp, _w in present[1:]: + if lp.is_conv != is_conv or lp.in_hw != in_hw: + raise ValueError(f"Incompatible shapes for module '{prefix}' in block-mix weighted mode") + + if method_norm == "stack": + compute_dtype = _resolve_working_dtype( + [t for lp, _w in present for t in (lp.down, lp.up)], + explicit_compute_dtype, + ) + + downs_list: List[torch.Tensor] = [] + ups_list: List[torch.Tensor] = [] + total_rank = 0 + for lp, w in present: + lp = _cast_pair(lp, compute_dtype) + rank_i = lp.down.shape[0] + scale_i = float(w) * (lp.alpha / max(rank_i, 1)) + downs_list.append(lp.down) + ups_list.append(lp.up * float(scale_i)) + total_rank += int(rank_i) + + new_down = torch.cat(downs_list, dim=0) + new_up = torch.cat(ups_list, dim=1) + new_alpha = float(total_rank) + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + continue + + compute_dtype = _resolve_working_dtype( + [t for lp, _w in present for t in (lp.down, lp.up)], + explicit_compute_dtype, + ) + + merged_delta = None + for lp, w in present: + d = compute_delta(lp, compute_dtype) * float(w) + merged_delta = d if merged_delta is None else (merged_delta + d) + + if merged_delta is None: + continue + + svd_rank = rank_value + if svd_rank == 0: + svd_rank = -1 + + new_down, new_up = delta_to_svd_factors( + merged_delta, + rank=svd_rank, + is_conv=is_conv, + auto_rank_threshold=auto_rank_threshold, + out_dtype=compute_dtype, + ) + new_alpha = float(new_down.shape[0]) + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + + return merged + + +def block_mix_routing_report( + a_pairs: Dict[str, LoraPair], + b_pairs: Dict[str, LoraPair], + preset: str, + recipe: str, + include_patterns: List[str], + exclude_patterns: List[str], +) -> Dict[str, Any]: + all_prefixes = set(a_pairs.keys()) | set(b_pairs.keys()) + prefixes_sorted = sorted(all_prefixes) + + preset_norm = (preset or "auto").lower().replace("_", "-") + if preset_norm == "auto": + preset_norm = _infer_block_mix_preset(prefixes_sorted) + + recipe_norm = (recipe or "").lower().replace("-", "_") + + routed_a = 0 + routed_b = 0 + fallback_to_a = 0 + fallback_to_b = 0 + missing_both = 0 + skipped_filter = 0 + requested_a = 0 + requested_b = 0 + + concept = 0 + style = 0 + + for prefix in prefixes_sorted: + if not module_is_included(prefix, include_patterns, exclude_patterns): + skipped_filter += 1 + continue + + role = _block_mix_role(prefix, preset_norm) + if role == "style": + style += 1 + else: + concept += 1 + + idx = _block_mix_pick_lora_index(prefix, preset_norm, recipe_norm) + if idx == 0: + requested_a += 1 + if prefix in a_pairs: + routed_a += 1 + elif prefix in b_pairs: + routed_b += 1 + fallback_to_b += 1 + else: + missing_both += 1 + else: + requested_b += 1 + if prefix in b_pairs: + routed_b += 1 + elif prefix in a_pairs: + routed_a += 1 + fallback_to_a += 1 + else: + missing_both += 1 + + included_total = (len(prefixes_sorted) - skipped_filter) + + return { + "preset": preset_norm, + "recipe": recipe_norm, + "modules_total": len(prefixes_sorted), + "modules_included": included_total, + "modules_skipped_filter": skipped_filter, + "requested_a": requested_a, + "requested_b": requested_b, + "routed_a": routed_a, + "routed_b": routed_b, + "fallback_to_a": fallback_to_a, + "fallback_to_b": fallback_to_b, + "missing_both": missing_both, + "role_concept": concept, + "role_style": style, + } + + +def _infer_block_mix_preset(prefixes: List[str]) -> str: + for p in prefixes: + if p.startswith("diffusion_model.layers."): + return "zimg-turbo" + if p.startswith("lora_unet_double_blocks_") or p.startswith("lora_unet_single_blocks_"): + return "flux" + if p.startswith("lora_unet_blocks_"): + return "wan" + if p.startswith("transformer_blocks."): + return "qwen" + if p.startswith("lora_te1_") or p.startswith("lora_te2_"): + return "sdxl" + if p.startswith("lora_unet_"): + return "sd" + return "generic" + + +def _block_mix_role(prefix: str, preset: str) -> str: + p = prefix.lower() + + if preset == "zimg-turbo": + if ".attention." in p: + return "concept" + if ".feed_forward." in p or ".adaln_modulation." in p: + return "style" + return "concept" + + if preset == "flux": + if "_attn_" in p: + return "concept" + if "_mlp_" in p or "_mod_" in p or "_mod_lin" in p: + return "style" + return "concept" + + if preset == "wan": + if "_attn_" in p: + return "concept" + if "_ffn_" in p or "_mlp_" in p: + return "style" + return "concept" + + if preset == "qwen": + if ".attn." in p: + return "concept" + if ".mlp." in p or "_mlp" in p: + return "style" + return "concept" + + attn_markers = ( + "attn", + "to_q", + "to_k", + "to_v", + "to_out", + "q_proj", + "k_proj", + "v_proj", + "out_proj", + "proj", + ) + if any(m in p for m in attn_markers): + return "concept" + + style_markers = ("mlp", "ff", "feed_forward", "res", "conv", "adaln", "norm") + if any(m in p for m in style_markers): + return "style" + + return "concept" + + +def _block_mix_pick_lora_index(prefix: str, preset: str, recipe: str) -> int: + r = (recipe or "").lower().replace("-", "_") + + if r in ("all_a", "only_a"): + return 0 + if r in ("all_b", "only_b"): + return 1 + + if preset == "flux" and r in ("img_a_txt_b", "img_b_txt_a"): + p = prefix.lower() + if "_img_" in p: + return 0 if r == "img_a_txt_b" else 1 + if "_txt_" in p: + return 1 if r == "img_a_txt_b" else 0 + return 0 + + role = _block_mix_role(prefix, preset) + if r in ("concept_a_style_b", "concept_from_a_style_from_b", "attn_a_ffn_b"): + return 0 if role == "concept" else 1 + if r in ("concept_b_style_a", "concept_from_b_style_from_a", "attn_b_ffn_a"): + return 1 if role == "concept" else 0 + + return 0 + + +def merge_mode_block_mix( + loras: List[Tuple[Dict[str, LoraPair], float]], + method: str, + rank_value: int, + preset: str, + recipe: str, + include_patterns: List[str], + exclude_patterns: List[str], + auto_rank_threshold: float, + explicit_compute_dtype: Optional[torch.dtype], + progress_cb: Optional[ProgressCallback] = None, +) -> Dict[str, LoraPair]: + if len(loras) != 2: + raise ValueError("block-mix mode requires exactly two LoRAs (A and B).") + + method_norm = (method or "svd").lower() + if method_norm not in ("stack", "svd"): + raise ValueError("block-mix method must be 'stack' or 'svd'.") + + a_pairs, a_weight = loras[0] + b_pairs, b_weight = loras[1] + + all_prefixes = set(a_pairs.keys()) | set(b_pairs.keys()) + prefixes_sorted = sorted(all_prefixes) + + preset_norm = (preset or "auto").lower().replace("_", "-") + if preset_norm == "auto": + preset_norm = _infer_block_mix_preset(prefixes_sorted) + + merged: Dict[str, LoraPair] = {} + total_modules = len(prefixes_sorted) + for i, prefix in enumerate( + tqdm(prefixes_sorted, desc=f"Merging (block-mix:{method_norm})", unit="module", disable=progress_cb is not None), + start=1, + ): + _progress(progress_cb, f"merge.block-mix.{method_norm}", i, total_modules, prefix) + if not module_is_included(prefix, include_patterns, exclude_patterns): + continue + + idx = _block_mix_pick_lora_index(prefix, preset_norm, recipe) + + present: List[Tuple[LoraPair, float]] = [] + if idx == 0: + if prefix in a_pairs: + present.append((a_pairs[prefix], float(a_weight))) + elif prefix in b_pairs: + present.append((b_pairs[prefix], float(b_weight))) + else: + if prefix in b_pairs: + present.append((b_pairs[prefix], float(b_weight))) + elif prefix in a_pairs: + present.append((a_pairs[prefix], float(a_weight))) + + if not present: + continue + + ref_lp = present[0][0] + is_conv = ref_lp.is_conv + in_hw = ref_lp.in_hw + for lp, _w in present[1:]: + if lp.is_conv != is_conv or lp.in_hw != in_hw: + raise ValueError(f"Incompatible shapes for module '{prefix}' in block-mix mode") + + if method_norm == "stack": + compute_dtype = _resolve_working_dtype( + [t for lp, _w in present for t in (lp.down, lp.up)], + explicit_compute_dtype, + ) + + downs_list: List[torch.Tensor] = [] + ups_list: List[torch.Tensor] = [] + total_rank = 0 + for lp, w in present: + lp = _cast_pair(lp, compute_dtype) + rank_i = lp.down.shape[0] + scale_i = float(w) * (lp.alpha / max(rank_i, 1)) + downs_list.append(lp.down) + ups_list.append(lp.up * float(scale_i)) + total_rank += int(rank_i) + + new_down = torch.cat(downs_list, dim=0) + new_up = torch.cat(ups_list, dim=1) + new_alpha = float(total_rank) + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + continue + + compute_dtype = _resolve_working_dtype( + [t for lp, _w in present for t in (lp.down, lp.up)], + explicit_compute_dtype, + ) + merged_delta = None + for lp, w in present: + d = compute_delta(lp, compute_dtype) * float(w) + merged_delta = d if merged_delta is None else (merged_delta + d) + + if merged_delta is None: + continue + + svd_rank = rank_value + if svd_rank == 0: + svd_rank = -1 + + new_down, new_up = delta_to_svd_factors( + merged_delta, + rank=svd_rank, + is_conv=is_conv, + auto_rank_threshold=auto_rank_threshold, + out_dtype=compute_dtype, + ) + new_alpha = float(new_down.shape[0]) + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + + return merged + + +def merge_mode_add_diff( + loras: List[Tuple[Dict[str, LoraPair], float]], + rank_value: int, + include_patterns: List[str], + exclude_patterns: List[str], + auto_rank_threshold: float, + explicit_compute_dtype: Optional[torch.dtype], + progress_cb: Optional[ProgressCallback] = None, +) -> Dict[str, LoraPair]: + if len(loras) < 2: + raise ValueError("add-diff mode requires at least two LoRAs (base + one or more others).") + + base_pairs, base_weight = loras[0] + + all_prefixes = set() + for pairs, _w in loras: + all_prefixes.update(pairs.keys()) + + merged: Dict[str, LoraPair] = {} + + prefixes_sorted = sorted(all_prefixes) + total_modules = len(prefixes_sorted) + for i, prefix in enumerate( + tqdm(prefixes_sorted, desc="Merging (add-diff)", unit="module", disable=progress_cb is not None), + start=1, + ): + _progress(progress_cb, "merge.add-diff", i, total_modules, prefix) + if not module_is_included(prefix, include_patterns, exclude_patterns): + continue + + base_lp = base_pairs.get(prefix, None) + + others: List[Tuple[LoraPair, float]] = [] + for pairs, w in loras[1:]: + if prefix in pairs: + others.append((pairs[prefix], w)) + + if base_lp is None and not others: + continue + + ref_lp = base_lp if base_lp is not None else others[0][0] + is_conv = ref_lp.is_conv + in_hw = ref_lp.in_hw + + if base_lp is not None and (base_lp.is_conv != is_conv or base_lp.in_hw != in_hw): + raise ValueError(f"Incompatible shapes for module '{prefix}' in add-diff (base).") + for lp, _w in others: + if lp.is_conv != is_conv or lp.in_hw != in_hw: + raise ValueError(f"Incompatible shapes for module '{prefix}' in add-diff (other).") + + compute_dtype = _resolve_working_dtype( + ([base_lp.down, base_lp.up] if base_lp is not None else []) + + [t for lp, _w in others for t in (lp.down, lp.up)], + explicit_compute_dtype, + ) + + if base_lp is not None: + delta_base_raw = compute_delta(base_lp, compute_dtype) + delta_base = delta_base_raw * float(base_weight) + else: + delta_base_raw = None + delta_base = torch.zeros_like(compute_delta(ref_lp, compute_dtype)) + + merged_delta = delta_base.clone() + + for lp_other, weight_other in others: + delta_other = compute_delta(lp_other, compute_dtype) + if delta_base_raw is None: + merged_delta = merged_delta + float(weight_other) * delta_other + else: + merged_delta = merged_delta + float(weight_other) * (delta_other - delta_base_raw) + + svd_rank = rank_value + if svd_rank == 0: + svd_rank = -1 + + new_down, new_up = delta_to_svd_factors( + merged_delta, + rank=svd_rank, + is_conv=is_conv, + auto_rank_threshold=auto_rank_threshold, + out_dtype=compute_dtype, + ) + new_alpha = float(new_down.shape[0]) + + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + + return merged + + +def merge_mode_svd( + loras: List[Tuple[Dict[str, LoraPair], float]], + rank_value: int, + preserve_norm: bool, + cap_mult: Optional[float], + include_patterns: List[str], + exclude_patterns: List[str], + auto_rank_threshold: float, + explicit_compute_dtype: Optional[torch.dtype], + progress_cb: Optional[ProgressCallback] = None, +) -> Dict[str, LoraPair]: + all_prefixes = set() + for pairs, _w in loras: + all_prefixes.update(pairs.keys()) + + merged: Dict[str, LoraPair] = {} + + prefixes_sorted = sorted(all_prefixes) + total_modules = len(prefixes_sorted) + for i, prefix in enumerate( + tqdm(prefixes_sorted, desc="Merging (svd)", unit="module", disable=progress_cb is not None), + start=1, + ): + _progress(progress_cb, "merge.svd", i, total_modules, prefix) + if not module_is_included(prefix, include_patterns, exclude_patterns): + continue + + present: List[Tuple[LoraPair, float]] = [] + for pairs, weight in loras: + if prefix in pairs: + present.append((pairs[prefix], weight)) + + if not present: + continue + + ref_lp = present[0][0] + is_conv = ref_lp.is_conv + in_hw = ref_lp.in_hw + + for lp, _w in present[1:]: + if lp.is_conv != is_conv or lp.in_hw != in_hw: + raise ValueError(f"Incompatible shapes for module '{prefix}' between LoRAs") + + compute_dtype = _resolve_working_dtype( + [t for lp, _w in present for t in (lp.down, lp.up)], + explicit_compute_dtype, + ) + + deltas: List[torch.Tensor] = [] + norms: List[float] = [] + + for lp, w in present: + d = compute_delta(lp, compute_dtype) + deltas.append(d * float(w)) + norms.append(tensor_norm(d)) + + base_norm = sum(norms) / max(len(norms), 1) + merged_delta = torch.zeros_like(deltas[0]) + for d in deltas: + merged_delta = merged_delta + d + + if preserve_norm: + merged_norm = tensor_norm(merged_delta) + if merged_norm > 1e-8 and base_norm > 1e-8: + merged_delta = merged_delta * (base_norm / merged_norm) + + if cap_mult is not None and cap_mult > 0.0 and base_norm > 0.0: + merged_norm = tensor_norm(merged_delta) + max_norm = base_norm * cap_mult + if merged_norm > max_norm and merged_norm > 0.0: + merged_delta = merged_delta * (max_norm / merged_norm) + + svd_rank = rank_value + if svd_rank == 0: + svd_rank = -1 + + new_down, new_up = delta_to_svd_factors( + merged_delta, + rank=svd_rank, + is_conv=is_conv, + auto_rank_threshold=auto_rank_threshold, + out_dtype=compute_dtype, + ) + new_alpha = float(new_down.shape[0]) + + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + + return merged + + +def merge_mode_add_orth( + loras: List[Tuple[Dict[str, LoraPair], float]], + rank_value: int, + include_patterns: List[str], + exclude_patterns: List[str], + auto_rank_threshold: float, + explicit_compute_dtype: Optional[torch.dtype], + progress_cb: Optional[ProgressCallback] = None, +) -> Dict[str, LoraPair]: + if len(loras) < 2: + raise ValueError("add-orth mode requires at least two LoRAs (base + one or more others).") + + base_pairs, base_weight = loras[0] + + all_prefixes = set() + for pairs, _w in loras: + all_prefixes.update(pairs.keys()) + + merged: Dict[str, LoraPair] = {} + + prefixes_sorted = sorted(all_prefixes) + total_modules = len(prefixes_sorted) + for i, prefix in enumerate( + tqdm(prefixes_sorted, desc="Merging (add-orth)", unit="module", disable=progress_cb is not None), + start=1, + ): + _progress(progress_cb, "merge.add-orth", i, total_modules, prefix) + if not module_is_included(prefix, include_patterns, exclude_patterns): + continue + + base_lp = base_pairs.get(prefix, None) + present_others: List[Tuple[LoraPair, float]] = [] + + for pairs, weight in loras[1:]: + if prefix in pairs: + present_others.append((pairs[prefix], weight)) + + if base_lp is None and not present_others: + continue + + if base_lp is None: + temp_loras = [(pairs, weight) for (pairs, weight) in loras if prefix in pairs] + temp_pairs_list = [] + for pairs, weight in temp_loras: + temp_pairs_list.append(({prefix: pairs[prefix]}, weight)) + return_fallback = merge_mode_add(temp_pairs_list, include_patterns, exclude_patterns, explicit_compute_dtype) + merged.update(return_fallback) + continue + + is_conv = base_lp.is_conv + in_hw = base_lp.in_hw + + for lp, _w in present_others: + if lp.is_conv != is_conv or lp.in_hw != in_hw: + raise ValueError(f"Incompatible shapes for module '{prefix}' in add-orth mode") + + compute_dtype = _resolve_working_dtype( + [base_lp.down, base_lp.up] + [t for lp, _w in present_others for t in (lp.down, lp.up)], + explicit_compute_dtype, + ) + + delta_base = compute_delta(base_lp, compute_dtype) * float(base_weight) + base_flat = delta_base.reshape(-1) + base_norm_sq = float(_safe_dot(base_flat, base_flat).item()) + + merged_delta = delta_base.clone() + eps = 1e-12 + + for lp_other, weight_other in present_others: + delta_other = compute_delta(lp_other, compute_dtype) * float(weight_other) + other_flat = delta_other.reshape(-1) + + if base_norm_sq < eps: + merged_delta = merged_delta + delta_other + continue + + coeff = float(_safe_dot(other_flat, base_flat).item()) / (base_norm_sq + eps) + orth_flat = other_flat - coeff * base_flat + orth_delta = orth_flat.view_as(delta_other) + merged_delta = merged_delta + orth_delta + + svd_rank = rank_value + if svd_rank == 0: + svd_rank = -1 + + new_down, new_up = delta_to_svd_factors( + merged_delta, + rank=svd_rank, + is_conv=is_conv, + auto_rank_threshold=auto_rank_threshold, + out_dtype=compute_dtype, + ) + new_alpha = float(new_down.shape[0]) + + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + + return merged + + +def merge_mode_diff_export( + loras: List[Tuple[Dict[str, LoraPair], float]], + rank_value: int, + include_patterns: List[str], + exclude_patterns: List[str], + auto_rank_threshold: float, + explicit_compute_dtype: Optional[torch.dtype], + progress_cb: Optional[ProgressCallback] = None, +) -> Dict[str, LoraPair]: + if len(loras) < 2: + raise ValueError("diff-export mode requires at least two LoRAs (base and capability).") + + base_pairs, base_weight = loras[0] + cap_pairs, cap_weight = loras[1] + + all_prefixes = set(base_pairs.keys()) | set(cap_pairs.keys()) + merged: Dict[str, LoraPair] = {} + + prefixes_sorted = sorted(all_prefixes) + total_modules = len(prefixes_sorted) + for i, prefix in enumerate( + tqdm(prefixes_sorted, desc="Merging (diff-export)", unit="module", disable=progress_cb is not None), + start=1, + ): + _progress(progress_cb, "merge.diff-export", i, total_modules, prefix) + if not module_is_included(prefix, include_patterns, exclude_patterns): + continue + + base_lp = base_pairs.get(prefix, None) + cap_lp = cap_pairs.get(prefix, None) + + if base_lp is None and cap_lp is None: + continue + + ref_lp = cap_lp if cap_lp is not None else base_lp + is_conv = ref_lp.is_conv + in_hw = ref_lp.in_hw + + if base_lp is not None and (base_lp.is_conv != is_conv or base_lp.in_hw != in_hw): + raise ValueError(f"Incompatible shapes for module '{prefix}' in base LoRA (diff-export).") + if cap_lp is not None and (cap_lp.is_conv != is_conv or cap_lp.in_hw != in_hw): + raise ValueError(f"Incompatible shapes for module '{prefix}' in capability LoRA (diff-export).") + + compute_dtype = _resolve_working_dtype( + ([base_lp.down, base_lp.up] if base_lp is not None else []) + + ([cap_lp.down, cap_lp.up] if cap_lp is not None else []), + explicit_compute_dtype, + ) + + if base_lp is not None: + delta_base = compute_delta(base_lp, compute_dtype) * float(base_weight) + else: + delta_base = torch.zeros_like(compute_delta(ref_lp, compute_dtype)) + + if cap_lp is not None: + delta_cap = compute_delta(cap_lp, compute_dtype) * float(cap_weight) + else: + delta_cap = torch.zeros_like(compute_delta(ref_lp, compute_dtype)) + + delta_diff = delta_cap - delta_base + + if torch.allclose(delta_diff, torch.zeros_like(delta_diff)): + continue + + svd_rank = rank_value + if svd_rank == 0: + svd_rank = -1 + + new_down, new_up = delta_to_svd_factors( + delta_diff, + rank=svd_rank, + is_conv=is_conv, + auto_rank_threshold=auto_rank_threshold, + out_dtype=compute_dtype, + ) + new_alpha = float(new_down.shape[0]) + + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + + return merged + + +def merge_mode_moe( + loras: List[Tuple[Dict[str, LoraPair], float]], + rank_value: int, + moe_temperature: float, + moe_hard: bool, + include_patterns: List[str], + exclude_patterns: List[str], + auto_rank_threshold: float, + explicit_compute_dtype: Optional[torch.dtype], + progress_cb: Optional[ProgressCallback] = None, +) -> Dict[str, LoraPair]: + if len(loras) < 1: + raise ValueError("moe mode requires at least one LoRA.") + + base_pairs, base_weight = loras[0] + + all_prefixes = set() + for pairs, _w in loras: + all_prefixes.update(pairs.keys()) + + merged: Dict[str, LoraPair] = {} + + prefixes_sorted = sorted(all_prefixes) + total_modules = len(prefixes_sorted) + for i, prefix in enumerate( + tqdm(prefixes_sorted, desc="Merging (moe)", unit="module", disable=progress_cb is not None), + start=1, + ): + _progress(progress_cb, "merge.moe", i, total_modules, prefix) + if not module_is_included(prefix, include_patterns, exclude_patterns): + continue + + base_lp = base_pairs.get(prefix, None) + experts: List[Tuple[LoraPair, float]] = [] + + for pairs, weight in loras[1:]: + if prefix in pairs: + experts.append((pairs[prefix], weight)) + + if base_lp is None and not experts: + continue + + ref_lp = base_lp if base_lp is not None else (experts[0][0] if experts else None) + if ref_lp is None: + continue + + is_conv = ref_lp.is_conv + in_hw = ref_lp.in_hw + + if base_lp is not None and (base_lp.is_conv != is_conv or base_lp.in_hw != in_hw): + raise ValueError(f"Incompatible shapes for module '{prefix}' in base LoRA (moe).") + for lp, _w in experts: + if lp.is_conv != is_conv or lp.in_hw != in_hw: + raise ValueError(f"Incompatible shapes for module '{prefix}' in moe experts.") + + compute_dtype = _resolve_working_dtype( + ([base_lp.down, base_lp.up] if base_lp is not None else []) + + [t for lp, _w in experts for t in (lp.down, lp.up)], + explicit_compute_dtype, + ) + + if base_lp is not None: + delta_base = compute_delta(base_lp, compute_dtype) * float(base_weight) + else: + delta_base = torch.zeros_like(compute_delta(ref_lp, compute_dtype)) + + if not experts: + merged_delta = delta_base + else: + scores: List[float] = [] + deltas_expert: List[torch.Tensor] = [] + + for lp, weight in experts: + d = compute_delta(lp, compute_dtype) * float(weight) + deltas_expert.append(d) + scores.append(tensor_norm(d)) + + scores_tensor = torch.tensor(scores, dtype=torch.float32, device=delta_base.device) + + if moe_hard: + best_index = int(torch.argmax(scores_tensor).item()) + gate = torch.zeros_like(scores_tensor) + gate[best_index] = 1.0 + else: + temp = max(moe_temperature, 1e-6) + scaled = scores_tensor / temp + gate = torch.softmax(scaled, dim=0) + + merged_delta = delta_base.clone() + for g, d in zip(gate, deltas_expert): + merged_delta = merged_delta + float(g.item()) * d + + svd_rank = rank_value + if svd_rank == 0: + svd_rank = -1 + + new_down, new_up = delta_to_svd_factors( + merged_delta, + rank=svd_rank, + is_conv=is_conv, + auto_rank_threshold=auto_rank_threshold, + out_dtype=compute_dtype, + ) + new_alpha = float(new_down.shape[0]) + + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=is_conv, + in_hw=in_hw, + ) + + return merged + + +def merge_mode_obfuscate( + loras: List[Tuple[Dict[str, LoraPair], float]], + include_patterns: List[str], + exclude_patterns: List[str], + explicit_compute_dtype: Optional[torch.dtype], + progress_cb: Optional[ProgressCallback] = None, +) -> Dict[str, LoraPair]: + if len(loras) < 1: + raise ValueError("obfuscate mode requires at least one LoRA.") + + all_prefixes = set() + for pairs, _w in loras: + all_prefixes.update(pairs.keys()) + + merged: Dict[str, LoraPair] = {} + + prefixes_sorted = sorted(all_prefixes) + total_modules = len(prefixes_sorted) + for i, prefix in enumerate( + tqdm(prefixes_sorted, desc="Merging (obfuscate)", unit="module", disable=progress_cb is not None), + start=1, + ): + _progress(progress_cb, "merge.obfuscate", i, total_modules, prefix) + if not module_is_included(prefix, include_patterns, exclude_patterns): + continue + + present: List[Tuple[LoraPair, float]] = [] + for pairs, weight in loras: + if prefix in pairs: + present.append((pairs[prefix], weight)) + + if not present: + continue + + ref_lp = present[0][0] + is_conv = ref_lp.is_conv + in_hw = ref_lp.in_hw + device = ref_lp.down.device + + for lp, _w in present[1:]: + if lp.is_conv != is_conv or lp.in_hw != in_hw: + raise ValueError(f"Incompatible shapes for module '{prefix}' between LoRAs in obfuscate mode") + + compute_dtype = _resolve_working_dtype( + [t for lp, _w in present for t in (lp.down, lp.up)], + explicit_compute_dtype, + ) + + down_blocks: List[torch.Tensor] = [] + up_blocks_flat: List[torch.Tensor] = [] + total_rank = 0 + c_in = None + k_h = None + k_w = None + + for lp, w in present: + lp = _cast_pair(lp, compute_dtype) + r_i = lp.down.shape[0] + if r_i == 0: + continue + + scale_i = float(w) * (lp.alpha / max(r_i, 1)) + + if is_conv: + down_i = lp.down.to(torch.float32) + up_i = lp.up.to(torch.float32) + c_in = down_i.shape[1] + k_h = down_i.shape[2] + k_w = down_i.shape[3] + + up_i_flat = up_i.view(up_i.shape[0], r_i) + up_i_flat = up_i_flat * scale_i + + down_blocks.append(down_i) + up_blocks_flat.append(up_i_flat) + else: + down_i = lp.down.to(torch.float32) + up_i = (lp.up.to(torch.float32) * scale_i) + + down_blocks.append(down_i) + up_blocks_flat.append(up_i) + + total_rank += r_i + + if total_rank == 0: + continue + + if is_conv: + down_cat = torch.cat(down_blocks, dim=0) + up_cat_flat = torch.cat(up_blocks_flat, dim=1) + r_total = down_cat.shape[0] + + rand = torch.randn((r_total, r_total), device=device, dtype=torch.float32) + q, _ = torch.linalg.qr(rand) + + down_flat = down_cat.view(r_total, -1) + down_flat = q @ down_flat + down_new32 = down_flat.view(r_total, c_in, k_h, k_w) + + up_cat_flat = up_cat_flat @ q.T + up_new32 = up_cat_flat.view(up_cat_flat.shape[0], r_total, 1, 1) + else: + down_cat = torch.cat(down_blocks, dim=0) + up_cat = torch.cat(up_blocks_flat, dim=1) + r_total = down_cat.shape[0] + + rand = torch.randn((r_total, r_total), device=device, dtype=torch.float32) + q, _ = torch.linalg.qr(rand) + + down_new32 = q @ down_cat + up_new32 = up_cat @ q.T + + down_new = down_new32.to(compute_dtype) + up_new = up_new32.to(compute_dtype) + alpha_new = float(r_total) + + merged[prefix] = LoraPair( + down=down_new.contiguous(), + up=up_new.contiguous(), + alpha=alpha_new, + is_conv=is_conv, + in_hw=in_hw, + ) + + return merged + + +def merge_mode_rebase( + lora: Tuple[Dict[str, LoraPair], float], + rank_value: int, + include_patterns: List[str], + exclude_patterns: List[str], + auto_rank_threshold: float, + explicit_compute_dtype: Optional[torch.dtype], + progress_cb: Optional[ProgressCallback] = None, +) -> Dict[str, LoraPair]: + pairs, weight = lora + merged: Dict[str, LoraPair] = {} + + prefixes_sorted = sorted(pairs.keys()) + total_modules = len(prefixes_sorted) + for i, prefix in enumerate( + tqdm(prefixes_sorted, desc="Rebasing (rebase)", unit="module", disable=progress_cb is not None), + start=1, + ): + _progress(progress_cb, "merge.rebase", i, total_modules, prefix) + if not module_is_included(prefix, include_patterns, exclude_patterns): + continue + + lp = pairs[prefix] + compute_dtype = _resolve_working_dtype([lp.down, lp.up], explicit_compute_dtype) + + delta = compute_delta(lp, compute_dtype) * float(weight) + original_rank = lp.down.shape[0] + + target_rank = -1 if rank_value == 0 else rank_value + + new_down, new_up = delta_to_svd_factors( + delta, + rank=target_rank, + is_conv=lp.is_conv, + auto_rank_threshold=auto_rank_threshold, + out_dtype=compute_dtype, + ) + + new_rank = new_down.shape[0] + if new_rank > original_rank: + if lp.is_conv: + new_down = new_down[:original_rank].contiguous() + new_up = new_up[:, :original_rank, ...].contiguous() + else: + new_down = new_down[:original_rank].contiguous() + new_up = new_up[:, :original_rank].contiguous() + + new_alpha = float(new_down.shape[0]) + + merged[prefix] = LoraPair( + down=new_down, + up=new_up, + alpha=new_alpha, + is_conv=lp.is_conv, + in_hw=lp.in_hw, + ) + + return merged + + +def build_state_dict( + merged: Dict[str, LoraPair], + metadata_ref: Dict[str, str], + dtype: torch.dtype, + device: torch.device, + progress_cb: Optional[ProgressCallback] = None, +) -> Tuple[Dict[str, torch.Tensor], Dict[str, str]]: + tensors: Dict[str, torch.Tensor] = {} + + key_style = _get_save_key_style(metadata_ref) + + items = list(merged.items()) + total_modules = len(items) + for i, (prefix, lp) in enumerate(items, start=1): + _progress(progress_cb, "build.state_dict", i, total_modules, prefix) + down = lp.down.to(dtype=dtype, device=device) + up = lp.up.to(dtype=dtype, device=device) + alpha_value = torch.tensor(lp.alpha, dtype=torch.float32, device=device) + + tensors[_down_key(prefix, key_style)] = down + tensors[_up_key(prefix, key_style)] = up + tensors[get_alpha_key(prefix)] = alpha_value + + meta = dict(metadata_ref or {}) + meta.setdefault("format", "kohya-lora") + + return tensors, meta + + +def summarize_pairs( + pairs: Dict[str, LoraPair], + weight: float, + explicit_compute_dtype: Optional[torch.dtype], + include_patterns: List[str], + exclude_patterns: List[str], + max_modules: int = 25, +) -> Dict[str, float]: + prefixes = [p for p in sorted(pairs.keys()) if module_is_included(p, include_patterns, exclude_patterns)] + if not prefixes: + return { + "modules": 0.0, + "avg_rank": 0.0, + "avg_alpha": 0.0, + "avg_scale": 0.0, + "mean_delta_norm": 0.0, + "max_delta_norm": 0.0, + } + + norms: List[float] = [] + ranks: List[float] = [] + alphas: List[float] = [] + scales: List[float] = [] + + for prefix in prefixes[: max_modules]: + lp = pairs[prefix] + compute_dtype = _resolve_working_dtype([lp.down, lp.up], explicit_compute_dtype) + delta = compute_delta(lp, compute_dtype) + norms.append(tensor_norm(delta) * float(abs(weight))) + r = float(lp.down.shape[0]) + a = float(lp.alpha) + ranks.append(r) + alphas.append(a) + scales.append(float(abs(weight)) * (a / max(r, 1.0))) + + return { + "modules": float(len(prefixes)), + "avg_rank": float(sum(ranks) / max(len(ranks), 1)), + "avg_alpha": float(sum(alphas) / max(len(alphas), 1)), + "avg_scale": float(sum(scales) / max(len(scales), 1)), + "mean_delta_norm": float(sum(norms) / max(len(norms), 1)), + "max_delta_norm": float(max(norms) if norms else 0.0), + } + + +def main(): + parser = argparse.ArgumentParser( + description=( + "Comprehensive LoRA merge tool with modes: " + "svd, rebase, add, add-diff, add-orth, diff-export, moe, obfuscate." + ) + ) + parser.add_argument( + "inputs", + nargs="+", + help="Input LoRAs as 'path[@weight]'. Example: A.safetensors@1.0 B.safetensors@-0.5", + ) + parser.add_argument("--out", required=True, help="Output .safetensors path") + parser.add_argument( + "--mode", + type=str, + default="svd", + choices=["svd", "rebase", "add", "add-diff", "add-orth", "diff-export", "moe", "obfuscate", "block-mix"], + help=( + "Merge mode: " + "svd = SVD rank-compressed merge, " + "rebase = single-LoRA SVD rank-compression, " + "add = exact linear stack (matches Comfy stacking behavior), " + "add-diff = base + weighted differences toward others, " + "add-orth = base + orthogonalized contributions, " + "diff-export = pure difference LoRA, " + "moe = mixture-of-experts per module (base + gated experts), " + "obfuscate = stack-equivalent factor-space rebasis without SVD, " + "block-mix = preset router that routes modules to LoRA A or B and merges via stack or svd." + ), + ) + parser.add_argument( + "--rank", + type=int, + default=64, + help=( + "Rank for SVD-based modes. " + "Use 0 to trigger auto rank in svd/add-orth/diff-export/moe/add-diff/rebase." + ), + ) + parser.add_argument( + "--auto-rank-threshold", + type=float, + default=0.99, + help="Energy ratio threshold for auto rank in SVD-based modes.", + ) + parser.add_argument( + "--preserve-norm", + action="store_true", + help="In svd mode, preserve average per-module norm.", + ) + parser.add_argument( + "--cap-mult", + type=float, + default=None, + help="In svd mode, cap merged norm at cap_mult * mean(source_norms).", + ) + parser.add_argument( + "--dtype", + type=str, + default="fp16", + help="Output dtype: fp16, fp32, bf16", + ) + parser.add_argument( + "--compute-dtype", + type=str, + default="auto", + help=( + "Internal merge dtype alignment: auto, bf16, fp16, fp32. " + "auto => if mixed dtypes encountered, favor bf16." + ), + ) + parser.add_argument( + "--cpu", + action="store_true", + help="Force CPU even if CUDA is available", + ) + parser.add_argument( + "--include-pattern", + action="append", + default=None, + help="Only merge modules whose prefix contains this substring. Can be used multiple times.", + ) + parser.add_argument( + "--exclude-pattern", + action="append", + default=None, + help="Exclude modules whose prefix contains this substring. Can be used multiple times.", + ) + parser.add_argument( + "--moe-temperature", + type=float, + default=1.0, + help="Softmax temperature for moe mode (lower = sharper gating).", + ) + parser.add_argument( + "--moe-hard", + action="store_true", + help="Use hard gating in moe mode (pick single best expert per module).", + ) + + parser.add_argument( + "--block-mix-method", + type=str, + default="svd", + choices=["svd", "stack"], + help="block-mix only: svd (delta merge then SVD) or stack (exact rank stacking).", + ) + parser.add_argument( + "--block-mix-preset", + type=str, + default="auto", + choices=["auto", "zimg-turbo", "flux", "wan", "qwen", "sd", "sdxl", "generic"], + help="block-mix only: routing preset family.", + ) + parser.add_argument( + "--block-mix-recipe", + type=str, + default="concept_a_style_b", + help=( + "block-mix only: recipe string. Examples: concept_a_style_b, concept_b_style_a, " + "attn_a_ffn_b, img_a_txt_b, img_b_txt_a." + ), + ) + + parser.add_argument( + "--report", + action="store_true", + help="Print a small diagnostic summary (module counts, ranks, scaling, delta norms).", + ) + + args = parser.parse_args() + + weighted_paths = parse_weighted_paths(args.inputs) + if not weighted_paths: + raise SystemExit("No input LoRA files provided.") + + device = get_device(args.cpu) + dtype = get_dtype(args.dtype) + explicit_compute_dtype = get_compute_dtype(args.compute_dtype) + + include_patterns = args.include_pattern or [] + exclude_patterns = args.exclude_pattern or [] + + loaded: List[Tuple[Dict[str, LoraPair], float]] = [] + ref_meta: Dict[str, str] = {} + + for index, (path, weight) in enumerate(weighted_paths): + pairs, meta = load_lora_pairs(path, device=device) + print(f"Loaded {len(pairs)} modules from {path} (weight {weight})") + if index == 0: + ref_meta = meta + loaded.append((pairs, weight)) + + if args.report: + stats = summarize_pairs( + pairs, + weight=weight, + explicit_compute_dtype=explicit_compute_dtype, + include_patterns=include_patterns, + exclude_patterns=exclude_patterns, + ) + print( + "Report(input): " + + json.dumps( + { + "path": path, + "weight": weight, + **stats, + }, + indent=2, + ) + ) + + if args.mode == "add": + merged_pairs = merge_mode_add(loaded, include_patterns, exclude_patterns, explicit_compute_dtype) + elif args.mode == "block-mix": + merged_pairs = merge_mode_block_mix( + loaded, + method=args.block_mix_method, + rank_value=args.rank, + preset=args.block_mix_preset, + recipe=args.block_mix_recipe, + include_patterns=include_patterns, + exclude_patterns=exclude_patterns, + auto_rank_threshold=args.auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + ) + elif args.mode == "add-diff": + merged_pairs = merge_mode_add_diff( + loaded, + rank_value=args.rank, + include_patterns=include_patterns, + exclude_patterns=exclude_patterns, + auto_rank_threshold=args.auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + ) + elif args.mode == "add-orth": + merged_pairs = merge_mode_add_orth( + loaded, + rank_value=args.rank, + include_patterns=include_patterns, + exclude_patterns=exclude_patterns, + auto_rank_threshold=args.auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + ) + elif args.mode == "diff-export": + merged_pairs = merge_mode_diff_export( + loaded, + rank_value=args.rank, + include_patterns=include_patterns, + exclude_patterns=exclude_patterns, + auto_rank_threshold=args.auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + ) + elif args.mode == "moe": + merged_pairs = merge_mode_moe( + loaded, + rank_value=args.rank, + moe_temperature=args.moe_temperature, + moe_hard=args.moe_hard, + include_patterns=include_patterns, + exclude_patterns=exclude_patterns, + auto_rank_threshold=args.auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + ) + elif args.mode == "obfuscate": + merged_pairs = merge_mode_obfuscate( + loaded, + include_patterns=include_patterns, + exclude_patterns=exclude_patterns, + explicit_compute_dtype=explicit_compute_dtype, + ) + elif args.mode == "rebase": + if len(loaded) != 1: + raise ValueError("rebase mode requires exactly one input LoRA.") + merged_pairs = merge_mode_rebase( + loaded[0], + rank_value=args.rank, + include_patterns=include_patterns, + exclude_patterns=exclude_patterns, + auto_rank_threshold=args.auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + ) + else: + merged_pairs = merge_mode_svd( + loaded, + rank_value=args.rank, + preserve_norm=args.preserve_norm, + cap_mult=args.cap_mult if args.cap_mult is not None else None, + include_patterns=include_patterns, + exclude_patterns=exclude_patterns, + auto_rank_threshold=args.auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + ) + + state, meta = build_state_dict( + merged_pairs, + metadata_ref=ref_meta, + dtype=dtype, + device=device, + ) + + if args.report: + merged_stats = summarize_pairs( + merged_pairs, + weight=1.0, + explicit_compute_dtype=explicit_compute_dtype, + include_patterns=include_patterns, + exclude_patterns=exclude_patterns, + ) + print("Report(merged): " + json.dumps(merged_stats, indent=2)) + + note_str = "Created by WAS Lora Merger" if args.mode == "obfuscate" else f"Created by WAS Lora Merger - Merging Mode: {args.mode}" + + meta.update( + { + "format": meta.get("format", "kohya-lora"), + "merged_from": str(weighted_paths), + "mode": args.mode, + "rank": str(args.rank), + "dtype": str(dtype).replace("torch.", ""), + "compute_dtype": ( + str(explicit_compute_dtype).replace("torch.", "") + if explicit_compute_dtype is not None + else "auto(bf16-if-mixed)" + ), + "preserve_norm": str(args.preserve_norm), + "cap_mult": str(args.cap_mult) if args.cap_mult is not None else "None", + "include_patterns": str(include_patterns), + "exclude_patterns": str(exclude_patterns), + "moe_temperature": str(args.moe_temperature), + "moe_hard": str(args.moe_hard), + "auto_rank_threshold": str(args.auto_rank_threshold), + "tool": "WAS LoRA Merger", + "note": note_str, + } + ) + + if args.mode == "obfuscate": + meta.pop("mode", None) + meta.pop("merged_from", None) + meta.pop("rank", None) + meta.pop("compute_dtype", None) + meta.pop("include_patterns", None) + meta.pop("exclude_patterns", None) + meta.pop("moe_temperature", None) + meta.pop("moe_hard", None) + meta.pop("auto_rank_threshold", None) + meta.pop("preserve_norm", None) + meta.pop("cap_mult", None) + + out_dir = os.path.dirname(os.path.abspath(args.out)) or "." + os.makedirs(out_dir, exist_ok=True) + save_file(state, args.out, metadata=meta) + print(f"Saved merged LoRA to: {args.out}") + print(json.dumps(meta, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/nodes/WASEdgeSafeLatentUpscale.py b/nodes/WASEdgeSafeLatentUpscale.py new file mode 100644 index 0000000..9e97042 --- /dev/null +++ b/nodes/WASEdgeSafeLatentUpscale.py @@ -0,0 +1,397 @@ +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import Tuple + +import torch +import torch.nn.functional as F + + +def is_latent_5d(samples: torch.Tensor) -> bool: + return samples.dim() == 5 + + +def flatten_5d_to_4d(samples: torch.Tensor) -> Tuple[torch.Tensor, int, int]: + b, c, t, h, w = samples.shape + flat = samples.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w).contiguous() + return flat, b, t + + +def unflatten_4d_to_5d(samples_btchw: torch.Tensor, b: int, t: int) -> torch.Tensor: + bt, c, h, w = samples_btchw.shape + if bt != b * t: + raise ValueError(f"Shape mismatch: bt={bt} != b*t={b*t}") + return samples_btchw.reshape(b, t, c, h, w).permute(0, 2, 1, 3, 4).contiguous() + + +def resize_tensor(x: torch.Tensor, size_hw: Tuple[int, int], mode: str) -> torch.Tensor: + if mode == "bilinear": + return F.interpolate(x, size=size_hw, mode="bilinear", align_corners=False) + if mode == "bicubic": + return F.interpolate(x, size=size_hw, mode="bicubic", align_corners=False) + if mode == "area": + return F.interpolate(x, size=size_hw, mode="area") + if mode == "nearest-exact": + return F.interpolate(x, size=size_hw, mode="nearest-exact") + return F.interpolate(x, size=size_hw, mode=mode) + + +def gaussian_blur_depthwise(x: torch.Tensor, sigma: float) -> torch.Tensor: + if sigma <= 0.0: + return x + + radius = max(1, int(math.ceil(3.0 * sigma))) + ksize = 2 * radius + 1 + + device = x.device + dtype = x.dtype + + coords = torch.arange(-radius, radius + 1, device=device, dtype=dtype) + kernel_1d = torch.exp(-(coords * coords) / (2.0 * sigma * sigma)) + kernel_1d = kernel_1d / kernel_1d.sum() + + kernel_x = kernel_1d.view(1, 1, 1, ksize) + kernel_y = kernel_1d.view(1, 1, ksize, 1) + + c = x.shape[1] + x_pad = F.pad(x, (radius, radius, radius, radius), mode="reflect") + x_blur = F.conv2d(x_pad, kernel_x.expand(c, 1, 1, ksize), groups=c) + x_blur = F.conv2d(x_blur, kernel_y.expand(c, 1, ksize, 1), groups=c) + return x_blur + + +def sigmoid_weight(d: torch.Tensor, threshold: float, softness: float) -> torch.Tensor: + t = float(threshold) + s = max(1e-8, float(softness)) + return torch.sigmoid((d - t) / s) + + +def apply_temporal_ema(weight_bt1hw: torch.Tensor, b: int, t: int, ema: float) -> torch.Tensor: + if ema <= 0.0 or t <= 1: + return weight_bt1hw + + ema = float(ema) + bt, c, h, w = weight_bt1hw.shape + wgt = weight_bt1hw.reshape(b, t, c, h, w).contiguous() + + for bi in range(b): + prev = wgt[bi, 0] + for ti in range(1, t): + cur = wgt[bi, ti] + prev = prev * ema + cur * (1.0 - ema) + wgt[bi, ti] = prev + + return wgt.reshape(bt, c, h, w).contiguous() + + +def sobel_grad_mag(x_bt1hw: torch.Tensor) -> torch.Tensor: + device = x_bt1hw.device + dtype = x_bt1hw.dtype + + kx = torch.tensor( + [[-1.0, 0.0, 1.0], + [-2.0, 0.0, 2.0], + [-1.0, 0.0, 1.0]], + device=device, + dtype=dtype, + ).view(1, 1, 3, 3) + + ky = torch.tensor( + [[-1.0, -2.0, -1.0], + [0.0, 0.0, 0.0], + [1.0, 2.0, 1.0]], + device=device, + dtype=dtype, + ).view(1, 1, 3, 3) + + gx = F.conv2d(x_bt1hw, kx, padding=1) + gy = F.conv2d(x_bt1hw, ky, padding=1) + mag = torch.sqrt(gx * gx + gy * gy + 1e-12) + return mag + + +def normalize_map(x: torch.Tensor) -> torch.Tensor: + mx = x.amax(dim=(2, 3), keepdim=True) + return x / (mx + 1e-8) + + +def clamp01(x: torch.Tensor) -> torch.Tensor: + return torch.clamp(x, 0.0, 1.0) + + +def compute_damp_mask( + latent_btchw_fp32: torch.Tensor, + weight_bt1hw_fp32: torch.Tensor, + gate_mode: str, + grad_blur_sigma: float, + damp_threshold: float, + damp_softness: float, + damp_power: float, + damp_mask_blur_sigma: float, +) -> torch.Tensor: + energy = latent_btchw_fp32.abs().mean(dim=1, keepdim=True) + + if grad_blur_sigma > 0.0: + energy = gaussian_blur_depthwise(energy, float(grad_blur_sigma)) + + grad = sobel_grad_mag(energy) + grad = normalize_map(grad) + + mask = sigmoid_weight(grad, float(damp_threshold), float(damp_softness)) + mask = clamp01(mask) + + if gate_mode == "weight": + mask = mask * clamp01(weight_bt1hw_fp32) + elif gate_mode == "weight_sqrt": + mask = mask * torch.sqrt(clamp01(weight_bt1hw_fp32) + 1e-8) + elif gate_mode == "none": + pass + else: + pass + + if damp_power != 1.0: + mask = clamp01(mask).pow(float(damp_power)) + + if damp_mask_blur_sigma > 0.0: + mask = gaussian_blur_depthwise(mask, float(damp_mask_blur_sigma)) + mask = clamp01(mask) + + return mask + + +def apply_highpass_damping( + latent_btchw_fp32: torch.Tensor, + damp_mask_bt1hw_fp32: torch.Tensor, + strength: float, + highpass_sigma: float, +) -> torch.Tensor: + s = float(strength) + if s <= 0.0: + return latent_btchw_fp32 + + low = gaussian_blur_depthwise(latent_btchw_fp32, float(highpass_sigma)) if highpass_sigma > 0.0 else latent_btchw_fp32 + high = latent_btchw_fp32 - low + + m = clamp01(damp_mask_bt1hw_fp32) + return low + high * (1.0 - s * m) + + +@dataclass +class AdaptiveBlendConfig: + scale: float = 2.0 + smooth_mode: str = "bilinear" + diff_blur_sigma: float = 0.6 + threshold: float = 0.12 + softness: float = 0.05 + weight_power: float = 1.0 + weight_blur_sigma: float = 0.0 + temporal_ema: float = 0.0 + + enable_directional_damping: bool = True + damping_strength: float = 0.35 + damping_gate_mode: str = "weight_sqrt" # none|weight|weight_sqrt + damping_grad_blur_sigma: float = 0.0 + damping_threshold: float = 0.25 + damping_softness: float = 0.08 + damping_power: float = 1.0 + damping_mask_blur_sigma: float = 0.6 + damping_highpass_sigma: float = 1.0 + damping_temporal_ema: float = 0.25 + + preview_mode: str = "both" # weight|damp|both + output_mask_pixel_scale: int = 8 + + +class WASAdaptiveDifferenceLatentUpscale: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "latent": ("LATENT",), + "scale": ("FLOAT", {"default": 2.0, "min": 1.0, "max": 8.0, "step": 0.05}), + "smooth_mode": (["bilinear", "bicubic", "area"], {"default": "bilinear"}), + + "diff_blur_sigma": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 8.0, "step": 0.05}), + "threshold": ("FLOAT", {"default": 0.12, "min": 0.0, "max": 1.0, "step": 0.005}), + "softness": ("FLOAT", {"default": 0.05, "min": 0.0005, "max": 1.0, "step": 0.001}), + "weight_power": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 6.0, "step": 0.05}), + "weight_blur_sigma": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 8.0, "step": 0.05}), + "temporal_ema": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 0.99, "step": 0.01}), + + "enable_directional_damping": ("BOOLEAN", {"default": True}), + "damping_strength": ("FLOAT", {"default": 0.35, "min": 0.0, "max": 1.0, "step": 0.01}), + "damping_gate_mode": (["none", "weight", "weight_sqrt"], {"default": "weight_sqrt"}), + "damping_grad_blur_sigma": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 8.0, "step": 0.05}), + "damping_threshold": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.005}), + "damping_softness": ("FLOAT", {"default": 0.08, "min": 0.0005, "max": 1.0, "step": 0.001}), + "damping_power": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 6.0, "step": 0.05}), + "damping_mask_blur_sigma": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 8.0, "step": 0.05}), + "damping_highpass_sigma": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 8.0, "step": 0.05}), + "damping_temporal_ema": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 0.99, "step": 0.01}), + + "preview_mode": (["weight", "damp", "both"], {"default": "both"}), + "output_mask_pixel_scale": ("INT", {"default": 8, "min": 1, "max": 16, "step": 1}), + } + } + + RETURN_TYPES = ("LATENT", "MASK", "IMAGE") + RETURN_NAMES = ("latent", "mask", "mask_preview") + FUNCTION = "upscale" + CATEGORY = "WAS/Latent" + + def upscale( + self, + latent, + scale: float, + smooth_mode: str, + diff_blur_sigma: float, + threshold: float, + softness: float, + weight_power: float, + weight_blur_sigma: float, + temporal_ema: float, + enable_directional_damping: bool, + damping_strength: float, + damping_gate_mode: str, + damping_grad_blur_sigma: float, + damping_threshold: float, + damping_softness: float, + damping_power: float, + damping_mask_blur_sigma: float, + damping_highpass_sigma: float, + damping_temporal_ema: float, + preview_mode: str, + output_mask_pixel_scale: int, + ): + if "samples" not in latent: + raise ValueError("LATENT input must be a dict containing key 'samples'.") + + samples: torch.Tensor = latent["samples"] + orig_dtype = samples.dtype + + cfg = AdaptiveBlendConfig( + scale=float(scale), + smooth_mode=str(smooth_mode), + diff_blur_sigma=float(diff_blur_sigma), + threshold=float(threshold), + softness=float(softness), + weight_power=float(weight_power), + weight_blur_sigma=float(weight_blur_sigma), + temporal_ema=float(temporal_ema), + enable_directional_damping=bool(enable_directional_damping), + damping_strength=float(damping_strength), + damping_gate_mode=str(damping_gate_mode), + damping_grad_blur_sigma=float(damping_grad_blur_sigma), + damping_threshold=float(damping_threshold), + damping_softness=float(damping_softness), + damping_power=float(damping_power), + damping_mask_blur_sigma=float(damping_mask_blur_sigma), + damping_highpass_sigma=float(damping_highpass_sigma), + damping_temporal_ema=float(damping_temporal_ema), + preview_mode=str(preview_mode), + output_mask_pixel_scale=int(output_mask_pixel_scale), + ) + + was_5d = is_latent_5d(samples) + if was_5d: + samples_4d, b, t = flatten_5d_to_4d(samples) + h, w = samples.shape[-2], samples.shape[-1] + else: + b = samples.shape[0] + t = 1 + samples_4d = samples + h, w = samples.shape[-2], samples.shape[-1] + + target_h = int(round(h * cfg.scale)) + target_w = int(round(w * cfg.scale)) + if target_h < 1 or target_w < 1: + raise ValueError("Invalid target size computed from scale.") + + base = resize_tensor(samples_4d, (target_h, target_w), mode="nearest-exact") + smooth = resize_tensor(samples_4d, (target_h, target_w), mode=cfg.smooth_mode) + + base_fp = base.float() + smooth_fp = smooth.float() + + diff = (base_fp - smooth_fp).abs().mean(dim=1, keepdim=True) + if cfg.diff_blur_sigma > 0.0: + diff = gaussian_blur_depthwise(diff, cfg.diff_blur_sigma) + + wgt = sigmoid_weight(diff, cfg.threshold, cfg.softness) + wgt = clamp01(wgt) + + if cfg.weight_power != 1.0: + wgt = clamp01(wgt).pow(cfg.weight_power) + + if cfg.weight_blur_sigma > 0.0: + wgt = gaussian_blur_depthwise(wgt, cfg.weight_blur_sigma) + wgt = clamp01(wgt) + + if was_5d and cfg.temporal_ema > 0.0: + wgt = apply_temporal_ema(wgt, b=b, t=t, ema=cfg.temporal_ema) + + out_fp = base_fp * (1.0 - wgt) + smooth_fp * wgt + + damp_mask = torch.zeros_like(wgt) + if cfg.enable_directional_damping and cfg.damping_strength > 0.0: + damp_mask = compute_damp_mask( + latent_btchw_fp32=out_fp, + weight_bt1hw_fp32=wgt, + gate_mode=cfg.damping_gate_mode, + grad_blur_sigma=cfg.damping_grad_blur_sigma, + damp_threshold=cfg.damping_threshold, + damp_softness=cfg.damping_softness, + damp_power=cfg.damping_power, + damp_mask_blur_sigma=cfg.damping_mask_blur_sigma, + ) + + if was_5d and cfg.damping_temporal_ema > 0.0: + damp_mask = apply_temporal_ema(damp_mask, b=b, t=t, ema=cfg.damping_temporal_ema) + + out_fp = apply_highpass_damping( + latent_btchw_fp32=out_fp, + damp_mask_bt1hw_fp32=damp_mask, + strength=cfg.damping_strength, + highpass_sigma=cfg.damping_highpass_sigma, + ) + + out = out_fp.to(dtype=orig_dtype) + + if was_5d: + out = unflatten_4d_to_5d(out, b=b, t=t) + + out_latent = dict(latent) + out_latent["samples"] = out + + mask_for_output = damp_mask if cfg.preview_mode in ("damp", "both") else wgt + mask_out = mask_for_output[:, 0, :, :].contiguous() + + ps = max(1, int(cfg.output_mask_pixel_scale)) + + def to_preview_image(m_bt1hw: torch.Tensor) -> torch.Tensor: + m = m_bt1hw + if ps != 1: + m = resize_tensor(m, (target_h * ps, target_w * ps), mode="nearest-exact") + img = m[:, 0:1, :, :].repeat(1, 3, 1, 1).permute(0, 2, 3, 1).contiguous() + return clamp01(img) + + if cfg.preview_mode == "weight": + prev_img = to_preview_image(wgt) + elif cfg.preview_mode == "damp": + prev_img = to_preview_image(damp_mask) + else: + a = to_preview_image(wgt) + bimg = to_preview_image(damp_mask) + prev_img = torch.cat([a, bimg], dim=2) + + return (out_latent, mask_out, prev_img) + + +NODE_CLASS_MAPPINGS = { + "WAS_AdaptiveDifferenceLatentUpscale": WASAdaptiveDifferenceLatentUpscale, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "WAS_AdaptiveDifferenceLatentUpscale": "WAS Adaptive Difference Latent Upscale (Damped)", +} diff --git a/nodes/WASLatentContrastLimitedDetailBoost.py b/nodes/WASLatentContrastLimitedDetailBoost.py new file mode 100644 index 0000000..8890386 --- /dev/null +++ b/nodes/WASLatentContrastLimitedDetailBoost.py @@ -0,0 +1,232 @@ +from __future__ import annotations + +import math +from typing import Tuple + +import torch +import torch.nn.functional as F + + +def is_latent_5d(samples: torch.Tensor) -> bool: + return samples.dim() == 5 + + +def flatten_5d_to_4d(samples: torch.Tensor) -> Tuple[torch.Tensor, int, int]: + b, c, t, h, w = samples.shape + flat = samples.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w).contiguous() + return flat, b, t + + +def unflatten_4d_to_5d(samples_btchw: torch.Tensor, b: int, t: int) -> torch.Tensor: + bt, c, h, w = samples_btchw.shape + if bt != b * t: + raise ValueError(f"Shape mismatch: bt={bt} != b*t={b*t}") + return samples_btchw.reshape(b, t, c, h, w).permute(0, 2, 1, 3, 4).contiguous() + + +def gaussian_blur_depthwise(x: torch.Tensor, sigma: float) -> torch.Tensor: + if sigma <= 0.0: + return x + + radius = max(1, int(math.ceil(3.0 * sigma))) + ksize = 2 * radius + 1 + + device = x.device + dtype = x.dtype + + coords = torch.arange(-radius, radius + 1, device=device, dtype=dtype) + kernel_1d = torch.exp(-(coords * coords) / (2.0 * sigma * sigma)) + kernel_1d = kernel_1d / kernel_1d.sum() + + kernel_x = kernel_1d.view(1, 1, 1, ksize) + kernel_y = kernel_1d.view(1, 1, ksize, 1) + + c = x.shape[1] + x_pad = F.pad(x, (radius, radius, radius, radius), mode="reflect") + x_blur = F.conv2d(x_pad, kernel_x.expand(c, 1, 1, ksize), groups=c) + x_blur = F.conv2d(x_blur, kernel_y.expand(c, 1, ksize, 1), groups=c) + return x_blur + + +def sobel_grad_mag(x_bt1hw: torch.Tensor) -> torch.Tensor: + device = x_bt1hw.device + dtype = x_bt1hw.dtype + + kx = torch.tensor( + [[-1.0, 0.0, 1.0], + [-2.0, 0.0, 2.0], + [-1.0, 0.0, 1.0]], + device=device, dtype=dtype + ).view(1, 1, 3, 3) + + ky = torch.tensor( + [[-1.0, -2.0, -1.0], + [0.0, 0.0, 0.0], + [1.0, 2.0, 1.0]], + device=device, dtype=dtype + ).view(1, 1, 3, 3) + + gx = F.conv2d(x_bt1hw, kx, padding=1) + gy = F.conv2d(x_bt1hw, ky, padding=1) + return torch.sqrt(gx * gx + gy * gy + 1e-12) + + +def clamp01(x: torch.Tensor) -> torch.Tensor: + return torch.clamp(x, 0.0, 1.0) + + +def resize_mask_for_preview(mask_bt1hw: torch.Tensor, pixel_scale: int) -> torch.Tensor: + ps = max(1, int(pixel_scale)) + if ps == 1: + m = mask_bt1hw + else: + h, w = mask_bt1hw.shape[-2], mask_bt1hw.shape[-1] + m = F.interpolate(mask_bt1hw, size=(h * ps, w * ps), mode="nearest-exact") + img = m[:, 0:1].repeat(1, 3, 1, 1).permute(0, 2, 3, 1).contiguous() + return clamp01(img) + + +class WASLatentContrastLimitedDetailBoost: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "latent": ("LATENT",), + + # Band-pass (DoG) parameters + "sigma_small": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 8.0, "step": 0.05}), + "sigma_large": ("FLOAT", {"default": 1.4, "min": 0.0, "max": 16.0, "step": 0.05}), + + # Strength and limiting + "gain": ("FLOAT", {"default": 0.35, "min": 0.0, "max": 2.0, "step": 0.01}), + "limit": ("FLOAT", {"default": 1.25, "min": 0.1, "max": 8.0, "step": 0.05}), + + # Local energy normalization (prevents halos) + "rms_sigma": ("FLOAT", {"default": 1.2, "min": 0.0, "max": 16.0, "step": 0.05}), + "rms_floor": ("FLOAT", {"default": 0.06, "min": 0.0, "max": 1.0, "step": 0.005}), + + # Optional edge protection (reduces dark outlines at strong boundaries) + "edge_protect": ("FLOAT", {"default": 0.45, "min": 0.0, "max": 1.0, "step": 0.01}), + "edge_sigma": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 8.0, "step": 0.05}), + "edge_threshold": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}), + "edge_softness": ("FLOAT", {"default": 0.10, "min": 0.0005, "max": 1.0, "step": 0.01}), + + # Preview + "preview_mask_scale": ("INT", {"default": 8, "min": 1, "max": 16, "step": 1}), + "preview_mode": (["edge_mask", "detail_mask"], {"default": "detail_mask"}), + } + } + + RETURN_TYPES = ("LATENT", "MASK", "IMAGE") + RETURN_NAMES = ("latent", "mask", "mask_preview") + FUNCTION = "boost" + CATEGORY = "WAS/Latent" + + def boost( + self, + latent, + sigma_small: float, + sigma_large: float, + gain: float, + limit: float, + rms_sigma: float, + rms_floor: float, + edge_protect: float, + edge_sigma: float, + edge_threshold: float, + edge_softness: float, + preview_mask_scale: int, + preview_mode: str, + ): + if "samples" not in latent: + raise ValueError("LATENT input must be a dict containing key 'samples'.") + + samples: torch.Tensor = latent["samples"] + orig_dtype = samples.dtype + + was_5d = is_latent_5d(samples) + if was_5d: + x, b, t = flatten_5d_to_4d(samples) + else: + x = samples + b, t = x.shape[0], 1 + + x_fp = x.float() + + s_small = float(sigma_small) + s_large = float(sigma_large) + if s_large < s_small: + s_small, s_large = s_large, s_small + + # DoG band-pass: emphasizes microdetail without classic unsharp overshoot tendencies + low_small = gaussian_blur_depthwise(x_fp, s_small) if s_small > 0.0 else x_fp + low_large = gaussian_blur_depthwise(x_fp, s_large) if s_large > 0.0 else x_fp + dog = low_small - low_large # band-pass + + # Local RMS normalization (contrast-limited): prevents dark halos / emboss + if float(rms_sigma) > 0.0: + rms = gaussian_blur_depthwise(dog * dog, float(rms_sigma)) + rms = torch.sqrt(torch.clamp(rms, min=0.0) + 1e-8) + else: + rms = torch.sqrt(torch.mean(dog * dog, dim=(2, 3), keepdim=True) + 1e-8) + + rms = rms + float(rms_floor) + dog_n = dog / rms + + # Soft limiter on normalized detail (prevents ringing) + lim = float(limit) + dog_l = torch.tanh(dog_n * lim) / max(1e-6, lim) + + # Detail magnitude mask (for inspection / optional gating) + detail_mag = dog_l.abs().mean(dim=1, keepdim=True) + detail_mag = detail_mag / (detail_mag.amax(dim=(2, 3), keepdim=True) + 1e-8) + detail_mag = clamp01(detail_mag) + + # Edge protection mask (reduce enhancement at strong boundaries) + if float(edge_protect) > 0.0: + energy = x_fp.abs().mean(dim=1, keepdim=True) + if float(edge_sigma) > 0.0: + energy = gaussian_blur_depthwise(energy, float(edge_sigma)) + gmag = sobel_grad_mag(energy) + gmag = gmag / (gmag.amax(dim=(2, 3), keepdim=True) + 1e-8) + + t0 = float(edge_threshold) + s0 = max(1e-6, float(edge_softness)) + edge = torch.sigmoid((gmag - t0) / s0) # 0..1 strong edges -> 1 + edge = clamp01(edge) + + protect = float(edge_protect) + edge_gate = 1.0 - protect * edge + edge_gate = torch.clamp(edge_gate, 0.0, 1.0) + else: + edge = torch.zeros_like(detail_mag) + edge_gate = 1.0 + + # Apply enhancement + out = x_fp + float(gain) * dog_l * edge_gate + out = out.to(dtype=orig_dtype) + + if was_5d: + out = unflatten_4d_to_5d(out, b=b, t=t) + + out_latent = dict(latent) + out_latent["samples"] = out + + # ComfyUI MASK output: [BT,H,W] + if preview_mode == "edge_mask": + mask_bt = edge[:, 0] + prev_img = resize_mask_for_preview(edge, preview_mask_scale) + else: + mask_bt = detail_mag[:, 0] + prev_img = resize_mask_for_preview(detail_mag, preview_mask_scale) + + return (out_latent, mask_bt.contiguous(), prev_img) + + +NODE_CLASS_MAPPINGS = { + "WASLatentContrastLimitedDetailBoost": WASLatentContrastLimitedDetailBoost, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "WASLatentContrastLimitedDetailBoost": "WAS Latent Detail Boost", +} diff --git a/nodes/WASPowerLoraMerger.py b/nodes/WASPowerLoraMerger.py new file mode 100644 index 0000000..ff38d43 --- /dev/null +++ b/nodes/WASPowerLoraMerger.py @@ -0,0 +1,862 @@ +import os +import sys +import json +import re + +from safetensors.torch import save_file +from typing import Any, Dict, List, Optional, Tuple +from tqdm import tqdm + + +import folder_paths +from comfy.utils import ProgressBar + +from ..modules.cli import merge_loras_z +from nodes import LoraLoader + +class WASProgress: + def __init__(self, total: int, desc: str): + self.total = int(total) + self.comfy = ProgressBar(self.total) + self.console = tqdm( + total=self.total, + desc=desc, + unit="step", + dynamic_ncols=True, + miniters=1, + mininterval=0.1, + file=sys.stderr, + ) + + def set_total(self, total: int): + self.total = int(total) + try: + self.console.total = self.total + except Exception: + pass + + def update_absolute(self, value: int): + v = int(value) + try: + if hasattr(self.comfy, "update_absolute"): + self.comfy.update_absolute(v, self.total) + else: + self.comfy.update(1) + except Exception: + pass + + try: + delta = v - int(self.console.n) + if delta > 0: + self.console.update(delta) + self.console.refresh() + except Exception: + pass + + def close(self): + try: + try: + self.console.refresh() + except Exception: + pass + self.console.close() + except Exception: + pass + + +class AnyType(str): + def __ne__(self, __value: object) -> bool: + return False + +any_type = AnyType("*") + + +class FlexibleOptionalInputType(dict): + def __init__(self, type, data: Optional[dict] = None): + self.type = type + self.data = data + if self.data is not None: + for k, v in self.data.items(): + self[k] = v + + def __getitem__(self, key): + if self.data is not None and key in self.data: + return self.data[key] + return (self.type,) + + def __contains__(self, key): + return True + + +class WASPowerLoraMerger: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL", {"tooltip": "Base MODEL. The merged LoRA will be applied to this model after saving."}), + "clip": ("CLIP", {"tooltip": "Base CLIP. The merged LoRA will be applied to this CLIP after saving."}), + "output_filename": ( + "STRING", + { + "default": "merged_lora.safetensors", + "tooltip": "Output filename (relative to ComfyUI models/loras). Must be a relative path. '.safetensors' is appended if missing.", + }, + ), + "output_model_strength": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 100.0, + "step": 0.01, + "tooltip": "Strength used when applying the newly-created LoRA to the output unet model.", + }, + ), + "output_clip_strength": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 100.0, + "step": 0.01, + "tooltip": "Strength used when applying the newly-created LoRA to the output clip model.", + }, + ), + + "mode": ( + ["svd", "rebase", "add", "add-diff", "add-orth", "diff-export", "moe", "obfuscate", "block-mix"], + { + "default": "svd", + "tooltip": "Merge mode. svd=recompress merged delta via SVD; rebase=single-LoRA SVD recompress; add=exact linear stacking; add-diff=base + weighted diffs toward others; add-orth=base + orthogonalized contributions; diff-export=export only the difference between first two; moe=mixture-of-experts per module; obfuscate=stack-equivalent factor rebasis without SVD; block-mix=route modules to LoRA A or B using preset/recipe, merge via stack or svd.", + }, + ), + "block_mix_recipe": ( + [ + "all_a", + "all_b", + "concept_a_style_b", + "concept_b_style_a", + "attn_a_ffn_b", + "attn_b_ffn_a", + "img_a_txt_b", + "img_b_txt_a", + ], + { + "default": "concept_a_style_b", + "tooltip": "block-mix mode only: routing recipe.", + }, + ), + }, + "optional": FlexibleOptionalInputType( + any_type, + { + "options": ( + "WAS_LORA_MERGE_OPTIONS", + { + "tooltip": "Optional advanced merge options dict (use a companion options node).", + }, + ), + }, + ), + "hidden": { + "was_lora_catalog": (["None"] + folder_paths.get_filename_list("loras"),), + }, + } + + RETURN_TYPES = ("MODEL", "CLIP", "STRING") + RETURN_NAMES = ("model", "clip", "lora_path") + FUNCTION = "merge" + CATEGORY = "WAS Extras" + + def merge( + self, + model, + clip, + output_filename: str, + output_model_strength: float, + output_clip_strength: float, + mode: str, + block_mix_recipe: str, + options: Any = None, + **kwargs: Any, + ): + loras: List[Tuple[str, float]] = [] + + def _to_bool(v: Any, default: bool = True) -> bool: + if v is None: + return default + if isinstance(v, bool): + return v + if isinstance(v, (int, float)): + return bool(int(v)) + if isinstance(v, str): + s = v.strip().lower() + if s in ("true", "1", "yes", "y", "on"): + return True + if s in ("false", "0", "no", "n", "off", ""): + return False + return default + return bool(v) + + def _coerce_payload(v: Any) -> Optional[dict]: + if isinstance(v, dict): + return v + if isinstance(v, str): + s = v.strip() + if not s: + return None + try: + obj = json.loads(s) + return obj if isinstance(obj, dict) else None + except Exception: + return None + return None + + payload_rows: List[dict] = [] + flat_enabled: Dict[int, bool] = {} + flat_lora: Dict[int, Optional[str]] = {} + flat_weight: Dict[int, float] = {} + + for key, value in kwargs.items(): + if not isinstance(key, str): + continue + k = key.lower() + if not (k.startswith("lora_") or k.startswith("lora_payload_")): + continue + + if k.startswith("lora_payload_"): + payload = _coerce_payload(value) + if payload is not None: + payload_rows.append(payload) + continue + + m_enabled = re.fullmatch(r"lora_(\d+)_enabled", k) + if m_enabled: + idx = int(m_enabled.group(1)) + flat_enabled[idx] = _to_bool(value, default=True) + continue + + m_weight = re.fullmatch(r"lora_(\d+)_weight", k) + if m_weight: + idx = int(m_weight.group(1)) + try: + flat_weight[idx] = float(value) + except Exception: + flat_weight[idx] = 1.0 + continue + + m_lora = re.fullmatch(r"lora_(\d+)", k) + if m_lora: + idx = int(m_lora.group(1)) + if isinstance(value, str) and value and value != "None": + flat_lora[idx] = value + else: + flat_lora[idx] = None + continue + + if payload_rows: + for payload in payload_rows: + on = _to_bool(payload.get("on", True), default=True) + lora_name = payload.get("lora", None) + weight = float(payload.get("weight", 1.0)) + + if not on: + continue + if lora_name is None or lora_name == "" or lora_name == "None": + continue + if weight == 0.0: + continue + + full_path = folder_paths.get_full_path("loras", lora_name) + if full_path is None: + raise ValueError(f"LoRA not found: {lora_name}") + loras.append((full_path, weight)) + else: + all_indices = sorted(set(flat_enabled.keys()) | set(flat_lora.keys()) | set(flat_weight.keys())) + for idx in all_indices: + on = _to_bool(flat_enabled.get(idx, True), default=True) + lora_name = flat_lora.get(idx, None) + weight = float(flat_weight.get(idx, 1.0)) + + if not on: + continue + if lora_name is None or lora_name == "" or lora_name == "None": + continue + if weight == 0.0: + continue + + full_path = folder_paths.get_full_path("loras", lora_name) + if full_path is None: + raise ValueError(f"LoRA not found: {lora_name}") + loras.append((full_path, weight)) + + if not loras: + raise ValueError("At least one LoRA must be provided.") + + progress = None + progress_step = 0 + + def progress_cb(stage: str, current: int, total: int, message: str | None = None): + nonlocal progress, progress_step + + if progress is None: + progress = WASProgress(1, desc="WAS LoRA Merger") + + stage_total = max(int(total), 1) + stage_current = max(0, min(int(current), stage_total)) + + try: + progress.update_absolute(progress_step + stage_current) + except Exception: + pass + + if mode == "rebase" and len(loras) != 1: + raise ValueError("rebase mode requires exactly one LoRA.") + + if mode == "block-mix" and len(loras) != 2: + raise ValueError("block-mix mode requires exactly two LoRAs (A and B).") + + if mode in ("add-diff", "add-orth", "diff-export") and len(loras) < 2: + raise ValueError(f"{mode} mode requires at least two LoRAs.") + + opt = options if isinstance(options, dict) else {} + + def _opt_bool(key: str, default: bool) -> bool: + return _to_bool(opt.get(key, default), default=default) + + def _opt_int(key: str, default: int) -> int: + try: + return int(opt.get(key, default)) + except Exception: + return int(default) + + def _opt_float(key: str, default: float) -> float: + try: + return float(opt.get(key, default)) + except Exception: + return float(default) + + def _opt_str(key: str, default: str) -> str: + v = opt.get(key, default) + if v is None: + return str(default) + return str(v) + + rank = _opt_int("rank", 32) + auto_rank_threshold = _opt_float("auto_rank_threshold", 0.99) + preserve_norm = _opt_bool("preserve_norm", False) + cap_mult_enable = _opt_bool("cap_mult_enable", False) + cap_mult = _opt_float("cap_mult", 1.0) + dtype = _opt_str("dtype", "bf16") + compute_dtype = _opt_str("compute_dtype", "auto") + cpu = _opt_bool("cpu", False) + include_patterns = _opt_str("include_patterns", "") + exclude_patterns = _opt_str("exclude_patterns", "") + moe_temperature = _opt_float("moe_temperature", 1.0) + moe_hard = _opt_bool("moe_hard", False) + block_mix_method = _opt_str("block_mix_method", "svd") + block_mix_preset = _opt_str("block_mix_preset", "auto") + + include_list = [p.strip() for p in include_patterns.split(",") if p.strip()] + exclude_list = [p.strip() for p in exclude_patterns.split(",") if p.strip()] + + device = merge_loras_z.get_device(force_cpu=cpu) + out_dtype = merge_loras_z.get_dtype(dtype) + explicit_compute_dtype = merge_loras_z.get_compute_dtype(compute_dtype) + + loaded: List[Tuple[Dict[str, merge_loras_z.LoraPair], float]] = [] + ref_meta: Dict[str, str] = {} + + for idx, (path, weight) in enumerate(loras): + pairs, meta = merge_loras_z.load_lora_pairs(path, device=device, progress_cb=progress_cb) + if idx == 0: + ref_meta = meta if isinstance(meta, dict) else {} + loaded.append((pairs, weight)) + + try: + stats = merge_loras_z.summarize_pairs( + pairs, + weight=weight, + explicit_compute_dtype=explicit_compute_dtype, + include_patterns=include_list, + exclude_patterns=exclude_list, + ) + print( + "WASPowerLoraMerger Report(input): " + + json.dumps( + { + "path": path, + "weight": float(weight), + "mode": mode, + "stats": stats, + }, + indent=2, + ) + ) + except Exception as e: + print(f"WASPowerLoraMerger Report(input) failed: {e}") + + all_prefixes = set() + for pairs, _w in loaded: + all_prefixes.update(pairs.keys()) + module_count = max(1, len(all_prefixes)) + + progress_total = (module_count * 2) + 5 + if progress is None: + progress = WASProgress(progress_total, desc="WAS LoRA Merger") + else: + progress.set_total(int(progress_total)) + try: + progress.comfy = ProgressBar(progress.total) + except Exception: + pass + + progress_step = 0 + progress.update_absolute(progress_step) + + def merge_progress_cb(stage: str, current: int, total: int, message: str | None = None): + nonlocal progress_step + if progress is None: + return + t = max(int(total), 1) + c = max(0, min(int(current), t)) + abs_value = min(module_count, int(round((c / t) * module_count))) + progress.update_absolute(abs_value) + + if mode == "add": + merged_pairs = merge_loras_z.merge_mode_add(loaded, include_list, exclude_list, explicit_compute_dtype, progress_cb=merge_progress_cb) + elif mode == "block-mix": + try: + a_pairs, _a_weight = loaded[0] + b_pairs, _b_weight = loaded[1] + routing = merge_loras_z.block_mix_routing_report( + a_pairs=a_pairs, + b_pairs=b_pairs, + preset=block_mix_preset, + recipe=block_mix_recipe, + include_patterns=include_list, + exclude_patterns=exclude_list, + ) + print("WASPowerLoraMerger Report(block-mix.routing): " + json.dumps(routing, indent=2)) + except Exception as e: + print(f"WASPowerLoraMerger Report(block-mix.routing) failed: {e}") + + weighted = bool(opt.get("block_mix_weighted", False)) + if weighted: + concept_mix = float(opt.get("block_mix_concept_mix", 0.5)) + style_mix = float(opt.get("block_mix_style_mix", 0.5)) + merged_pairs = merge_loras_z.merge_mode_block_mix_weighted( + loaded, + method=block_mix_method, + rank_value=rank, + preset=block_mix_preset, + recipe=block_mix_recipe, + concept_mix=concept_mix, + style_mix=style_mix, + include_patterns=include_list, + exclude_patterns=exclude_list, + auto_rank_threshold=auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + progress_cb=merge_progress_cb, + ) + else: + merged_pairs = merge_loras_z.merge_mode_block_mix( + loaded, + method=block_mix_method, + rank_value=rank, + preset=block_mix_preset, + recipe=block_mix_recipe, + include_patterns=include_list, + exclude_patterns=exclude_list, + auto_rank_threshold=auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + progress_cb=merge_progress_cb, + ) + elif mode == "add-diff": + merged_pairs = merge_loras_z.merge_mode_add_diff( + loaded, + rank_value=rank, + include_patterns=include_list, + exclude_patterns=exclude_list, + auto_rank_threshold=auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + progress_cb=merge_progress_cb, + ) + elif mode == "add-orth": + merged_pairs = merge_loras_z.merge_mode_add_orth( + loaded, + rank_value=rank, + include_patterns=include_list, + exclude_patterns=exclude_list, + auto_rank_threshold=auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + progress_cb=merge_progress_cb, + ) + elif mode == "diff-export": + merged_pairs = merge_loras_z.merge_mode_diff_export( + loaded, + rank_value=rank, + include_patterns=include_list, + exclude_patterns=exclude_list, + auto_rank_threshold=auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + progress_cb=merge_progress_cb, + ) + elif mode == "moe": + merged_pairs = merge_loras_z.merge_mode_moe( + loaded, + rank_value=rank, + moe_temperature=moe_temperature, + moe_hard=moe_hard, + include_patterns=include_list, + exclude_patterns=exclude_list, + auto_rank_threshold=auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + progress_cb=merge_progress_cb, + ) + elif mode == "obfuscate": + merged_pairs = merge_loras_z.merge_mode_obfuscate( + loaded, + include_patterns=include_list, + exclude_patterns=exclude_list, + explicit_compute_dtype=explicit_compute_dtype, + progress_cb=merge_progress_cb, + ) + elif mode == "rebase": + merged_pairs = merge_loras_z.merge_mode_rebase( + loaded[0], + rank_value=rank, + include_patterns=include_list, + exclude_patterns=exclude_list, + auto_rank_threshold=auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + progress_cb=merge_progress_cb, + ) + else: + merged_pairs = merge_loras_z.merge_mode_svd( + loaded, + rank_value=rank, + preserve_norm=preserve_norm, + cap_mult=(cap_mult if cap_mult_enable else None), + include_patterns=include_list, + exclude_patterns=exclude_list, + auto_rank_threshold=auto_rank_threshold, + explicit_compute_dtype=explicit_compute_dtype, + progress_cb=merge_progress_cb, + ) + + def build_progress_cb(stage: str, current: int, total: int, message: str | None = None): + nonlocal progress_step + if progress is None: + return + t = max(int(total), 1) + c = max(0, min(int(current), t)) + abs_value = module_count + min(module_count, int(round((c / t) * module_count))) + progress.update_absolute(abs_value) + + state, meta = merge_loras_z.build_state_dict( + merged_pairs, + metadata_ref=ref_meta, + dtype=out_dtype, + device=device, + progress_cb=build_progress_cb, + ) + + try: + merged_stats = merge_loras_z.summarize_pairs( + merged_pairs, + weight=1.0, + explicit_compute_dtype=explicit_compute_dtype, + include_patterns=include_list, + exclude_patterns=exclude_list, + ) + print("WASPowerLoraMerger Report(merged): " + json.dumps(merged_stats, indent=2)) + except Exception as e: + print(f"WASPowerLoraMerger Report(merged) failed: {e}") + + if progress is not None: + progress.update_absolute((module_count * 2) + 1) + + if mode == "obfuscate": + note_str = "Created by WAS Merge Loras" + else: + note_str = f"Created by WAS Merge Loras - Merging Mode: {mode}" + + merged_from = [(p, w) for p, w in loras] if mode != "obfuscate" else None + + meta.update( + { + "merged_from": str(merged_from), + "mode": mode, + "rank": str(rank), + "dtype": str(out_dtype).replace("torch.", ""), + "compute_dtype": ( + str(explicit_compute_dtype).replace("torch.", "") + if explicit_compute_dtype is not None + else "auto(bf16-if-mixed)" + ), + "preserve_norm": str(preserve_norm), + "cap_mult": str(cap_mult) if cap_mult_enable else "None", + "include_patterns": str(include_list), + "exclude_patterns": str(exclude_list), + "moe_temperature": str(moe_temperature), + "moe_hard": str(moe_hard), + "auto_rank_threshold": str(auto_rank_threshold), + "tool": "WAS Merge Loras", + "note": note_str, + } + ) + + if mode == "obfuscate": + meta.pop("mode", None) + meta.pop("merged_from", None) + meta.pop("include_patterns", None) + meta.pop("exclude_patterns", None) + + try: + from folder_paths import folder_names_and_paths + + lora_dirs = folder_names_and_paths.get("loras", [[], []])[0] + lora_dir = lora_dirs[0] if lora_dirs else None + except Exception: + lora_dir = None + + if not lora_dir: + raise RuntimeError("Unable to resolve the LoRA output directory.") + + rel_out = (output_filename or "").strip().replace("\\\\", "/") + while rel_out.startswith("/"): + rel_out = rel_out[1:] + if not rel_out: + rel_out = "merged_lora.safetensors" + if rel_out.startswith("../") or "/../" in rel_out or rel_out == "..": + raise ValueError("output_filename must not contain '..' path traversal.") + if ":" in rel_out: + raise ValueError("output_filename must be a relative path under models/loras.") + + if not rel_out.lower().endswith(".safetensors"): + rel_out = rel_out + ".safetensors" + + lora_root_abs = os.path.abspath(lora_dir) + out_path = os.path.abspath(os.path.join(lora_root_abs, *rel_out.split("/"))) + if os.path.commonpath([lora_root_abs, out_path]) != lora_root_abs: + raise ValueError("`output_filename` resolves outside models/loras directory.\nFor security reasons, LoRA cannot be saved to this location.") + + os.makedirs(os.path.dirname(out_path), exist_ok=True) + + save_file(state, out_path, metadata=meta) + if os.path.exists(out_path) and os.path.getsize(out_path) > 0: + print(f"Saved merged LoRA to {out_path}") + + model, clip = LoraLoader().load_lora(model, clip, rel_out, output_model_strength, output_clip_strength) + + if progress is not None: + progress.update_absolute(progress.total) + progress.close() + + return (model, clip, rel_out) + + +class WASPowerLoraMergerOptions: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "rank": ( + "INT", + { + "default": 32, + "min": 0, + "max": 4096, + "step": 1, + "tooltip": "Rank for SVD-based modes (svd, rebase, add-diff, add-orth, diff-export, moe). Set to 0 to use auto-rank (energy threshold).", + }, + ), + "auto_rank_threshold": ( + "FLOAT", + { + "default": 0.99, + "min": 0.5, + "max": 1.0, + "step": 0.0001, + "tooltip": "Auto-rank energy threshold for SVD-based modes when rank=0. Higher values keep more singular-value energy (larger rank).", + }, + ), + "preserve_norm": ( + "BOOLEAN", + { + "default": False, + "tooltip": "svd mode only: preserve average per-module delta norm after merging (helps prevent overall strength drift).", + }, + ), + "cap_mult_enable": ( + "BOOLEAN", + { + "default": False, + "tooltip": "svd mode only: enable capping merged per-module norm (cap_mult ร— mean(source_norms)).", + }, + ), + "cap_mult": ( + "FLOAT", + { + "default": 1.0, + "min": 0.0, + "max": 100.0, + "step": 0.01, + "tooltip": "svd mode only (when cap_mult_enable is on): cap merged module norm to cap_mult ร— mean(source module norms).", + }, + ), + "dtype": ( + ["fp16", "fp32", "bf16"], + { + "default": "bf16", + "tooltip": "Output dtype for the saved LoRA tensors.", + }, + ), + "compute_dtype": ( + ["auto", "bf16", "fp16", "fp32"], + { + "default": "auto", + "tooltip": "Internal merge dtype alignment. auto => if mixed dtypes are encountered, prefer bf16; otherwise use existing dtype.", + }, + ), + "cpu": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Force CPU even if CUDA is available (slower but avoids VRAM usage).", + }, + ), + "include_patterns": ( + "STRING", + { + "default": "", + "multiline": True, + "tooltip": "Optional filter: only merge modules whose prefix contains any of these substrings. Comma-separated list (whitespace is trimmed).", + }, + ), + "exclude_patterns": ( + "STRING", + { + "default": "", + "multiline": True, + "tooltip": "Optional filter: exclude modules whose prefix contains any of these substrings. Comma-separated list (whitespace is trimmed).", + }, + ), + "moe_temperature": ( + "FLOAT", + { + "default": 1.0, + "min": 1e-6, + "max": 100.0, + "step": 0.01, + "tooltip": "moe mode only: softmax temperature for expert gating (lower = sharper selection).", + }, + ), + "moe_hard": ( + "BOOLEAN", + { + "default": False, + "tooltip": "moe mode only: hard gating (pick a single best expert per module) instead of soft mixture.", + }, + ), + "block_mix_method": ( + ["svd", "stack"], + { + "default": "svd", + "tooltip": "block-mix mode only: svd (delta merge then SVD) or stack (exact rank stacking).", + }, + ), + "block_mix_preset": ( + ["auto", "zimg-turbo", "flux", "wan", "qwen", "sd", "sdxl", "generic"], + { + "default": "auto", + "tooltip": "block-mix mode only: routing preset family.", + }, + ), + "block_mix_weighted": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Enable weighted block-mix (blend A/B per module using role-specific mix ratios).", + }, + ), + "block_mix_concept_mix": ( + "FLOAT", + { + "default": 0.5, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "When weighted block-mix is enabled: fraction of LoRA A to use for concept/attention modules (B = 1 - A).", + }, + ), + "block_mix_style_mix": ( + "FLOAT", + { + "default": 0.5, + "min": 0.0, + "max": 1.0, + "step": 0.01, + "tooltip": "When weighted block-mix is enabled: fraction of LoRA A to use for style/FFN modules (B = 1 - A).", + }, + ), + } + } + + RETURN_TYPES = ("WAS_LORA_MERGE_OPTIONS",) + RETURN_NAMES = ("options",) + FUNCTION = "build" + CATEGORY = "WAS Extras" + + def build( + self, + rank: int, + auto_rank_threshold: float, + preserve_norm: bool, + cap_mult_enable: bool, + cap_mult: float, + dtype: str, + compute_dtype: str, + cpu: bool, + include_patterns: str, + exclude_patterns: str, + moe_temperature: float, + moe_hard: bool, + block_mix_method: str, + block_mix_preset: str, + block_mix_weighted: bool, + block_mix_concept_mix: float, + block_mix_style_mix: float, + ): + options = { + "rank": int(rank), + "auto_rank_threshold": float(auto_rank_threshold), + "preserve_norm": bool(preserve_norm), + "cap_mult_enable": bool(cap_mult_enable), + "cap_mult": float(cap_mult), + "dtype": str(dtype), + "compute_dtype": str(compute_dtype), + "cpu": bool(cpu), + "include_patterns": str(include_patterns), + "exclude_patterns": str(exclude_patterns), + "moe_temperature": float(moe_temperature), + "moe_hard": bool(moe_hard), + "block_mix_method": str(block_mix_method), + "block_mix_preset": str(block_mix_preset), + "block_mix_weighted": bool(block_mix_weighted), + "block_mix_concept_mix": float(block_mix_concept_mix), + "block_mix_style_mix": float(block_mix_style_mix), + } + return (options,) + + +NODE_CLASS_MAPPINGS = { + "WASPowerLoraMerger": WASPowerLoraMerger, + "WASPowerLoraMergerOptions": WASPowerLoraMergerOptions, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "WASPowerLoraMerger": "WAS Power LoRA Merger", + "WASPowerLoraMergerOptions": "WAS Power LoRA Merger Options", +} diff --git a/requirements.txt b/requirements.txt index 558d8fd..0165ea8 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,2 +1,4 @@ rich -numpy \ No newline at end of file +numpy +tqdm +safetensors \ No newline at end of file diff --git a/web/was_power_lora_merger.js b/web/was_power_lora_merger.js new file mode 100644 index 0000000..a0516e8 --- /dev/null +++ b/web/was_power_lora_merger.js @@ -0,0 +1,719 @@ +import { app } from "../../scripts/app.js"; +import { api } from "../../scripts/api.js"; + +const EXT_NAME = "WAS_Extras.PowerLoraMergerUI"; +const NODE_NAME = "WASPowerLoraMerger"; + +function toBool(v, defaultValue = true) { + if (v == null) return defaultValue; + if (typeof v === "boolean") return v; + if (typeof v === "number") return !!v; + if (typeof v === "string") { + const s = v.trim().toLowerCase(); + if (s === "true" || s === "1" || s === "yes" || s === "y" || s === "on") return true; + if (s === "false" || s === "0" || s === "no" || s === "n" || s === "off" || s === "") return false; + return defaultValue; + } + return !!v; +} + +function sleep(ms) { + return new Promise((resolve) => setTimeout(resolve, ms)); +} + +async function refreshNodeDefsAndUpdate(node) { + try { + const command = app?.extensionManager?.command; + if (command && typeof command.execute === "function") { + try { + await command.execute("Comfy.RefreshNodeDefinitions"); + } catch (e) { + } + } + + const prev = getLoraOptions(node); + let fromBackend = null; + + for (let attempt = 0; attempt < 8; attempt++) { + const defs = await api.getNodeDefs(); + const def = defs?.[NODE_NAME] ?? (Array.isArray(defs) ? defs.find((d) => d?.name === NODE_NAME) : null); + const cat = def?.input?.hidden?.was_lora_catalog; + if (Array.isArray(cat) && cat.length) { + const next = normalizeLoraOptions(cat); + const changed = JSON.stringify(next) !== JSON.stringify(prev); + if (changed || attempt === 7) { + fromBackend = next; + break; + } + } + + await sleep(150); + } + + if (Array.isArray(fromBackend) && fromBackend.length) { + setLoraOptions(node, fromBackend); + rebuildLoraRows(node, fromBackend); + node.setDirtyCanvas(true, true); + return; + } + } catch (e) { + } +} + +function showLoraChooser(event, callback, parentMenu, loras) { + const canvas = app.canvas; + const safeLoras = normalizeLoraOptions(loras); + const nestedMenuValues = buildNestedLoraMenuValues(safeLoras, callback); + + let safeEvent = event; + if (!(safeEvent instanceof MouseEvent) && !(safeEvent instanceof CustomEvent)) { + try { + const canvasEl = canvas?.canvas; + const rect = canvasEl?.getBoundingClientRect?.(); + const mx = canvas?.mouse?.[0] ?? canvas?.last_mouse?.[0] ?? 0; + const my = canvas?.mouse?.[1] ?? canvas?.last_mouse?.[1] ?? 0; + const clientX = (rect?.left ?? 0) + mx; + const clientY = (rect?.top ?? 0) + my; + safeEvent = new MouseEvent("contextmenu", { + bubbles: true, + cancelable: true, + clientX, + clientY, + }); + } catch (e) { + safeEvent = new MouseEvent("contextmenu", { bubbles: true, cancelable: true }); + } + } + + new LiteGraph.ContextMenu(nestedMenuValues, { + event: safeEvent, + parentMenu: parentMenu ?? undefined, + title: "WAS LoRA Picker", + scale: Math.max(1, canvas?.ds?.scale ?? 1), + className: "dark", + callback, + }); +} + +function splitLoraPath(path) { + const normalized = String(path ?? "").replace(/\\/g, "/"); + return normalized.split("/").filter((p) => p.length > 0); +} + +function buildLoraTree(loras) { + const root = { folders: new Map(), files: [], all: [] }; + + for (const raw of loras) { + if (typeof raw !== "string") continue; + if (raw === "None") continue; + + const parts = splitLoraPath(raw); + if (!parts.length) continue; + + let node = root; + node.all.push(raw); + for (let i = 0; i < parts.length; i++) { + const part = parts[i]; + const isLeaf = i === parts.length - 1; + if (isLeaf) { + node.files.push({ name: part, full: raw }); + } else { + if (!node.folders.has(part)) { + node.folders.set(part, { folders: new Map(), files: [], all: [] }); + } + node = node.folders.get(part); + node.all.push(raw); + } + } + } + + return root; +} + +function filterLoras(loras, query) { + const q = String(query ?? "").trim().toLowerCase(); + if (!q) return Array.isArray(loras) ? loras.slice() : []; + const arr = Array.isArray(loras) ? loras : []; + return arr.filter((x) => typeof x === "string" && x.toLowerCase().includes(q)); +} + +function makeSearchMenuItem(title, allLoras, onPick) { + return { + content: title, + callback: (_value, _options, event, parentMenu, node) => { + let q = ""; + try { + q = window?.prompt?.("Filter LoRAs", "") ?? ""; + } catch (e) { + q = ""; + } + + const filtered = filterLoras(allLoras, q); + const t = buildLoraTree(filtered); + const menuValues = treeToMenuValues(t, onPick, filtered); + + new LiteGraph.ContextMenu(menuValues, { + event, + parentMenu: parentMenu ?? undefined, + title: "WAS LoRA Picker", + scale: Math.max(1, app?.canvas?.ds?.scale ?? 1), + className: "dark", + callback: (_v, _o, _e, _pm, _n) => { + }, + }, node); + }, + }; +} + +function treeToMenuValues(treeNode, onPick, allLorasForNode) { + const values = []; + + // Folder-aware search: show a filter action for the current folder/subtree. + const folderAll = Array.isArray(allLorasForNode) ? allLorasForNode : (treeNode?.all ?? []); + if (Array.isArray(folderAll) && folderAll.length) { + values.push(makeSearchMenuItem("๐Ÿ”Ž Filter in this folder", folderAll, onPick)); + values.push(null); + } + + const folderNames = Array.from(treeNode.folders.keys()).sort((a, b) => + a.localeCompare(b, undefined, { numeric: true, sensitivity: "base" }), + ); + + for (const folderName of folderNames) { + const child = treeNode.folders.get(folderName); + values.push({ + content: `๐Ÿ“ ${folderName}`, + has_submenu: true, + callback: () => { + }, + submenu: { + options: treeToMenuValues(child, onPick, child?.all ?? []), + }, + }); + } + + const files = Array.isArray(treeNode.files) ? treeNode.files.slice() : []; + files.sort((a, b) => a.name.localeCompare(b.name, undefined, { numeric: true, sensitivity: "base" })); + for (const f of files) { + values.push({ + content: f.name, + rgthree_originalValue: f.full, + callback: (_value, options, event, parentMenu, node) => { + onPick?.(f.full, options, event, parentMenu, node); + }, + }); + } + + return values; +} + +function buildNestedLoraMenuValues(loras, onPick) { + const safe = normalizeLoraOptions(loras); + + const out = []; + + // Root-level search across the entire LoRA catalog. + const safeNoNone = safe.filter((x) => typeof x === "string" && x !== "None"); + if (safeNoNone.length) { + out.push(makeSearchMenuItem("๐Ÿ”Ž Filter LoRAs", safeNoNone, onPick)); + out.push(null); + } + + if (safe.includes("None")) { + out.push({ + content: "None", + rgthree_originalValue: "None", + callback: (_value, options, event, parentMenu, node) => { + onPick?.("None", options, event, parentMenu, node); + }, + }); + } + + const tree = buildLoraTree(safe); + const nested = treeToMenuValues(tree, onPick, tree?.all ?? []); + out.push(...nested); + + return out; +} + +function pickedLoraValue(value) { + if (typeof value === "string") return value; + if (value && typeof value === "object") { + // Folder items should never be treated as a picked LoRA. + if (value.has_submenu || value.submenu) return null; + const orig = value.rgthree_originalValue; + if (typeof orig === "string") return orig; + const c = value.content; + if (typeof c === "string") { + if (c.startsWith("๐Ÿ“ ")) return null; + return c; + } + } + return null; +} + +function ensureState(node) { + if (!node.properties) node.properties = {}; + if (!Array.isArray(node.properties.was_lora_rows)) { + node.properties.was_lora_rows = []; + } + return node.properties.was_lora_rows; +} + +function normalizeLoraOptions(options) { + let out = options; + if (Array.isArray(out) && out.length === 1 && Array.isArray(out[0])) { + out = out[0]; + } + if (!Array.isArray(out)) return ["None"]; + const filtered = out.filter((x) => typeof x === "string"); + if (!filtered.length) return ["None"]; + if (filtered[0] !== "None") { + if (!filtered.includes("None")) filtered.unshift("None"); + } + return filtered; +} + +function setLoraOptions(node, options) { + if (!node.properties) node.properties = {}; + node.properties.was_lora_options = normalizeLoraOptions(options); +} + +function getLoraOptions(node) { + const opts = node?.properties?.was_lora_options; + return normalizeLoraOptions(opts); +} + +function makeHiddenPayloadWidget(name, stateRef) { + return { + name, + type: "custom", + value: stateRef, + computeSize() { + return [0, 0]; + }, + draw() { + }, + serializeValue() { + return { ...stateRef }; + }, + }; +} + +function makeSectionHeaderWidget(name, title) { + return { + name, + type: "custom", + value: title, + computeSize(width) { + return [width ?? 0, 22]; + }, + draw(ctx, node, width, y, height) { + try { + const h = height ?? 22; + ctx.save(); + ctx.globalAlpha = 1.0; + ctx.fillStyle = "rgba(120, 170, 255, 0.15)"; + ctx.fillRect(0, y, width, h); + ctx.fillStyle = "rgba(210, 230, 255, 0.95)"; + ctx.font = "12px sans-serif"; + ctx.textBaseline = "middle"; + ctx.fillText(title, 10, y + h / 2); + ctx.restore(); + } catch (e) { + } + }, + serializeValue() { + return title; + }, + }; +} + +function clearWasWidgets(node) { + if (!node.widgets) return; + node.widgets = node.widgets.filter((w) => { + const n = w?.name; + const rowId = w?.options?.id; + if (rowId === "row_add") return false; + if (rowId === "row_remove_last") return false; + if (rowId === "row_clear") return false; + if (rowId === "row_refresh") return false; + if (!n) return true; + return !(n.startsWith("was_row_") || n.startsWith("lora_")); + }); +} + +function syncRowsFromWidgets(node) { + const rows = ensureState(node); + const widgets = Array.isArray(node?.widgets) ? node.widgets : []; + + let maxIdxSeen = -1; + + for (const w of widgets) { + const name = w?.name; + if (typeof name !== "string") continue; + + const mCombo = /^lora_(\d+)$/.exec(name); + if (mCombo) { + const idx = Math.max(0, Number(mCombo[1]) - 1); + maxIdxSeen = Math.max(maxIdxSeen, idx); + while (rows.length <= idx) rows.push({ on: true, lora: null, weight: 1.0 }); + const v = w?.value; + rows[idx].lora = typeof v === "string" && v !== "None" ? v : null; + continue; + } + + const mWeight = /^lora_(\d+)_weight$/.exec(name); + if (mWeight) { + const idx = Math.max(0, Number(mWeight[1]) - 1); + maxIdxSeen = Math.max(maxIdxSeen, idx); + while (rows.length <= idx) rows.push({ on: true, lora: null, weight: 1.0 }); + const n = Number(w?.value); + rows[idx].weight = Number.isFinite(n) ? n : 1.0; + } + + const mEnabled = /^lora_(\d+)_enabled$/.exec(name); + if (mEnabled) { + const idx = Math.max(0, Number(mEnabled[1]) - 1); + maxIdxSeen = Math.max(maxIdxSeen, idx); + while (rows.length <= idx) rows.push({ on: true, lora: null, weight: 1.0 }); + rows[idx].on = toBool(w?.value, true); + } + } + + const targetLen = Math.max(1, maxIdxSeen + 1); + if (rows.length > targetLen) { + rows.length = targetLen; + } + + return rows; +} + +function rebuildLoraRows(node, loraOptions, sync = true) { + const rows = sync ? syncRowsFromWidgets(node) : ensureState(node); + clearWasWidgets(node); + + node.addCustomWidget(makeSectionHeaderWidget("was_row_header", "Selected LoRA's")); + + const resolvedOptions = normalizeLoraOptions(loraOptions); + + for (const row of rows) { + if (row && typeof row.lora === "string" && row.lora !== "None" && !resolvedOptions.includes(row.lora)) { + resolvedOptions.push(row.lora); + } + } + setLoraOptions(node, resolvedOptions); + + rows.forEach((row, idx) => { + const rowIndex = idx + 1; + const payloadName = `lora_payload_${rowIndex}`; + node.addCustomWidget(makeHiddenPayloadWidget(payloadName, row)); + + if (typeof row.lora !== "string") row.lora = row.lora == null ? null : String(row.lora); + if (!Number.isFinite(row.weight)) row.weight = 1.0; + if (typeof row.on !== "boolean") row.on = toBool(row.on, true); + + node.addWidget( + "toggle", + `lora_${rowIndex}_enabled`, + !!row.on, + (v) => { + row.on = !!v; + }, + { label: `lora_${rowIndex}_enabled` }, + ); + + const comboWidget = node.addWidget( + "combo", + `lora_${rowIndex}`, + row.lora ?? "None", + (v) => { + row.lora = v === "None" ? null : v; + }, + { values: resolvedOptions, label: `lora_${rowIndex}` }, + ); + + if (comboWidget && comboWidget.options) { + comboWidget.options.values = resolvedOptions; + } + + node.addWidget( + "number", + `lora_${rowIndex}_weight`, + Number.isFinite(row.weight) ? row.weight : 1.0, + (v) => { + const n = Number(v); + row.weight = Number.isFinite(n) ? n : 1.0; + }, + { min: -10.0, max: 10.0, step: 0.01, precision: 3, label: `lora_${rowIndex}_strength` }, + ); + }); + + node.addWidget( + "button", + "โž• Add LoRA", + null, + (...args) => { + const event = args?.[1] ?? args?.[0]; + const opts = getLoraOptions(node); + showLoraChooser( + event, + (value) => { + const picked = pickedLoraValue(value); + if (typeof picked === "string" && picked !== "None") { + syncRowsFromWidgets(node); + const curRows = ensureState(node); + curRows.push({ on: true, lora: picked, weight: 1.0 }); + rebuildLoraRows(node, getLoraOptions(node), false); + const computed = node.computeSize?.() ?? [node.size[0], node.size[1]]; + node.size[1] = Math.max(node.size[1], computed[1]); + node.setDirtyCanvas(true, true); + } + }, + null, + opts, + ); + }, + { id: "row_add" }, + ); + + node.addWidget( + "button", + "โž– Remove Last LoRA", + null, + () => { + syncRowsFromWidgets(node); + const curRows = ensureState(node); + if (!curRows.length) return; + curRows.pop(); + rebuildLoraRows(node, getLoraOptions(node), false); + const computed = node.computeSize?.() ?? [node.size[0], node.size[1]]; + node.size[1] = Math.max(node.size[1], computed[1]); + node.setDirtyCanvas(true, true); + }, + { id: "row_remove_last" }, + ); + + node.addWidget( + "button", + "๐Ÿงน Clear LoRAs", + null, + () => { + syncRowsFromWidgets(node); + const curRows = ensureState(node); + if (!curRows.length) return; + curRows.length = 0; + curRows.push({ on: true, lora: null, weight: 1.0 }); + rebuildLoraRows(node, getLoraOptions(node), false); + const computed = node.computeSize?.() ?? [node.size[0], node.size[1]]; + node.size[1] = Math.max(node.size[1], computed[1]); + node.setDirtyCanvas(true, true); + }, + { id: "row_clear" }, + ); + + node.addWidget( + "button", + "โ™ป๏ธ Refresh LoRA List", + null, + async () => { + await refreshNodeDefsAndUpdate(node); + }, + { id: "row_refresh" }, + ); + + try { + const computed = node.computeSize?.() ?? null; + if (computed && Array.isArray(computed) && computed.length >= 2) { + node.size[1] = computed[1]; + } + } catch (e) { + } +} + +app.registerExtension({ + name: EXT_NAME, + getNodeMenuItems(node) { + if (node?.comfyClass !== NODE_NAME) return []; + + const rows = ensureState(node); + const opts = getLoraOptions(node); + + return [ + null, + { + content: "โž• Add LoRA", + callback: (_value, _options, event) => { + showLoraChooser( + event, + (picked) => { + const lora = pickedLoraValue(picked); + if (typeof lora === "string" && lora !== "None") { + syncRowsFromWidgets(node); + const curRows = ensureState(node); + curRows.push({ on: true, lora: lora, weight: 1.0 }); + rebuildLoraRows(node, getLoraOptions(node), false); + node.setDirtyCanvas(true, true); + } + }, + null, + opts, + ); + }, + }, + { + content: "โž– Remove Last LoRA", + disabled: rows.length === 0, + callback: () => { + syncRowsFromWidgets(node); + const curRows = ensureState(node); + curRows.pop(); + rebuildLoraRows(node, getLoraOptions(node), false); + node.setDirtyCanvas(true, true); + }, + }, + { + content: "โŒ Clear LoRAs", + disabled: rows.length === 0, + callback: () => { + syncRowsFromWidgets(node); + const curRows = ensureState(node); + curRows.length = 0; + curRows.push({ on: true, lora: null, weight: 1.0 }); + rebuildLoraRows(node, getLoraOptions(node), false); + node.setDirtyCanvas(true, true); + }, + }, + { + content: "โ™ป๏ธ Refresh LoRA List", + callback: async () => { + await refreshNodeDefsAndUpdate(node); + }, + }, + ]; + }, + async beforeRegisterNodeDef(nodeType, nodeData) { + if (nodeData?.name !== NODE_NAME) return; + + const backendCatalog = nodeData?.input?.hidden?.was_lora_catalog; + + const onNodeCreated = nodeType.prototype.onNodeCreated; + nodeType.prototype.onNodeCreated = function () { + onNodeCreated?.apply(this, arguments); + + if (!this.widgets) this.widgets = []; + + if (Array.isArray(backendCatalog) && backendCatalog.length) { + setLoraOptions(this, backendCatalog); + } + + const rows = ensureState(this); + if (!rows.length) { + rows.push({ on: true, lora: null, weight: 1.0 }); + } + + const opts = getLoraOptions(this); + rebuildLoraRows(this, opts); + const computed = this.computeSize?.() ?? [this.size[0], this.size[1]]; + this.size[1] = computed[1]; + this.setDirtyCanvas(true, true); + + try { + if (!this.properties) this.properties = {}; + if (!this.properties._was_lora_catalog_refresh_pending) { + this.properties._was_lora_catalog_refresh_pending = true; + setTimeout(async () => { + try { + await refreshNodeDefsAndUpdate(this); + } catch (e) { + } + try { + this.properties._was_lora_catalog_refresh_pending = false; + } catch (e) { + } + }, 0); + } + } catch (e) { + } + }; + + const configure = nodeType.prototype.configure; + nodeType.prototype.configure = function (info) { + const widgetValues = info?.widgets_values || []; + const rows = ensureState(this); + rows.length = 0; + + for (const v of widgetValues) { + if (v && typeof v === "object" && Object.prototype.hasOwnProperty.call(v, "lora")) { + rows.push({ + on: toBool(v.on, true), + lora: v.lora ?? null, + weight: Number.isFinite(v.weight) ? v.weight : 1.0, + }); + } + } + + if (!rows.length) { + rows.push({ on: true, lora: null, weight: 1.0 }); + } + + if (Array.isArray(backendCatalog) && backendCatalog.length) { + setLoraOptions(this, backendCatalog); + } + + const opts = getLoraOptions(this); + rebuildLoraRows(this, opts); + const computed = this.computeSize?.() ?? [this.size[0], this.size[1]]; + this.size[1] = computed[1]; + this.setDirtyCanvas(true, true); + + try { + if (!this.properties) this.properties = {}; + if (!this.properties._was_lora_catalog_refresh_pending) { + this.properties._was_lora_catalog_refresh_pending = true; + setTimeout(async () => { + try { + await refreshNodeDefsAndUpdate(this); + } catch (e) { + } + try { + this.properties._was_lora_catalog_refresh_pending = false; + } catch (e) { + } + }, 0); + } + } catch (e) { + } + + configure?.apply(this, arguments); + }; + + const onSerialize = nodeType.prototype.onSerialize; + nodeType.prototype.onSerialize = function (o) { + try { + const rows = ensureState(this); + const safeRows = rows.map((r) => { + const lora = typeof r?.lora === "string" ? r.lora : null; + const weight = Number.isFinite(Number(r?.weight)) ? Number(r.weight) : 1.0; + return { on: !!r?.on, lora, weight }; + }); + o.properties = o.properties || {}; + o.properties.was_lora_rows = safeRows; + } catch (e) { + } + + return onSerialize?.apply(this, arguments); + }; + + const refreshComboInNode = nodeType.prototype.refreshComboInNode; + nodeType.prototype.refreshComboInNode = function (defs) { + const fromBackend = defs?.input?.hidden?.was_lora_catalog; + if (Array.isArray(fromBackend) && fromBackend.length) { + setLoraOptions(this, fromBackend); + rebuildLoraRows(this, fromBackend); + this.setDirtyCanvas(true, true); + } + return refreshComboInNode?.apply(this, arguments); + }; + }, +});