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:
co-authored by
Claude Opus 4.8
parent
036934f911
commit
030ab02fe7
+76
-28
@@ -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)
|
||||
@@ -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()
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user