import json, math, random import numpy as np import torch from torch import nn SEED = 17 def seed_all(seed=SEED): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def entropy(logits): p = torch.softmax(logits, dim=-1) return -(p * torch.log(p.clamp_min(1e-8))).sum(dim=-1) class MLP(nn.Module): def __init__(self): super().__init__() self.net = nn.Sequential(nn.Linear(2, 24), nn.Tanh(), nn.Linear(24, 2)) def forward(self, x): return self.net(x) class Joint(nn.Module): def __init__(self): super().__init__() self.classifier = MLP() self.a = nn.Parameter(torch.tensor(-0.2)) self.beta = nn.Parameter(torch.tensor(0.2)) def forward(self, x): logits = self.classifier(x) h = entropy(logits) b = torch.nn.functional.softplus(self.beta) q = torch.sigmoid(self.a + b*h) return logits, h, q, b def make_data(n, seed, missing_mode='informative'): g = torch.Generator().manual_seed(seed) y = torch.randint(0, 2, (n,), generator=g) means = torch.where(y[:, None] == 0, torch.tensor([-1.2, 0.0]), torch.tensor([1.2, 0.0])) x = means + 0.95 * torch.randn(n, 2, generator=g) # Bayes posterior entropy is the teacher uncertainty signal. teacher_logits = torch.stack([-2.4*x[:, 0], 2.4*x[:, 0]], dim=1) h = entropy(teacher_logits) if missing_mode == 'informative': q = torch.sigmoid(-1.35 + 3.0*h) elif missing_mode == 'random': q = torch.full((n,), 0.35) else: raise ValueError(missing_mode) m = torch.bernoulli(q, generator=g) return x, y, m, q, h def math_check(): torch.manual_seed(3) logits = torch.tensor([[1.2, -0.4], [0.1, 0.2], [-.7, .9]], dtype=torch.double, requires_grad=True) a = torch.tensor(-0.4, dtype=torch.double, requires_grad=True) b = torch.tensor(1.7, dtype=torch.double, requires_grad=True) h = entropy(logits) q = torch.sigmoid(a + b*h) m = torch.tensor([1., 0., 1.], dtype=torch.double) y = torch.tensor([0, 1, 1]) ce = nn.functional.cross_entropy(logits, y, reduction='none') loss = (m*(-torch.log(q)) + (1-m)*(-torch.log(1-q) + ce)).mean() loss.backward() # For each item, d[-log likelihood]/d eta = q-m, eta=a+bH. expected_a = (q-m).mean().item() expected_b = ((q-m)*h).mean().item() finite = [] for i in range(3): z = logits.detach().double().clone().requires_grad_(True) hh = entropy(z)[i] qq = torch.sigmoid(a.detach()+b.detach()*hh) 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]) finite.append(torch.autograd.grad(li, z)[0][i].norm().item()) return {'a_grad': float(a.grad), 'a_expected': expected_a, 'b_grad': float(b.grad), 'b_expected': expected_b, 'nonzero_entropy_to_logit_grad_norms': finite, '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)} def train(x, y, m, mode, steps=500, seed=0): seed_all(seed) device = 'cuda' if torch.cuda.is_available() else 'cpu' try: x, y, m = x.to(device), y.to(device), m.to(device) if mode == 'joint': model = Joint().to(device) else: model = MLP().to(device) opt = torch.optim.Adam(model.parameters(), lr=0.012) for _ in range(steps): logits = model(x) if mode == 'supervised' else model.classifier(x)[0] if False else None if mode == 'supervised': loss = nn.functional.cross_entropy(logits[m < 0.5], y[m < 0.5]) else: logits, h, q, b = model(x) if mode == 'detached': q = torch.sigmoid(model.a + b*h.detach()) ce = nn.functional.cross_entropy(logits, y, reduction='none') loss = ((1-m)*ce - m*torch.log(q.clamp_min(1e-7)) - (1-m)*torch.log((1-q).clamp_min(1e-7))).mean() loss = loss + 1e-3*(model.a.square()+b.square()) opt.zero_grad(); loss.backward(); opt.step() return model, device except Exception: device = 'cpu'; x, y, m = x.cpu(), y.cpu(), m.cpu() return train_cpu(x,y,m,mode,steps,seed) def train_cpu(x,y,m,mode,steps,seed): seed_all(seed); model = Joint() if mode != 'supervised' else MLP(); opt=torch.optim.Adam(model.parameters(),lr=.012) for _ in range(steps): if mode=='supervised': loss=nn.functional.cross_entropy(model(x)[m<.5],y[m<.5]) else: 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') 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()) opt.zero_grad();loss.backward();opt.step() return model,'cpu' def evaluate(model, device, x, y, m, qtrue): with torch.no_grad(): z,h,q,b=model(x.to(device)) if isinstance(model,Joint) else (model(x.to(device)),None,None,None) p=torch.softmax(z,-1); pred=p.argmax(-1).cpu(); yy=y acc=(pred==yy).float().mean().item(); nll=nn.functional.cross_entropy(z,y.to(device)).item() ent=(-(p*torch.log(p.clamp_min(1e-8))).sum(-1)).cpu() if q is not None: corr=np.corrcoef(q.cpu().numpy(),qtrue.numpy())[0,1] return acc,nll,float(ent.mean()),float(b.cpu()),float(corr) return acc,nll,float(ent.mean()),None,None def main(): seed_all(); check=math_check(); results={'math_check':check} x,y,m,q,h=make_data(900, 41, 'informative'); xt,yt,mt,qt,ht=make_data(2200, 42, 'informative') results['missing_fraction']=float(m.mean()); results['train_entropy_missing_corr']=float(np.corrcoef(m.numpy(),h.numpy())[0,1]) for mode in ['supervised','detached','joint']: model,dev=train(x,y,m,mode,500,seed=55) results[mode]=evaluate(model,dev,xt,yt,mt,qt) xr,yr,mr,qr,hr=make_data(900, 43, 'random'); xrt,yrt,mrt,qrt,hrt=make_data(2200,44,'random') for mode in ['supervised','joint']: model,dev=train(xr,yr,mr,mode,500,seed=56) results['random_'+mode]=evaluate(model,dev,xrt,yrt,mrt,qrt) with open('results.json','w') as f: json.dump(results,f,indent=2) print(json.dumps(results,indent=2)) if __name__=='__main__': main()