Add Power Lora Merger

This commit is contained in:
Jordan Thompson
2025-12-23 23:51:14 -08:00
parent ac1ab31df6
commit 932d928847
10 changed files with 4338 additions and 2 deletions
+4 -1
View File
@@ -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"]
View File
View File
+148
View File
@@ -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()
File diff suppressed because it is too large Load Diff
+397
View File
@@ -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)",
}
@@ -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",
}
+862
View File
@@ -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",
}
+3 -1
View File
@@ -1,2 +1,4 @@
rich
numpy
numpy
tqdm
safetensors
+719
View File
@@ -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);
};
},
});