| """PANDA lineage distribution across LARRY time points d2/d9/d16 + clonal purity + per-lineage DE."""
|
| from __future__ import annotations
|
| from pathlib import Path
|
| import warnings, json, sys
|
| warnings.filterwarnings("ignore")
|
| import numpy as np, pandas as pd, anndata as ad, scanpy as sc, torch
|
| from scipy import stats
|
|
|
| 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"))
|
| CORPUS = Path(str(PANDA_ROOT / "data/corpus/hematopoiesis/harmonized/corpus.h5ad"))
|
| OUT = Path(str(PANDA_ROOT / "discovery/hematopoiesis/marker"))
|
| OUT.mkdir(parents=True, exist_ok=True)
|
|
|
|
|
| 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)
|
|
|
|
|
| a = ad.read_h5ad(CORPUS)
|
| print(f"[hsc-mech] corpus: {a.shape}", flush=True)
|
| X = np.asarray(a.obsm["X_pca"]).astype(np.float32)
|
|
|
| all_z = []
|
| with torch.no_grad():
|
| for i in range(0, X.shape[0], 8192):
|
| xb = torch.from_numpy(X[i:i+8192]).to(DEVICE)
|
| aux = torch.zeros(len(xb), 2, device=DEVICE)
|
| all_z.append(model(xb, aux, lam_dann=0.0)["z"].cpu().numpy())
|
| Z = np.concatenate(all_z, axis=0)
|
| cos = Z @ protos.T
|
| pred_ix = cos.argmax(axis=1)
|
| pred = np.array([classes[i] for i in pred_ix], dtype=object)
|
| a.obs["pred_label"] = pred
|
| a.obs["pred_conf"] = cos.max(axis=1)
|
|
|
| if "Time point" not in a.obs.columns:
|
| print("[hsc-mech] no Time point column; skipping time analysis")
|
| else:
|
| tp = a.obs["Time point"].astype(int)
|
| xt = pd.crosstab(pred, tp, normalize="columns")
|
| print("\n[hsc-mech] Fraction per predicted class per time point:")
|
| print(xt.round(3))
|
| xt.to_csv(OUT / "62_time_course_class_fractions.csv")
|
|
|
| rows = []
|
| n_d2 = int((tp == 2).sum()); n_d16 = int((tp == 16).sum())
|
| for c in classes:
|
| n_c_d16 = int(((pred == c) & (tp == 16)).sum())
|
| n_c_d2 = int(((pred == c) & (tp == 2)).sum())
|
| contingency = np.array([[n_c_d16, n_d16 - n_c_d16], [n_c_d2, n_d2 - n_c_d2]])
|
| odds, p = stats.fisher_exact(contingency)
|
| f16 = (n_c_d16 + 1) / (n_d16 + 2); f2 = (n_c_d2 + 1) / (n_d2 + 2)
|
| rows.append({"class": c, "n_d16": n_c_d16, "n_d2": n_c_d2,
|
| "log2_fold_d16_vs_d2": round(np.log2(f16 / f2), 3),
|
| "fisher_p": p})
|
| df = pd.DataFrame(rows).sort_values("log2_fold_d16_vs_d2", ascending=False)
|
| print("\n[hsc-mech] class enrichment d16 vs d2 (positive = expanded at late time):")
|
| print(df.to_string(index=False))
|
| df.to_csv(OUT / "62_time_course_enrichment.csv", index=False)
|
|
|
|
|
| if "Library" in a.obs.columns:
|
| libs = a.obs["Library"].astype(str)
|
| top_lib = libs.value_counts().head(200).index
|
| clone_purity = []
|
| for L in top_lib:
|
| m = libs == L
|
| if m.sum() < 3: continue
|
| pl = pd.Series(pred[m.values]).value_counts(normalize=True)
|
| clone_purity.append({
|
| "library": L, "n": int(m.sum()),
|
| "dominant_class": pl.index[0],
|
| "purity": float(pl.iloc[0]),
|
| })
|
| cp = pd.DataFrame(clone_purity)
|
| print(f"\n[hsc-mech] clonal purity (dominant-class fraction) — {len(cp)} clones:")
|
| print(f" median: {cp['purity'].median():.3f}, mean: {cp['purity'].mean():.3f}, "
|
| f"n_clones_pure_>0.9: {(cp['purity'] > 0.9).sum()}/{len(cp)}")
|
| cp.to_csv(OUT / "62_clonal_purity.csv", index=False)
|
|
|
| a.obs["pred_label"] = pd.Categorical(pred)
|
| keep_classes = [c for c in classes if (pred == c).sum() >= 100]
|
| a_sub = a[np.isin(pred, keep_classes)].copy()
|
| if a_sub.n_obs >= 500:
|
| sc.tl.rank_genes_groups(a_sub, "pred_label", method="wilcoxon",
|
| n_genes=30, use_raw=False)
|
| rows = []
|
| for cls in keep_classes:
|
| try:
|
| names = a_sub.uns["rank_genes_groups"]["names"][cls]
|
| lfc = a_sub.uns["rank_genes_groups"]["logfoldchanges"][cls]
|
| for g, l in zip(names[:15], lfc[:15]):
|
| rows.append({"class": cls, "gene": g, "logfc": round(float(l), 3)})
|
| except Exception: pass
|
| pd.DataFrame(rows).to_csv(OUT / "62_lineage_markers.csv", index=False)
|
| print(f"\n[hsc-mech] wrote lineage markers to 62_lineage_markers.csv")
|
|
|
| print(f"\n[hsc-mech] complete. Outputs in {OUT}/")
|
|
|
|
|
| if __name__ == "__main__":
|
| main()
|
|
|