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

Mechanism confirmed, baseline not beaten

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