File size: 7,483 Bytes
99db05a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 | """DELLA magnitude pruning + task-vector combine over voice head tensors.
Clean-room port of the DELLA spec (arXiv:2406.11617) in numpy. No mergekit
dependency, no vendored code: mergekit is LGPL-3.0 and the authors' snapshot
carries no license grant, so both served as read-only references only.
One deliberate fork from upstream mergekit's ``della_magprune``, recorded
here rather than collapsed: upstream rescales L1 against the whole-tensor
sum, while ``magprune`` rescales each row to its own pre-prune ``sum|x|``.
Per-row keeps every output neuron's scale (each head row is one token's
direction), stays bounded at aggressive sparsity, and keeps row-blocking
exactly correct (a global rescale would need the full tensor at once).
Equivalence spot-checks against mergekit therefore expect "close", not
bit-exact: same family, different rescale locality.
Determinism: ``seed`` spawns one child stream per voice
(``SeedSequence(seed)``), so a seeded run is bit-identical regardless of
``block_rows`` — each voice's Bernoulli draws are consumed in row order on
its own stream. ``seed=None`` gives fresh entropy per voice
(nondeterministic, as upstream).
"""
VALID_RESCALES = ("l1", "inv_p")
def _validate_density_epsilon(density, epsilon):
if density + epsilon >= 1 or density - epsilon <= 0:
raise ValueError(
"della density±epsilon must stay in (0,1): "
f"got density={density} epsilon={epsilon}"
)
def _check_rescale(rescale):
if rescale not in VALID_RESCALES:
raise ValueError(f"della rescale must be 'l1' or 'inv_p', got {rescale!r}")
def magprune(delta, density=0.5, epsilon=0.1, rescale="l1", rng=None):
"""Rank-weighted Bernoulli prune of one delta, row-wise. Always F32 out.
Keep-probability ramps linearly with within-row magnitude rank, from
``density-epsilon`` (smallest) to ``density+epsilon`` (largest), then a
seeded Bernoulli draw decides survivors. ``rescale="l1"`` restores each
surviving row's ``sum|x|``; ``rescale="inv_p"`` divides each survivor by
its own keep-prob (paper-faithful, heavy-tailed).
"""
import numpy as np
F32 = np.float32
_check_rescale(rescale)
a = np.asarray(delta, dtype=F32)
if density >= 1:
return a
if density <= 0:
return np.zeros_like(a)
_validate_density_epsilon(density, epsilon)
orig_shape = a.shape
w = np.atleast_2d(a)
order = np.argsort(np.abs(w), axis=-1, kind="stable")
ranks = np.argsort(order, axis=-1).astype(F32) + F32(1.0)
rmin = ranks.min(axis=-1, keepdims=True)
rmax = ranks.max(axis=-1, keepdims=True)
denom = rmax - rmin
# Single-column rows have no ordering; fall back to flat density.
rank_norm = np.divide(ranks - rmin, denom,
out=np.full_like(ranks, F32(0.5)),
where=denom > 0)
rank_norm = np.clip(rank_norm, F32(0.0), F32(1.0))
d = F32(density)
e = F32(epsilon)
probs = (d - e) + rank_norm * (F32(2.0) * e)
if rng is None:
rng = np.random.default_rng()
mask = rng.binomial(1, probs).astype(F32)
masked = w * mask
if rescale == "l1":
before = np.abs(w).sum(axis=-1, keepdims=True)
after = np.abs(masked).sum(axis=-1, keepdims=True)
ok = (before >= 1e-7) & (after >= 1e-7)
scale = np.where(ok, before / np.maximum(after, F32(1e-30)), F32(1.0))
out = masked * scale
else: # inv_p
out = masked / probs
if not np.all(np.isfinite(out)):
raise ValueError("della inv_p rescale produced NaN/Inf "
"- refusing to write a corrupt voice")
return np.asarray(out.reshape(orig_shape), dtype=F32)
def combine(base, deltas, weights, lambda_=1.0, normalize=True):
"""TIES-sum combine of pruned deltas onto base. F32 in, F32 out.
Weight-scale, elect the majority sign per position (ties go positive),
average the elected (``normalize``) or sum them, scale by ``lambda_``,
add base. Positions every voice pruned divide by 1, not 0.
"""
import numpy as np
F32 = np.float32
b = np.asarray(base, dtype=F32)
ds = [np.asarray(d, dtype=F32) for d in deltas]
if not ds:
return b
if len(weights) != len(ds):
raise ValueError(f"della weights ({len(weights)}) != deltas ({len(ds)})")
for i, d in enumerate(ds):
if d.shape != b.shape:
raise ValueError(f"della delta {i} shape {d.shape} != base {b.shape}")
wcol = np.asarray(list(weights), dtype=F32).reshape((-1,) + (1,) * b.ndim)
W = np.stack(ds).astype(F32, copy=False) * wcol
majority = np.where(W.sum(axis=0) >= 0, F32(1.0), F32(-1.0))
elected = np.sign(W) == majority
mixed = np.where(elected, W, F32(0.0)).sum(axis=0)
divisor = np.where(elected, wcol, F32(0.0)).sum(axis=0)
divisor = np.where(divisor == 0, F32(1.0), divisor)
if normalize:
mixed = mixed / divisor
if lambda_ != 1:
mixed = mixed * F32(lambda_)
return b + mixed
def merge(base, voices, weights=None, density=0.5, epsilon=0.1, lambda_=1.0,
rescale="l1", normalize=True, seed=None, block_rows=4096):
"""Stream one DELLA merge: per-block (voice-base) prune, then combine.
Only ``block_rows`` rows per voice are pruned/combined at once, so the
full ``n_voices x rows x cols`` delta stack is never materialized; peak
besides the inputs/output is ``O(block_rows x cols x n_voices)``.
Blocking is exactly correct: pruning ranks within rows and election
votes per position, neither crosses a block boundary. Each voice draws
on its own ``SeedSequence(seed)`` child stream, so repartitioning blocks
never changes the result. Output is cast back to the base input's dtype.
"""
import numpy as np
F32 = np.float32
_check_rescale(rescale)
b = np.asarray(base)
out_dtype = b.dtype
bw = np.asarray(b, dtype=F32)
vs = [np.asarray(v) for v in voices]
if not vs:
return np.asarray(b, dtype=out_dtype)
wts = [1.0] * len(vs) if weights is None else list(weights)
if len(wts) != len(vs):
raise ValueError(f"della weights ({len(wts)}) != voices ({len(vs)})")
for i, v in enumerate(vs):
if v.shape != bw.shape:
raise ValueError(f"della voice {i} shape {v.shape} != base {bw.shape}")
if not (density >= 1 or density <= 0):
_validate_density_epsilon(density, epsilon)
squeeze = False
if bw.ndim == 1:
squeeze = True
bw = bw[None, :]
vs = [v[None, :] for v in vs]
if bw.ndim < 1:
raise ValueError("della merge needs at least 1-D tensors")
if seed is None:
rngs = [np.random.default_rng() for _ in vs]
else:
rngs = [np.random.default_rng(s)
for s in np.random.SeedSequence(seed).spawn(len(vs))]
vw = [np.asarray(v, dtype=F32) for v in vs]
out = np.empty_like(bw)
for start in range(0, bw.shape[0], block_rows):
sl = slice(start, min(start + block_rows, bw.shape[0]))
bb = bw[sl]
pruned = [magprune(v[sl] - bb, density=density, epsilon=epsilon,
rescale=rescale, rng=rngs[vi])
for vi, v in enumerate(vw)]
out[sl] = combine(bb, pruned, wts, lambda_=lambda_, normalize=normalize)
if squeeze:
out = out.reshape(b.shape)
if out_dtype != F32:
return out.astype(out_dtype, copy=False)
return out
|