import json, math, os, random, time import numpy as np import torch import torch.nn as nn SEED = 442 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(min(8, os.cpu_count() or 1)) try: device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') except Exception: device = torch.device('cpu') def normalized_silu(t, r): s = torch.nn.functional.softplus(r) + 1e-3 return torch.nn.functional.silu(s * t) / s, s class VarMLP(nn.Module): def __init__(self, width=32, depth=2): super().__init__() self.width, self.depth = width, depth self.weights = nn.ParameterList() self.biases = nn.ParameterList() self.scales = nn.ParameterList() in_dim = 1 for _ in range(depth): self.weights.append(nn.Parameter(torch.randn(width, in_dim) * math.sqrt(2 / in_dim))) self.biases.append(nn.Parameter(torch.zeros(width))) self.scales.append(nn.Parameter(torch.zeros(width))) in_dim = width self.out_w = nn.Parameter(torch.randn(1, width) * math.sqrt(2 / width)) self.out_b = nn.Parameter(torch.zeros(1)) def forward(self, x, return_v=False): h = x q = torch.zeros(self.width, device=x.device) for W, b, r in zip(self.weights, self.biases, self.scales): h = h @ W.t() + b s = torch.nn.functional.softplus(r) + 1e-3 h = torch.nn.functional.silu(s * h) / s q = q + torch.sqrt(W.square() + 1e-8).sum(1) + torch.sqrt(b.square() + 1e-8) y = h @ self.out_w.t() + self.out_b if return_v: V = (torch.sqrt(self.out_w.square() + 1e-8) * (1 + q[None, :])).sum() + torch.sqrt(self.out_b.square() + 1e-8).sum() return y, V, q return y def l2(self): return sum(p.square().sum() for p in self.parameters()) @torch.no_grad() def math_checks(): t = torch.linspace(-3, 3, 1001) r = torch.tensor(0.37) s = torch.nn.functional.softplus(r) + 1e-3 direct = torch.nn.functional.silu(s * t) / s impl, _ = normalized_silu(t, r) formula_err = float((direct - impl).abs().max()) m = VarMLP(width=1, depth=2) expected = sum(torch.sqrt(W.square() + 1e-8).sum() + torch.sqrt(b.square() + 1e-8).sum() for W, b in zip(m.weights, m.biases)) _, _, q = m(torch.zeros(1, 1), True) q_err = float((q[0] - expected).abs()) tt = torch.tensor([-2., -0.5, 0.5, 2.]) vals = [] for ss in [0.25, 1., 4.]: rr = torch.log(torch.expm1(torch.tensor(ss - 1e-3))) vals.append((torch.nn.functional.silu(ss * tt) / ss).numpy()) shape_delta = float(np.max(np.abs(vals[0] - vals[2]))) return {'normalized_formula_max_error': formula_err, 'recursive_q_max_error': q_err, 'silu_shape_delta_s025_vs_s4': shape_delta} def train(kind, lam, steps=700): torch.manual_seed(SEED + (0 if kind == 'l2' else 100) + int(lam * 1e6)) ntr, nva = 128, 256 xtr = torch.linspace(-1, 1, ntr, device=device).unsqueeze(1) ytr = torch.sin(2 * math.pi * 3 * xtr) g = torch.Generator(device='cpu').manual_seed(SEED) xv = torch.rand(nva, 1, generator=g).to(device) * 2 - 1 yv = torch.sin(2 * math.pi * 3 * xv) model = VarMLP(width=32, depth=2).to(device) if kind == 'l2': opt = torch.optim.AdamW(model.parameters(), lr=3e-3, weight_decay=lam) else: opt = torch.optim.Adam(model.parameters(), lr=3e-3) t0 = time.time() for _ in range(steps): opt.zero_grad(set_to_none=True) pred, V, _ = model(xtr, True) mse = ((pred - ytr) ** 2).mean() loss = mse + (lam * V if kind == 'var' else 0) loss.backward() opt.step() with torch.no_grad(): pv, V, _ = model(xv, True) val = float(((pv - yv) ** 2).mean()) train_mse = float(((model(xtr) - ytr) ** 2).mean()) l2 = float(model.l2()) maxpred = float(pv.abs().max()) return {'kind': kind, 'lambda': lam, 'train_mse': train_mse, 'val_mse': val, 'V': float(V), 'param_l2': l2, 'max_prediction': maxpred, 'seconds': time.time() - t0, 'steps': steps} def main(): global device checks = math_checks() results = [] for lam in [1e-4, 1e-3, 1e-2]: for kind in ['l2', 'var']: try: results.append(train(kind, lam)) except Exception: if device.type == 'cuda': device = torch.device('cpu') results.append(train(kind, lam)) else: raise best_l2 = min((r for r in results if r['kind'] == 'l2'), key=lambda r: r['val_mse']) best_var = min((r for r in results if r['kind'] == 'var'), key=lambda r: r['val_mse']) out = {'device': str(device), 'checks': checks, 'results': results, 'best_l2': best_l2, 'best_var': best_var} with open('results.json', 'w') as f: json.dump(out, f, indent=2) print(json.dumps(out, indent=2)) if __name__ == '__main__': main()