File size: 19,051 Bytes
2e3eac1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c08a743
 
 
2e3eac1
 
 
c08a743
2e3eac1
c08a743
 
 
 
 
 
 
2e3eac1
 
c08a743
 
2e3eac1
 
 
 
 
 
 
 
 
 
 
 
c08a743
2e3eac1
 
c08a743
2e3eac1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f609d01
 
2e3eac1
f609d01
 
2e3eac1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c08a743
2e3eac1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c08a743
2e3eac1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c08a743
2e3eac1
 
 
 
 
 
 
 
 
 
 
 
 
 
c08a743
2e3eac1
 
 
c08a743
2e3eac1
 
 
 
 
c08a743
2e3eac1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c08a743
 
 
 
 
 
 
 
 
 
 
 
2e3eac1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c08a743
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
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
"""PhenoSeq β€” Cell Painting morphology β†’ single-cell transcriptomic embeddings.

Gradio Space for Sentinal4D/PhenoSeq: a conditional Gaussian-diffusion model that
generates 512-d scGPT RNA-seq embeddings from ViT-L Cell Painting imaging features.
"""

import os

os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")

import spaces  # noqa: E402  β€” must precede torch / CUDA-touching imports

import json  # noqa: E402
import tempfile  # noqa: E402
import time  # noqa: E402

import matplotlib  # noqa: E402

matplotlib.use("Agg")

import gradio as gr  # noqa: E402
import matplotlib.pyplot as plt  # noqa: E402
import numpy as np  # noqa: E402
import torch  # noqa: E402

from pipeline import PhenoSeqPipeline  # noqa: E402

MODEL_ID = "Sentinal4D/PhenoSeq"
ASSETS = os.path.join(os.path.dirname(os.path.abspath(__file__)), "assets")

# ──────────────────────────────────────────────────────────────────────────────
# Model β€” loaded once at module scope, moved to CUDA eagerly (ZeroGPU packs it)
# ──────────────────────────────────────────────────────────────────────────────

pipe = PhenoSeqPipeline.from_pretrained(MODEL_ID, device="cuda", ddim_steps=50)
print(f"Loaded {pipe}", flush=True)

# ──────────────────────────────────────────────────────────────────────────────
# Reference bundle (derived from altoslabs/scGeneScope, held-out val wells)
# ──────────────────────────────────────────────────────────────────────────────

SAMPLES: dict = json.load(open(os.path.join(ASSETS, "samples.json")))
_IMG = np.load(os.path.join(ASSETS, "imaging.npz"))
_RNA = np.load(os.path.join(ASSETS, "rna.npz"))

IMAGING = {k: _IMG[k].astype(np.float32) for k in _IMG.files}
REAL = {k[6:]: _RNA[k].astype(np.float32) for k in _RNA.files if k.startswith("real__")}
CENTROID_NAMES = sorted(k[6:] for k in _RNA.files if k.startswith("cent__"))
CENTROIDS = np.stack([_RNA[f"cent__{n}"] for n in CENTROID_NAMES]).astype(np.float32)
PCA_MEAN = _RNA["pca_mean"].astype(np.float32)
PCA_BASIS = _RNA["pca_basis"].astype(np.float32)  # (3, 512)
BG = _RNA["bg"].astype(np.float32)  # background cloud of real scGPT cells

# scGPT embeddings share a large common mean component, so every similarity below
# is computed on mean-centred vectors (centre = mean of the real training cells).
CENTRE = PCA_MEAN
CENTROIDS_C = CENTROIDS - CENTRE
CENTROIDS_C = CENTROIDS_C / (np.linalg.norm(CENTROIDS_C, axis=1, keepdims=True) + 1e-9)

SORTED_IDS = sorted(SAMPLES, key=lambda s: (SAMPLES[s]["treatment"].lower(), s))
CHOICES = [
    (
        f"{SAMPLES[s]['treatment']} Β· {SAMPLES[s]['dosage']} β€” "
        f"plate {SAMPLES[s]['plate']} well {SAMPLES[s]['well']} ({s})",
        s,
    )
    for s in SORTED_IDS
]
DEFAULT_ID = "202502PA_86"  # DMSO vehicle control


