Add Power Lora Merger
This commit is contained in:
+4
-1
@@ -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"]
|
||||
|
||||
@@ -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
@@ -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",
|
||||
}
|
||||
@@ -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
@@ -1,2 +1,4 @@
|
||||
rich
|
||||
numpy
|
||||
numpy
|
||||
tqdm
|
||||
safetensors
|
||||
@@ -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);
|
||||
};
|
||||
},
|
||||
});
|
||||
Reference in New Issue
Block a user