Missingness-as-a-Label Signal / missingness_signal.py
Mechanism confirmed, baseline not beaten
1import json, math, random
2import numpy as np
3import torch
4from torch import nn
5
6SEED = 17
7
8def seed_all(seed=SEED):
9 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
10 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
11
12def entropy(logits):
13 p = torch.softmax(logits, dim=-1)
14 return -(p * torch.log(p.clamp_min(1e-8))).sum(dim=-1)
15
16class MLP(nn.Module):
17 def __init__(self):
18 super().__init__()
19 self.net = nn.Sequential(nn.Linear(2, 24), nn.Tanh(), nn.Linear(24, 2))
20 def forward(self, x): return self.net(x)
21
22class Joint(nn.Module):
23 def __init__(self):
24 super().__init__()
25 self.classifier = MLP()
26 self.a = nn.Parameter(torch.tensor(-0.2))
27 self.beta = nn.Parameter(torch.tensor(0.2))
28 def forward(self, x):
29 logits = self.classifier(x)
30 h = entropy(logits)
31 b = torch.nn.functional.softplus(self.beta)
32 q = torch.sigmoid(self.a + b*h)
33 return logits, h, q, b
34
35def make_data(n, seed, missing_mode='informative'):
36 g = torch.Generator().manual_seed(seed)
37 y = torch.randint(0, 2, (n,), generator=g)
38 means = torch.where(y[:, None] == 0, torch.tensor([-1.2, 0.0]), torch.tensor([1.2, 0.0]))
39 x = means + 0.95 * torch.randn(n, 2, generator=g)
40 # Bayes posterior entropy is the teacher uncertainty signal.
41 teacher_logits = torch.stack([-2.4*x[:, 0], 2.4*x[:, 0]], dim=1)
42 h = entropy(teacher_logits)
43 if missing_mode == 'informative':
44 q = torch.sigmoid(-1.35 + 3.0*h)
45 elif missing_mode == 'random':
46 q = torch.full((n,), 0.35)
47 else:
48 raise ValueError(missing_mode)
49 m = torch.bernoulli(q, generator=g)
50 return x, y, m, q, h
51
52def math_check():
53 torch.manual_seed(3)
54 logits = torch.tensor([[1.2, -0.4], [0.1, 0.2], [-.7, .9]], dtype=torch.double, requires_grad=True)
55 a = torch.tensor(-0.4, dtype=torch.double, requires_grad=True)
56 b = torch.tensor(1.7, dtype=torch.double, requires_grad=True)
57 h = entropy(logits)
58 q = torch.sigmoid(a + b*h)
59 m = torch.tensor([1., 0., 1.], dtype=torch.double)
60 y = torch.tensor([0, 1, 1])
61 ce = nn.functional.cross_entropy(logits, y, reduction='none')
62 loss = (m*(-torch.log(q)) + (1-m)*(-torch.log(1-q) + ce)).mean()
63 loss.backward()
64 # For each item, d[-log likelihood]/d eta = q-m, eta=a+bH.
65 expected_a = (q-m).mean().item()
66 expected_b = ((q-m)*h).mean().item()
67 finite = []
68 for i in range(3):
69 z = logits.detach().double().clone().requires_grad_(True)
70 hh = entropy(z)[i]
71 qq = torch.sigmoid(a.detach()+b.detach()*hh)
72 li = -torch.log(qq) if bool(m[i].item()) else -torch.log(1-qq)+nn.functional.cross_entropy(z[i:i+1], y[i:i+1])
73 finite.append(torch.autograd.grad(li, z)[0][i].norm().item())
74 return {'a_grad': float(a.grad), 'a_expected': expected_a,
75 'b_grad': float(b.grad), 'b_expected': expected_b,
76 'nonzero_entropy_to_logit_grad_norms': finite,
77 'gradient_match': bool(abs(a.grad.item()-expected_a)<1e-10 and abs(b.grad.item()-expected_b)<1e-10 and max(finite)>1e-6)}
78
79def train(x, y, m, mode, steps=500, seed=0):
80 seed_all(seed)
81 device = 'cuda' if torch.cuda.is_available() else 'cpu'
82 try:
83 x, y, m = x.to(device), y.to(device), m.to(device)
84 if mode == 'joint': model = Joint().to(device)
85 else: model = MLP().to(device)
86 opt = torch.optim.Adam(model.parameters(), lr=0.012)
87 for _ in range(steps):
88 logits = model(x) if mode == 'supervised' else model.classifier(x)[0] if False else None
89 if mode == 'supervised':
90 loss = nn.functional.cross_entropy(logits[m < 0.5], y[m < 0.5])
91 else:
92 logits, h, q, b = model(x)
93 if mode == 'detached':
94 q = torch.sigmoid(model.a + b*h.detach())
95 ce = nn.functional.cross_entropy(logits, y, reduction='none')
96 loss = ((1-m)*ce - m*torch.log(q.clamp_min(1e-7)) - (1-m)*torch.log((1-q).clamp_min(1e-7))).mean()
97 loss = loss + 1e-3*(model.a.square()+b.square())
98 opt.zero_grad(); loss.backward(); opt.step()
99 return model, device
100 except Exception:
101 device = 'cpu'; x, y, m = x.cpu(), y.cpu(), m.cpu()
102 return train_cpu(x,y,m,mode,steps,seed)
103
104def train_cpu(x,y,m,mode,steps,seed):
105 seed_all(seed); model = Joint() if mode != 'supervised' else MLP(); opt=torch.optim.Adam(model.parameters(),lr=.012)
106 for _ in range(steps):
107 if mode=='supervised': loss=nn.functional.cross_entropy(model(x)[m<.5],y[m<.5])
108 else:
109 z,h,q,b=model(x); q=torch.sigmoid(model.a+b*(h.detach() if mode=='detached' else h)); ce=nn.functional.cross_entropy(z,y,reduction='none')
110 loss=((1-m)*ce-m*torch.log(q.clamp_min(1e-7))-(1-m)*torch.log((1-q).clamp_min(1e-7))).mean()+1e-3*(model.a.square()+b.square())
111 opt.zero_grad();loss.backward();opt.step()
112 return model,'cpu'
113
114def evaluate(model, device, x, y, m, qtrue):
115 with torch.no_grad():
116 z,h,q,b=model(x.to(device)) if isinstance(model,Joint) else (model(x.to(device)),None,None,None)
117 p=torch.softmax(z,-1); pred=p.argmax(-1).cpu(); yy=y
118 acc=(pred==yy).float().mean().item(); nll=nn.functional.cross_entropy(z,y.to(device)).item()
119 ent=(-(p*torch.log(p.clamp_min(1e-8))).sum(-1)).cpu()
120 if q is not None:
121 corr=np.corrcoef(q.cpu().numpy(),qtrue.numpy())[0,1]
122 return acc,nll,float(ent.mean()),float(b.cpu()),float(corr)
123 return acc,nll,float(ent.mean()),None,None
124
125def main():
126 seed_all(); check=math_check(); results={'math_check':check}
127 x,y,m,q,h=make_data(900, 41, 'informative'); xt,yt,mt,qt,ht=make_data(2200, 42, 'informative')
128 results['missing_fraction']=float(m.mean()); results['train_entropy_missing_corr']=float(np.corrcoef(m.numpy(),h.numpy())[0,1])
129 for mode in ['supervised','detached','joint']:
130 model,dev=train(x,y,m,mode,500,seed=55)
131 results[mode]=evaluate(model,dev,xt,yt,mt,qt)
132 xr,yr,mr,qr,hr=make_data(900, 43, 'random'); xrt,yrt,mrt,qrt,hrt=make_data(2200,44,'random')
133 for mode in ['supervised','joint']:
134 model,dev=train(xr,yr,mr,mode,500,seed=56)
135 results['random_'+mode]=evaluate(model,dev,xrt,yrt,mrt,qrt)
136 with open('results.json','w') as f: json.dump(results,f,indent=2)
137 print(json.dumps(results,indent=2))
138
139if __name__=='__main__': main()