phenoseq / app.py
multimodalart's picture
multimodalart HF Staff
Size ZeroGPU duration to measured sampling time (was ~5x over)
f609d01 verified
Raw History Blame Contribute Delete
19.1 kB
"""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)