Moment-Tuple Propagation for Compressed Neural Inference / experiment.py

Mechanism failed

Raw ⬇ ZIP
  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()