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