rwbaber's picture
Upload app.py with huggingface_hub
5b1e7f4 verified
Raw
History Blame
12.5 kB
import os
import warnings
import gradio as gr
import numpy as np
import plotly.graph_objects as go
from plotly.subplots import make_subplots
import scanpy as sc
warnings.filterwarnings('ignore')
sc.settings.verbosity = 0
DATASET_REPO = "minuttilab/cDCFUN-data"
_DIR = os.path.dirname(__file__)
# ── Load data ─────────────────────────────────────────────────────────────────
_LOCAL_H5AD = os.path.join(_DIR, "adata_annotated.h5ad")
if os.path.exists(_LOCAL_H5AD):
print("Loading local h5ad…")
_data_path = _LOCAL_H5AD
else:
print("Fetching dataset from HuggingFace…")
from huggingface_hub import hf_hub_download
_data_path = hf_hub_download(
repo_id=DATASET_REPO,
filename="adata_annotated.h5ad",
repo_type="dataset",
token=os.environ.get("cDCFUN"),
)
print("Loading AnnData…")
adata = sc.read_h5ad(_data_path)
print(f"Ready — {adata.n_obs:,} cells · {adata.n_vars:,} genes.")
# ── Gene name lookup (case-insensitive) ───────────────────────────────────────
_var_lower = {g.lower(): g for g in adata.var_names}
# ── UMAP coordinates ──────────────────────────────────────────────────────────
if 'X_umap' in adata.obsm:
_umap = np.asarray(adata.obsm['X_umap'])
else:
raise RuntimeError("AnnData object does not contain UMAP coordinates (adata.obsm['X_umap']).")
def _get_expr(gene: str) -> np.ndarray:
try:
sub = adata[:, gene].X
if hasattr(sub, 'toarray'):
sub = sub.toarray()
return np.asarray(sub).ravel()
except Exception:
return np.zeros(adata.n_obs, dtype=float)
# ── Layout constants ──────────────────────────────────────────────────────────
_MODEBAR_REMOVE = [
'select2d', 'lasso2d', 'zoomIn2d', 'zoomOut2d',
'toggleSpikelines', 'hoverClosestCartesian', 'hoverCompareCartesian',
]
CELL_SIZE = 260
CB_PAD = 60
_FONT = dict(family='Roboto, sans-serif', size=11)
UMAP_STYLE = {
'expr': {
'size': 2,
'opacity': 0.85,
'colorscale': 'Spectral',
'reversescale': True,
}
}
# ── Gene expression grid ──────────────────────────────────────────────────────
def _make_expr_grid(genes: list, shared_scale: bool = False, n_cols: int = 3, size: int = None, opacity: float = None, colorscale: str = None, reversescale: bool = None) -> go.Figure:
if not genes:
return go.Figure()
if size is None: size = UMAP_STYLE['expr']['size']
if opacity is None: opacity = UMAP_STYLE['expr']['opacity']
if colorscale is None: colorscale = UMAP_STYLE['expr']['colorscale']
if reversescale is None: reversescale = UMAP_STYLE['expr']['reversescale']
n_rows = (len(genes) + n_cols - 1) // n_cols
H_GAP = 0.04
GAP_PX = 30
MARGIN_T = 30
MARGIN_B = 5
paper_h = n_rows * CELL_SIZE + max(0, n_rows - 1) * GAP_PX
total_h = paper_h + MARGIN_T + MARGIN_B
V_GAP = (GAP_PX / paper_h) if n_rows > 1 else 0.0
subplot_w = (1.0 - H_GAP * (n_cols - 1)) / n_cols
subplot_h = CELL_SIZE / paper_h
all_vals = [_get_expr(g) for g in genes]
global_max = float(max(v.max() for v in all_vals)) if shared_scale else None
subtitles = genes + [''] * (n_rows * n_cols - len(genes))
fig = make_subplots(
rows=n_rows, cols=n_cols,
subplot_titles=subtitles,
horizontal_spacing=H_GAP,
vertical_spacing=V_GAP,
)
for i, (gene, vals) in enumerate(zip(genes, all_vals)):
r, c = divmod(i, n_cols)
vmin = 0.0 if shared_scale else float(vals.min())
vmax = global_max if shared_scale else float(vals.max())
cb_x = c * (subplot_w + H_GAP) + subplot_w + 0.01
cb_y = 1.0 - r * (subplot_h + V_GAP) - subplot_h / 2
fig.add_trace(go.Scattergl(
x=_umap[:, 0], y=_umap[:, 1],
mode='markers',
marker=dict(
color=vals,
colorscale=colorscale,
reversescale=reversescale,
cmin=vmin, cmax=vmax,
size=size,
opacity=opacity,
colorbar=dict(
x=cb_x, xanchor='left',
y=cb_y, yanchor='middle',
len=subplot_h,
thickness=10,
outlinewidth=0,
tickfont=dict(size=9),
nticks=4,
),
),
hoverinfo='skip',
showlegend=False,
), row=r + 1, col=c + 1)
fig.update_xaxes(visible=False, row=r + 1, col=c + 1)
fig.update_yaxes(visible=False, row=r + 1, col=c + 1)
fig.update_layout(
height=total_h,
margin=dict(l=5, r=CB_PAD, t=MARGIN_T, b=MARGIN_B),
plot_bgcolor='rgba(0,0,0,0)',
paper_bgcolor='rgba(0,0,0,0)',
dragmode='pan',
font=_FONT,
modebar_remove=_MODEBAR_REMOVE,
)
return fig
# ── UI ────────────────────────────────────────────────────────────────────────
INSTRUCTIONS = """
Type gene name(s) separated by commas — e.g. `Clec7a, Cd274, Ido1`
Press **Enter** or click **Plot genes**.
"""
CSS = """
* { font-family: Roboto, sans-serif !important; }
footer, header { display: none !important; }
.gradio-container { max-width: 100% !important; padding: 8px !important; }
@media (prefers-color-scheme: light) {
.js-plotly-plot text { fill: #1e293b !important; }
}
@media (prefers-color-scheme: dark) {
.js-plotly-plot text { fill: #e2e8f0 !important; }
}
.js-plotly-plot { width: 100% !important; }
.gradio-container .plotly-graph-div { width: 100% !important; overflow: visible !important; }
.js-plotly-plot svg { max-width: 100% !important; }
.gradio-container .plot-container { overflow: visible !important; height: auto !important; }
/* Responsive expression grid: desktop shows 3-col, mobile shows 1-col */
#expr-desktop { display: block; }
#expr-mobile { display: none; }
@media (max-width: 768px) {
#expr-desktop { display: none !important; }
#expr-mobile { display: block !important; }
}
.js-plotly-plot .modebar-container {
right: auto !important;
left: -25px !important;
top: 0px !important;
width: auto !important;
}
.js-plotly-plot .modebar {
left: 28px !important;
right: auto !important;
transform: none !important;
display: flex !important;
flex-direction: row !important;
flex-wrap: wrap !important;
width: 86px !important;
justify-content: flex-start !important;
}
.js-plotly-plot .modebar-btn {
opacity: 1 !important;
display: flex !important;
align-items: center !important;
justify-content: center !important;
width: 28px !important;
height: 28px !important;
background: transparent !important;
margin: 0 !important;
padding: 0 !important;
}
.js-plotly-plot .modebar-btn svg {
width: 16px !important;
height: 16px !important;
overflow: visible !important;
}
.js-plotly-plot .modebar-btn svg path,
.js-plotly-plot .modebar-btn svg rect,
.js-plotly-plot .modebar-btn svg polygon {
fill: #888888 !important;
stroke: none !important;
}
.js-plotly-plot .modebar-btn:hover svg path,
.js-plotly-plot .modebar-btn:hover svg rect,
.js-plotly-plot .modebar-btn:hover svg polygon {
fill: #2196F3 !important;
}
.js-plotly-plot .modebar-group {
display: contents !important;
}
.js-plotly-plot .modebar-btn { order: 10 !important; }
.js-plotly-plot .modebar-btn[data-title*="Plotly"] { order: 1 !important; }
.js-plotly-plot .modebar-btn[data-title="Reset axes"] { order: 2 !important; }
.js-plotly-plot .modebar-btn[data-title="Autoscale"] { order: 3 !important; }
.js-plotly-plot .modebar-btn[data-title="Pan"] { order: 4 !important; }
.js-plotly-plot .modebar-btn[data-title="Zoom"] { order: 5 !important; }
.js-plotly-plot .modebar-btn[data-title*="Download"] { order: 6 !important; }
.js-plotly-plot .modebar-btn[data-title]::before { display: none !important; }
.js-plotly-plot .modebar-btn[data-title]::after {
right: auto !important;
left: 0px !important;
margin-right: 0 !important;
text-align: left !important;
}
.js-plotly-plot .modebar-btn[data-title="Zoom"]::after { content: "Draw Box to Zoom" !important; }
.js-plotly-plot .modebar-btn[data-title="Pan"]::after { content: "Drag to Pan" !important; }
.js-plotly-plot .modebar-btn[data-title="Autoscale"]::after { content: "Auto Scale" !important; }
.js-plotly-plot .modebar-btn[data-title="Reset axes"]::after { content: "Reset View" !important; }
.js-plotly-plot .modebar-btn[data-title*="Download"]::after { content: "Save as PNG" !important; }
.js-plotly-plot .modebar-btn[data-title*="Plotly"]::after { content: "About Plotly" !important; }
"""
with gr.Blocks(theme=gr.themes.Soft(primary_hue="indigo"), title="Gene Expression Explorer", css=CSS, fill_width=True) as demo:
with gr.Row():
with gr.Column(scale=1):
gr.Markdown("# Gene Expression Explorer")
gr.Markdown("Explore gene expression across dendritic cell populations. Interactive Plotly UMAPs are shown.")
gene_box = gr.Textbox(placeholder="e.g. Clec7a, Cd274, Ido1, Mafb", label="Gene names (comma-separated)", lines=1)
with gr.Row():
plot_btn = gr.Button("Plot genes", variant="primary")
shared_chk = gr.Checkbox(label="Shared scale (0 → max)", value=False)
error_md = gr.Markdown("")
gr.Markdown(INSTRUCTIONS)
gr.Markdown("**Reference UMAPs**")
gr.Markdown("Cell type")
gr.Image(value=os.path.join(_DIR, "ref_celltype.png"), show_label=False,
interactive=False, show_download_button=False, show_fullscreen_button=False,
height=CELL_SIZE)
gr.Markdown("Lineage")
gr.Image(value=os.path.join(_DIR, "ref_lineage.png"), show_label=False,
interactive=False, show_download_button=False, show_fullscreen_button=False,
height=CELL_SIZE)
gr.Markdown("Condition")
gr.Image(value=os.path.join(_DIR, "ref_condition.png"), show_label=False,
interactive=False, show_download_button=False, show_fullscreen_button=False,
height=CELL_SIZE)
with gr.Column(scale=3):
gr.Markdown("**Gene expression**")
out_plot_desktop = gr.Plot(show_label=False, elem_id="expr-desktop")
out_plot_mobile = gr.Plot(show_label=False, elem_id="expr-mobile")
def plot_genes_interactive(gene_input: str, shared_scale: bool):
if not gene_input or not gene_input.strip():
return go.Figure(), go.Figure(), ""
raw = [g.strip() for g in gene_input.replace(';', ',').split(',') if g.strip()]
found = [_var_lower[g.lower()] for g in raw if g.lower() in _var_lower]
missed = [g for g in raw if g.lower() not in _var_lower]
msg = ""
if missed:
msg = f"⚠️ Not found in dataset: {', '.join(missed)}"
if not found:
return go.Figure(), go.Figure(), msg or "No valid gene names entered."
fig_desktop = _make_expr_grid(found, shared_scale=shared_scale, n_cols=3, **UMAP_STYLE['expr'])
fig_mobile = _make_expr_grid(found, shared_scale=shared_scale, n_cols=1, **UMAP_STYLE['expr'])
plotted = f"**Plotting:** {', '.join(found)}"
msg = plotted + ("\n\n" + msg if msg else "")
return fig_desktop, fig_mobile, msg
plot_btn.click(plot_genes_interactive, inputs=[gene_box, shared_chk], outputs=[out_plot_desktop, out_plot_mobile, error_md], api_name="plot")
gene_box.submit(plot_genes_interactive, inputs=[gene_box, shared_chk], outputs=[out_plot_desktop, out_plot_mobile, error_md])
demo.queue()
demo.launch()