| """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:
|
|
|
| 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:
|
|
|
| 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)
|
| 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)
|
|
|