PANDA / scripts /analysis /84_adversary_purification.py
bryan7264's picture
Correction pass: gate-matched Dahlin, retracted unsupported claims, complete HF-placode DEG set, restyled figures
141bacd verified
Raw
History Blame Contribute Delete
5.8 kB
"""probe dataset + depth adversary heads on training data to check trunk invariance."""
from pathlib import Path
import json, warnings, pickle, sys, numpy as np, pandas as pd, torch
warnings.filterwarnings("ignore")
import anndata as ad, scanpy as sc, scipy.sparse as sp
from pathlib import Path as _P_root
ROOT = _P_root(__file__).resolve().parents[2]
ROOT_STR = str(ROOT)
sys.path.insert(0, ROOT_STR)
from panda import PANDAEncoder
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
CKPT = Path(f"{ROOT_STR}/checkpoints")
CORP = Path(f"{ROOT_STR}/data/corpus")
OUT = Path(f"{ROOT_STR}/discovery"); OUT.mkdir(exist_ok=True)
def project_corpus(sys, subsample=20000):
stats = np.load(CORP / sys / "harmonized/corpus_stats.npz", allow_pickle=True)
hvgs = [str(g) for g in stats["shared_hvgs"]]
mu = np.asarray(stats["mean"], dtype=np.float32)
sig = np.asarray(stats["std"], dtype=np.float32)
pca = pickle.load(open(CORP / sys / "harmonized/pca_basis.pkl", "rb"))
ck = torch.load(CKPT / sys / "panda_final.pt", map_location=DEVICE, weights_only=False)
m = PANDAEncoder(n_pca=50, n_classes=len(ck["classes"]),
n_datasets=len(ck["datasets"])).to(DEVICE).eval()
m.load_state_dict(ck["model"])
corp = ad.read_h5ad(CORP / sys / "harmonized/corpus.h5ad")
print(f"[{sys}] corpus {corp.shape}", flush=True)
dcol_candidates = ["dataset_id", "dataset", "sample", "batch"]
dcol = None
for c in dcol_candidates:
if c in corp.obs.columns:
dcol = c; break
if dcol is None:
# HSC corpus is Weinreb-only (single dataset)
corp.obs["dataset_id"] = "single"
dcol = "dataset_id"
print(f"[{sys}] dataset col={dcol} n_unique={corp.obs[dcol].nunique()} train datasets in ckpt={len(ck['datasets'])}", flush=True)
if corp.n_obs > subsample:
rng = np.random.default_rng(0)
idx = rng.choice(corp.n_obs, size=subsample, replace=False)
corp = corp[idx].copy()
print(f"[{sys}] subsampled to {corp.n_obs} cells", flush=True)
hvg2i = {g: i for i, g in enumerate(hvgs)}
common = [g for g in corp.var_names.astype(str) if g in hvg2i]
a_c = corp[:, common].copy()
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((corp.n_obs, len(hvgs)), dtype=np.float32)
cols = np.array([hvg2i[g] for g in common])
Xf[:, cols] = X
Xz = np.clip((Xf - mu) / sig, -10, 10)
Xpca = pca.transform(Xz).astype(np.float32)
if "n_counts" in corp.obs.columns:
n_counts = corp.obs["n_counts"].astype(float).values
elif "total_counts" in corp.obs.columns:
n_counts = corp.obs["total_counts"].astype(float).values
else:
# X is log-normalised; expm1 first for approximate raw sum
n_counts = np.expm1(X.sum(axis=1))
log10c = np.log10(n_counts + 1)
log10c_z = (log10c - log10c.mean()) / (log10c.std() + 1e-8)
dset = corp.obs[dcol].astype(str).values
train_dsets = ck["datasets"]
dset_to_idx = {d: i for i, d in enumerate(train_dsets)}
y_dset = np.array([dset_to_idx.get(d, -1) for d in dset])
valid = y_dset >= 0
print(f"[{sys}] valid rows for dataset-eval: {valid.sum()} / {len(y_dset)}", flush=True)
dom_preds = np.zeros((corp.n_obs, len(train_dsets)), dtype=np.float32)
depth_preds = np.zeros(corp.n_obs, dtype=np.float32)
with torch.no_grad():
for i in range(0, corp.n_obs, 2048):
xb = torch.from_numpy(Xpca[i:i+2048]).to(DEVICE)
aux = torch.zeros(len(xb), 2, device=DEVICE)
out = m(xb, aux, lam_dann=0.0) # lam=0 disables GRL at inference
dom_preds[i:i+2048] = out["dom"].cpu().numpy()
depth_preds[i:i+2048] = out["depth"].cpu().numpy().squeeze()
from sklearn.metrics import accuracy_score, top_k_accuracy_score, mean_squared_error, r2_score
result = {"system": sys, "n_cells": int(corp.n_obs), "n_train_datasets": len(train_dsets)}
if valid.sum() > 0 and len(train_dsets) > 1:
dom_pred_argmax = dom_preds[valid].argmax(axis=1)
acc = float(accuracy_score(y_dset[valid], dom_pred_argmax))
chance = 1.0 / len(train_dsets)
result.update({
"dataset_adv_accuracy": acc,
"dataset_adv_chance": chance,
"dataset_adv_above_chance": acc - chance,
"dataset_adv_random_baseline_test": "None (single-dataset)" if len(train_dsets) == 1 else f"n_datasets={len(train_dsets)}, chance={chance:.3f}",
})
mse_depth = float(mean_squared_error(log10c_z, depth_preds))
r2_depth = float(r2_score(log10c_z, depth_preds))
result.update({
"depth_adv_mse_z": mse_depth,
"depth_adv_r2_z": r2_depth,
"depth_target_std_z": float(log10c_z.std()),
})
print(f"[{sys}] dom_adv_acc={result.get('dataset_adv_accuracy', 'NA')} vs chance={result.get('dataset_adv_chance', 'NA')}", flush=True)
print(f"[{sys}] depth_adv MSE_z={mse_depth:.4f} R²_z={r2_depth:.4f} (R²≤0 ⇒ trunk fully depth-invariant)", flush=True)
return result
all_results = {}
for sys in ["pan_skin", "hematopoiesis", "pancreas"]:
try:
all_results[sys] = project_corpus(sys)
except Exception as e:
import traceback; traceback.print_exc()
all_results[sys] = {"error": str(e)}
json.dump(all_results, open(OUT / "84_adversary_purification.json", "w"), indent=2, default=str)
print(f"\nwrote {OUT}/84_adversary_purification.json", flush=True)
print(json.dumps(all_results, indent=2, default=str), flush=True)