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 <noreply@anthropic.com>
This commit is contained in:
larsupb
2026-07-19 02:31:16 +02:00
co-authored by Claude Opus 4.8
parent 036934f911
commit 030ab02fe7
2 changed files with 77 additions and 29 deletions
+76 -28
View File
@@ -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
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)
+1 -1
View File
@@ -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()