PANDA / scripts /analysis /81_counterfactual_knockouts.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.04 kB
"""per-class gene knockouts: zero each candidate gene, remeasure prototype cosine."""
from __future__ import annotations
from pathlib import Path
import warnings, json, pickle, sys, numpy as np, pandas as pd
warnings.filterwarnings("ignore")
import torch
import anndata as ad
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
CKPT_ROOT = Path(f"{ROOT_STR}/checkpoints")
CORP_ROOT = Path(f"{ROOT_STR}/data/corpus")
OUT = Path(f"{ROOT_STR}/discovery"); OUT.mkdir(exist_ok=True)
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
TOP_KO_GENES = 100 # only test top-100 attributed genes for KO
def compute_ko(sys):
print(f"\n===== {sys} =====", flush=True)
stats = np.load(CORP_ROOT / sys / "harmonized/corpus_stats.npz", allow_pickle=True)
hvgs = [str(g) for g in stats["shared_hvgs"]]
pca = pickle.load(open(CORP_ROOT / sys / "harmonized/pca_basis.pkl", "rb"))
mu = np.asarray(stats["mean"], dtype=np.float32)
sig = np.asarray(stats["std"], dtype=np.float32)
ck = torch.load(CKPT_ROOT / sys / "marker" / "panda_final.pt", map_location=DEVICE, weights_only=False)
classes = ck["classes"]
model = PANDAEncoder(variant="marker", n_pca=50, n_markers=len(ck.get("marker_genes",[])), n_classes=len(classes), n_sub=3,
n_datasets=len(ck["datasets"])).to(DEVICE).eval()
model.load_state_dict(ck["model"])
protos = torch.from_numpy(ck["prototypes"]).to(DEVICE)
protos = protos / (protos.norm(dim=1, keepdim=True) + 1e-8)
import scipy.sparse as sp
import scanpy as sc
corp = ad.read_h5ad(CORP_ROOT / sys / "harmonized/corpus.h5ad")
label_key = "canonical_label" if "canonical_label" in corp.obs else "cell_type"
y = corp.obs[label_key].astype(str).values
hvg2i = {g: i for i, g in enumerate(hvgs)}
common = [g for g in corp.var_names.astype(str) if g in hvg2i]
a = corp[:, common].copy()
sc.pp.normalize_total(a, target_sum=1e4); sc.pp.log1p(a)
X = a.X.toarray().astype(np.float32) if sp.issparse(a.X) else a.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
att = np.load(OUT / f"80_{sys}_gene_attribution_full.npy") # (K, G)
rows = []
for ci, cls_name in enumerate(classes):
mask = y == cls_name
if mask.sum() == 0:
print(f"[!!] {cls_name}: no cells in corpus", flush=True)
continue
x_gene_log_mean = Xf[mask].mean(axis=0)
x_z = np.clip((x_gene_log_mean - mu) / sig, -10, 10)
x_pca0 = pca.transform(x_z.reshape(1, -1))[0].astype(np.float32)
with torch.no_grad():
xt = torch.from_numpy(x_pca0).unsqueeze(0).to(DEVICE)
aux = torch.zeros(1, 2, device=DEVICE)
out = model(xt, aux, lam_dann=0.0)
z0 = out["z"]
s0 = float((z0 * protos[ci].unsqueeze(0)).sum())
att_c = att[ci]
cand_idx = np.argsort(-np.abs(att_c))[:TOP_KO_GENES]
# zero each gene in log space, reproject; z is clipped so recompute per KO
deltas = []
for gi in cand_idx:
x_gene_ko = x_gene_log_mean.copy()
x_gene_ko[gi] = 0.0
x_z_ko = np.clip((x_gene_ko - mu) / sig, -10, 10)
x_pca_ko = pca.transform(x_z_ko.reshape(1, -1))[0].astype(np.float32)
with torch.no_grad():
xtko = torch.from_numpy(x_pca_ko).unsqueeze(0).to(DEVICE)
zko = model(xtko, aux, lam_dann=0.0)["z"]
s_ko = float((zko * protos[ci].unsqueeze(0)).sum())
deltas.append(s0 - s_ko) # positive delta = drop when KO'd
deltas = np.array(deltas)
rank = np.argsort(-deltas)
top30 = rank[:30]
rows.append({
"class": cls_name,
"baseline_cos": s0,
"n_cells_class": int(mask.sum()),
"top_essential_genes": ",".join([hvgs[cand_idx[r]] for r in top30[:20]]),
"top_essential_deltas": ",".join([f"{deltas[r]:+.4f}" for r in top30[:20]]),
"top_essential_baseline_expression": ",".join([f"{x_gene_log_mean[cand_idx[r]]:.2f}" for r in top30[:20]]),
})
print(f"[{cls_name}] baseline_cos={s0:.4f} top-5 essentials: "
f"{', '.join([f'{hvgs[cand_idx[r]]}{deltas[r]:+.3f})' for r in top30[:5]])}",
flush=True)
df = pd.DataFrame(rows)
df.to_csv(OUT / f"81_{sys}_ko_essentials.csv", index=False)
print(f"[wrote] {OUT}/81_{sys}_ko_essentials.csv", flush=True)
for sys in ["pan_skin", "hematopoiesis", "pancreas"]:
try:
compute_ko(sys)
except Exception as e:
import traceback; traceback.print_exc()
print(f"[!] {sys} failed: {e}", flush=True)
print("\n=== DONE ===", flush=True)