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