def _montage(sample_id: str) -> str:
    return os.path.join(ASSETS, "cells", f"{sample_id}.png")


def _well_header(sid: str) -> str:
    info = SAMPLES[sid]
    return (
        f"### {info['treatment']} β€” {info['dosage']} ({info['group']})\n"
        f"scGeneScope well `{sid}` Β· plate {info['plate']} well {info['well']} Β· "
        f"replicate {info['replicate']} Β· {info['n_cells_well']:,} imaged cells in the well Β· "
        f"**held out of PhenoSeq training**"
    )


# ──────────────────────────────────────────────────────────────────────────────
# Plot / metric helpers (figures are rendered to PNG so nothing unpicklable
# crosses the ZeroGPU fork boundary)
# ──────────────────────────────────────────────────────────────────────────────


def _png(fig) -> str:
    path = os.path.join(tempfile.mkdtemp(), "plot.png")
    fig.savefig(path, dpi=110, bbox_inches="tight", facecolor="white")
    plt.close(fig)
    return path


def _project(x: np.ndarray) -> np.ndarray:
    return (x - PCA_MEAN) @ PCA_BASIS[:2].T


def _scatter(gen: np.ndarray, real, title: str) -> str:
    fig, ax = plt.subplots(figsize=(5.6, 5.0))
    b = _project(BG)
    ax.scatter(b[:, 0], b[:, 1], s=3, c="#d8d8de", alpha=0.6, linewidths=0,
               label="real scGPT cells (all wells)")
    if real is not None and len(real):
        r = _project(real)
        ax.scatter(r[:, 0], r[:, 1], s=12, c="#2563eb", alpha=0.6, linewidths=0,
                   label="real cells β€” this well")
    g = _project(gen)
    ax.scatter(g[:, 0], g[:, 1], s=14, c="#ea580c", alpha=0.8, linewidths=0,
               label="PhenoSeq generated")
    ax.set_xlabel("scGPT PC1")
    ax.set_ylabel("scGPT PC2")
    ax.set_title(title, fontsize=10)
    ax.legend(loc="best", fontsize=8)
    ax.spines[["top", "right"]].set_visible(False)
    return _png(fig)


def _profile(gen: np.ndarray, real) -> str:
    fig, ax = plt.subplots(figsize=(9.0, 2.8))
    gm = gen.mean(0) - CENTRE
    order = np.argsort(-np.abs(gm))[:64]
    idx = np.arange(len(order))
    if real is not None and len(real):
        ax.bar(idx - 0.21, (real.mean(0) - CENTRE)[order], width=0.42,
               color="#2563eb", label="real (well mean)")
    ax.bar(idx + 0.21, gm[order], width=0.42, color="#ea580c",
           label="PhenoSeq (generated mean)")
    ax.axhline(0, color="#999", lw=0.6)
    ax.set_xlabel("64 largest-magnitude scGPT dimensions (mean-centred)")
    ax.set_ylabel("value")
    ax.set_xticks([])
    ax.legend(fontsize=8, frameon=False)
    ax.spines[["top", "right"]].set_visible(False)
    return _png(fig)


def _cos(a: np.ndarray, b: np.ndarray) -> float:
    return float(a @ b / ((np.linalg.norm(a) * np.linalg.norm(b)) + 1e-9))


