CELL-FM / app.py
BoHuangLab's picture
Import the CondenSeq app modules from cell_fm.apps.condenseq; add vit_cls
a46a831 verified
Raw History Blame Contribute Delete
10.3 kB
"""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 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
# 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()
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} |
| c<sub>sat</sub> (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_<name>_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],
)
with gr.Blocks(title="CELL-FM", theme=gr.themes.Soft()) 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()
# 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)