File size: 5,796 Bytes
141bacd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
"""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)