Full-Likelihood Auxiliary Representation Training / experiment.py
Mechanism works
1import json, math, random
2import numpy as np
3import torch
4from torch import nn
5import torch.nn.functional as F
6
7SEED = 204
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9torch.set_num_threads(4)
10
11
12def math_checks():
13 rng = np.random.default_rng(SEED)
14 # Check Bayes odds identity for two Gaussian populations at random x values.
15 n = 20000
16 x = rng.normal(size=n)
17 mu0, mu1, pi = -0.7, 1.1, 0.3
18 p0 = np.exp(-0.5*(x-mu0)**2)/np.sqrt(2*np.pi)
19 p1 = np.exp(-0.5*(x-mu1)**2)/np.sqrt(2*np.pi)
20 q = pi*p1/(pi*p1+(1-pi)*p0)
21 lhs = q/(1-q)
22 rhs = (pi/(1-pi)) * p1/p0
23 odds_rel_err = float(np.max(np.abs(lhs-rhs)/(np.abs(rhs)+1e-12)))
24 # Check total-variance/Fisher decomposition on simulated scalar scores.
25 a = rng.integers(0, 2, size=n)
26 scores = rng.normal(loc=np.where(a == 0, -0.8, 1.4), scale=np.where(a == 0, .7, 1.1))
27 total = np.var(scores)
28 within = sum(np.mean(a == k)*np.var(scores[a == k]) for k in (0,1))
29 means = np.array([np.mean(scores[a == k]) for k in (0,1)])
30 probs = np.array([np.mean(a == k) for k in (0,1)])
31 between = np.sum(probs*(means - np.sum(probs*means))**2)
32 fisher_abs_err = float(abs(total - within - between))
33 return {"odds_max_relative_error": odds_rel_err, "fisher_decomposition_abs_error": fisher_abs_err,
34 "total_variance": float(total), "within": float(within), "between": float(between)}
35
36
37def make_data(n_target=64, n_aux=4000, n_test=4000):
38 rng = np.random.default_rng(SEED + 1)
39 # Target labels depend on x0. Auxiliary examples are unlabeled and shifted;
40 # the domain head can learn population structure from many examples.
41 xt = rng.normal(size=(n_target, 2)).astype(np.float32)
42 yt = (xt[:, 0] + 0.35*xt[:, 1] > 0).astype(np.int64)
43 xa = rng.normal(size=(n_aux, 2)).astype(np.float32)
44 xa[:, 0] += 1.5
45 xa[:, 1] += 0.25
46 xe = rng.normal(size=(n_test, 2)).astype(np.float32)
47 ye = (xe[:, 0] + 0.35*xe[:, 1] > 0).astype(np.int64)
48 return tuple(torch.tensor(v) for v in (xt, yt, xa, xe, ye))
49
50
51class Net(nn.Module):
52 def __init__(self):
53 super().__init__()
54 self.encoder = nn.Sequential(nn.Linear(2, 12), nn.Tanh(), nn.Linear(12, 8), nn.Tanh())
55 self.target = nn.Linear(8, 2)
56 self.domain = nn.Linear(8, 1)
57 def forward(self, x):
58 z = self.encoder(x)
59 return self.target(z), self.domain(z).squeeze(-1)
60
61
62def train(kind, data, lam=0.1, epochs=50, batch=32):
63 xt, yt, xa, xe, ye = data
64 torch.manual_seed(SEED + (0 if kind == 'baseline' else 10))
65 model = Net()
66 opt = torch.optim.AdamW(model.parameters(), lr=3e-3, weight_decay=1e-4)
67 gen = torch.Generator().manual_seed(SEED + 20)
68 for _ in range(epochs):
69 perm = torch.randperm(len(xt), generator=gen)
70 for start in range(0, len(xt), batch):
71 ids = perm[start:start+batch]
72 xb, yb = xt[ids], yt[ids]
73 logits, _ = model(xb)
74 loss = F.cross_entropy(logits, yb)
75 if kind == 'joint':
76 aid = torch.randint(0, len(xa), (len(ids),), generator=gen)
77 xd = torch.cat([xb, xa[aid]])
78 ad = torch.cat([torch.zeros(len(ids)), torch.ones(len(ids))])
79 _, dl = model(xd)
80 loss = loss + lam * F.binary_cross_entropy_with_logits(dl, ad)
81 opt.zero_grad(); loss.backward(); opt.step()
82 with torch.no_grad():
83 logits, dl = model(xe)
84 nll = float(F.cross_entropy(logits, ye))
85 acc = float((logits.argmax(1) == ye).float().mean())
86 # Domain accuracy is diagnostic, not a target metric.
87 xd = torch.cat([xt, xa]); ad = torch.cat([torch.zeros(len(xt)), torch.ones(len(xa))])
88 _, dl = model(xd)
89 da = float(((dl > 0).float() == ad).float().mean())
90 return {"test_nll": nll, "test_accuracy": acc, "domain_accuracy": da}
91
92
93def main():
94 checks = math_checks()
95 rows = []
96 for nt in (16, 64, 256):
97 full = make_data(n_target=nt, n_aux=2000)
98 for kind, lam in (("baseline", 0.0), ("joint", 0.01), ("joint", 0.1), ("joint", 0.5)):
99 # Keep the same auxiliary pool and target test distribution; only target labels vary.
100 out = train(kind, full, lam=lam)
101 rows.append({"n_target": nt, "method": kind, "lambda": lam, **out})
102 print(json.dumps({"math_checks": checks, "results": rows}, indent=2))
103
104if __name__ == '__main__':
105 main()