SabaPivot's picture
download
raw
23.6 kB
"""
Independent re-derivation of the constructions behind arXiv:2605.01702 (OpenReview g89qqA6qmD).
We do not transcribe the paper's proof. We build our own floating-point sigma-networks
that satisfy the *statements*:
Lemma 3.4 : f = f* and D_{f,x}(y) = 0 (values kept, AD gradient killed)
Lemma 3.5 : f = 0 and D_{f,x}(h*(x)) = g*(x) (values killed, AD gradient arbitrary)
Theorem 3.1: f = f* and D_{f,x}(h*(x)) = g*(x) (f = f2 # f1)
for every sigma in {ReLU, ELU, GELU, Swish, Sigmoid, tanh}, in genuine IEEE-754 arithmetic.
Both halves rest on the non-associativity of floating-point summation:
wipe identity : (t (+) C) (-) C = 0 whenever |t| < ulp(C)/2 ,
while 0 (+) C (-) C (+) u = u leaves a *later* summand untouched.
* the +C/-C pair inside the AD accumulation grad(x_k) = (+)_r s_r (x) A_1[r,k]
annihilates the value path's gradient; the forward pass never sees the pair
because the two units are identical and their downstream weights are +/-w.
* the +C/-C pair inside a forward accumulation annihilates the value; AD never
sees it because the backward pass multiplies (it does not sum) along that edge.
* for the saturating activations sigma' = round(sigma_hat') underflows to exactly 0
where sigma is exactly +-1 : values survive, gradients do not.
Layer-1 unit order (this order *is* the summation order of the AD accumulation):
[f1 features][f2 features][P Q][f2 carriers][K1 K2]
^^^ annihilates everything to its left
"""
import numpy as np
from fpnet import Act, Builder
SAT = {"sigmoid", "tanh"}
ACT_NAMES = ["relu", "elu", "gelu", "swish", "sigmoid", "tanh"]
def acfg(name, dtype, act=None, kmax=1024):
"""Operating points, all found numerically (never assumed).
T : smallest power of two such that on every t = k*T, 1 <= k <= kmax,
non-saturating : sigma(t) = t and sigma'(t) = 1 exactly,
saturating : sigma(t) = 1 and sigma'(t) = 0 exactly,
and on t = -k*T sigma is a single constant and sigma' = 0 exactly.
(This is the paper's Condition-1 regime |sigma'(gamma)| << |sigma(gamma)|.)
S = 2T is the layer-1 grid scale, BP = T the detector scale.
"""
a = act if act is not None else Act(name, dtype)
sat = name in SAT
ks = np.arange(1, kmax + 1, dtype=np.float64)
T = None
for e in range(0, 60):
P = float(2.0**e)
if P * kmax > float(np.finfo(dtype).max) / 8:
break
t = np.asarray(ks * P, dtype=dtype)
sp, sn, dp, dn = a.s(t), a.s(-t), a.sp(t), a.sp(-t)
ok_neg = bool(np.all(sn == sn[0]) and np.all(dn == 0))
ok_pos = (
bool(np.all(sp == 1) and np.all(dp == 0))
if sat
else bool(np.all(sp == t) and np.all(dp == 1))
)
if ok_pos and ok_neg:
T = dtype(P)
break
assert T is not None, f"no operating scale found for {name}/{dtype}"
c = dict(
name=name,
sat=sat,
act=a,
dtype=dtype,
T=T,
S=dtype(2.0) * T,
BP=T,
NEG=a.s(np.asarray([-T], dtype=dtype))[0],
)
if sat:
GP = None
for e in range(1, 200):
P = dtype(2.0) ** e
if (
a.s(np.array([P], dtype=dtype))[0] == 1
and a.sp(np.array([P], dtype=dtype))[0] != 0
):
GP = P
break
assert GP is not None
c["GP"] = GP
c["U"] = dtype(1.0) # well-conditioned: sigma'(1) is O(1)
c["ONV"] = T
else:
U = dtype(2.0) ** max(20, int(np.log2(float(T))) + 2)
c["U"] = U
c["ONV"] = U * (dtype(2.0) ** 20)
for p in (c["U"], c["ONV"]):
assert a.s(np.array([p], dtype=dtype))[0] == p
assert a.sp(np.array([p], dtype=dtype))[0] == 1
c["const_pre"] = c["U"]
c["const_val"] = a.s(np.array([c["U"]], dtype=dtype))[0]
return c
def _pow2_ceil(v, dtype):
v = float(abs(v))
return dtype(1.0) if v == 0 else dtype(2.0) ** int(np.ceil(np.log2(v)))
class Construction:
"""f = f2 # f1 on a finite uniform grid domain, for one activation / dtype."""
def __init__(
self, name, dtype, hp, Z, d=1, mode="thm31", n_carrier=4, margin=None,
kmax=None, n_extra=0
):
self.name, self.dtype, self.hp, self.d = name, dtype, hp, d
self.Z = np.asarray(Z, dtype=np.int64).reshape(-1, d)
self.m = self.Z.shape[0]
self.cfg = acfg(
name,
dtype,
kmax=kmax or max(8, 4 * int(self.Z.max() - self.Z.min()) + 8),
)
self.mode = mode
self.nc = n_carrier
self.margin = int(np.finfo(dtype).nmant + 6) if margin is None else int(margin)
self.h = dtype(2.0) ** hp
self.X = (np.asarray(self.Z, dtype=np.float64) * float(self.h)).astype(dtype)
if self.cfg["sat"]:
# saturating sigma uses two threshold knots that must fall strictly
# between grid points and stay exactly representable: the domain is
# the even multiples of h, the knots the odd ones.
assert np.all(self.Z % 2 == 0), "saturating sigma needs an even grid"
self.use_f1 = mode in ("thm31", "lem34")
self.use_f2 = mode in ("thm31", "lem35")
self.alpha = dtype(1.0)
self.n_extra = int(n_extra)
self.L = 5 + self.n_extra
# ---------------- targets ------------------------------------------
def draw_targets(
self, seed=0, fexp=(-8, 8), gexp=(-8, 8), hexp=(-4, 4), zero_frac=0.15
):
rng = np.random.default_rng(seed)
dt, m, d = self.dtype, self.m, self.d
def rf(lo, hi, size):
e = rng.integers(lo, hi + 1, size=size)
man = 1.0 + rng.integers(0, 2**20, size=size) / 2.0**20
s = rng.choice([-1.0, 1.0], size=size)
return (s * man * 2.0**e).astype(dt)
self.fstar = rf(*fexp, m)
self.hstar = rf(*hexp, m)
zero = rng.random(m) < zero_frac
self.hstar[zero] = dt(0.0)
self.gstar = rf(*gexp, (m, d)).astype(dt)
self.gstar[zero, :] = dt(0.0)
self.zero_mask = zero
return self
def set_targets(self, fstar=None, hstar=None, gstar=None):
dt = self.dtype
if fstar is not None:
self.fstar = np.asarray(fstar, dtype=dt)
if hstar is not None:
self.hstar = np.asarray(hstar, dtype=dt)
if gstar is not None:
self.gstar = np.asarray(gstar, dtype=dt).reshape(self.m, self.d)
self.zero_mask = self.hstar == 0
return self
# ---------------- skeleton -----------------------------------------
def _skeleton(self):
dt, cfg, d = self.dtype, self.cfg, self.d
sat, S, U = cfg["sat"], cfg["S"], cfg["U"]
W1 = dt(2.0) ** (int(np.log2(float(S))) - self.hp)
bld = Builder(d, self.L)
i1 = list(range(self.m)) if self.use_f1 else []
i2 = list(range(self.m)) if self.use_f2 else []
self.i1, self.i2 = i1, i2
# ---- layer 1 ----
self.hat1 = {i: self._features(bld, i, W1, S) for i in i1}
self.hat2 = {i: self._features(bld, i, W1, S) for i in i2}
self.pair = None
if not sat:
P = bld.add(0, {k: 1.0 for k in range(d)}, float(cfg["ONV"]))
Q = bld.add(0, {k: 1.0 for k in range(d)}, float(cfg["ONV"]))
self.pair = (P, Q)
self.carr = {
i: [
[bld.add(0, {k: 0.0}, float(U)) for _ in range(self.nc)]
for k in range(d)
]
for i in i2
}
self.k12 = (bld.add(0, {}, float(U)), bld.add(0, {}, float(U))) if i2 else None
# ---- layer 2 ----
self.sink = None
if self.pair is not None:
P, Q = self.pair
self.sink = (
bld.add(1, {P: 1.0, Q: -1.0}, float(U)),
bld.add(1, {}, float(U)),
)
self.det1 = {i: self._detector(bld, self.hat1[i], S) for i in i1}
self.det2 = {i: self._detector(bld, self.hat2[i], S) for i in i2}
self.wipe = {}
for i in i2:
ins = {c: float(self.alpha) for k in range(d) for c in self.carr[i][k]}
ins[self.k12[0]] = 0.0
ins[self.k12[1]] = 0.0
self.wipe[i] = bld.add(1, ins, float(U))
# ---- padding relay layers (identity in value, extra depth) ----
self.L3, self.L4, self.L5 = 2 + self.n_extra, 3 + self.n_extra, 4 + self.n_extra
self.relay_units = []
n2 = bld.width(1)
for r in range(self.n_extra):
self.relay_units.append([bld.add(2 + r, {j: 1.0}, 0.0) for j in range(n2)])
# ---- layer 3 ----
self.T = None
if self.sink is not None:
s1, s2 = self.sink
self.T = (
bld.add(self.L3, {s1: 1.0, s2: -1.0}, float(U)),
bld.add(self.L3, {}, float(U)),
)
self.V = {i: bld.add(self.L3, {self.det1[i]: 0.0}, 0.0) for i in i1}
self.gate = {
i: bld.add(self.L3, {self.det2[i]: 0.0, self.wipe[i]: 0.0}, 0.0) for i in i2
}
# ---- layer 4 ----
self.U2 = None
if self.T is not None:
t1, t2 = self.T
self.U2 = (
bld.add(self.L4, {t1: 1.0, t2: -1.0}, float(U)),
bld.add(self.L4, {}, float(U)),
)
self.R = {i: bld.add(self.L4, {self.gate[i]: 0.0}, 0.0) for i in i2}
self.kout = (bld.add(self.L4, {}, float(U)), bld.add(self.L4, {}, float(U))) if i2 else None
self.A = {i: bld.add(self.L4, {self.V[i]: 0.0}, 0.0) for i in i1}
# ---- layer 5 : output ----
ins = {}
if self.U2 is not None:
ins[self.U2[0]] = 0.0
ins[self.U2[1]] = 0.0
for i in i2:
ins[self.R[i]] = 1.0
if self.kout is not None:
ins[self.kout[0]] = 0.0
ins[self.kout[1]] = 0.0
for i in i1:
ins[self.A[i]] = 0.0
bld.add(self.L5, ins, 0.0)
self.bld = bld
return bld
def _features(self, bld, i, W1, S):
out = []
for k in range(self.d):
nz = int(self.Z[i, k])
offs = [1.0, -1.0] if self.cfg["sat"] else [1.0, 0.0, -1.0]
out.append(
[
bld.add(0, {k: float(W1)}, float((np.float64(o) - nz) * float(S)))
for o in offs
]
)
return out
def _detector(self, bld, hats, S):
ins = {}
S = float(S)
if self.cfg["sat"]:
sc = S if self.name == "sigmoid" else S / 2.0
for k in range(self.d):
ins[hats[k][0]] = sc
ins[hats[k][1]] = -sc
else:
for k in range(self.d):
ins[hats[k][0]] = 1.0
ins[hats[k][1]] = -2.0
ins[hats[k][2]] = 1.0
return bld.add(1, ins, -(float(self.d) - 0.5) * S)
# ---------------- measurement helpers -------------------------------
def _vals(self, net, layer):
out = []
for x in self.X:
net.forward(x)
out.append(net.act.s(net._pres[layer]))
return np.asarray(out)
def _pres(self, net, layer):
out = []
for x in self.X:
net.forward(x)
out.append(net._pres[layer].copy())
return np.asarray(out)
def _grad_l1(self, net, x, gin):
net.forward(x)
_, tr = net.backward(np.array([gin], dtype=self.dtype), debug=True)
return tr[("post_act", 0)]
def _pair_values(self, vals, idx, i):
col = vals[:, idx]
off = np.unique(np.delete(col, i))
assert off.size <= 1, (self.name, "off-value not unique", off[:4])
return col[i], (off[0] if off.size else self.dtype(0.0))
def _exact_relay(self, v_on, v_off, p_on, p_off):
"""w,b with w (x) v_off (+) b == p_off EXACTLY and w (x) v_on (+) b ~= p_on."""
dt = self.dtype
v_on, v_off = dt(v_on), dt(v_off)
w = dt(np.float64(dt(p_on) - dt(p_off)) / np.float64(v_on - v_off))
b = dt(dt(p_off) - dt(w * v_off))
return w, b
def _p_off(self):
return -self.cfg["BP"] if self.name == "sigmoid" else self.dtype(0.0)
# ---------------- build ---------------------------------------------
def build(self, verbose=False):
dt, cfg = self.dtype, self.cfg
act = cfg["act"]
bld = self._skeleton()
if self.use_f2: # wipe constants first: they fix v_W
self._calibrate_wipe_l2(bld, bld.build(act))
# relay layers: reproduce layer-2's values exactly, adding depth
if self.n_extra:
net = bld.build(act)
v2m, p2m = self._vals(net, 1), self._pres(net, 1)
for j in range(bld.width(1)):
vals, pres = np.unique(v2m[:, j]), np.unique(p2m[:, j])
assert vals.size <= 2, ("relay: unit takes >2 values", j, vals[:4])
if vals.size == 1:
w, b = dt(1.0), dt(dt(pres.max()) - dt(vals[0]))
else:
ia = int(np.argmax(np.abs(v2m[:, j] - vals[0]) > 0))
on, off = dt(vals[1]), dt(vals[0])
pon = dt(p2m[np.argmax(v2m[:, j] == on), j])
poff = dt(p2m[np.argmax(v2m[:, j] == off), j])
w, b = self._exact_relay(on, off, pon, poff)
for r in range(self.n_extra):
bld.set_weight(2 + r, j, j, float(w))
bld.set_bias(2 + r, j, float(b))
chk = bld.build(act)
assert np.array_equal(self._vals(chk, 1 + self.n_extra), v2m), \
"relay layers must reproduce the layer-2 values exactly"
net = bld.build(act)
v2 = self._vals(net, 1 + self.n_extra)
for i in self.i1:
on, off = self._pair_values(v2, self.det1[i], i)
w, b = self._exact_relay(on, off, cfg["ONV"], self._p_off())
bld.set_weight(self.L3, self.V[i], self.det1[i], float(w))
bld.set_bias(self.L3, self.V[i], float(b))
# aW is the AD path gate -> wipe-unit; its forward contribution aW (x) v_W is a
# constant and is kept small enough to be irrelevant (verified below).
aW = dt(2.0) ** -30 if cfg["sat"] else dt(2.0) ** -20
for i in self.i2:
on, off = self._pair_values(v2, self.det2[i], i)
p_off = -max(dt(cfg["BP"]), dt(cfg["U"]))
w, b = self._exact_relay(on, off, cfg["U"], p_off)
bld.set_weight(self.L3, self.gate[i], self.det2[i], float(w))
bld.set_weight(self.L3, self.gate[i], self.wipe[i], float(aW))
bld.set_bias(self.L3, self.gate[i], float(b))
net = bld.build(act)
v3 = self._vals(net, self.L3)
p3 = self._pres(net, self.L3)
for i in self.i2: # the gate must be shut off its point
spg = act.sp(np.delete(p3[:, self.gate[i]], i))
assert np.all(spg == 0), (self.name, "gate leaks", spg[spg != 0][:3])
for i in self.i1:
on, off = self._pair_values(v3, self.V[i], i)
w, b = self._exact_relay(on, off, cfg["ONV"], self._p_off())
bld.set_weight(self.L4, self.A[i], self.V[i], float(w))
bld.set_bias(self.L4, self.A[i], float(b))
for i in self.i2:
on, off = self._pair_values(v3, self.gate[i], i)
w, b = self._exact_relay(on, off, cfg["U"], self._p_off())
bld.set_weight(self.L4, self.R[i], self.gate[i], float(w))
bld.set_bias(self.L4, self.R[i], float(b))
net = bld.build(act)
v4 = self._vals(net, self.L4)
self.ONA = {}
for i in self.i1:
on, off = self._pair_values(v4, self.A[i], i)
assert off == 0, (self.name, "f1 off value must be exactly 0", off)
self.ONA[i] = on
bld.set_weight(self.L5, 0, self.A[i], float(dt(self.fstar[i]) / dt(on)))
if self.use_f2:
for _ in range(5):
net = bld.build(act)
if self._calibrate_scale(bld, net):
break
for _ in range(8):
net = bld.build(act)
if self._calibrate_carriers(bld, net):
break
net = bld.build(act)
self._calibrate_wipes(bld, net)
if self.pair is not None:
net = bld.build(act)
self._calibrate_pair(bld, net)
if self.use_f2:
for _ in range(5):
net = bld.build(act)
if self._calibrate_carriers(bld, net):
break
net = bld.build(act)
self._calibrate_wipes(bld, net)
self.net = bld.build(act)
return self.net
# ---------------- calibration ---------------------------------------
def _calibrate_scale(self, bld, net):
dt, cfg = self.dtype, self.cfg
xmax = max(1.0, float(np.abs(self.X).max()))
ulpU = float(np.spacing(np.asarray(cfg["U"], dtype=dt)))
w_target = dt(2.0) ** int(np.floor(np.log2(ulpU / (8.0 * xmax))))
stable = True
for i in self.i2:
if self.hstar[i] == 0:
continue
gl1 = self._grad_l1(net, self.X[i], self.hstar[i])
mu = abs(dt(gl1[self.carr[i][0][0]]))
gmax = _pow2_ceil(np.max(np.abs(self.gstar[i])), dt)
desired = dt(gmax) / dt(w_target)
if mu == 0:
stable = False
mu = dt(2.0) ** -60
e = int(round(np.log2(float(desired) / float(mu))))
if e != 0:
stable = False
cur = dt(net.As[self.L5][0, self.R[i]])
bld.set_weight(self.L5, 0, self.R[i], float(cur * dt(2.0) ** e))
return stable
def _calibrate_carriers(self, bld, net):
dt, d = self.dtype, self.d
done = True
for i in self.i2:
if self.hstar[i] == 0:
for k in range(d):
for c in self.carr[i][k]:
bld.set_weight(0, c, k, 0.0)
continue
gl1 = self._grad_l1(net, self.X[i], self.hstar[i])
for k in range(d):
mult = [dt(gl1[c]) for c in self.carr[i][k]]
target = dt(self.gstar[i, k])
if all(mu == 0 for mu in mult):
done = False
continue
ws, acc = [], dt(0.0)
for mu in mult:
if mu == 0:
ws.append(dt(0.0))
continue
r = dt(target - acc)
w = dt(np.float64(r) / np.float64(mu))
ws.append(w)
acc = dt(acc + dt(mu * w))
if acc != target:
done = False
for j, c in enumerate(self.carr[i][k]):
bld.set_weight(0, c, k, float(ws[j]))
return done
def _calibrate_wipes(self, bld, net):
self._calibrate_wipe_l2(bld, net)
self._calibrate_wipe_out(bld)
def _calibrate_wipe_l2(self, bld, net):
dt, cfg = self.dtype, self.cfg
vconst = dt(cfg["const_val"])
v1 = self._vals(net, 0)
for i in self.i2:
worst = dt(0.0)
for r in range(self.m):
s = dt(0.0)
for k in range(self.d):
for c in self.carr[i][k]:
s = dt(s + dt(self.alpha * dt(v1[r, c])))
worst = max(worst, abs(s))
C0 = _pow2_ceil(worst, dt) * (dt(2.0) ** self.margin)
w = float(dt(C0) / vconst)
bld.set_weight(1, self.wipe[i], self.k12[0], w)
bld.set_weight(1, self.wipe[i], self.k12[1], -w)
def _calibrate_wipe_out(self, bld):
dt, cfg = self.dtype, self.cfg
vconst = dt(cfg["const_val"])
net2 = bld.build(cfg["act"])
v4 = self._vals(net2, self.L4)
worst_o = dt(0.0)
for r in range(self.m):
s = dt(0.0)
for i in self.i2:
s = dt(s + dt(dt(net2.As[self.L5][0, self.R[i]]) * dt(v4[r, self.R[i]])))
worst_o = max(worst_o, abs(s))
C1 = _pow2_ceil(worst_o, dt) * (dt(2.0) ** self.margin)
w = float(dt(C1) / vconst)
bld.set_weight(self.L5, 0, self.kout[0], w)
bld.set_weight(self.L5, 0, self.kout[1], -w)
def _calibrate_pair(self, bld, net):
"""scale the P/Q wipe pair until its +C/-C annihilates the detector-path
gradient. The gain is spread over the four edges of the pair's dummy
chain (P -> sink -> T -> U -> output); every edge cancels exactly in the
forward pass, so the gain is invisible to f."""
dt = self.dtype
P = self.pair[0]
saved = {}
for i in self.i2: # silence the carriers: measure leakage only
for k in range(self.d):
for c in self.carr[i][k]:
saved[(c, k)] = bld.units[0][c][0].get(k, 0.0)
bld.set_weight(0, c, k, 0.0)
edges = [(1, self.sink[0], self.pair[0], self.pair[1]),
(self.L3, self.T[0], self.sink[0], self.sink[1]),
(self.L4, self.U2[0], self.T[0], self.T[1]),
(self.L5, 0, self.U2[0], self.U2[1])]
for (lay, u, a, b) in edges: # unit gain probe
bld.set_weight(lay, u, a, 1.0)
bld.set_weight(lay, u, b, -1.0)
probe = bld.build(net.act)
bld.set_weight(self.L5, 0, self.U2[0], 0.0)
bld.set_weight(self.L5, 0, self.U2[1], 0.0)
base = bld.build(net.act)
ys = sorted({float(v) for v in self.hstar} | {1.0, -1.0, 2.0 ** 12, 2.0 ** -12})
worst = 0.0
for r in range(self.m):
for y in ys:
y = dt(y)
if y == 0:
continue
base.forward(self.X[r])
t = float(np.max(np.abs(base.backward(np.array([y], dtype=dt)))))
probe.forward(self.X[r])
_, tr = probe.backward(np.array([y], dtype=dt), debug=True)
sP = float(abs(dt(tr[("post_act", 0)][P])))
if sP > 0 and t > 0:
worst = max(worst, t / sP)
for (c, k), w in saved.items():
bld.set_weight(0, c, k, w)
need = int(np.ceil(np.log2(worst))) + self.margin if worst > 0 else self.margin
# per-edge headroom so that w (x) value stays finite in the forward pass
big = float(np.finfo(dt).max) / 8.0
caps = []
for (lay, u, a, b) in edges:
v = float(np.abs(self._vals(probe, lay - 1)[:, a]).max())
caps.append(int(np.floor(np.log2(big / max(v, 1e-300)))))
assert need <= sum(caps), ("pair gain does not fit", need, caps)
rem, exps = need, []
for cap in caps:
e = int(np.clip(rem, -cap, cap))
exps.append(e)
rem -= e
self.pair_gain_exps = exps
for (lay, u, a, b), e in zip(edges, exps):
w = float(dt(2.0) ** e)
bld.set_weight(lay, u, a, w)
bld.set_weight(lay, u, b, -w)

Xet Storage Details

Size:
23.6 kB
·
Xet hash:
160992f278b5d4343640e2077a6042bc3370a58328a099133b82ed155870462c

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.