def _report(gen, real, truth, sid, elapsed, n_cells, ddim_steps) -> str:
    gmean = gen.mean(0)
    gm_c = gmean - CENTRE
    rows = ["| metric | value |", "|---|---|"]
    if real is not None and len(real):
        rm_c = real.mean(0) - CENTRE
        others = [_cos(gm_c, REAL[o].mean(0) - CENTRE) for o in REAL if o != sid]
        rows += [
            f"| **cosine to THIS well's real transcriptome** (mean-centred) | "
            f"**{_cos(gm_c, rm_c):+.3f}** |",
            f"| cosine to the other {len(others)} held-out wells (median Β· max) | "
            f"{float(np.median(others)):+.3f} Β· {float(np.max(others)):+.3f} |",
            f"| MSE vs. real well mean (raw scGPT scale) | "
            f"{float(((gmean - real.mean(0)) ** 2).mean()):.2e} |",
            f"| real held-out cells in the comparison | {len(real)} |",
        ]
    rows += [
        f"| generated cells Β· DDIM steps | {n_cells} Β· "
        f"{ddim_steps if ddim_steps else '1000 (full DDPM)'} |",
        f"| GPU sampling time | {elapsed:.2f} s |",
    ]

    sim = CENTROIDS_C @ (gm_c / (np.linalg.norm(gm_c) + 1e-9))
    order = np.argsort(-sim)
    out = ["\n".join(rows), "",
           "**Perturbation retrieval** β€” nearest real-transcriptome treatment centroid "
           f"({len(CENTROID_NAMES)} candidates, centroids built from the *training* wells only):",
           "", "| rank | treatment | cosine |", "|---|---|---|"]
    for k, i in enumerate(order[:5], 1):
        mark = " βœ…" if truth and CENTROID_NAMES[i] == truth else ""
        out.append(f"| {k} | {CENTROID_NAMES[i]}{mark} | {float(sim[i]):+.3f} |")
    if truth and truth in CENTROID_NAMES:
        rank = int(np.where(np.array(CENTROID_NAMES)[order] == truth)[0][0]) + 1
        out += ["", f"Ground truth for this well is **{truth}** β†’ retrieved at "
                    f"**rank {rank} / {len(CENTROID_NAMES)}** (chance β‰ˆ "
                    f"{len(CENTROID_NAMES) // 2})."]
    return "\n".join(out)


def _load_uploaded(path: str) -> np.ndarray:
    if str(path).endswith(".npz"):
        z = np.load(path)
        arr = z["X"] if "X" in z.files else z[z.files[0]]
    else:
        arr = np.load(path)
    arr = np.asarray(arr, dtype=np.float32)
    if arr.ndim == 3:
        arr = arr[0]
    if arr.ndim != 2 or arr.shape[-1] != 5120:
        raise gr.Error(
            f"Expected ViT-L imaging features of shape (N, 5120), got {tuple(arr.shape)}. "
            "Features are 5 fluorescence channels Γ— 1024 ViT-L/14 dims."
        )
    if arr.shape[0] > 16:
        arr = arr[np.linspace(0, arr.shape[0] - 1, 16, dtype=int)]
    return np.ascontiguousarray(arr)


# ──────────────────────────────────────────────────────────────────────────────
# Inference
# ──────────────────────────────────────────────────────────────────────────────


def _duration(sample_id=None, features_file=None, n_cells=256, ddim_steps=50, seed=0,
              *args, **kwargs) -> int:
    # Measured on ZeroGPU: 1.5 s for 256 cells x 50 steps, 16.3 s for 1024 x 200,
    # plus ~5 s of weight streaming / plotting / file IO inside the GPU context.
    steps = 1000 if not int(ddim_steps or 0) else int(ddim_steps)
    n = max(16, min(1024, int(n_cells or 256)))
    return int(min(180, 8 + 9.5e-5 * n * steps))


