Buckets:
| """ | |
| 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.