Joint Modeling for Stochastic Interventions / run_custom_bench.py

Mechanism confirmed, baseline not beaten

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