import json, time import numpy as np import torch import torch.nn as nn SEED = 214 np.random.seed(SEED) torch.manual_seed(SEED) def tuple_affine(t, W, d): mu, s2, b, v, c = t return (mu @ W.T + d, s2 @ (W * W).T, b @ W.T, v @ (W * W).T, c @ (W * W).T) def tuple_activation(t, kind="tanh"): mu, s2, b, v, c = t if kind == "tanh": f = torch.tanh(mu) fp = 1 - f*f fpp = -2*f*fp elif kind == "relu": f = torch.relu(mu) fp = (mu > 0).to(mu.dtype) fpp = torch.zeros_like(mu) else: raise ValueError(kind) sz = torch.clamp(s2 + v + 2*c, min=0) m2 = sz + b*b out_mu = f + fp*b + 0.5*fpp*m2 out_var = (fp + fpp*b)**2 * sz + 0.5*fpp*fpp*sz*sz return out_mu, torch.zeros_like(out_var), torch.zeros_like(out_var), out_var, torch.zeros_like(out_var) class MLP(nn.Module): def __init__(self): super().__init__() self.l1 = nn.Linear(2, 24) self.l2 = nn.Linear(24, 1) def forward(self, x): return self.l2(torch.tanh(self.l1(x))) def math_check(): # Affine identities are checked against empirical moments, including correlation. rng = np.random.default_rng(SEED) n = 500000 x = rng.normal(0.7, 0.8, n) u = rng.normal(size=n) e = -0.12 + 0.35*(0.55*(x-0.7)/0.8 + np.sqrt(1-0.55**2)*u) a, k = 0.4, -1.7 xa, xk = x+a, k*x ea, ek = e, k*e affine_err = max(abs(np.mean(xa)-(np.mean(x)+a)), abs(np.var(xa)-(np.var(x))), abs(np.mean(ek)-k*np.mean(e)), abs(np.var(ek)-k*k*np.var(e)), abs(np.cov(xk, ek, bias=True)[0,1]-k*k*np.cov(x,e,bias=True)[0,1])) # Smooth activation: compare Taylor tuple mean/variance to high-sample truth. 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] sz = s2+v+2*c f = np.tanh(mu); fp=1-f*f; fpp=-2*f*fp approx_mean=f+fp*b+0.5*fpp*(sz+b*b) approx_var=(fp+fpp*b)**2*sz+0.5*fpp*fpp*sz*sz truth=np.tanh(x+e) # A local regime check: second-order Taylor should improve when total spread is small. xl = rng.normal(0.7, 0.06, n) el = rng.normal(-0.01, 0.025, n) 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] zl = sl2+vl+2*cl fl=np.tanh(ml); fpl=1-fl*fl; fppl=-2*fl*fpl local_mean=fl+fpl*bl+0.5*fppl*(zl+bl*bl) local_var=(fpl+fppl*bl)**2*zl+0.5*fppl*fppl*zl*zl local_truth=np.tanh(xl+el) return {"affine_max_abs_error": float(affine_err), "tanh_mean_abs_error": float(abs(approx_mean-np.mean(truth))), "tanh_variance_abs_error": float(abs(approx_var-np.var(truth))), "local_tanh_mean_abs_error": float(abs(local_mean-np.mean(local_truth))), "local_tanh_variance_abs_error": float(abs(local_var-np.var(local_truth))), "tanh_true_mean": float(np.mean(truth)), "tanh_approx_mean": float(approx_mean)} def train_model(): rng=np.random.default_rng(SEED) x=rng.normal(size=(10000,2)).astype(np.float32) y=(np.sin(x[:,0])+0.25*x[:,1]**2+0.05*x[:,0]*x[:,1]).astype(np.float32)[:,None] xt=torch.tensor(x[:8000]); yt=torch.tensor(y[:8000]) xv=torch.tensor(x[8000:]); yv=torch.tensor(y[8000:]) model=MLP() opt=torch.optim.Adam(model.parameters(),lr=0.015) for _ in range(220): idx=torch.randperm(len(xt))[:128] pred=model(xt[idx]); loss=((pred-yt[idx])**2).mean() opt.zero_grad(); loss.backward(); opt.step() return model, xv, yv def tuple_predict(model, xhat, qvar): # Decode is the tuple center; calibrated quantization error is zero mean. z=(xhat, torch.zeros_like(xhat), torch.zeros_like(xhat), torch.full_like(xhat, qvar), torch.zeros_like(xhat)) z=tuple_affine(z, model.l1.weight, model.l1.bias) z=tuple_activation(z, "tanh") z=tuple_affine(z, model.l2.weight, model.l2.bias) return z[0], torch.clamp(z[1]+z[3]+2*z[4],min=0) def experiment(): model,x,y=train_model(); model.eval() # Uniform scalar quantizer; noise variance q^2/12 is a deliberately cheap calibration. q=0.40 xhat=(torch.round(x/q)*q) qvar=q*q/12 with torch.no_grad(): t0=time.perf_counter(); base=model(xhat); base_time=time.perf_counter()-t0 t0=time.perf_counter(); tm,tv=tuple_predict(model,xhat,qvar); tuple_time=time.perf_counter()-t0 t0=time.perf_counter() mcs=[] for _ in range(32): # Independent calibrated compression perturbations. mcs.append(model(xhat + (torch.rand_like(xhat)-0.5)*q)) mc=torch.stack(mcs); mcmean=mc.mean(0); mcstd=mc.std(0,unbiased=False) mc_time=time.perf_counter()-t0 clean=model(x) def rmse(a,b): return float(torch.sqrt(torch.mean((a-b)**2))) def coverage(mean,std,target): return float((((target >= mean-1.645*std) & (target <= mean+1.645*std)).float().mean())) return { "n_test":len(x), "quant_step":q, "baseline_rmse_to_clean":rmse(base,clean), "tuple_rmse_to_clean":rmse(tm,clean), "mc_mean_rmse_to_clean":rmse(mcmean,clean), "tuple_vs_mc_mean_rmse":rmse(tm,mcmean), "tuple_90pct_coverage_clean":coverage(tm,torch.sqrt(tv+1e-12),clean), "mc_90pct_coverage_clean":coverage(mcmean,mcstd+1e-12,clean), "baseline_ms_per_sample":1000*base_time/len(x), "tuple_ms_per_sample":1000*tuple_time/len(x), "mc32_ms_per_sample":1000*mc_time/len(x), "tuple_mean_std":float(torch.sqrt(tv).mean()), "mc_mean_std":float(mcstd.mean()) } def main(): out={"math_check":math_check(),"experiment":experiment()} print(json.dumps(out,indent=2,sort_keys=True)) if __name__ == "__main__": main()