File size: 6,707 Bytes
4e2940e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
#!/usr/bin/env python
"""Empirical identifiability: negative controls for disentanglement.

Shows the z_T/z_PT disentanglement is real by comparing:
- Real data: z_T captures cell type, z_PT captures different structure
- Permuted data: both latents are random noise

If disentanglement disappears on permuted data, it's a real signal.
"""
from _common import *
from sklearn.metrics import silhouette_score, adjusted_rand_score
from sklearn.cluster import KMeans

OUT = output_dir("21_identifiability")


def run_permutation_test(adata_base, cluster_key, n_perm=3):
    """Run DeepPTR on real vs permuted data."""
    from scipy.sparse import issparse

    # ── Real data ────────────────────────────────────────────────────
    print("  Real data...")
    adata_real = adata_base.copy()
    for key in ("spliced", "unspliced"):
        if key in adata_real.layers and issparse(adata_real.layers[key]):
            adata_real.layers[key] = np.asarray(adata_real.layers[key].todense())

    torch.set_num_threads(4)
    scptr.deep.fit_deepptr(adata_real, verbose=False, **DEEP_HP)

    labels = adata_real.obs[cluster_key].astype("category").cat.codes.values
    n_sample = min(2000, len(labels))

    sil_T_real = silhouette_score(adata_real.obsm["X_z_T"], labels, sample_size=n_sample)
    sil_PT_real = silhouette_score(adata_real.obsm["X_z_PT"], labels, sample_size=n_sample)

    # PT cluster vs expression cluster ARI
    km = KMeans(n_clusters=5, random_state=0, n_init=10)
    pt_labels = km.fit_predict(adata_real.obsm["X_z_PT"])
    ari_real = adjusted_rand_score(labels, pt_labels)

    # Count PT-specific genes
    gamma = adata_real.layers["gamma"]
    z_T = adata_real.obsm["X_z_T"]
    z_PT = adata_real.obsm["X_z_PT"]
    n_pt_genes = 0
    for g in range(adata_real.n_vars):
        gv = gamma[:, g]
        if gv.std() < 1e-8:
            continue
        r_T = max(abs(stats.spearmanr(gv, z_T[:, d]).statistic) for d in range(z_T.shape[1]))
        r_PT = max(abs(stats.spearmanr(gv, z_PT[:, d]).statistic) for d in range(z_PT.shape[1]))
        if r_PT > 0.3 and r_PT > r_T * 1.5:
            n_pt_genes += 1

    print(f"    sil_T={sil_T_real:.4f}, sil_PT={sil_PT_real:.4f}, ARI={ari_real:.4f}, PT_genes={n_pt_genes}")

    # ── Permuted data ────────────────────────────────────────────────
    perm_results = []
    for p in range(n_perm):
        print(f"  Permutation {p+1}/{n_perm}...")
        adata_perm = adata_base.copy()
        for key in ("spliced", "unspliced"):
            if key in adata_perm.layers and issparse(adata_perm.layers[key]):
                adata_perm.layers[key] = np.asarray(adata_perm.layers[key].todense())

        # Permute cells independently per gene (destroy cell-gene structure)
        rng = np.random.RandomState(p)
        s_perm = adata_perm.layers["spliced"].copy()
        u_perm = adata_perm.layers["unspliced"].copy()
        for g in range(s_perm.shape[1]):
            s_perm[:, g] = rng.permutation(s_perm[:, g])
            u_perm[:, g] = rng.permutation(u_perm[:, g])
        adata_perm.layers["spliced"] = s_perm
        adata_perm.layers["unspliced"] = u_perm

        torch.set_num_threads(4)
        try:
            scptr.deep.fit_deepptr(adata_perm, verbose=False, **{**DEEP_HP, "seed": p})
        except Exception as e:
            print(f"    Failed: {e}")
            continue

        sil_T_perm = silhouette_score(adata_perm.obsm["X_z_T"], labels, sample_size=n_sample)
        sil_PT_perm = silhouette_score(adata_perm.obsm["X_z_PT"], labels, sample_size=n_sample)

        pt_labels_perm = km.fit_predict(adata_perm.obsm["X_z_PT"])
        ari_perm = adjusted_rand_score(labels, pt_labels_perm)

        perm_results.append({
            "sil_T": float(sil_T_perm),
            "sil_PT": float(sil_PT_perm),
            "ari": float(ari_perm),
        })
        print(f"    sil_T={sil_T_perm:.4f}, sil_PT={sil_PT_perm:.4f}, ARI={ari_perm:.4f}")

    return {
        "real": {
            "sil_T": float(sil_T_real),
            "sil_PT": float(sil_PT_real),
            "ari": float(ari_real),
            "n_pt_genes": n_pt_genes,
        },
        "permuted": perm_results,
    }


def main():
    set_figure_style()
    all_results = {}

    for ds_name, loader, ck in DATASETS:
        print(f"\n{'=' * 60}\n{ds_name.upper()}: Identifiability test\n{'=' * 60}")

        adata_raw = loader()
        scptr.pp.filter_genes(adata_raw)
        scptr.pp.normalize_layers(adata_raw)
        scptr.pp.neighbors(adata_raw, n_neighbors=30)
        scptr.pp.smooth_layers(adata_raw)
        scptr.tl.estimate_beta(adata_raw)
        adata_base = select_top_genes(adata_raw, n_top=300)

        results = run_permutation_test(adata_base, ck, n_perm=3)
        all_results[ds_name] = results

        # Summary
        mean_perm_sil_T = np.mean([p["sil_T"] for p in results["permuted"]])
        mean_perm_ari = np.mean([p["ari"] for p in results["permuted"]])

        print(f"\n  SUMMARY:")
        print(f"    z_T silhouette: real={results['real']['sil_T']:.4f}, permuted={mean_perm_sil_T:.4f}")
        print(f"    PT-expr ARI:    real={results['real']['ari']:.4f}, permuted={mean_perm_ari:.4f}")
        print(f"    PT-specific genes: real={results['real']['n_pt_genes']}")

    save_json(all_results, "identifiability", OUT)

    # Figure
    fig, axes = plt.subplots(1, 2, figsize=(12, 5))
    for ax_idx, (ds_name, res) in enumerate(all_results.items()):
        if ax_idx >= 2:
            break
        ax = axes[ax_idx]

        metrics = ["sil_T", "sil_PT", "ari"]
        labels = ["Silhouette z_T", "Silhouette z_PT", "ARI (PT vs expr)"]
        real_vals = [res["real"][m] for m in metrics]
        perm_vals = [np.mean([p[m] for p in res["permuted"]]) for m in metrics]
        perm_stds = [np.std([p[m] for p in res["permuted"]]) for m in metrics]

        x = np.arange(len(metrics))
        ax.bar(x - 0.2, real_vals, 0.35, label="Real", color="darkorange", alpha=0.8)
        ax.bar(x + 0.2, perm_vals, 0.35, yerr=perm_stds, label="Permuted",
               color="gray", alpha=0.6, capsize=4)
        ax.set_xticks(x)
        ax.set_xticklabels(labels, fontsize=8)
        ax.set_ylabel("Score")
        ax.set_title(f"{ds_name}: Real vs permuted")
        ax.legend()
        ax.axhline(0, color="k", lw=0.5)

    fig.suptitle("Empirical identifiability: disentanglement is real", y=1.02)
    fig.tight_layout()
    save_fig(fig, "identifiability", OUT)


if __name__ == "__main__":
    main()