cfREF / src /cfref /_core.py
qiongli0705's picture
Upload 10 files
703e278 verified
Raw History Blame Contribute Delete
2.34 kB
"""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