"""CELL-FM CondenSeq demo: protein sequence -> condensate titration curve.
Enter an IDP sequence; the app generates a concentration ladder of images with
CELL-FM, classifies each as condensed or diffuse, smooths the calls into a
titration curve, and integrates it into AUC (condensation propensity) and AAC
(reentrant dissolution).
"""
import os
import shutil
import tempfile
import gradio as gr
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from cell_fm.apps.condenseq import metrics, pipeline
from cell_fm.apps.virtual_opencell import viewer
# gradio_client walks component API schemas to build /gradio_api/info. JSON Schema
# allows `additionalProperties: true`, but the walker recurses into that bool as if
# it were a schema and then does `"const" in schema`, raising
# TypeError: argument of type 'bool' is not iterable
# which 500s /gradio_api/info -- the API panel and any gradio_client call -- while
# the UI itself keeps serving. Short-circuit bool schemas before they reach get_type().
# The recursive calls resolve this name through the module global, so rebinding it
# covers the nested case that actually triggers here.
try:
import gradio_client.utils as _gc_utils
_orig_schema_to_type = _gc_utils._json_schema_to_python_type
def _schema_to_type(schema, defs=None):
if isinstance(schema, bool):
return "Any" if schema else "None"
return _orig_schema_to_type(schema, defs)
_gc_utils._json_schema_to_python_type = _schema_to_type
except Exception as exc: # never let the shim take the app down
print(f"[startup] gradio_client schema patch skipped: {exc}")
# ZeroGPU: only present on HF Spaces hardware, so the import is optional.
try:
import spaces
GPU_DECORATOR = spaces.GPU(duration=int(os.environ.get("ZEROGPU_DURATION", "300")))
except Exception: # running locally or on a dedicated GPU
def GPU_DECORATOR(fn):
return fn
HERE = os.path.dirname(os.path.abspath(__file__))
# palette: categorical slots 1 and 2 (validated, CVD dE 24.7 / normal dE 33.6)
CURVE = "#2a78d6" # the titration curve and the area under it
GAP = "#eb6834" # the area above the curve
INK = "#0b0b0b"
MUTED = "#52514e"
SURFACE = "#fcfcfb"
def load_default_sequence() -> str:
"""The NUP98 IDP, wild type -- what the box is pre-filled with."""
df = pd.read_csv(os.path.join(HERE, "assets", "nup98_mutation_seqs.csv"), dtype=str)
return df.loc[df.mutation == "WT", "sequence"].iloc[0]
DEFAULT_SEQUENCE = load_default_sequence()
# The proteins with virtual staining, and OpenCell's annotations of their real cell lines
PROTEINS = viewer.load_catalog(os.path.join(HERE, "assets", "virtual_opencell_proteins.csv"))
def plot_curve(result: pipeline.Result):
"""The titration curve, with the two integrated areas drawn as the areas they are."""
x, y, curve = result.intensities, result.predictions, result.curve
flat = metrics.no_reentrant_curve(curve)
fig, ax = plt.subplots(figsize=(8.5, 4.6), facecolor=SURFACE)
ax.set_facecolor(SURFACE)
# the raw per-image calls, kept visible but recessive: the curve is a summary of these
ax.scatter(x, y, s=7, alpha=0.16, color=MUTED, linewidths=0, zorder=1)
# AUC: what the sequence actually does
ax.fill_between(x, 0, curve, color=CURVE, alpha=0.18, linewidth=0, zorder=2,
label=f"AUC {result.auc:.3f}")
# AAC: the condensation lost to reentrant dissolution at high concentration
if result.aac > 1e-9:
ax.fill_between(x, curve, flat, where=flat > curve, color=GAP, alpha=0.30,
linewidth=0, zorder=3, label=f"AAC {result.aac:.3f}")
ax.plot(x, flat, color=GAP, linewidth=1.2, linestyle="--", alpha=0.8, zorder=4)
ax.plot(x, curve, color=CURVE, linewidth=2.0, zorder=5)
ax.set_xscale("log")
ax.set_xlim(x.min(), x.max())
ax.set_ylim(-0.02, 1.02)
ax.set_xlabel("Protein intensity level (a.u., log scale)", fontsize=11, color=INK)
ax.set_ylabel("Condensate probability", fontsize=11, color=INK)
ax.tick_params(labelsize=10, colors=MUTED)
ax.grid(True, alpha=0.18, linewidth=0.8)
ax.set_axisbelow(True)
for side in ("top", "right"):
ax.spines[side].set_visible(False)
for side in ("left", "bottom"):
ax.spines[side].set_color(MUTED)
ax.spines[side].set_linewidth(0.8)
if result.c_sat != float("inf"):
ax.axvline(result.c_sat, color=MUTED, linewidth=1.0, linestyle=":", zorder=4)
ax.annotate(
f"$c_{{sat}}$ {result.c_sat:.0f}",
xy=(result.c_sat, 1.0), xytext=(6, -13), textcoords="offset points",
fontsize=9, color=MUTED, ha="left", va="top",
)
ax.legend(loc="upper left", frameon=False, fontsize=10, labelcolor=INK)
fig.tight_layout()
return fig
def plot_samples(result: pipeline.Result):
"""A strip of generated images across the ladder, so the curve is checkable by eye."""
imgs, xs = result.images, result.image_intensities
fig, axes = plt.subplots(1, len(imgs), figsize=(2.0 * len(imgs), 2.3), facecolor=SURFACE)
for ax, img, level in zip(np.atleast_1d(axes), imgs, xs):
ax.imshow(img, cmap="magma", vmin=0, vmax=1)
ax.set_title(f"{level:.0f}", fontsize=10, color=INK)
ax.set_xticks([]); ax.set_yticks([])
for s in ax.spines.values():
s.set_visible(False)
fig.suptitle("Generated protein channel across the concentration ladder (a.u.)",
fontsize=10, color=MUTED, y=0.04)
fig.tight_layout(rect=(0, 0.06, 1, 1))
return fig
def summary_markdown(result: pipeline.Result) -> str:
csat = "not reached" if result.c_sat == float("inf") else f"{result.c_sat:.0f}"
return f"""
| | |
|---|---|
| **AUC** — area under the curve, condensation propensity | **{result.auc:.4f}** |
| **AAC** — area above the curve, reentrant dissolution | **{result.aac:.4f}** |
| No-reentrant AUC | {result.no_reentrant_auc:.4f} |
| csat (first intensity at P ≥ 0.8) | {csat} |
| Condensed calls | {int(result.predictions.sum())} / {len(result.predictions)} |
| Moving-average window | {result.window} |
| Sequence length | {len(result.sequence)} aa |
Both areas are normalised by the width of the log-intensity range, so they lie in
[0, 1] and are comparable across sequences.
"""
def results_table(result: pipeline.Result) -> pd.DataFrame:
"""The numbers behind the plot, so the chart is never the only way to read them."""
return pd.DataFrame({
"protein_intensity_level": np.round(result.intensities, 4),
"predicted_class": result.predictions,
"condensate_probability_ma": np.round(result.curve, 6),
})
@GPU_DECORATOR
def analyse(sequence, progress=gr.Progress()):
try:
sequence = pipeline.clean_sequence(sequence)
except ValueError as exc:
raise gr.Error(str(exc)) from exc
result = pipeline.run(sequence, progress=progress)
table = results_table(result)
csv_path = os.path.join("/tmp", "pred_vs_intensity.csv")
table.to_csv(csv_path, index=False)
return plot_curve(result), plot_samples(result), summary_markdown(result), table, csv_path
# ---------------------------------------------------------------------------
# Sections
#
# Each application of CELL-FM is one tab. To add another, write a
# build__section() that lays out its own components and wires its own
# events, then give it a gr.Tab in the Blocks at the bottom. Sections are
# independent: they share the loaded models via pipeline.py and nothing else.
# ---------------------------------------------------------------------------
def build_condensate_titration_section():
"""Sequence in, condensate titration curve out, summarised as AUC and AAC."""
gr.Markdown(
"""
Predict how an Intrinsically Disordered Peptide (IDP) behaves as its concentration rises.
"""
)
gr.Image(
value=os.path.join(HERE, "images", "cellfm_flow_diagram.png"),
show_label=False, show_download_button=False, container=False,
interactive=False,
)
with gr.Row():
with gr.Column(scale=2):
sequence = gr.Textbox(
value=DEFAULT_SEQUENCE, lines=4,
label="IDP sequence",
info=(
f"Exactly {pipeline.SEQUENCE_LENGTH} single-letter amino acids — the whole "
"CondenSeq library is 66-mers. FASTA headers are stripped."
),
)
run_button = gr.Button("Run analysis", variant="primary")
with gr.Column(scale=3):
summary = gr.Markdown()
curve_plot = gr.Plot(label="Condensate titration curve")
sample_plot = gr.Plot(label="Generated images")
with gr.Accordion("Per-image results", open=False):
table = gr.Dataframe(label="pred_vs_intensity", wrap=True)
download = gr.File(label="Download CSV")
gr.Markdown(
"""
---
Generation follows `evaluate_seq2img_dict.py`, classification `evaluate_single_img.py`,
smoothing `ma_plot.py`, and the areas `analysis/ana_all_mutation_log_scale.py` from the
CELL-FM repository. Every image is conditioned on one fixed reference nucleus, so curves
are comparable across sequences.
"""
)
run_button.click(
analyse,
inputs=sequence,
outputs=[curve_plot, sample_plot, summary, table, download],
)
def build_virtual_opencell_section():
"""Pick an OpenCell protein and browse CELL-FM's virtual staining of it, laid out like an
OpenCell target page: the real cell line's annotations on the left, the images on the right."""
def load(gene):
row = PROTEINS.loc[gene]
stack, paths = viewer.load_samples(gene, int(row["n_samples"]))
return row, stack, paths
def show(gene, index, channel, n_min, n_max, n_gamma, t_min, t_max, t_gamma):
_, stack, _ = load(gene)
return viewer.render(stack[index], channel, (n_min, n_max, n_gamma), (t_min, t_max, t_gamma))
def sample_file(gene, index):
"""The raw sample under a readable name, copied where Gradio is allowed to serve it."""
_, _, paths = load(gene)
path = os.path.join(tempfile.gettempdir(), f"virtual_opencell_{gene}_{index + 1:04d}.tif")
shutil.copyfile(paths[index], path)
return path
def protein_view(gene, *settings):
"""Panel, thumbnails, first sample and its file for one protein."""
row, stack, _ = load(gene)
image = show(gene, 0, *settings)
return viewer.target_html(gene, row), viewer.thumbnails(stack), image, sample_file(gene, 0)
def umap_plot(gene):
try:
return viewer.umap_figure(gene)
except Exception as exc: # the map is an extra; the images work without it
print(f"[virtual opencell] map unavailable: {exc}")
return None
def select_protein(gene, channel, n_min, n_max, n_gamma, t_min, t_max, t_gamma):
try:
info, thumbs, image, path = protein_view(gene, channel, n_min, n_max, n_gamma, t_min, t_max, t_gamma)
except Exception as exc:
raise gr.Error(f"Could not load the images for {gene}: {viewer.load_error(exc)}") from exc
return info, gr.Gallery(value=thumbs, selected_index=0), image, 0, path, umap_plot(gene)
def select_sample(gene, channel, n_min, n_max, n_gamma, t_min, t_max, t_gamma, evt: gr.SelectData):
image = show(gene, evt.index, channel, n_min, n_max, n_gamma, t_min, t_max, t_gamma)
return image, evt.index, sample_file(gene, evt.index)
# Shown on first load: the default protein, unless its images cannot be fetched
defaults = ("Both", 0, 100, 1.0, 0, 100, 1.0)
try:
info0, thumbs0, image0, file0 = protein_view(viewer.DEFAULT_GENE, *defaults)
except Exception as exc: # e.g. no access to the dataset; the error is shown in the panel
print(f"[startup] virtual OpenCell preload failed: {exc}")
info0, thumbs0, image0, file0 = viewer.error_html(viewer.DEFAULT_GENE, exc), None, None, None
umap0 = umap_plot(viewer.DEFAULT_GENE)
with gr.Column(elem_id="virtual-opencell"):
with gr.Row(elem_classes="vo-navbar", equal_height=True):
gr.HTML(viewer.NAVBAR_HTML)
protein = gr.Dropdown(
choices=list(PROTEINS.index), value=viewer.DEFAULT_GENE, filterable=True,
show_label=False, container=False, scale=0, min_width=220, elem_classes="vo-search",
)
with gr.Row():
with gr.Column(scale=4):
info = gr.HTML(info0)
gr.HTML(viewer.UMAP_HTML, padding=False)
umap = gr.Plot(value=umap0, show_label=False, elem_classes="vo-umap")
with gr.Column(scale=6):
with gr.Row():
channel = gr.Radio(viewer.CHANNELS, value=defaults[0], label="Channel",
elem_classes="vo-buttons")
download = gr.DownloadButton("Download TIFF", value=file0, size="sm", variant="secondary",
scale=0, min_width=140, elem_classes="vo-download")
image = gr.Image(
value=image0, show_label=False, interactive=False, height=600,
show_download_button=True, elem_classes="vo-viewer",
)
gr.HTML('Samples
', padding=False)
samples = gr.Gallery(
value=thumbs0, selected_index=0 if thumbs0 else None, show_label=False,
columns=8, height=230, allow_preview=False, object_fit="contain",
show_share_button=False, show_download_button=False, elem_classes="vo-thumbnails",
)
with gr.Row(elem_classes="vo-settings"):
with gr.Column():
gr.HTML('Nucleus image settings
', padding=False)
n_min = gr.Slider(0, 100, value=defaults[1], step=1, label="Intensity min (%)")
n_max = gr.Slider(1, 150, value=defaults[2], step=1, label="Intensity max (%)")
n_gamma = gr.Slider(0.5, 1.5, value=defaults[3], step=0.05, label="Gamma")
with gr.Column():
gr.HTML('Target image settings
', padding=False)
t_min = gr.Slider(0, 100, value=defaults[4], step=1, label="Intensity min (%)")
t_max = gr.Slider(1, 150, value=defaults[5], step=1, label="Intensity max (%)")
t_gamma = gr.Slider(0.5, 1.5, value=defaults[6], step=0.05, label="Gamma")
gr.HTML(viewer.FOOTER_HTML)
index = gr.State(0)
settings = [channel, n_min, n_max, n_gamma, t_min, t_max, t_gamma]
protein.change(select_protein, inputs=[protein] + settings,
outputs=[info, samples, image, index, download, umap])
# Redraws take milliseconds: no progress overlay, and a dragged slider only renders its latest value
samples.select(select_sample, inputs=[protein] + settings, outputs=[image, index, download],
show_progress="hidden")
for control in settings:
control.change(show, inputs=[protein, index] + settings, outputs=image,
show_progress="hidden", trigger_mode="always_last")
with gr.Blocks(title="CELL-FM", theme=gr.themes.Soft(), css=viewer.CSS, head=viewer.HEAD) as demo:
gr.Markdown(
"""
# CELL-FM
A virtual microscopy model that bridges microscopy images and protein sequences.
"""
)
with gr.Tabs():
with gr.Tab("Condensate Titration"):
build_condensate_titration_section()
with gr.Tab("Virtual OpenCell"):
build_virtual_opencell_section()
# Download weights and build the models at startup, on CPU. Under ZeroGPU the GPU
# budget only covers the decorated call, so this must not happen inside it.
try:
pipeline.load_models()
except Exception as exc: # surfaced in the Space logs; the UI still loads
print(f"[startup] model preload failed: {exc}")
if __name__ == "__main__":
# Only pass share when explicitly asked (run_local.sh). On Spaces the kwarg must
# be omitted entirely: an explicit share=False makes Gradio demand a share link
# when it cannot reach localhost, which kills the app at startup.
launch_kwargs = {"share": True} if os.environ.get("GRADIO_SHARE") == "1" else {}
demo.queue().launch(**launch_kwargs)