Spaces:
Running on Zero
Running on Zero
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)
|