Moment-Tuple Propagation for Compressed Neural Inference / experiment.py
Mechanism failed
1import json, time
2import numpy as np
3import torch
4import torch.nn as nn
5
6SEED = 214
7np.random.seed(SEED)
8torch.manual_seed(SEED)
9
10def tuple_affine(t, W, d):
11 mu, s2, b, v, c = t
12 return (mu @ W.T + d, s2 @ (W * W).T, b @ W.T,
13 v @ (W * W).T, c @ (W * W).T)
14
15def tuple_activation(t, kind="tanh"):
16 mu, s2, b, v, c = t
17 if kind == "tanh":
18 f = torch.tanh(mu)
19 fp = 1 - f*f
20 fpp = -2*f*fp
21 elif kind == "relu":
22 f = torch.relu(mu)
23 fp = (mu > 0).to(mu.dtype)
24 fpp = torch.zeros_like(mu)
25 else:
26 raise ValueError(kind)
27 sz = torch.clamp(s2 + v + 2*c, min=0)
28 m2 = sz + b*b
29 out_mu = f + fp*b + 0.5*fpp*m2
30 out_var = (fp + fpp*b)**2 * sz + 0.5*fpp*fpp*sz*sz
31 return out_mu, torch.zeros_like(out_var), torch.zeros_like(out_var), out_var, torch.zeros_like(out_var)
32
33class MLP(nn.Module):
34 def __init__(self):
35 super().__init__()
36 self.l1 = nn.Linear(2, 24)
37 self.l2 = nn.Linear(24, 1)
38 def forward(self, x):
39 return self.l2(torch.tanh(self.l1(x)))
40
41def math_check():
42 # Affine identities are checked against empirical moments, including correlation.
43 rng = np.random.default_rng(SEED)
44 n = 500000
45 x = rng.normal(0.7, 0.8, n)
46 u = rng.normal(size=n)
47 e = -0.12 + 0.35*(0.55*(x-0.7)/0.8 + np.sqrt(1-0.55**2)*u)
48 a, k = 0.4, -1.7
49 xa, xk = x+a, k*x
50 ea, ek = e, k*e
51 affine_err = max(abs(np.mean(xa)-(np.mean(x)+a)),
52 abs(np.var(xa)-(np.var(x))),
53 abs(np.mean(ek)-k*np.mean(e)),
54 abs(np.var(ek)-k*k*np.var(e)),
55 abs(np.cov(xk, ek, bias=True)[0,1]-k*k*np.cov(x,e,bias=True)[0,1]))
56 # Smooth activation: compare Taylor tuple mean/variance to high-sample truth.
57 mu, s2, b, v, c = np.mean(x), np.var(x), np.mean(e), np.var(e), np.cov(x,e,bias=True)[0,1]
58 sz = s2+v+2*c
59 f = np.tanh(mu); fp=1-f*f; fpp=-2*f*fp
60 approx_mean=f+fp*b+0.5*fpp*(sz+b*b)
61 approx_var=(fp+fpp*b)**2*sz+0.5*fpp*fpp*sz*sz
62 truth=np.tanh(x+e)
63 # A local regime check: second-order Taylor should improve when total spread is small.
64 xl = rng.normal(0.7, 0.06, n)
65 el = rng.normal(-0.01, 0.025, n)
66 ml, sl2, bl, vl, cl = np.mean(xl), np.var(xl), np.mean(el), np.var(el), np.cov(xl,el,bias=True)[0,1]
67 zl = sl2+vl+2*cl
68 fl=np.tanh(ml); fpl=1-fl*fl; fppl=-2*fl*fpl
69 local_mean=fl+fpl*bl+0.5*fppl*(zl+bl*bl)
70 local_var=(fpl+fppl*bl)**2*zl+0.5*fppl*fppl*zl*zl
71 local_truth=np.tanh(xl+el)
72 return {"affine_max_abs_error": float(affine_err),
73 "tanh_mean_abs_error": float(abs(approx_mean-np.mean(truth))),
74 "tanh_variance_abs_error": float(abs(approx_var-np.var(truth))),
75 "local_tanh_mean_abs_error": float(abs(local_mean-np.mean(local_truth))),
76 "local_tanh_variance_abs_error": float(abs(local_var-np.var(local_truth))),
77 "tanh_true_mean": float(np.mean(truth)), "tanh_approx_mean": float(approx_mean)}
78
79def train_model():
80 rng=np.random.default_rng(SEED)
81 x=rng.normal(size=(10000,2)).astype(np.float32)
82 y=(np.sin(x[:,0])+0.25*x[:,1]**2+0.05*x[:,0]*x[:,1]).astype(np.float32)[:,None]
83 xt=torch.tensor(x[:8000]); yt=torch.tensor(y[:8000])
84 xv=torch.tensor(x[8000:]); yv=torch.tensor(y[8000:])
85 model=MLP()
86 opt=torch.optim.Adam(model.parameters(),lr=0.015)
87 for _ in range(220):
88 idx=torch.randperm(len(xt))[:128]
89 pred=model(xt[idx]); loss=((pred-yt[idx])**2).mean()
90 opt.zero_grad(); loss.backward(); opt.step()
91 return model, xv, yv
92
93def tuple_predict(model, xhat, qvar):
94 # Decode is the tuple center; calibrated quantization error is zero mean.
95 z=(xhat, torch.zeros_like(xhat), torch.zeros_like(xhat),
96 torch.full_like(xhat, qvar), torch.zeros_like(xhat))
97 z=tuple_affine(z, model.l1.weight, model.l1.bias)
98 z=tuple_activation(z, "tanh")
99 z=tuple_affine(z, model.l2.weight, model.l2.bias)
100 return z[0], torch.clamp(z[1]+z[3]+2*z[4],min=0)
101
102def experiment():
103 model,x,y=train_model(); model.eval()
104 # Uniform scalar quantizer; noise variance q^2/12 is a deliberately cheap calibration.
105 q=0.40
106 xhat=(torch.round(x/q)*q)
107 qvar=q*q/12
108 with torch.no_grad():
109 t0=time.perf_counter(); base=model(xhat); base_time=time.perf_counter()-t0
110 t0=time.perf_counter(); tm,tv=tuple_predict(model,xhat,qvar); tuple_time=time.perf_counter()-t0
111 t0=time.perf_counter()
112 mcs=[]
113 for _ in range(32):
114 # Independent calibrated compression perturbations.
115 mcs.append(model(xhat + (torch.rand_like(xhat)-0.5)*q))
116 mc=torch.stack(mcs); mcmean=mc.mean(0); mcstd=mc.std(0,unbiased=False)
117 mc_time=time.perf_counter()-t0
118 clean=model(x)
119 def rmse(a,b): return float(torch.sqrt(torch.mean((a-b)**2)))
120 def coverage(mean,std,target): return float((((target >= mean-1.645*std) & (target <= mean+1.645*std)).float().mean()))
121 return {
122 "n_test":len(x), "quant_step":q,
123 "baseline_rmse_to_clean":rmse(base,clean),
124 "tuple_rmse_to_clean":rmse(tm,clean), "mc_mean_rmse_to_clean":rmse(mcmean,clean),
125 "tuple_vs_mc_mean_rmse":rmse(tm,mcmean),
126 "tuple_90pct_coverage_clean":coverage(tm,torch.sqrt(tv+1e-12),clean),
127 "mc_90pct_coverage_clean":coverage(mcmean,mcstd+1e-12,clean),
128 "baseline_ms_per_sample":1000*base_time/len(x),
129 "tuple_ms_per_sample":1000*tuple_time/len(x),
130 "mc32_ms_per_sample":1000*mc_time/len(x),
131 "tuple_mean_std":float(torch.sqrt(tv).mean()), "mc_mean_std":float(mcstd.mean())
132 }
133
134def main():
135 out={"math_check":math_check(),"experiment":experiment()}
136 print(json.dumps(out,indent=2,sort_keys=True))
137
138if __name__ == "__main__": main()