#!/usr/bin/env python
"""轨迹/拟时序样例(公开数据:Paul et al. 2015 小鼠骨髓造血)。

对应价目表里「高级模块」那一档($900–3,200/项目)。之前那档只有价格没有实物,
客户凭什么信我们会做?所以真跑一遍:PAGA 画出细胞群之间的拓扑关系,
再用扩散拟时序给细胞排序,最后看关键基因是否沿着已知的分化方向变化。

自检:造血是教科书级已知的 —— 从祖细胞出发应分出红系(Hba-a2/Klf1)与髓系(Elane/Mpo/Prtn3)两支。
如果算出来不是这个拓扑,就是流程有问题,而不是发现了新生物学。
"""
import json, os, warnings
warnings.filterwarnings("ignore")
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import scanpy as sc

OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "out_traj")
os.makedirs(OUT, exist_ok=True)
sc.settings.figdir = OUT
sc.settings.set_figure_params(dpi=140, frameon=False, figsize=(5.2, 4.2), facecolor="white")
M = {"dataset": "Paul et al. 2015 — mouse bone marrow haematopoiesis (public, via scanpy)"}

adata = sc.datasets.paul15()
adata.X = adata.X.astype("float64")
M["cells"], M["genes"] = int(adata.n_obs), int(adata.n_vars)
M["annotated_groups"] = int(adata.obs["paul15_clusters"].nunique())

sc.pp.recipe_zheng17(adata)
sc.tl.pca(adata, svd_solver="arpack", n_comps=50)
sc.pp.neighbors(adata, n_neighbors=15, n_pcs=30)
sc.tl.draw_graph(adata)
sc.tl.leiden(adata, resolution=1.0, key_added="leiden", flavor="igraph", n_iterations=2, directed=False)
M["leiden_clusters"] = int(adata.obs["leiden"].nunique())

# PAGA:先算群与群之间的连接强度,再据此画出可信的拓扑(而不是硬把 UMAP 上的形状当轨迹)
sc.tl.paga(adata, groups="leiden")
sc.pl.paga(adata, color="leiden", save="_topology.png", show=False)
conn = adata.uns["paga"]["connectivities"].toarray()
M["paga_strong_edges"] = int((np.triu(conn) > 0.1).sum())

# 起点:取 Paul 的原始注释里最像祖细胞的那群,而不是随手指一个细胞
prog = [c for c in adata.obs["paul15_clusters"].cat.categories if "MEP" in c or "GMP" in c or "Ery" not in c]
root_group = None
for cand in ["7MEP", "8Mk", "1Ery"]:
    if cand in list(adata.obs["paul15_clusters"].cat.categories):
        root_group = cand; break
if root_group is None:
    root_group = str(adata.obs["paul15_clusters"].cat.categories[0])
mask = (adata.obs["paul15_clusters"] == root_group).values
adata.uns["iroot"] = int(np.flatnonzero(mask)[0])
M["root_group"] = root_group

sc.tl.diffmap(adata)
sc.tl.dpt(adata)
M["pseudotime_range"] = [round(float(adata.obs["dpt_pseudotime"].min()), 3),
                         round(float(adata.obs["dpt_pseudotime"].max()), 3)]

sc.tl.draw_graph(adata, init_pos="paga")
sc.pl.draw_graph(adata, color=["paul15_clusters"], legend_loc="on data",
                 title="Published annotation", save="_groups.png", show=False)
sc.pl.draw_graph(adata, color=["dpt_pseudotime"], title="Diffusion pseudotime",
                 save="_pseudotime.png", show=False)

# 自检:两条已知分支的标志基因,沿拟时序应当分开
ERY = [g for g in ["Hba-a2", "Klf1", "Gata1"] if g in adata.var_names]
MYE = [g for g in ["Elane", "Mpo", "Prtn3", "Cebpe"] if g in adata.var_names]
M["marker_genes_checked"] = {"erythroid": ERY, "myeloid": MYE}
if ERY and MYE:
    sc.pl.draw_graph(adata, color=ERY + MYE, ncols=4, save="_markers.png", show=False)
    X = adata[:, ERY + MYE].X
    X = np.asarray(X.todense()) if hasattr(X, "todense") else np.asarray(X)
    pt = adata.obs["dpt_pseudotime"].values
    late = pt > np.quantile(pt, 0.75)
    ery_late = float(X[late][:, :len(ERY)].mean())
    mye_late = float(X[late][:, len(ERY):].mean())
    M["late_pseudotime_expression"] = {"erythroid_mean": round(ery_late, 3), "myeloid_mean": round(mye_late, 3)}
    # 相关性:红系与髓系标志基因在细胞层面应当互斥(负相关),这是分支存在的直接证据
    e = X[:, :len(ERY)].mean(axis=1); m = X[:, len(ERY):].mean(axis=1)
    M["ery_vs_mye_correlation"] = round(float(np.corrcoef(e, m)[0, 1]), 3)
    M["branch_check_passed"] = bool(M["ery_vs_mye_correlation"] < 0)

with open(f"{OUT}/metrics.json", "w") as f:
    json.dump(M, f, indent=2, ensure_ascii=False)
print(json.dumps(M, indent=2, ensure_ascii=False))
