Full-Likelihood Auxiliary Representation Training / experiment.py

Mechanism works

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