import json, math, random import numpy as np import torch from torch import nn import torch.nn.functional as F SEED = 204 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) def math_checks(): rng = np.random.default_rng(SEED) # Check Bayes odds identity for two Gaussian populations at random x values. n = 20000 x = rng.normal(size=n) mu0, mu1, pi = -0.7, 1.1, 0.3 p0 = np.exp(-0.5*(x-mu0)**2)/np.sqrt(2*np.pi) p1 = np.exp(-0.5*(x-mu1)**2)/np.sqrt(2*np.pi) q = pi*p1/(pi*p1+(1-pi)*p0) lhs = q/(1-q) rhs = (pi/(1-pi)) * p1/p0 odds_rel_err = float(np.max(np.abs(lhs-rhs)/(np.abs(rhs)+1e-12))) # Check total-variance/Fisher decomposition on simulated scalar scores. a = rng.integers(0, 2, size=n) scores = rng.normal(loc=np.where(a == 0, -0.8, 1.4), scale=np.where(a == 0, .7, 1.1)) total = np.var(scores) within = sum(np.mean(a == k)*np.var(scores[a == k]) for k in (0,1)) means = np.array([np.mean(scores[a == k]) for k in (0,1)]) probs = np.array([np.mean(a == k) for k in (0,1)]) between = np.sum(probs*(means - np.sum(probs*means))**2) fisher_abs_err = float(abs(total - within - between)) return {"odds_max_relative_error": odds_rel_err, "fisher_decomposition_abs_error": fisher_abs_err, "total_variance": float(total), "within": float(within), "between": float(between)} def make_data(n_target=64, n_aux=4000, n_test=4000): rng = np.random.default_rng(SEED + 1) # Target labels depend on x0. Auxiliary examples are unlabeled and shifted; # the domain head can learn population structure from many examples. xt = rng.normal(size=(n_target, 2)).astype(np.float32) yt = (xt[:, 0] + 0.35*xt[:, 1] > 0).astype(np.int64) xa = rng.normal(size=(n_aux, 2)).astype(np.float32) xa[:, 0] += 1.5 xa[:, 1] += 0.25 xe = rng.normal(size=(n_test, 2)).astype(np.float32) ye = (xe[:, 0] + 0.35*xe[:, 1] > 0).astype(np.int64) return tuple(torch.tensor(v) for v in (xt, yt, xa, xe, ye)) class Net(nn.Module): def __init__(self): super().__init__() self.encoder = nn.Sequential(nn.Linear(2, 12), nn.Tanh(), nn.Linear(12, 8), nn.Tanh()) self.target = nn.Linear(8, 2) self.domain = nn.Linear(8, 1) def forward(self, x): z = self.encoder(x) return self.target(z), self.domain(z).squeeze(-1) def train(kind, data, lam=0.1, epochs=50, batch=32): xt, yt, xa, xe, ye = data torch.manual_seed(SEED + (0 if kind == 'baseline' else 10)) model = Net() opt = torch.optim.AdamW(model.parameters(), lr=3e-3, weight_decay=1e-4) gen = torch.Generator().manual_seed(SEED + 20) for _ in range(epochs): perm = torch.randperm(len(xt), generator=gen) for start in range(0, len(xt), batch): ids = perm[start:start+batch] xb, yb = xt[ids], yt[ids] logits, _ = model(xb) loss = F.cross_entropy(logits, yb) if kind == 'joint': aid = torch.randint(0, len(xa), (len(ids),), generator=gen) xd = torch.cat([xb, xa[aid]]) ad = torch.cat([torch.zeros(len(ids)), torch.ones(len(ids))]) _, dl = model(xd) loss = loss + lam * F.binary_cross_entropy_with_logits(dl, ad) opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): logits, dl = model(xe) nll = float(F.cross_entropy(logits, ye)) acc = float((logits.argmax(1) == ye).float().mean()) # Domain accuracy is diagnostic, not a target metric. xd = torch.cat([xt, xa]); ad = torch.cat([torch.zeros(len(xt)), torch.ones(len(xa))]) _, dl = model(xd) da = float(((dl > 0).float() == ad).float().mean()) return {"test_nll": nll, "test_accuracy": acc, "domain_accuracy": da} def main(): checks = math_checks() rows = [] for nt in (16, 64, 256): full = make_data(n_target=nt, n_aux=2000) for kind, lam in (("baseline", 0.0), ("joint", 0.01), ("joint", 0.1), ("joint", 0.5)): # Keep the same auxiliary pool and target test distribution; only target labels vary. out = train(kind, full, lam=lam) rows.append({"n_target": nt, "method": kind, "lambda": lam, **out}) print(json.dumps({"math_checks": checks, "results": rows}, indent=2)) if __name__ == '__main__': main()