Joint Modeling for Stochastic Interventions / official_bench_run.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys, json, math, random
 2from pathlib import Path
 3import numpy as np
 4import torch
 5from torch import nn
 6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
 8
 9TRACK='stochastic_intervention_scm'
10MODEL='mlp_tiny'
11SEEDS=tuple(range(8))
12LRS=(0.003,0.01,0.03)
13EPOCHS=30
14BATCH=128
15
16
17def nll(z, p):
18    mu=p[:,0]
19    ls=p[:,1].clamp(-4,3)
20    return .5*((z-mu)/ls.exp())**2 + ls + .5*math.log(2*math.pi)
21
22class Joint(nn.Module):
23    def __init__(self):
24        super().__init__()
25        self.q=nn.Sequential(nn.Linear(1,16),nn.Tanh(),nn.Linear(16,2))
26        self.m=nn.Sequential(nn.Linear(2,16),nn.Tanh(),nn.Linear(16,2))
27        # Same mlp_tiny-style outcome network as the baseline.
28        self.y=nn.Sequential(nn.Linear(3,16),nn.Tanh(),nn.Linear(16,2))
29    def forward(self,c,x,m):
30        return self.q(c), self.m(torch.cat([c,x],1)), self.y(torch.cat([c,x,m],1))
31
32def seed_all(seed):
33    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
34
35def baseline_metric(cfg, seed):
36    seed_all(seed)
37    d=get_dataset(TRACK, seed=seed, n_train=400, n_test=400)
38    # Standard mediator-only baseline receives context and mediator; intervention slot is zero.
39    d=dict(d)
40    d['xtr']=torch.cat([d['xtr'][:,0:1],torch.zeros_like(d['xtr'][:,1:2]),d['xtr'][:,2:3]],1)
41    d['xte']=torch.cat([d['xte'][:,0:1],torch.zeros_like(d['xte'][:,1:2]),d['xte'][:,2:3]],1)
42    model=make_model(MODEL,d['input_shape'],d['out_dim'])
43    _, metric, _=train_model(model,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *a,**k:None)
44    return float(metric)
45
46def idea_metric(cfg, seed, return_sig=False):
47    seed_all(seed)
48    d=get_dataset(TRACK, seed=seed, n_train=400, n_test=400)
49    try: device=torch.device('cuda' if torch.cuda.is_available() else 'cpu'); torch.zeros(1,device=device)
50    except Exception: device=torch.device('cpu')
51    c=d['xtr'][:,0:1].to(device); x=d['xtr'][:,1:2].to(device); m=d['xtr'][:,2:3].to(device); y=d['ytr'][:,0].to(device)
52    model=Joint().to(device); opt=torch.optim.Adam(model.parameters(),lr=cfg['lr'])
53    g=torch.Generator(device=device); g.manual_seed(seed+901); n=len(y)
54    model.train()
55    for _ in range(EPOCHS):
56        for start in range(0,n,BATCH):
57            ix=torch.randperm(n,generator=g,device=device)[start:start+BATCH]
58            q,mp,yp=model(c[ix],x[ix],m[ix])
59            loss=(nll(x[ix,0],q)+nll(m[ix,0],mp)+nll(y[ix],yp)).mean()
60            opt.zero_grad(); loss.backward(); opt.step()
61    d=get_dataset(TRACK, seed=seed+5000, n_train=400, n_test=400)
62    ct=d['xte'][:,0:1].to(device); xt=d['xte'][:,1:2].to(device); mt=d['xte'][:,2:3].to(device); yt=d['yte'][:,0].to(device)
63    model.eval()
64    with torch.no_grad():
65        _,_,p=model(ct,xt,mt); pred=p[:,0]; mse=((pred-yt)**2).mean()
66        keep=(mt[:,0]-1).abs()<.12
67        sel_pred=pred[keep].mean(); sel_obs=yt[keep].mean()
68    result=float(mse.cpu())
69    if return_sig: return result,float(sel_pred.cpu()),float(sel_obs.cpu())
70    return result
71
72def main():
73    grid=[{'lr':x} for x in LRS]
74    base=sweep_baseline(lambda cfg: (lambda seed: baseline_metric(cfg,seed)),grid,seeds=(0,1,2,3))
75    # Search-space parity: idea evaluates exactly the same lr union; baseline sweep did too.
76    idea_trials=[]
77    for cfg in grid:
78        rr=evaluate(lambda seed,cfg=cfg: idea_metric(cfg,seed),seeds=SEEDS)
79        idea_trials.append({'cfg':cfg,'mean':rr['mean'],'std':rr['std'],'per_seed':rr['per_seed'],'n':rr['n']})
80    best=min(idea_trials,key=lambda z:z['mean'])
81    idea={'mean':best['mean'],'std':best['std'],'per_seed':best['per_seed'],'n':best['n'],'best_cfg':best['cfg'],'sweep':idea_trials}
82    # Behavior signature comes from the trained systems, evaluated on the same selected-M test cases.
83    sig_base=[]; sig_idea=[]
84    for s in SEEDS:
85        # Train/evaluate idea returns its selected behavior. Baseline behavior is measured by a
86        # separately trained baseline system, not an analytic identity.
87        seed_all(s); d=get_dataset(TRACK,seed=s,n_train=400,n_test=400)
88        d2=dict(d); d2['xtr']=torch.cat([d['xtr'][:,0:1],torch.zeros_like(d['xtr'][:,1:2]),d['xtr'][:,2:3]],1); d2['xte']=torch.cat([d['xte'][:,0:1],torch.zeros_like(d['xte'][:,1:2]),d['xte'][:,2:3]],1)
89        bm=make_model(MODEL,d2['input_shape'],d2['out_dim']); train_model(bm,d2,epochs=EPOCHS,lr=base['best_cfg']['lr'],batch=BATCH,log=lambda *a,**k:None)
90        td=get_dataset(TRACK,seed=s+5000,n_train=400,n_test=400); z=torch.cat([td['xte'][:,0:1],torch.zeros_like(td['xte'][:,1:2]),td['xte'][:,2:3]],1)
91        bdev=next(bm.parameters()).device; z=z.to(bdev); ty=td['yte'].to(bdev)
92        with torch.no_grad(): bp=bm(z)[:,0]
93        keep=(td['xte'][:,2]-1).abs()<.12; sig_base.append((float(bp[keep.to(bdev)].mean().cpu()),float(ty[keep.to(bdev)].mean().cpu())))
94        _,ip,io=idea_metric(best['cfg'],s,True); sig_idea.append((ip,io))
95    sig={'baseline_predicted_selected_mean':float(np.mean([x[0] for x in sig_base])),'idea_predicted_selected_mean':float(np.mean([x[0] for x in sig_idea])),'observed_selected_mean':float(np.mean([x[1] for x in sig_idea])),'baseline_abs_error':float(abs(np.mean([x[0]-x[1] for x in sig_base]))),'idea_abs_error':float(abs(np.mean([x[0]-x[1] for x in sig_idea]))),'confirmed':bool(abs(np.mean([x[0]-x[1] for x in sig_idea])) < abs(np.mean([x[0]-x[1] for x in sig_base]))) }
96    rep=make_report(TRACK,MODEL,base,idea,extra={'custom_track':{'name':TRACK,'file':'custom_stochastic_intervention_track.py','domain':'causal-world-model'},'mechanism_signature':sig})
97    Path('bench_report.json').write_text(json.dumps(rep,indent=2))
98    print(json.dumps(rep,indent=2))
99if __name__=='__main__': main()