@spaces.GPU(duration=_duration)
def generate(
    sample_id: str = DEFAULT_ID,
    features_file: str | None = None,
    n_cells: int = 256,
    ddim_steps: int = 50,
    seed: int = 0,
):
    """Generate synthetic single-cell scGPT transcriptomes from Cell Painting morphology.

    Args:
        sample_id: held-out scGeneScope well id whose Cell Painting features condition the model.
        features_file: optional .npy/.npz with your own ViT-L features of shape (N, 5120); overrides sample_id.
        n_cells: how many single-cell transcriptomes to synthesise for the well.
        ddim_steps: DDIM denoising steps; 0 runs the full 1000-step DDPM sampler.
        seed: RNG seed for the diffusion noise.

    Returns:
        A Cell Painting montage of the conditioning cells, a PCA scatter of generated vs.
        real scGPT cells, a per-dimension profile plot, a well description, a metrics
        report, and the generated embeddings as an .npy file of shape (n_cells, 512).
    """
    if features_file:
        feats = _load_uploaded(features_file)
        sid, montage, real, truth = None, None, None, None
        title = "uploaded ViT-L features"
        header = f"### Custom imaging features\n{feats.shape[0]} imaging cells as conditioning context."
    else:
        sid = sample_id if sample_id in IMAGING else DEFAULT_ID
        feats = IMAGING[sid]
        montage = _montage(sid)
        real = REAL.get(sid)
        truth = SAMPLES[sid]["treatment"]
        title = f"{truth} Β· {SAMPLES[sid]['dosage']}"
        header = _well_header(sid)

    n_cells = max(16, min(1024, int(n_cells)))
    ddim_steps = int(ddim_steps)

    torch.manual_seed(int(seed))
    ctx = np.repeat(feats[None], n_cells, axis=0)  # (n_cells, 16, 5120)
    t0 = time.perf_counter()
    gen = np.asarray(pipe(ctx, batch_size=n_cells, ddim_steps=ddim_steps), dtype=np.float32)
    elapsed = time.perf_counter() - t0
    print(f"sampled {gen.shape} in {elapsed:.2f}s (steps={ddim_steps})", flush=True)

    npy = os.path.join(tempfile.mkdtemp(), "phenoseq_scgpt.npy")
    np.save(npy, gen)

    return (
        montage,
        _scatter(gen, real, f"scGPT embedding space β€” {title}"),
        _profile(gen, real),
        header,
        _report(gen, real, truth, sid, elapsed, n_cells, ddim_steps),
        npy,
    )


def preview(sample_id: str):
    """Show the Cell Painting montage and metadata for a well without running the model."""
    sid = sample_id if sample_id in IMAGING else DEFAULT_ID
    return _montage(sid), _well_header(sid)


# ──────────────────────────────────────────────────────────────────────────────
# UI
# ──────────────────────────────────────────────────────────────────────────────

CSS = """
#col-container { max-width: 1200px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""

INTRO = """# 🧫 PhenoSeq β€” Cell Painting β†’ single-cell transcriptomes

