"""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) # predict on all corpus cells (this is training-set prediction — for analysis only) 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) # sibling-fate concordance: Library = clonal barcode if "Library" in a.obs.columns: libs = a.obs["Library"].astype(str) top_lib = libs.value_counts().head(200).index # top 200 largest clones 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()