From 030ab02fe7e7e59b8089db71da227b1fb68f8d5d Mon Sep 17 00:00:00 2001 From: larsupb Date: Sun, 19 Jul 2026 02:31:16 +0200 Subject: [PATCH] perf(gta): stream the merge instead of stacking (cut peak VRAM ~2x) gta_merge stacked all N deltas into [N,out,in] plus several full copies, peaking at ~3.5 GB for a single 21504x3072 FLUX layer (~13x the delta). It now streams over a list: one accumulator for the elected sign, then numerator/ divisor accumulators, freeing each input delta as consumed. Peak ~1.7 GB/key, independent of N. Preserves mergekit parity incl. the density<1 sparsified sign vote (3-state torch.sign so sparsified-out zeros are excluded). gta_merge now consumes its deltas list (documented); behavior test snapshots first. Co-Authored-By: Claude Opus 4.8 --- src/merge/gta.py | 104 +++++++++++++++++++++++++++---------- tests/test_gta_behavior.py | 2 +- 2 files changed, 77 insertions(+), 29 deletions(-) diff --git a/src/merge/gta.py b/src/merge/gta.py index 066d83a..a237190 100644 --- a/src/merge/gta.py +++ b/src/merge/gta.py @@ -154,7 +154,11 @@ def gta_merge(deltas: List[torch.Tensor], weights: torch.Tensor, *, mode: str, rescale_norm: str = "default") -> torch.Tensor: """Merge a list of full LoRA deltas with GTA semantics on the delta itself. - `weights` are the per-LoRA merge strengths (signed). Returns the merged delta.""" + `weights` are the per-LoRA merge strengths (signed). Returns the merged delta. + + NOTE: this **consumes** `deltas` -- entries are replaced/freed in place to keep + peak memory low (large FLUX deltas are hundreds of MB each). Snapshot anything + you need to reuse before calling.""" if mode not in _MODE_SPARSIFY: raise ValueError(f"unknown GTA mode {mode!r}") if mode in _MODE_ALWAYS_CONSENSUS: @@ -166,11 +170,65 @@ def gta_merge(deltas: List[torch.Tensor], weights: torch.Tensor, *, mode: str, sp = _MODE_SPARSIFY[mode] res = resolve_rescale_norm(mode, rescale_norm) - sparse = [sparsify(d, sp, density=density, gamma=gamma, epsilon=epsilon, - rescale_norm=res) for d in deltas] - stack = torch.stack(sparse, dim=0) - return disjoint_merge(stack, weights.to(stack.dtype), sign_consensus=consensus, - normalize=normalize) + # Sparsify in place (replace each entry) so we never hold both the original + # and sparsified copy of every delta at once. + for i in range(len(deltas)): + deltas[i] = sparsify(deltas[i], sp, density=density, gamma=gamma, + epsilon=epsilon, rescale_norm=res) + w = weights.to(deltas[0].dtype) + return _stream_merge(deltas, w, sign_consensus=consensus, normalize=normalize) + + +def _stream_merge(deltas: List[torch.Tensor], weights: torch.Tensor, *, + sign_consensus: bool, normalize: bool) -> torch.Tensor: + """Memory-frugal merge over a *list* of deltas. + + Never stacks into an ``[N, out, in]`` tensor (that plus mergekit-style + intermediate copies peaked at ~13x a single delta for large FLUX layers and + exhausted VRAM). Instead it streams: one accumulator for the elected sign, + then per-element accumulators for the numerator and divisor. Peak is a small + constant number of ``[out, in]`` buffers, independent of ``N``.""" + n = len(deltas) + + if sign_consensus: + acc = deltas[0] * weights[0] + for i in range(1, n): + acc = acc + deltas[i] * weights[i] + # Elected sign as +1/-1 (matches mergekit's TIES 'sum' method). Using a + # 3-state sign comparison below means sparsified-out elements (exact 0, + # sign 0) never match the elected +/-1 and so are excluded from both the + # numerator and the divisor -- essential when density < 1. + sign_pm = (acc >= 0).to(deltas[0].dtype) * 2 - 1 + del acc + + mixed = None + divisor = None + zero = None + for i in range(n): + wd = deltas[i] * weights[i] + deltas[i] = None # free the input delta as soon as used + if sign_consensus: + if zero is None: + zero = torch.zeros((), dtype=wd.dtype, device=wd.device) + agree = torch.sign(wd) == sign_pm + if normalize: + dcontrib = agree.to(wd.dtype) * weights[i].abs() + divisor = dcontrib if divisor is None else divisor + dcontrib + del dcontrib + wd = torch.where(agree, wd, zero) # reuse wd as the masked contribution + del agree + elif normalize: + divisor = weights[i] if divisor is None else divisor + weights[i] + mixed = wd if mixed is None else mixed + wd + del wd + + if normalize: + if not sign_consensus: + divisor = torch.as_tensor(divisor, dtype=mixed.dtype, device=mixed.device) + divisor = divisor.expand_as(mixed) if divisor.dim() == 0 else divisor + divisor = torch.where(divisor.abs() < 1e-8, torch.ones_like(divisor), divisor) + mixed = mixed / divisor + return mixed # --------------------------------------------------------- sign + merge @@ -181,27 +239,17 @@ def elect_sign(weighted_deltas: torch.Tensor) -> torch.Tensor: return (sign_weight >= 0).to(weighted_deltas.dtype) * 2 - 1 -def disjoint_merge(deltas: torch.Tensor, weights: torch.Tensor, *, +def disjoint_merge(deltas, weights: torch.Tensor, *, sign_consensus: bool, normalize: bool) -> torch.Tensor: - """Merge stacked deltas (shape [N, *]) with per-LoRA `weights` (shape [N]). + """Merge deltas with per-LoRA ``weights`` (shape [N]). - weighted_deltas = deltas * weights (broadcast). When sign_consensus, elect a - per-element sign and keep only agreeing contributions. `normalize` divides by - the per-element sum of surviving weights (weighted average).""" - w = weights.clone() - while w.dim() < deltas.dim(): - w = w.unsqueeze(-1) - weighted = deltas * w - - if sign_consensus: - sign = elect_sign(weighted) - agree = (torch.sign(weighted) == sign).to(deltas.dtype) - else: - agree = torch.ones_like(weighted) - - mixed = (weighted * agree).sum(dim=0) - divisor = (w.abs() * agree).sum(dim=0) if sign_consensus else (w * agree).sum(dim=0) - divisor = torch.where(divisor.abs() < 1e-8, torch.ones_like(divisor), divisor) - if normalize: - mixed = mixed / divisor - return mixed \ No newline at end of file + Accepts either a list of ``[out, in]`` deltas or a stacked ``[N, *]`` tensor + (the latter is unbound into views, no copy). Delegates to the memory-frugal + :func:`_stream_merge`; kept as the stable public entry used by the unit + tests. When ``sign_consensus``, elect a per-element sign and keep only + agreeing contributions; ``normalize`` divides by the per-element sum of + surviving weights.""" + if isinstance(deltas, torch.Tensor): + deltas = list(deltas.unbind(0)) + return _stream_merge(deltas, weights, sign_consensus=sign_consensus, + normalize=normalize) \ No newline at end of file diff --git a/tests/test_gta_behavior.py b/tests/test_gta_behavior.py index 809b51c..2084787 100644 --- a/tests/test_gta_behavior.py +++ b/tests/test_gta_behavior.py @@ -51,9 +51,9 @@ def test_style_plus_character_nonoverlap_keeps_strength(): def test_normalize_does_not_collapse_as_1_over_n_squared(): torch.manual_seed(0) deltas = [torch.randn(6, 6) for _ in range(4)] + avg = torch.stack(deltas).mean(0) # snapshot: gta_merge consumes `deltas` merged = gta.gta_merge(deltas, torch.ones(4), mode="ties", density=1.0, normalize=True) - avg = torch.stack(deltas).mean(0) assert merged.norm() > 0.5 * avg.norm()