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