Missingness-as-a-Label Signal / missingness_bench.py

Mechanism confirmed, baseline not beaten

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