Voice / vlib /della.py
Wiself's picture
Upload vlib/della.py with huggingface_hub
99db05a verified
Raw History Blame Contribute Delete
7.48 kB
"""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