Download src/cfref/_core.py from qiongli0705/cfREF: direct link, hf CLI and curl.
- Browser
- Download file 2.34 kB
-
https://huggingface.co/qiongli0705/cfREF/resolve/main/src/cfref/_core.py
- Command line
-
hf download hf://qiongli0705/cfREF/src/cfref/_core.py
-
curl -L -o _core.py https://huggingface.co/qiongli0705/cfREF/resolve/main/src/cfref/_core.py
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 | |