Download scripts/common/train_panda.py from bryan7264/PANDA: direct link, hf CLI and curl.
- Browser
- Download file 8.62 kB
-
https://huggingface.co/bryan7264/PANDA/resolve/main/scripts/common/train_panda.py
- Command line
-
hf download hf://bryan7264/PANDA/scripts/common/train_panda.py
-
curl -L -o train_panda.py https://huggingface.co/bryan7264/PANDA/resolve/main/scripts/common/train_panda.py
8.62 kB
| """train panda on the canonical paper-labeled corpus for one system + variant.""" | |
| from __future__ import annotations | |
| import argparse, sys, json, pickle, warnings, numpy as np, pandas as pd, torch, torch.nn.functional as F | |
| from pathlib import Path | |
| import anndata as ad, scanpy as sc, scipy.sparse as sp, yaml | |
| warnings.filterwarnings("ignore"); sc.settings.verbosity = 0 | |
| import os as _os | |
| from pathlib import Path as _Path | |
| PANDA_ROOT = _Path(_os.environ.get("PANDA_ROOT", str(_Path(__file__).resolve().parents[2]))) | |
| sys.path.insert(0, str(PANDA_ROOT)) | |
| from panda import ( | |
| PANDAEncoder, supcon_loss, vicreg_loss, hsic_biased, | |
| subcenter_angular_infonce, prototype_repulsion, | |
| ) | |
| ROOT = Path(str(PANDA_ROOT)) | |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| def load_corpus(system): | |
| """canonical loader; alias for load_corpus_v3 after finalize_rename.""" | |
| return load_corpus_v3(system) | |
| def load_corpus_v3(system): | |
| # kept for backward-compat with older scripts that import load_corpus_v3 | |
| p = ROOT / f"data/corpus/{system}/harmonized/corpus.h5ad" | |
| stats = np.load(ROOT / f"data/corpus/{system}/harmonized/corpus_stats.npz", allow_pickle=True) | |
| pca = pickle.load(open(ROOT / f"data/corpus/{system}/harmonized/pca_basis.pkl", "rb")) | |
| a = ad.read_h5ad(p) | |
| hvgs = [str(g) for g in stats["shared_hvgs"]] | |
| return a, hvgs, stats["mean"], stats["std"], pca | |
| def get_marker_gene_list(system): | |
| y = yaml.safe_load(open(ROOT / "panda/markers.yaml")) | |
| return y[system] | |
| def prepare_batches(adata, hvgs, mu, sig, pca, marker_genes=None, variant="pca", | |
| legacy_double_norm=True): | |
| # corpus.h5ad X is already normalize_total+log1p'd by the corpus builders, and the | |
| # published checkpoints were trained with a second normalize/log1p applied on top | |
| # (against per-gene stats computed from singly-normalized data). legacy_double_norm=True | |
| # reproduces that behaviour bit-for-bit; pass False to train on the corpus values the | |
| # stats were actually computed from. Do not mix: a checkpoint must be evaluated under | |
| # the same setting it was trained with. | |
| hvg2i = {g: i for i, g in enumerate(hvgs)} | |
| common = [g for g in adata.var_names.astype(str) if g in hvg2i] | |
| a_c = adata[:, common].copy() | |
| if legacy_double_norm: | |
| sc.pp.normalize_total(a_c, target_sum=1e4); sc.pp.log1p(a_c) | |
| X_ = a_c.X.toarray().astype(np.float32) if sp.issparse(a_c.X) else a_c.X.astype(np.float32) | |
| Xf = np.zeros((adata.n_obs, len(hvgs)), dtype=np.float32) | |
| cols = np.array([hvg2i[g] for g in common]) | |
| Xf[:, cols] = X_ | |
| Xz = np.clip((Xf - mu.astype(np.float32)) / sig.astype(np.float32), -10, 10) | |
| Xpca = pca.transform(Xz).astype(np.float32) | |
| Xmark = None | |
| if variant == "marker" and marker_genes: | |
| mvals = np.zeros((adata.n_obs, len(marker_genes)), dtype=np.float32) | |
| for j, g in enumerate(marker_genes): | |
| if g in adata.var_names: | |
| col = adata[:, g].X | |
| if sp.issparse(col): col = col.toarray() | |
| mvals[:, j] = col.flatten().astype(np.float32) | |
| mmu = mvals.mean(axis=0, keepdims=True); msig = mvals.std(axis=0, keepdims=True) + 1e-6 | |
| Xmark = np.clip((mvals - mmu) / msig, -5, 5).astype(np.float32) | |
| # expose the training-corpus marker stats so the checkpoint can carry them; | |
| # inference must reuse these rather than refit on the target dataset | |
| prepare_batches.last_marker_stats = (mmu, msig) | |
| labels = adata.obs["canonical_label"].astype(str).values | |
| classes = sorted(set(labels)) | |
| y = np.array([classes.index(l) for l in labels], dtype=np.int64) | |
| datasets = sorted(set(adata.obs["dataset"].astype(str).values)) | |
| y_dset = np.array([datasets.index(d) for d in adata.obs["dataset"].astype(str).values], dtype=np.int64) | |
| counts = np.asarray(adata.X.sum(axis=1)).ravel() | |
| log10cz = ((np.log10(counts + 1) - np.log10(counts + 1).mean()) / | |
| (np.log10(counts + 1).std() + 1e-6)).astype(np.float32) | |
| return Xpca, Xmark, y, classes, y_dset, datasets, log10cz | |
| def train(system, variant, epochs=8, batch=256, lr=1e-3): | |
| a, hvgs, mu, sig, pca = load_corpus_v3(system) | |
| marker_genes = get_marker_gene_list(system) if variant == "marker" else [] | |
| Xpca, Xmark, y, classes, y_dset, datasets, log10cz = prepare_batches( | |
| a, hvgs, mu, sig, pca, marker_genes, variant | |
| ) | |
| print(f"[train] {system}/{variant} n={a.n_obs} K={len(classes)} datasets={len(datasets)}", flush=True) | |
| print(f"[train] classes: {classes}", flush=True) | |
| n_markers = Xmark.shape[1] if Xmark is not None else 0 | |
| model = PANDAEncoder(variant=variant, n_pca=50, n_markers=n_markers, | |
| n_classes=len(classes), n_sub=3, n_datasets=len(datasets), dropout=0.2).to(DEVICE) | |
| opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4) | |
| rng = np.random.default_rng(0) | |
| for epoch in range(epochs): | |
| stage = 0 if epoch < 1 else 1 if epoch < 3 else 2 if epoch < 6 else 3 | |
| for g in opt.param_groups: g["lr"] = lr * (0.5 if epoch >= epochs - 1 else 1.0) | |
| perm = rng.permutation(a.n_obs) | |
| losses = [] | |
| for bstart in range(0, a.n_obs, batch): | |
| idx = perm[bstart:bstart+batch] | |
| x = torch.from_numpy(Xpca[idx]).to(DEVICE) | |
| xm = torch.from_numpy(Xmark[idx]).to(DEVICE) if Xmark is not None else None | |
| yy = torch.from_numpy(y[idx]).to(DEVICE) | |
| yd = torch.from_numpy(y_dset[idx]).to(DEVICE) | |
| dd = torch.from_numpy(log10cz[idx]).float().to(DEVICE).unsqueeze(1) | |
| aux = torch.zeros(len(idx), 2, device=DEVICE) | |
| lam = 0.1 if stage >= 2 else 0.0 | |
| out = model(x, aux, x_markers=xm, lam_dann=lam) | |
| z = out["z"] | |
| L = supcon_loss(z, yy, 0.1) + 1.0 * vicreg_loss(z) + 0.4 * F.cross_entropy(out["logits"], yy) | |
| if stage >= 1: | |
| L = L + 0.6 * subcenter_angular_infonce(z, yy, model.prototypes.detach().clone(), | |
| margin=0.15, temperature=0.07) | |
| if stage >= 2: | |
| L = L + F.cross_entropy(out["dom"], yd) + 0.3 * F.mse_loss(out["depth"], dd) + 0.05 * hsic_biased(out["repr"], dd) | |
| # NOTE: earlier revisions added 0.5 * prototype_repulsion(model.prototypes.detach()) | |
| # at stage 3. prototypes is a gradient-free EMA buffer and the tensor was detached, | |
| # so the term contributed exactly zero gradient — it only inflated the printed loss. | |
| # Removed from the objective (behaviour-preserving); prototype_repulsion() remains | |
| # in panda.model as an analysis metric. | |
| opt.zero_grad(); L.backward(); opt.step() | |
| if stage >= 1: | |
| with torch.no_grad(): model.update_prototypes(z.detach(), yy) | |
| losses.append(float(L)) | |
| print(f"[train {system}/{variant}] epoch {epoch}/{epochs} stage={stage} loss={np.mean(losses):.4f}", flush=True) | |
| # save to the same path every consumer (zero_shot, run_all_zero_shot, extract_prototypes) | |
| # loads from. Earlier revisions wrote to checkpoints/{system}_v3/ while consumers read | |
| # checkpoints/{system}/ — a retrain silently never propagated. | |
| ck_dir = ROOT / f"checkpoints/{system}/{variant}" | |
| ck_dir.mkdir(parents=True, exist_ok=True) | |
| mstats = getattr(prepare_batches, "last_marker_stats", None) if variant == "marker" else None | |
| torch.save({"model": model.state_dict(), "classes": classes, "datasets": datasets, | |
| "marker_genes": marker_genes if variant == "marker" else [], | |
| "marker_mu": mstats[0] if mstats else None, | |
| "marker_sig": mstats[1] if mstats else None, | |
| "legacy_double_norm": True, | |
| "prototypes": model.prototypes.detach().cpu().numpy()}, | |
| ck_dir / "panda_final.pt") | |
| print(f"[save] {ck_dir}/panda_final.pt", flush=True) | |
| if __name__ == "__main__": | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("system", choices=["pan_skin", "hematopoiesis", "pancreas"]) | |
| ap.add_argument("--variant", choices=["pca", "marker"], required=True) | |
| ap.add_argument("--epochs", type=int, default=8) | |
| ap.add_argument("--batch", type=int, default=256) | |
| ap.add_argument("--lr", type=float, default=1e-3) | |
| args = ap.parse_args() | |
| train(args.system, args.variant, args.epochs, args.batch, args.lr) | |