Download scripts/figures/generate_multi_umap.py from bryan7264/PANDA: direct link, hf CLI and curl.
- Browser
- Download file 9.29 kB
-
https://huggingface.co/bryan7264/PANDA/resolve/main/scripts/figures/generate_multi_umap.py
- Command line
-
hf download hf://bryan7264/PANDA/scripts/figures/generate_multi_umap.py
-
curl -L -o generate_multi_umap.py https://huggingface.co/bryan7264/PANDA/resolve/main/scripts/figures/generate_multi_umap.py
9.29 kB
| """3-panel UMAP: dingwall (En1 genotype), dahlin (Kit genotype), veres (stage). 8k cells/panel.""" | |
| from pathlib import Path | |
| import warnings, sys, pickle, json | |
| warnings.filterwarnings("ignore") | |
| import numpy as np, pandas as pd, anndata as ad, torch, scanpy as sc, scipy.sparse as sp | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import umap | |
| 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 | |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| FIG = Path(f"{ROOT_STR}/figures") | |
| FIG.mkdir(exist_ok=True) | |
| RS = 42 | |
| SAMPLE_N = 8000 | |
| def get_projection(ckpt_dir, data_a, shared_hvgs, mu, sig, pca): | |
| """128-d PANDA projections for the target AnnData.""" | |
| ck = torch.load(ckpt_dir / "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) | |
| G = len(shared_hvgs); hvg2i = {g: i for i, g in enumerate(shared_hvgs)} | |
| common = [g for g in data_a.var_names.astype(str) if g in hvg2i] | |
| a_c = data_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((data_a.n_obs, 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) | |
| Xpca = pca.transform(Xz).astype(np.float32) | |
| all_z = [] | |
| with torch.no_grad(): | |
| for i in range(0, data_a.n_obs, 4096): | |
| xb = torch.from_numpy(Xpca[i:i+4096]).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 = np.array([classes[i] for i in cos.argmax(axis=1)], dtype=object) | |
| return Z, pred | |
| def umap_it(Z, seed=RS): | |
| reducer = umap.UMAP(n_neighbors=30, min_dist=0.3, random_state=seed, | |
| metric="cosine", n_components=2) | |
| return reducer.fit_transform(Z) | |
| def load_sharon_stages(): | |
| from pathlib import Path as _P | |
| SHARON_DIR = _P(f"{ROOT_STR}/data/corpus/pancreas/held_out_unlabeled/sharon_extract") | |
| parts, stages = [], [] | |
| for meta_file in sorted(SHARON_DIR.glob("*.cell_metadata.tsv.gz")): | |
| counts_file = str(meta_file).replace("cell_metadata", "processed_counts") | |
| if not _P(counts_file).exists(): continue | |
| meta = pd.read_csv(meta_file, sep="\t", compression="gzip") | |
| counts = pd.read_csv(counts_file, sep="\t", compression="gzip", index_col=0) | |
| counts.columns = [c[0].upper() + c[1:].lower() if len(c) > 1 else c | |
| for c in counts.columns.astype(str)] | |
| counts = counts.T.groupby(level=0).sum().T | |
| obs = meta.set_index("library.barcode") | |
| obs = obs.loc[obs.index.intersection(counts.index)] | |
| counts_al = counts.loc[obs.index] | |
| X = sp.csr_matrix(counts_al.values.astype(np.float32)) | |
| obs["dataset"] = "sharon" | |
| var = pd.DataFrame({"gene_symbol": counts_al.columns}, index=counts_al.columns) | |
| a = ad.AnnData(X=X, obs=obs, var=var); a.var_names_make_unique() | |
| parts.append(a) | |
| return ad.concat(parts, join="outer", label="_batch") | |
| def load_dahlin(): | |
| from pathlib import Path as _P | |
| D_DIR = _P(f"{ROOT_STR}/data/corpus/hematopoiesis/held_out_unlabeled/dahlin_extract") | |
| GT = {"SIGAB1":"WT","SIGAC1":"WT","SIGAD1":"WT","SIGAF1":"WT","SIGAG1":"WT", | |
| "SIGAH1":"WT","SIGAG8":"Kit_W41","SIGAH8":"Kit_W41"} | |
| parts = [] | |
| for f in sorted(D_DIR.glob("*.txt.gz")): | |
| sample = f.name.split("_")[1].split(".")[0] | |
| df = pd.read_csv(f, sep="\t", compression="gzip", index_col=0) | |
| X = sp.csr_matrix(df.values.T.astype(np.float32)) | |
| obs = pd.DataFrame(index=[f"{sample}_{bc}" for bc in df.columns.astype(str)]) | |
| obs["sample"] = sample; obs["genotype"] = GT.get(sample, "unknown") | |
| var = pd.DataFrame(index=df.index.astype(str)) | |
| parts.append(ad.AnnData(X=X, obs=obs, var=var)) | |
| a = ad.concat(parts, join="outer", label="_batch") | |
| import mygene | |
| mg = mygene.MyGeneInfo() | |
| ids = a.var_names.astype(str).tolist() | |
| res = mg.querymany(ids, scopes="ensembl.gene", fields="symbol", species="mouse", | |
| verbose=False) | |
| id2sym = {r["query"]: r["symbol"] for r in res if "symbol" in r} | |
| syms = pd.Series(a.var_names.astype(str)).map(id2sym).values | |
| keep = pd.notna(syms) | |
| a = a[:, keep].copy(); a.var_names = syms[keep]; a.var_names_make_unique() | |
| return a | |
| def main(): | |
| fig, axes = plt.subplots(1, 3, figsize=(17, 5.3)) | |
| # ---- Panel A: Dingwall (skin) ---- | |
| print("[fig] Dingwall UMAP …", flush=True) | |
| p = ad.read_h5ad(f"{ROOT_STR}/discovery/pan_skin/marker/50_aldrich_projections.h5ad") | |
| Z = np.asarray(p.obsm["Z_projection"]) | |
| rng = np.random.default_rng(RS) | |
| idx = rng.choice(len(Z), size=min(SAMPLE_N, len(Z)), replace=False) | |
| Z_a = Z[idx] | |
| emb = umap_it(Z_a) | |
| genotype = p.obs["genotype"].values[idx] | |
| ax = axes[0] | |
| for g, c in zip(["WT", "En1-cKO"], ["#2b83ba", "#d7191c"]): | |
| m = genotype == g | |
| ax.scatter(emb[m, 0], emb[m, 1], s=2, alpha=0.5, c=c, | |
| label=f"{g} (n={int(m.sum())})") | |
| ax.set_title(f"(a) Dingwall skin (n={SAMPLE_N}) — En1 genotype", fontsize=10) | |
| ax.set_xlabel("UMAP 1"); ax.set_ylabel("UMAP 2") | |
| ax.legend(markerscale=6, frameon=False, fontsize=9) | |
| # ---- Panel B: Dahlin (HSC) ---- | |
| print("[fig] Dahlin UMAP …", flush=True) | |
| cache_d = Path(f"{ROOT_STR}/figures/_cache_dahlin_umap.npz") | |
| if cache_d.exists(): | |
| c = np.load(cache_d, allow_pickle=True) | |
| emb2 = c["emb"]; gt_d = c["gt"] | |
| else: | |
| stats_h = np.load(f"{ROOT_STR}/data/corpus/hematopoiesis/harmonized/corpus_stats.npz", | |
| allow_pickle=True) | |
| shared_hvgs_h = [str(g) for g in stats_h["shared_hvgs"]] | |
| pca_h = pickle.load(open(f"{ROOT_STR}/data/corpus/hematopoiesis/harmonized/pca_basis.pkl","rb")) | |
| a_d = load_dahlin() | |
| rng2 = np.random.default_rng(RS) | |
| idx2 = rng2.choice(a_d.n_obs, size=min(SAMPLE_N, a_d.n_obs), replace=False) | |
| a_d_sub = a_d[idx2].copy() | |
| Z_d, pred_d = get_projection(Path(f"{ROOT_STR}/checkpoints/hematopoiesis"), | |
| a_d_sub, shared_hvgs_h, stats_h["mean"], stats_h["std"], pca_h) | |
| emb2 = umap_it(Z_d) | |
| gt_d = a_d_sub.obs["genotype"].values | |
| np.savez(cache_d, emb=emb2, gt=np.asarray(gt_d, dtype=object)) | |
| ax = axes[1] | |
| for g, c in zip(["WT", "Kit_W41"], ["#2b83ba", "#d7191c"]): | |
| m = gt_d == g | |
| ax.scatter(emb2[m, 0], emb2[m, 1], s=2, alpha=0.5, c=c, | |
| label=f"{g} (n={int(m.sum())})") | |
| ax.set_title(f"(b) Dahlin HSC (n={SAMPLE_N}) — Kit genotype", fontsize=10) | |
| ax.set_xlabel("UMAP 1"); ax.set_ylabel("UMAP 2") | |
| ax.legend(markerscale=6, frameon=False, fontsize=9) | |
| # ---- Panel C: Veres (pancreas) ---- | |
| print("[fig] Veres UMAP …", flush=True) | |
| stats_p = np.load(f"{ROOT_STR}/data/corpus/pancreas/harmonized/corpus_stats.npz", | |
| allow_pickle=True) | |
| shared_hvgs_p = [str(g) for g in stats_p["shared_hvgs"]] | |
| pca_p = pickle.load(open(f"{ROOT_STR}/data/corpus/pancreas/harmonized/pca_basis.pkl","rb")) | |
| a_s = load_sharon_stages() | |
| stage_num = pd.to_numeric(a_s.obs["Stage"], errors="coerce") | |
| keep_st = stage_num.notna().values | |
| a_s = a_s[keep_st].copy() | |
| a_s.obs["Stage_int"] = stage_num[keep_st].astype(int).values | |
| rng3 = np.random.default_rng(RS) | |
| idx3 = rng3.choice(a_s.n_obs, size=min(SAMPLE_N, a_s.n_obs), replace=False) | |
| a_s_sub = a_s[idx3].copy() | |
| Z_s, pred_s = get_projection(Path(f"{ROOT_STR}/checkpoints/pancreas"), | |
| a_s_sub, shared_hvgs_p, stats_p["mean"], stats_p["std"], pca_p) | |
| emb3 = umap_it(Z_s) | |
| stage = a_s_sub.obs["Stage_int"].values | |
| ax = axes[2] | |
| stage_colors = {3: "#fdae61", 4: "#f8b0d1", 5: "#7570b3", 6: "#d7191c"} | |
| for s in sorted(np.unique(stage)): | |
| m = stage == s | |
| ax.scatter(emb3[m, 0], emb3[m, 1], s=2, alpha=0.5, | |
| c=stage_colors.get(s, "#666"), label=f"Stage {s} (n={int(m.sum())})") | |
| ax.set_title(f"(c) Veres hPSC (n={SAMPLE_N}) — differentiation stage", fontsize=10) | |
| ax.set_xlabel("UMAP 1"); ax.set_ylabel("UMAP 2") | |
| ax.legend(markerscale=6, frameon=False, fontsize=9) | |
| plt.tight_layout() | |
| plt.savefig(FIG / "fig6_multi_umap.pdf", bbox_inches="tight", dpi=100) | |
| plt.close() | |
| print(f"[fig] wrote {FIG}/fig6_multi_umap.pdf") | |
| if __name__ == "__main__": | |
| main() | |