feat(merge): interp_delta_merge helper for delta-space slerp/nuslerp/karcher/nearswap

This commit is contained in:
larsupb
2026-07-21 00:54:33 +02:00
parent fbaf1e9ae7
commit 48150ddd2e
3 changed files with 228 additions and 2 deletions
+7 -1
View File
@@ -19,7 +19,10 @@ from .utils import (
simple_weighted_average,
)
from .dispatcher import get_merge_method, prepare_method_args
from .algorithms import MERGE_ALGORITHMS, get_merge_algorithm
from .algorithms import (
MERGE_ALGORITHMS, get_merge_algorithm, sce_merge_deltas,
interp_delta_merge, INTERP_MODES,
)
from .base_node import BaseMergeMethodNode, BaseTaskArithmeticNode
__all__ = [
@@ -36,6 +39,9 @@ __all__ = [
# Algorithms
'MERGE_ALGORITHMS',
'get_merge_algorithm',
'sce_merge_deltas',
'interp_delta_merge',
'INTERP_MODES',
# Base classes
'BaseMergeMethodNode',
'BaseTaskArithmeticNode',
+94 -1
View File
@@ -29,7 +29,8 @@ from mergekit.merge_methods.slerp import SlerpTask
from mergekit.sparsify import RescaleNorm
import mergekit.sparsify as sparsify_module
from .utils import apply_weights_to_tensors
from .utils import apply_weights_to_tensors, create_map, create_tensor_param
from .gta import sce_delta_merge
def generalized_task_arithmetic_merge(
tensors: Dict[ModelReference, torch.Tensor],
@@ -230,6 +231,98 @@ def sce_merge(
return result
def sce_merge_deltas(
deltas: list,
weights: torch.Tensor,
select_topk: float = 0.1,
int8_mask: bool = False,
normalize: bool = False,
) -> torch.Tensor:
"""
Delta-space SCE merge.
Runs mergekit's SCE directly on the full per-LoRA deltas (out x in), the same
way the GTA family is merged in delta space, instead of merging the up/down
factors separately. Merging the factors independently injects meaningless
up_i @ down_j cross-terms and produces a near-noise result for SCE; operating
on the reconstructed deltas fixes that.
Args:
deltas: List of dense LoRA delta tensors (all the same shape).
weights: Per-LoRA strength weights (1-D tensor, aligned with `deltas`).
select_topk: Fraction of highest-variance elements to retain.
int8_mask: Use int8 masking inside mergekit's SCE.
Args:
normalize: False (default) -> additive SUM of the sign-agreeing selected
contributions (strengths act as gains, full stacked magnitude). True ->
mergekit's normalized weighted AVERAGE (strengths act as ratios; a
blend). See :func:`sce_delta_merge`.
Returns:
Merged dense delta (to be refactored back into a LoRA via
`merged_delta_to_lora`).
Uses a streaming implementation (:func:`sce_delta_merge`) that never stacks the
deltas into an ``[N, out, in]`` tensor, so peak VRAM stays bounded on large
FLUX layers (the naive mergekit stack OOMs an 8 GB card). ``int8_mask`` is
accepted for API parity but does not change the result and is ignored.
"""
if not deltas:
raise ValueError("sce_merge_deltas requires at least one delta")
# Bake the per-LoRA strength into each delta (the task vectors SCE operates on).
weighted = [float(w) * d for w, d in zip(weights.tolist(), deltas)]
# A single contributing LoRA has zero cross-tensor variance, so SCE's
# variance mask would zero the whole delta (no impact). Fall back to the
# single weighted delta in that case.
if len(weighted) == 1:
return weighted[0]
# select_topk<=0 retains no elements -> zero delta.
if select_topk <= 0:
return torch.zeros_like(deltas[0])
return sce_delta_merge(weighted, select_topk=select_topk, normalize=normalize)
INTERP_MODES = ("slerp", "nuslerp", "karcher", "nearswap")
def interp_delta_merge(
method,
deltas: list,
weights: torch.Tensor,
method_args: Dict[str, Any],
key: str = "merge",
*,
normalize: bool = False,
dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
if not deltas:
raise ValueError("interp_delta_merge requires at least one delta")
mag = weights.abs().to(torch.float32)
strength_scale = float(mag.sum()) if not normalize else float(mag.mean())
if len(deltas) == 1:
return deltas[0] * strength_scale
tensor_map, weight_map = {}, {}
weight_info = WeightInfo(name=f"{key}.merge", dtype=dtype, is_embed=False)
for i, d in enumerate(deltas):
ref = ModelReference(model=ModelPath(path=f"{key}.{i}"))
tensor_map[ref] = d
weight_map[ref] = torch.tensor(1.0)
gather = GatherTensors(weight_info=create_map(key, tensor_map, dtype))
params = ImmutableMap(
{r: ImmutableMap(create_tensor_param(weight_map[r], method_args))
for r in tensor_map})
blend = method(tensor_map, gather, weight_info, params, method_args)
return blend * strength_scale
def karcher_merge(
tensors: Dict[ModelReference, torch.Tensor],
gather_tensors: GatherTensors,
+127
View File
@@ -0,0 +1,127 @@
# tests/test_interp_delta_merge.py
# Standalone script test (repo pytest is broken). Verifies the delta-space
# interpolation merge helper: unit-weight blend + strength post-scale.
import os, sys, traceback
REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
PARENT = os.path.dirname(REPO) # .../custom_nodes
COMFY_ROOT = os.path.dirname(PARENT) # ComfyUI root, so `import comfy` resolves
sys.path.insert(0, COMFY_ROOT)
sys.path.insert(0, PARENT)
import torch
from mergekit.architecture import WeightInfo
from mergekit.common import ModelReference, ModelPath, ImmutableMap
from mergekit.io.tasks import GatherTensors
# import the package under a synthetic name so relative imports resolve
import importlib.util
PKG = "LoRA_Merger_ComfyUI_test"
spec = importlib.util.spec_from_file_location(
PKG, os.path.join(REPO, "__init__.py"), submodule_search_locations=[REPO])
pkg = importlib.util.module_from_spec(spec)
sys.modules[PKG] = pkg
spec.loader.exec_module(pkg)
from LoRA_Merger_ComfyUI_test.src.merge.algorithms import (
interp_delta_merge, INTERP_MODES,
slerp_merge, nuslerp_merge, karcher_merge, nearswap_merge,
)
from LoRA_Merger_ComfyUI_test.src.merge.utils import create_map, create_tensor_param
METHODS = {
"slerp": (slerp_merge, {"mode": "slerp", "t": 0.5, "lambda_": 1.0}),
"nuslerp": (nuslerp_merge, {"mode": "nuslerp", "nuslerp_flatten": True,
"nuslerp_row_wise": False, "lambda_": 1.0}),
"karcher": (karcher_merge, {"mode": "karcher", "max_iter": 10, "tol": 1e-5,
"lambda_": 1.0}),
"nearswap": (nearswap_merge, {"mode": "nearswap", "similarity_threshold": 0.001,
"lambda_": 1.0}),
}
def _delta(seed):
torch.manual_seed(seed)
return (torch.randn(64, 8) * 0.1) @ (torch.randn(8, 96) * 0.1)
def _blend_unit(method, deltas, margs):
"""Reference blend computed directly with unit weights (mirrors the helper)."""
tm, wm = {}, {}
wi = WeightInfo(name="k.merge", dtype=torch.float32, is_embed=False)
for i, d in enumerate(deltas):
ref = ModelReference(model=ModelPath(path=f"k.{i}"))
tm[ref] = d
wm[ref] = torch.tensor(1.0)
gt = GatherTensors(weight_info=create_map("k", tm, torch.float32))
tp = ImmutableMap({r: ImmutableMap(create_tensor_param(wm[r], margs)) for r in tm})
return method(tm, gt, wi, tp, margs)
def test_modes_constant():
assert INTERP_MODES == ("slerp", "nuslerp", "karcher", "nearswap"), INTERP_MODES
def test_additive_is_blend_times_sum():
for name, (fn, margs) in METHODS.items():
deltas = [_delta(1), _delta(2)]
B = _blend_unit(fn, [d.clone() for d in deltas], dict(margs))
w = torch.tensor([1.0, 0.5])
out = interp_delta_merge(fn, [d.clone() for d in deltas], w, dict(margs),
key="k", normalize=False)
exp = B * float(w.abs().sum()) # additive -> * sum(|s|)
assert torch.allclose(out, exp, atol=1e-5), f"{name}: additive != B*sum"
def test_average_is_blend_times_mean():
for name, (fn, margs) in METHODS.items():
deltas = [_delta(3), _delta(4)]
B = _blend_unit(fn, [d.clone() for d in deltas], dict(margs))
w = torch.tensor([1.0, 0.5])
out = interp_delta_merge(fn, [d.clone() for d in deltas], w, dict(margs),
key="k", normalize=True)
exp = B * float(w.abs().mean()) # average -> * mean(|s|)
assert torch.allclose(out, exp, atol=1e-5), f"{name}: average != B*mean"
def test_strength_is_linear_gain_additive():
fn, margs = METHODS["slerp"]
deltas = [_delta(5), _delta(6)]
o1 = interp_delta_merge(fn, [d.clone() for d in deltas], torch.tensor([1.0, 1.0]),
dict(margs), key="k", normalize=False)
o2 = interp_delta_merge(fn, [d.clone() for d in deltas], torch.tensor([2.0, 2.0]),
dict(margs), key="k", normalize=False)
# sum(|s|) doubles from 2 -> 4, so magnitude doubles
assert abs(o2.norm().item() / o1.norm().item() - 2.0) < 1e-4
def test_single_owner_fallback_not_zero():
fn, margs = METHODS["slerp"]
d = _delta(7)
out = interp_delta_merge(fn, [d.clone()], torch.tensor([1.0]), dict(margs),
key="k", normalize=False)
assert torch.allclose(out, d, atol=1e-6), "single-owner additive != delta*|s|"
out2 = interp_delta_merge(fn, [d.clone()], torch.tensor([0.5]), dict(margs),
key="k", normalize=False)
assert torch.allclose(out2, d * 0.5, atol=1e-6)
def run(tests):
failed = 0
for name, fn in tests:
try:
fn(); print(f"PASS {name}")
except Exception:
failed += 1; print(f"FAIL {name}"); traceback.print_exc()
if failed:
print(f"\n{failed} FAILED"); sys.exit(1)
print(f"\nAll {len(tests)} passed")
run([
("modes_constant", test_modes_constant),
("additive_is_blend_times_sum", test_additive_is_blend_times_sum),
("average_is_blend_times_mean", test_average_is_blend_times_mean),
("strength_linear_gain_additive", test_strength_is_linear_gain_additive),
("single_owner_fallback", test_single_owner_fallback_not_zero),
])