diff --git a/src/merge/__init__.py b/src/merge/__init__.py index aa70cfc..aa550fd 100644 --- a/src/merge/__init__.py +++ b/src/merge/__init__.py @@ -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', diff --git a/src/merge/algorithms.py b/src/merge/algorithms.py index 7eda351..1b36d2b 100644 --- a/src/merge/algorithms.py +++ b/src/merge/algorithms.py @@ -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, diff --git a/tests/test_interp_delta_merge.py b/tests/test_interp_delta_merge.py new file mode 100644 index 0000000..41a0528 --- /dev/null +++ b/tests/test_interp_delta_merge.py @@ -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), +]) \ No newline at end of file