[**PhenoSeq**](https://huggingface.co/Sentinal4D/PhenoSeq) is a conditional Gaussian-diffusion
model that synthesises **512-d scGPT RNA-seq embeddings** from **ViT-L Cell Painting imaging
features** β€” a transcriptome inferred from morphology alone, at single-cell resolution.

Pick a held-out [scGeneScope](https://huggingface.co/datasets/altoslabs/scGeneScope) well (or
upload your own ViT-L features): the model generates single-cell transcriptomes for it, and the
Space scores them against the **real matched RNA-seq** of that well.

[Paper β€” ICML 2026 FM4LS](https://openreview.net/forum?id=ACHMa1e8J1) Β· [GitHub](https://github.com/reednaidoo/PhenoSeq)
"""

with gr.Blocks(title="PhenoSeq") as demo:
    with gr.Column(elem_id="col-container"):
        gr.Markdown(INTRO)

        with gr.Row():
            sample_id = gr.Dropdown(
                CHOICES, value=DEFAULT_ID, scale=4, filterable=True,
                label="Cell Painting well β€” held-out scGeneScope validation samples",
            )
            run = gr.Button("Generate transcriptomes", variant="primary", scale=1)

        info_out = gr.Markdown(_well_header(DEFAULT_ID))

        with gr.Row():
            montage_out = gr.Image(
                value=_montage(DEFAULT_ID), height=440,
                label="Conditioning cells β€” Cell Painting composite (DNA Β· RNA Β· ER Β· AGP Β· Mito)",
            )
            scatter_out = gr.Image(
                height=440,
                label="Generated vs. real scGPT cells (PCA of the real manifold)",
            )

        report_out = gr.Markdown()
        profile_out = gr.Image(
            height=250,
            label="Mean embedding profile β€” generated vs. real",
        )

        with gr.Accordion("Advanced settings", open=False):
            with gr.Row():
                n_cells = gr.Slider(16, 1024, value=256, step=16,
                                    label="Single cells to generate")
                ddim_steps = gr.Slider(0, 200, value=50, step=10,
                                       label="DDIM steps (0 = full 1000-step DDPM)")
                seed = gr.Number(value=0, precision=0, label="Seed")
            features_file = gr.File(
                label="Optional β€” your own ViT-L imaging features (.npy / .npz, shape (N, 5120))",
                file_types=[".npy", ".npz"],
            )
            gr.Markdown(
                "Feature layout: 5 Γ— 1024 ViT-L/14 channel embeddings concatenated in the order "
                "**DNA, RNA, AGP, Mito, ER**. Pass raw (un-normalised) features β€” the pipeline "
                "applies the training-split `img_norm.npz` statistics for you. Up to 16 imaging "
                "cells are used as context, matching the training recipe."
            )

        npy_out = gr.File(label="Generated scGPT embeddings β€” .npy, shape (n_cells, 512)")

        gr.Examples(
            examples=[[s] for s in ["202502PA_86", "202502PA_95", "202502PA_23",
                                    "202502PA_57", "202502PA_12", "202502PA_31",
                                    "202502PA_78"]],
            inputs=[sample_id],
            outputs=[montage_out, scatter_out, profile_out, info_out, report_out, npy_out],
            fn=generate,
            cache_examples=True,
            cache_mode="lazy",
            label="Held-out wells: DMSO vehicle Β· translation inhibitor Β· NF-ΞΊB inhibitor Β· "
                  "CDK inhibitor Β· PKC activator Β· NAMPT inhibitor Β· inert analgesic",
        )

        gr.Markdown(
            "---\n"
            "**How to read this.** scGPT embeddings carry a large common mean component β€” raw "
            "cosine between *any* two of them is β‰ˆ0.99 β€” so every similarity here is computed on "
            "**mean-centred** vectors, and the row below the headline number shows how similar the "
            "generation is to the *other* held-out wells. That gap is the signal.\n\n"
            "**Honest baseline.** Sampling all 15 held-out wells (256 cells, 50 DDIM steps, seed 0): "
            "median centred cosine to the correct well **0.66** vs **0.40** to the others; the "
            "correct well is ranked #1 for 5/15 wells and top-3 for 9/15 (chance: 1/15 and 3/15). "
            "On the harder 29-way treatment retrieval the median rank is 9/29 (chance 15) with 3/15 "
            "hits at rank 1. So: clearly above chance, far from solved β€” some wells (BAY 11-7082, "
            "roscovitine) are recovered exactly, others (TPA at 202502PA_12) are missed. PhenoSeq is "
            "trained with population-level supervision, so the generated cells model the well's "
            "transcriptional distribution, not cell-to-cell pairings.\n\n"
            "**Model** [Sentinal4D/PhenoSeq](https://huggingface.co/Sentinal4D/PhenoSeq) "
            "(Apache-2.0) β€” 168 M-parameter cross-attention denoiser, 1000-step cosine schedule, "
            "EMA weights, DDIM sampling. **Reference data** derived from "
            "[altoslabs/scGeneScope](https://huggingface.co/datasets/altoslabs/scGeneScope) "
            "(CC-BY-NC-4.0, Β© Altos Labs): the bundled cell crops, ViT-L features and matched "
            "scGPT embeddings are small excerpts reused here for non-commercial research "
            "demonstration, with attribution."
        )

    ev = dict(
        fn=generate,
        inputs=[sample_id, features_file, n_cells, ddim_steps, seed],
        outputs=[montage_out, scatter_out, profile_out, info_out, report_out, npy_out],
        api_name="generate",
    )
    run.click(**ev)
    sample_id.change(fn=preview, inputs=[sample_id],
                     outputs=[montage_out, info_out], api_name=False)

if __name__ == "__main__":
    demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)