| """zero-shot HSC PANDA on nestorowa 2016 smart-seq2 as unlabeled discovery target."""
|
| from __future__ import annotations
|
| from pathlib import Path
|
| import warnings, json, sys, pickle
|
| warnings.filterwarnings("ignore")
|
| import numpy as np, pandas as pd, anndata as ad, scanpy as sc, scipy.sparse as sp
|
| import torch
|
| import os as _os
|
| from pathlib import Path as _Path
|
| PANDA_ROOT = _Path(_os.environ.get("PANDA_ROOT", str(_Path(__file__).resolve().parents[2])))
|
| sys.path.insert(0, str(PANDA_ROOT))
|
| from panda.model import PANDAEncoder
|
|
|
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| CKPT = Path(str(PANDA_ROOT / "checkpoints/hematopoiesis"))
|
| HARM = Path(str(PANDA_ROOT / "data/corpus/hematopoiesis/harmonized"))
|
| OUT = Path(str(PANDA_ROOT / "discovery/hematopoiesis/marker"))
|
|
|
|
|
| def load_nestorowa():
|
| p = Path(str(PANDA_ROOT / "data/corpus/hematopoiesis/tier_a/nestorowa_GSE81682_counts.txt.gz"))
|
| df = pd.read_csv(p, sep="\t", index_col=0, compression="gzip")
|
|
|
| import mygene
|
| print(f"[nestorowa] converting {len(df)} ENSMUSG IDs to symbols via mygene …", flush=True)
|
| mg = mygene.MyGeneInfo()
|
| ids = df.index.astype(str).tolist()
|
| res = mg.querymany(ids, scopes="ensembl.gene", fields="symbol",
|
| species="mouse", returnall=False, verbose=False)
|
| id2sym = {}
|
| for r in res:
|
| if "symbol" in r and "query" in r and not r.get("notfound", False):
|
| id2sym[r["query"]] = r["symbol"]
|
| syms_series = pd.Series(df.index.astype(str)).map(id2sym).values
|
| syms = pd.Series(syms_series, index=df.index)
|
| keep = syms.notna()
|
| print(f"[nestorowa] mapped {int(keep.sum())}/{len(df)} genes", flush=True)
|
| df = df.loc[keep]
|
| df.index = syms[keep].values
|
| df = df.groupby(df.index).sum()
|
| X = sp.csr_matrix(df.values.T.astype(np.float32))
|
| obs = pd.DataFrame(index=df.columns.astype(str))
|
| obs["dataset"] = "nestorowa_GSE81682"
|
| obs["sample"] = obs.index.str.split("_").str[0]
|
| var = pd.DataFrame({"gene_symbol": df.index.astype(str).values},
|
| index=df.index.astype(str))
|
| a = ad.AnnData(X=X, obs=obs, var=var)
|
| a.var_names_make_unique()
|
| return a
|
|
|
|
|
| def project(a, shared_hvgs, mu, sig):
|
| G = len(shared_hvgs); n = a.n_obs
|
| hvg2i = {g: i for i, g in enumerate(shared_hvgs)}
|
| common = [g for g in a.var_names.astype(str) if g in hvg2i]
|
| frac = len(common) / G
|
| print(f"[proj] {len(common)}/{G} HVGs present ({frac:.1%})", flush=True)
|
| a_c = a[:, 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((n, G), dtype=np.float32)
|
| cols = [hvg2i[g] for g in common]
|
| Xf[:, cols] = X
|
| Xz = np.clip((Xf - mu.astype(np.float32)) / sig.astype(np.float32), -10, 10)
|
| return Xz, frac
|
|
|
|
|
| def main():
|
| ck = torch.load(CKPT / "panda_final.pt", map_location=DEVICE, weights_only=False)
|
| classes = ck["classes"]; datasets = ck["datasets"]
|
| model = PANDAEncoder(n_pca=50, n_classes=len(classes),
|
| n_datasets=len(datasets)).to(DEVICE).eval()
|
| model.load_state_dict(ck["model"])
|
| protos = ck["prototypes"]
|
| protos = protos / (np.linalg.norm(protos, axis=1, keepdims=True) + 1e-8)
|
| print(f"[model] classes: {classes}", flush=True)
|
|
|
| stats = np.load(HARM / "corpus_stats.npz", allow_pickle=True)
|
| shared_hvgs = [str(g) for g in stats["shared_hvgs"]]
|
| mu, sig = stats["mean"], stats["std"]
|
| with open(HARM / "pca_basis.pkl", "rb") as f: pca = pickle.load(f)
|
|
|
| a = load_nestorowa()
|
| print(f"[target] Nestorowa shape: {a.shape}", flush=True)
|
|
|
| Xz, frac = project(a, shared_hvgs, mu, sig)
|
| Xpca = pca.transform(Xz).astype(np.float32)
|
|
|
| all_z = []
|
| with torch.no_grad():
|
| for i in range(0, a.n_obs, 4096):
|
| xb = torch.from_numpy(Xpca[i:i+4096]).to(DEVICE)
|
| aux = torch.zeros(len(xb), 2, device=DEVICE)
|
| out = model(xb, aux, lam_dann=0.0)
|
| all_z.append(out["z"].cpu().numpy())
|
| Z = np.concatenate(all_z, axis=0)
|
| cos = Z @ protos.T
|
| pred_ix = cos.argmax(axis=1)
|
| conf = cos.max(axis=1)
|
| entropy = -(np.exp(cos / 0.07) / np.exp(cos / 0.07).sum(axis=1, keepdims=True) *
|
| np.log(np.exp(cos / 0.07) / np.exp(cos / 0.07).sum(axis=1, keepdims=True) + 1e-12)
|
| ).sum(axis=1)
|
|
|
| a.obs["pred_label"] = np.array([classes[i] for i in pred_ix], dtype=object)
|
| a.obs["pred_conf"] = conf.astype(np.float32)
|
| a.obs["pred_entropy"] = entropy.astype(np.float32)
|
|
|
| print(f"\n[nestorowa] predicted class distribution:")
|
| print(a.obs["pred_label"].value_counts())
|
| print(f"\n[nestorowa] pred_conf: p10={np.percentile(conf,10):.3f}, "
|
| f"p50={np.percentile(conf,50):.3f}, p90={np.percentile(conf,90):.3f}")
|
| print(f"[nestorowa] shared-HVG fraction: {frac:.1%}")
|
|
|
| a.obs.to_csv(OUT / "63_nestorowa_predictions.csv")
|
| print(f"\n[nestorowa] wrote predictions to 63_nestorowa_predictions.csv")
|
|
|
|
|
| thr = np.percentile(conf, 10)
|
| mask = conf <= thr
|
| print(f"\n[nestorowa] bottom-decile confidence: {int(mask.sum())} cells (thr={thr:.3f})",
|
| flush=True)
|
| if mask.sum() >= 30:
|
| sub_low = ad.AnnData(X=Z[mask].astype(np.float32))
|
| sc.pp.neighbors(sub_low, use_rep="X", n_neighbors=10)
|
| sc.tl.leiden(sub_low, resolution=0.5, key_added="cluster")
|
| print(f"[nestorowa] novel clusters: {sub_low.obs['cluster'].nunique()}")
|
| a_low = a[mask].copy()
|
| a_low.obs["cluster"] = sub_low.obs["cluster"].values
|
| try:
|
| sc.tl.rank_genes_groups(a_low, "cluster", method="wilcoxon", n_genes=10, use_raw=False)
|
| rows = []
|
| for cl in sorted(a_low.obs["cluster"].unique()):
|
| names = a_low.uns["rank_genes_groups"]["names"][cl]
|
| lfc = a_low.uns["rank_genes_groups"]["logfoldchanges"][cl]
|
| for g, l in zip(names[:8], lfc[:8]):
|
| rows.append({"cluster": cl, "gene": g, "logfc": round(float(l), 3)})
|
| pd.DataFrame(rows).to_csv(OUT / "63_nestorowa_novel_markers.csv", index=False)
|
| print(f"[nestorowa] novel markers saved")
|
| except Exception as e:
|
| print(f"[nestorowa] DE failed: {e}")
|
|
|
|
|
| if __name__ == "__main__":
|
| main()
|
|
|