feat(merge): interp_delta_merge helper for delta-space slerp/nuslerp/karcher/nearswap
This commit is contained in:
@@ -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
@@ -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,
|
||||
|
||||
@@ -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),
|
||||
])
|
||||
Reference in New Issue
Block a user