Missingness-as-a-Label Signal / missingness_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6import torch.nn.functional as F
7
8sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
9from bench import make_model, make_report
10from bench.protocol import permutation_pvalue
11
12META = {'name': 'uncertainty_missingness_classification', 'domain': 'label_missingness',
13 'description': 'Binary classification with labels preferentially missing at high teacher posterior entropy.'}
14
15# Custom track contract: m is included as an auxiliary array for the local loss.
16def get_dataset(seed, n_train=400, n_test=400):
17 rng = np.random.RandomState(seed)
18 def sample(n, s):
19 r = np.random.RandomState(s)
20 y = r.randint(0, 2, size=n).astype(np.int64)
21 mu = np.where(y[:, None] == 0, [-1.15, 0.0], [1.15, 0.0])
22 x = (mu + .95 * r.randn(n, 2)).astype(np.float32)
23 teacher = np.stack([-2.3*x[:, 0], 2.3*x[:, 0]], axis=1)
24 pt = np.exp(teacher - teacher.max(1, keepdims=True)); pt /= pt.sum(1, keepdims=True)
25 h = -(pt*np.log(np.maximum(pt, 1e-8))).sum(1)
26 q = 1/(1+np.exp(-(-1.15 + 3.0*h)))
27 m = r.binomial(1, q).astype(np.float32)
28 return x, y, m, q, h
29 xtr,ytr,m,q,h = sample(n_train, seed)
30 xte,yte,mt,qt,ht = sample(n_test, seed + 5000)
31 return {'xtr':xtr, 'ytr':ytr, 'mtr':m, 'qtr':q, 'htr':h,
32 'xte':xte, 'yte':yte, 'mte':mt, 'qte':qt, 'hte':ht,
33 'task':'classification', 'metric':'nll'}
34
35def seed_all(seed):
36 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
37 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
38
39def device_model(model):
40 try:
41 dev = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
42 return model.to(dev), dev
43 except Exception:
44 return model.cpu(), torch.device('cpu')
45
46def train(d, lr, wd, joint, seed, epochs=35):
47 seed_all(seed)
48 net = make_model('mlp_tiny', (2,), 2)
49 net, dev = device_model(net)
50 a = nn.Parameter(torch.tensor(-0.2, device=dev))
51 beta = nn.Parameter(torch.tensor(0.2, device=dev))
52 params = list(net.parameters()) + ([a, beta] if joint else [])
53 opt = torch.optim.Adam(params, lr=lr, weight_decay=wd)
54 x=torch.tensor(d['xtr'], device=dev); y=torch.tensor(d['ytr'], device=dev)
55 m=torch.tensor(d['mtr'], device=dev)
56 n=x.shape[0]
57 for _ in range(epochs):
58 # Full-batch is intentionally identical across systems and within budget.
59 logits=net(x)
60 if joint:
61 p=logits.softmax(-1); h=-(p*torch.log(p.clamp_min(1e-8))).sum(-1)
62 b=F.softplus(beta)
63 q=torch.sigmoid(a+b*h)
64 ce=F.cross_entropy(logits,y,reduction='none')
65 loss=((1-m)*ce - m*torch.log(q.clamp_min(1e-7)) -(1-m)*torch.log((1-q).clamp_min(1e-7))).mean()
66 loss=loss+1e-3*(a*a+b*b)
67 else:
68 # Standard practice with selective labels: CE only on observed labels.
69 loss=F.cross_entropy(logits[m < .5], y[m < .5])
70 opt.zero_grad(); loss.backward(); opt.step()
71 return net, dev, (a,beta) if joint else None
72
73def evaluate(net, dev, d, aux):
74 with torch.no_grad():
75 x=torch.tensor(d['xte'],device=dev); y=torch.tensor(d['yte'],device=dev)
76 z=net(x); p=z.softmax(-1)
77 nll=F.cross_entropy(z,y).item(); err=(p.argmax(1)!=y).float().mean().item()
78 h=-(p*torch.log(p.clamp_min(1e-8))).sum(-1)
79 if aux is not None:
80 a,beta=aux; b=F.softplus(beta); q=torch.sigmoid(a+b*h)
81 qn=q.cpu().numpy(); obs=d['mte']; true=d['qte']
82 corr=float(np.corrcoef(qn,true)[0,1]) if np.std(qn)>1e-8 else 0.0
83 return {'nll':nll,'err':err,'missing_corr':corr,'b':float(b.cpu()),
84 'pred_missing':float(qn.mean()),'observed_missing':float(obs.mean())}
85 return {'nll':nll,'err':err}
86
87def run_config(lr,wd,joint,seeds):
88 vals=[]
89 for s in seeds:
90 d=get_dataset(s,400,400)
91 net,dev,aux=train(d,lr,wd,joint,s)
92 vals.append(evaluate(net,dev,d,aux))
93 return vals
94
95def mean_metric(vals): return float(np.mean([v['nll'] for v in vals]))
96
97def main():
98 # Math sanity: exact missingness score derivatives and nonzero entropy path.
99 torch.manual_seed(3)
100 z=torch.tensor([[1.2,-.4],[.1,.2],[-.7,.9]],dtype=torch.double,requires_grad=True)
101 a=torch.tensor(-.4,dtype=torch.double,requires_grad=True); b=torch.tensor(1.7,dtype=torch.double,requires_grad=True)
102 h=-(z.softmax(-1)*torch.log(z.softmax(-1))).sum(-1); q=torch.sigmoid(a+b*h); m=torch.tensor([1.,0.,1.],dtype=torch.double)
103 loss=(-m*torch.log(q)-(1-m)*torch.log(1-q)).mean(); loss.backward()
104 math_check={'a_grad':float(a.grad),'a_expected':float((q-m).mean()),'b_grad':float(b.grad),'b_expected':float(((q-m)*h).mean()),'entropy_grad_norm':float(z.grad.norm()),'gradient_match':abs(float(a.grad-(q-m).mean()))<1e-10 and abs(float(b.grad-((q-m)*h).mean()))<1e-10}
105 lrs=[0.003,0.01,0.03]; wds=[0.0,1e-4]
106 grid=[{'lr':lr,'wd':wd} for lr in lrs for wd in wds]
107 sweep_seeds=(0,1,2,3); seeds=tuple(range(8))
108 sweep=[]
109 for c in grid:
110 vals=run_config(c['lr'],c['wd'],False,sweep_seeds)
111 sweep.append({'config':c,'mean_nll':mean_metric(vals),'per_seed':vals})
112 best=min(sweep,key=lambda x:x['mean_nll'])['config']
113 idea=[]
114 for lr in lrs:
115 vals=run_config(lr,best['wd'],True,seeds)
116 idea.append({'config':{'lr':lr,'wd':best['wd']},'per_seed':vals,'mean_nll':mean_metric(vals)})
117 best_idea=min(idea,key=lambda x:x['mean_nll'])
118 # Full paired baseline at every idea lr, preserving union parity.
119 base_by_lr={}
120 for lr in lrs:
121 vals=run_config(lr,best['wd'],False,seeds); base_by_lr[lr]=vals
122 base_vals=base_by_lr[best_idea['config']['lr']]
123 diffs=[i['nll']-b['nll'] for i,b in zip(best_idea['per_seed'],base_vals)]
124 report=make_report('uncertainty_missingness_classification','mlp_tiny',
125 {'full': {'per_seed':[v['nll'] for v in base_vals], 'mean':mean_metric(base_vals), 'std':float(np.std([v['nll'] for v in base_vals])), 'n':len(base_vals)}, 'best_config':best, 'sweep':sweep, 'full_eval_config':best_idea['config']},
126 {'per_seed':[v['nll'] for v in best_idea['per_seed']], 'mean':mean_metric(best_idea['per_seed']), 'std':float(np.std([v['nll'] for v in best_idea['per_seed']])), 'n':len(best_idea['per_seed']), 'best_config':best_idea['config'], 'sweep':idea},
127 {'mechanism_signature': {'predicted': 'positive uncertainty-to-missingness slope and high predicted/observed missingness correlation', 'observed_b_mean':float(np.mean([v['b'] for v in best_idea['per_seed']])), 'observed_missing_corr_mean':float(np.mean([v['missing_corr'] for v in best_idea['per_seed']])), 'confirmed':bool(np.mean([v['b'] for v in best_idea['per_seed']])>0 and np.mean([v['missing_corr'] for v in best_idea['per_seed']])>.5)}, 'custom_track': {'name':META['name'],'file':'missingness_bench.py','domain':META['domain']}, 'math_check':math_check, 'protocol':{'baseline_grid':grid,'idea_grid':lrs,'paired_deltas_nll':diffs,'permutation_p':permutation_pvalue(diffs),'seeds':list(seeds)}})
128 report['mechanism_signature']={'predicted': 'positive uncertainty-to-missingness slope and high predicted/observed missingness correlation',
129 'observed_b_mean':float(np.mean([v['b'] for v in best_idea['per_seed']])),
130 'observed_missing_corr_mean':float(np.mean([v['missing_corr'] for v in best_idea['per_seed']])),
131 'confirmed':bool(np.mean([v['b'] for v in best_idea['per_seed']])>0 and np.mean([v['missing_corr'] for v in best_idea['per_seed']])>.5)}
132 report['custom_track']={'name':META['name'],'file':'missingness_bench.py','domain':META['domain']}
133 report['math_check']=math_check
134 report['protocol']={'baseline_grid':grid,'idea_grid':lrs,'paired_deltas_nll':diffs,'permutation_p':permutation_pvalue(diffs),'seeds':list(seeds)}
135 print(json.dumps(report,indent=2))
136 Path('bench_report.json').write_text(json.dumps(report,indent=2))
137
138if __name__=='__main__': main()