Joint Modeling for Stochastic Interventions / run_custom_bench.py
Mechanism confirmed, baseline not beaten
1import json, math, os, random
2import numpy as np
3import torch
4from torch import nn
5from custom_stochastic_intervention_track import get_dataset, META
6
7SEEDS = list(range(8))
8LRS = [0.003, 0.01, 0.03]
9EPOCHS = 90
10BATCH = 128
11
12def math_check(seed=123):
13 rng = np.random.RandomState(seed)
14 n = 500000
15 out = {}
16 vals = []
17 for sx in [1.0, 2.0]:
18 x = sx*rng.randn(n); u = rng.randn(n); v = rng.randn(n)
19 m = x + u; y = m + x + v
20 keep = np.abs(m-1.0) < .012
21 empirical = float(y[keep].mean())
22 analytic = 1.0 + sx*sx/(sx*sx+1.0)
23 vals.append((empirical, analytic, int(keep.sum())))
24 out['selected_m_equals_1'] = {'sigma1': vals[0], 'sigma2': vals[1]}
25 out['shift_observed'] = abs(vals[1][0]-vals[0][0]) > 0.15
26 return out
27
28def normal_nll(z, pars):
29 mu = pars[:, 0]
30 logsd = pars[:, 1].clamp(-4.0, 3.0)
31 return 0.5*((z-mu)/logsd.exp())**2 + logsd + 0.5*math.log(2*math.pi)
32
33class SharedMLP(nn.Module):
34 def __init__(self):
35 super().__init__()
36 self.net = nn.Sequential(nn.Linear(3, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 2))
37 def forward(self, z): return self.net(z)
38
39class JointMLP(nn.Module):
40 def __init__(self):
41 super().__init__()
42 self.q = nn.Sequential(nn.Linear(1,32), nn.Tanh(), nn.Linear(32,2))
43 self.m = nn.Sequential(nn.Linear(2,32), nn.Tanh(), nn.Linear(32,2))
44 self.y = SharedMLP()
45 def forward(self, c, x, m):
46 return self.q(c), self.m(torch.cat([c,x],1)), self.y(torch.cat([c,x,m],1))
47
48def train_one(seed, lr, idea, device):
49 torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
50 d = get_dataset(seed+1000, 800, 800)
51 c = torch.tensor(d['xtr'][:,0:1], device=device); x = torch.tensor(d['xtr'][:,1:2], device=device)
52 m = torch.tensor(d['xtr'][:,2:3], device=device); y = torch.tensor(d['ytr'][:,None], device=device)
53 if idea: model = JointMLP().to(device)
54 else: model = SharedMLP().to(device)
55 opt = torch.optim.Adam(model.parameters(), lr=lr)
56 g = torch.Generator(device=device); g.manual_seed(seed+55)
57 n = len(y)
58 model.train()
59 for _ in range(EPOCHS):
60 ix = torch.randperm(n, generator=g, device=device)
61 for start in range(0,n,BATCH):
62 j = ix[start:start+BATCH]
63 if idea:
64 q, mp, yp = model(c[j],x[j],m[j])
65 loss = (normal_nll(x[j].squeeze(1),q) + normal_nll(m[j].squeeze(1),mp) + normal_nll(y[j].squeeze(1),yp)).mean()
66 else:
67 # Standard mediator-only practice: the identical outcome MLP is trained without X.
68 z = torch.cat([c[j], torch.zeros_like(x[j]), m[j]], 1)
69 loss = normal_nll(y[j].squeeze(1), model(z)).mean()
70 opt.zero_grad(); loss.backward(); opt.step()
71 d = get_dataset(seed+2000, 12000, 12000)
72 ct = torch.tensor(d['xte'][:,0:1], device=device); xt = torch.tensor(d['xte'][:,1:2], device=device)
73 mt = torch.tensor(d['xte'][:,2:3], device=device); yt = torch.tensor(d['yte'], device=device)
74 model.eval()
75 with torch.no_grad():
76 if idea:
77 _,_,pars = model(ct,xt,mt)
78 pred = pars[:,0]
79 else:
80 pars = model(torch.cat([ct,torch.zeros_like(xt),mt],1)); pred = pars[:,0]
81 mse = float(((pred-yt)**2).mean().cpu())
82 # Re-test the claimed selection mechanism on model outputs. Selection is on observed mediator;
83 # compare predicted E[Y|M approximately 1] with the empirical selected outcome.
84 keep = (mt[:,0]-1.0).abs() < 0.035
85 selected_pred = float(pred[keep].mean().cpu())
86 selected_obs = float(yt[keep].mean().cpu())
87 return {'seed':seed,'lr':lr,'mse':mse,'selected_pred':selected_pred,'selected_obs':selected_obs,'n_selected':int(keep.sum().cpu())}, model
88
89def pvalue(deltas):
90 deltas=np.asarray(deltas); count=0; total=1<<len(deltas)
91 for mask in range(total):
92 signs=np.array([1 if (mask>>i)&1 else -1 for i in range(len(deltas))])
93 if (signs*deltas).mean() <= 0: count += 1
94 return count/total
95
96def main():
97 try: device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
98 except Exception: device=torch.device('cpu')
99 try:
100 # A CUDA allocation can fail in the shared environment; retry all work on CPU.
101 torch.zeros(1,device=device)
102 except Exception: device=torch.device('cpu')
103 math_result=math_check()
104 baseline={}; idea={}
105 for lr in LRS:
106 baseline[str(lr)] = [train_one(s,lr,False,device)[0] for s in SEEDS]
107 idea[str(lr)] = [train_one(s,lr,True,device)[0] for s in SEEDS]
108 bmeans={k:float(np.mean([r['mse'] for r in v])) for k,v in baseline.items()}
109 imeans={k:float(np.mean([r['mse'] for r in v])) for k,v in idea.items()}
110 best_lr=min(bmeans,key=bmeans.get); best_idea_lr=min(imeans,key=imeans.get)
111 br=baseline[best_lr]; ir=idea[best_idea_lr]
112 deltas=[ir[i]['mse']-br[i]['mse'] for i in range(8)]
113 sig={
114 'selected_m':1.0,
115 'baseline_predicted_mean':float(np.mean([r['selected_pred'] for r in br])),
116 'idea_predicted_mean':float(np.mean([r['selected_pred'] for r in ir])),
117 'observed_mean':float(np.mean([r['selected_obs'] for r in ir])),
118 'baseline_abs_error':float(abs(np.mean([r['selected_pred'] for r in br])-np.mean([r['selected_obs'] for r in br]))),
119 'idea_abs_error':float(abs(np.mean([r['selected_pred'] for r in ir])-np.mean([r['selected_obs'] for r in ir]))),
120 'confirmed': bool(math_result['shift_observed'] and abs(np.mean([r['selected_pred'] for r in ir])-np.mean([r['selected_obs'] for r in ir])) < 0.35)
121 }
122 report={'track':META,'official_bench_available':False,'infrastructure_note':'Specified /home/maxwelhelp/all/math2nn/bench and README.md were absent; this is the required-contract local fallback, not bench.make_report output.','math_check':math_result,'budget':{'seeds':SEEDS,'epochs':EPOCHS,'batch':BATCH,'lr_union':LRS},'baseline_sweep':bmeans,'idea_sweep':imeans,'best_baseline_lr':float(best_lr),'best_idea_lr':float(best_idea_lr),'baseline_per_seed':br,'idea_per_seed':ir,'paired_delta_mean':float(np.mean(deltas)),'paired_deltas':deltas,'permutation_p':pvalue(deltas),'mechanism_signature':sig,'bench_report':{'custom_track':{'name':META['name'],'file':'custom_stochastic_intervention_track.py','domain':META['domain']},'verdict':'idea better (significant)' if np.mean(deltas)<0 and pvalue(deltas)<.05 else 'no significant win'}}
123 with open('custom_bench_results.json','w') as f: json.dump(report,f,indent=2)
124 print(json.dumps(report,indent=2))
125if __name__=='__main__': main()