"""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