"""Network and episodic training adapted from the manuscript implementation.""" import numpy as np import torch import torch.nn as nn def _mlp(d_in, d_h, d_out): return nn.Sequential(nn.Linear(d_in, d_h), nn.GELU(), nn.Linear(d_h, d_out)) class CfRefNet(nn.Module): def __init__(self, d_in, d_h=128, d_e=48): super().__init__() self.phi = _mlp(d_in, d_h, d_e) self.head = _mlp(3 * d_e, d_h, 1) def forward(self, x, R): """x: (n, d) queries; R: (m, d) reference controls from the same centre.""" e = self.phi(x) er = self.phi(R) mu = er.mean(0, keepdim=True) sd = er.std(0, keepdim=True) + 1e-6 feats = torch.cat([e, e - mu, ((e - mu) / sd)], dim=-1) return self.head(feats).squeeze(-1) def train_cfref(X, y, cohort, n_ref_range=(5, 30), episodes=3000, d_h=128, d_e=48, lr=1e-3, weight_decay=1e-4, batch=32, seed=0, device="cpu"): """Episodic training: each episode draws a centre, a reference set of that centre's controls, and a query minibatch from the remaining samples.""" rng = np.random.RandomState(seed) torch.manual_seed(seed) X = np.asarray(X, dtype=np.float32); y = np.asarray(y); cohort = np.asarray(cohort) net = CfRefNet(X.shape[1], d_h, d_e).to(device) opt = torch.optim.AdamW(net.parameters(), lr=lr, weight_decay=weight_decay) Xt = torch.tensor(X, device=device) yt = torch.tensor(y, dtype=torch.float32, device=device) centres = [c for c in np.unique(cohort) if ((cohort == c) & (y == 0)).sum() >= n_ref_range[0] + 3 and ((cohort == c) & (y == 1)).sum() >= 1] if not centres: raise ValueError("no training centre has enough controls") for ep in range(episodes): c = centres[rng.randint(len(centres))] idx_c = np.where(cohort == c)[0] ctrl = idx_c[y[idx_c] == 0] m = rng.randint(n_ref_range[0], min(n_ref_range[1], len(ctrl) - 2) + 1) ref = rng.choice(ctrl, m, replace=False) pool = np.setdiff1d(idx_c, ref) q = rng.choice(pool, min(batch, len(pool)), replace=False) logit = net(Xt[q], Xt[ref]) loss = nn.functional.binary_cross_entropy_with_logits(logit, yt[q]) opt.zero_grad(); loss.backward(); opt.step() net.eval() return net