bozdaglab's picture F4RUKYLDRM's picture
Redesign workbench dashboard shell (#1)
2e6a2a4
Raw History Blame Contribute Delete
90.2 kB
from __future__ import annotations
import hashlib
import html
import math
import platform
import secrets
import tempfile
import time
import zipfile
from datetime import datetime
from functools import lru_cache, partial
from pathlib import Path
import numpy as np
import spaces
import gradio as gr
import pandas as pd
import plotly.graph_objects as go
import torch
from sklearn.metrics import roc_auc_score
from biolmnet.artifacts import EXPECTED_FILES, load_bundle, save_bundle
from biolmnet.data import (
PreparedWorkspace,
UPSTREAM_REPOSITORY,
attach_embeddings_and_pathways,
build_biological_mask,
github_dataset_sources,
load_genept_embeddings,
read_csv,
upstream_example_sources,
upstream_interaction_sources,
validate_and_align_omics,
)
from biolmnet.training import (
Hyperparameters,
ModelBundle,
pathway_importance,
predict,
train,
)
from biolmnet.ui import components as ui
from biolmnet.ui.styles import CSS
GENEPT_OPTIONS = {
"Auto — bulk or single-cell based on source": "auto",
"Bulk · large-3 context": "embedding_original_large_3.parquet",
"Bulk · original ada-text context": "embedding_original_ada_text.parquet",
"Single-cell · cell type, tissue, drug & pathway": (
"embedding_associations_cell_type_tissue_drug_pathway_openai_large.parquet"
),
"Single-cell · age, cell type, drugs & pathways": (
"embedding_associations_age_cell_type_drugs_pathways_openai_large.parquet"
),
}
STAGE_KEYS = ["data", "train", "export", "predict", "results"]
STAGE_LABELS = ["Data & Priors", "Train Model", "Export Artifacts", "Predict", "Results"]
STAGE_TOTAL = len(STAGE_KEYS)
# ═══════════════════════════════════════════════════════════════════════
# Small presentation-only helpers (no scientific logic below this line
# calls into biolmnet.model / biolmnet.data / biolmnet.training / the
# artifacts module for anything other than reading the values they
# already computed).
# ═══════════════════════════════════════════════════════════════════════
def _status(message: str, error: bool = False) -> str:
return ui.simple_status_html(html.escape(message), error=error)
def _source_visibility(mode: str):
return (
gr.update(visible=mode == "BioLM-NET examples"),
gr.update(visible=mode == "GitHub folder"),
gr.update(visible=mode == "Upload files"),
)
def _workflow_visibility(active_page: str):
return tuple(gr.update(visible=page == active_page) for page in STAGE_KEYS)
@lru_cache(maxsize=1)
def _upstream_interactions() -> tuple[pd.DataFrame, pd.DataFrame]:
pdi_url, ppi_url = upstream_interaction_sources()
return read_csv(pdi_url), read_csv(ppi_url)
def _read_required_upload(path: str | None, label: str) -> pd.DataFrame:
if not path:
raise ValueError(f"Upload {label}.")
return read_csv(path)
def _resolve_embedding_file(option: str, source_mode: str, example_dataset: str) -> str:
selected = GENEPT_OPTIONS[option]
if selected != "auto":
return selected
if source_mode == "BioLM-NET examples" and example_dataset == "scTrioseq2":
return "embedding_associations_cell_type_tissue_drug_pathway_openai_large.parquet"
return "embedding_original_large_3.parquet"
def _graph_preview_html(mask_density: str, gene_nodes: str, pathway_units: str, genept_dim: str) -> str:
return ui.blueprint_div(
'<div style="display:grid;grid-template-columns:1fr 1fr;gap:14px 10px">'
+ ui.mini_stat_html(mask_density, "Mask density")
+ ui.mini_stat_html(gene_nodes, "Gene nodes")
+ ui.mini_stat_html(pathway_units, "Pathway units")
+ ui.mini_stat_html(genept_dim, "GenePT dim")
+ "</div>"
)
# ── rail / chrome state machine (presentation only) ─────────────────────
def _rail_states(active: str, workspace_ready: bool, model_ready: bool, predicted_ready: bool):
states = []
states.append(
("on", "Active") if active == "data" else (("done", "Done") if workspace_ready else ("next", "Next"))
)
if active == "train":
states.append(("on", "Active"))
elif model_ready:
states.append(("done", "Done"))
elif workspace_ready:
states.append(("next", "Next"))
else:
states.append(("off", "Locked"))
if active == "export":
states.append(("on", "Active"))
elif model_ready:
states.append(("done", "Done"))
else:
states.append(("off", "Locked"))
if active == "predict":
states.append(("on", "Active"))
elif predicted_ready:
states.append(("done", "Done"))
elif model_ready:
states.append(("done", "Ready"))
else:
states.append(("off", "Available"))
if active == "results":
states.append(("on", "Active"))
elif model_ready:
states.append(("done", "Ready"))
elif workspace_ready:
states.append(("off", "Partial"))
else:
states.append(("off", "Available"))
return states
def _run_state_rows(active: str, workspace, bundle, run_meta: dict, align: dict, predicted: dict):
dataset = (run_meta or {}).get("source_name") or "—"
arch_tone = "accent" if workspace else "muted"
arch_text = "Built" if workspace else "None"
model_tone = "accent" if bundle else "muted"
model_text = "Trained" if bundle else "Untrained"
if active == "export":
return [
("Dataset", dataset, None),
("Architecture", arch_text, arch_tone),
("Model", model_text, model_tone),
("Val accuracy", f"{bundle.metrics['accuracy']:.3f}" if bundle else "—", None),
("Artifact", "Ready" if bundle else "None", "accent" if bundle else "muted"),
]
if active == "predict":
align = align or {}
source_label = align.get("source_label", "Session model")
required = align.get("required")
matched = align.get("matched")
ok = align.get("ok")
return [
("Scoring with", source_label, None),
("Model", model_text, model_tone),
("Features required", f"{required:,}" if required is not None else "—", None),
(
"Features matched",
f"{matched:,}" if matched is not None else "—",
None if ok in (None, True) else "error",
),
("Inference", "Ready" if ok else ("Blocked" if ok is not None else "—"), "accent" if ok else ("error" if ok is not None else None)),
]
if active == "results":
predicted = predicted or {}
rows = [
("Run", (run_meta or {}).get("session_id", "—"), None),
(
"Epochs",
f"{bundle.metrics['epochs_completed']}" if bundle else "—",
None,
),
("Val accuracy", f"{bundle.metrics['accuracy']:.3f}" if bundle else "—", "accent" if bundle else None),
("Macro F1", f"{bundle.metrics['f1_macro']:.3f}" if bundle else "—", None),
(
"Predictions",
f"{predicted['n_samples']:,} scored" if predicted.get("n_samples") else "None yet",
None,
),
]
return rows
if active == "data":
if workspace:
genes_aligned = f"{len(workspace.gene_branch.input_genes):,} · {len(workspace.dna_branch.input_genes):,}"
pathways_kept = f"{len(workspace.gene_branch.pathways) + len(workspace.dna_branch.pathways):,}"
else:
genes_aligned = "—"
pathways_kept = "—"
return [
("Dataset", dataset, None),
("Genes aligned", genes_aligned, None),
("Pathways kept", pathways_kept, None),
("Architecture", arch_text, arch_tone),
("Model", model_text, model_tone),
]
# train
rows = [
("Dataset", dataset, None),
("Architecture", arch_text, arch_tone),
("Model", model_text, model_tone),
]
if bundle:
rows.append(("Epochs", f"{bundle.metrics['epochs_completed']}", None))
rows.append(("Device", (bundle.metrics.get("device") if bundle else "ZeroGPU") or "ZeroGPU", None))
return rows
def _rail_row_updates(active, workspace, bundle, predicted_ready):
states = _rail_states(active, bool(workspace), bool(bundle), predicted_ready)
return [ui.rail_row_html(i + 1, label, state, word) for i, (label, (state, word)) in enumerate(zip(STAGE_LABELS, states))]
def _session_note(bundle) -> str:
if bundle:
device = bundle.metrics.get("device", "cpu")
return "ZeroGPU idle" if device == "cuda" else "CPU"
return "ZeroGPU idle"
def _refresh_chrome(active, workspace, bundle, run_meta, align, predicted, session_id):
run_meta = run_meta or {}
predicted = predicted or {}
predicted_ready = bool(predicted.get("n_samples"))
rail_html = _rail_row_updates(active, workspace, bundle, predicted_ready)
plate = ui.run_state_plate(_run_state_rows(active, workspace, bundle, run_meta, align, predicted))
stage_index = STAGE_KEYS.index(active) + 1
topbar = ui.topbar_html(
stage_index, STAGE_TOTAL, STAGE_LABELS[stage_index - 1], session_id or "——————", _session_note(bundle)
)
train_nav_ok = bool(workspace)
export_nav_ok = bool(bundle)
return (*rail_html, plate, topbar, gr.update(interactive=train_nav_ok), gr.update(interactive=export_nav_ok))
CHROME_OUTPUT_NAMES = [
"rail_1", "rail_2", "rail_3", "rail_4", "rail_5", "run_plate", "topbar", "train_click", "export_click",
]
def _enter_page(page_key, workspace, bundle, run_meta, align, predicted, session_id):
visibility = tuple(gr.update(visible=page == page_key) for page in STAGE_KEYS)
chrome = _refresh_chrome(page_key, workspace, bundle, run_meta, align, predicted, session_id)
return (page_key, *visibility, *chrome)
# ═══════════════════════════════════════════════════════════════════════
# Data & Priors
# ═══════════════════════════════════════════════════════════════════════
def prepare_workspace(
source_mode: str,
example_dataset: str,
github_folder: str,
uploaded_gene: str | None,
uploaded_dna: str | None,
uploaded_labels: str | None,
uploaded_gene_pathway: str | None,
uploaded_dna_pathway: str | None,
uploaded_pdi: str | None,
uploaded_ppi: str | None,
pathway_files_are_significant: bool,
embedding_option: str,
progress=gr.Progress(track_tqdm=False),
):
empty = pd.DataFrame()
try:
progress(0.03, desc="Reading omics data")
if source_mode == "BioLM-NET examples":
sources = upstream_example_sources(example_dataset)
frames = {name: read_csv(url) for name, url in sources.items()}
file_labels = {name: Path(url).name for name, url in sources.items()}
source_name = f"BioLM-NET / {example_dataset}"
precomputed_significant = True
allow_preset_trim = True
elif source_mode == "GitHub folder":
if not github_folder.strip():
raise ValueError("Enter a GitHub dataset folder URL.")
sources = github_dataset_sources(github_folder)
frames = {name: read_csv(url) for name, url in sources.items()}
file_labels = {name: Path(url).name for name, url in sources.items()}
source_name = github_folder.strip()
precomputed_significant = pathway_files_are_significant
allow_preset_trim = False
else:
uploads = {
"gene": (uploaded_gene, "Gene_Expression.csv"),
"dna": (uploaded_dna, "DNA_Methylation.csv"),
"labels": (uploaded_labels, "label.csv"),
"gene_pathways": (uploaded_gene_pathway, "the gene-expression pathway mapping CSV"),
"dna_pathways": (uploaded_dna_pathway, "the DNA-methylation pathway mapping CSV"),
}
frames = {key: _read_required_upload(path, label) for key, (path, label) in uploads.items()}
file_labels = {key: Path(path).name for key, (path, _) in uploads.items()}
source_name = "Uploaded dataset"
precomputed_significant = pathway_files_are_significant
allow_preset_trim = False
(gene_frame, dna_frame, labels, label_names, warnings) = validate_and_align_omics(
frames["gene"], frames["dna"], frames["labels"], allow_preset_trim=allow_preset_trim
)
progress(0.18, desc="Loading PDI and PPI priors")
if uploaded_pdi or uploaded_ppi:
if not uploaded_pdi or not uploaded_ppi:
raise ValueError("To override the repository priors, upload both PDI and PPI files.")
pdi_frame = read_csv(uploaded_pdi)
ppi_frame = read_csv(uploaded_ppi)
else:
pdi_frame, ppi_frame = _upstream_interactions()
progress(0.38, desc="Constructing sparse biological masks")
gene_branch = build_biological_mask(list(gene_frame.columns), pdi_frame, ppi_frame)
dna_branch = build_biological_mask(list(dna_frame.columns), pdi_frame, ppi_frame)
embedding_file = _resolve_embedding_file(embedding_option, source_mode, example_dataset)
progress(0.53, desc="Retrieving GenePT embeddings")
embeddings = load_genept_embeddings(embedding_file)
progress(0.72, desc="Building enriched pathway connections")
gene_enrichment = attach_embeddings_and_pathways(
gene_branch, embeddings, frames["gene_pathways"], precomputed_significant=precomputed_significant
)
# Snapshot the unmatched pathway symbols before dna's own attach call
# mutates dna_branch.hidden_genes in place — presentation only, used
# by the "Show symbols" reveal.
dna_hidden_before = list(dna_branch.hidden_genes)
dna_enrichment = attach_embeddings_and_pathways(
dna_branch, embeddings, frames["dna_pathways"], precomputed_significant=precomputed_significant
)
dna_pathway_symbols = set(frames["dna_pathways"]["SYMBOL"].astype(str).str.strip())
unmatched_symbols = sorted(set(dna_hidden_before) - dna_pathway_symbols)
workspace = PreparedWorkspace(
gene_expression=gene_frame.to_numpy(dtype="float32"),
dna_methylation=dna_frame.to_numpy(dtype="float32"),
labels=labels,
label_names=label_names,
gene_branch=gene_branch,
dna_branch=dna_branch,
source_name=source_name,
warnings=warnings,
)
architecture = pd.DataFrame(
[
{
"branch": "Gene expression",
"samples": len(gene_frame),
"input genes": len(gene_branch.input_genes),
"PDI edges": gene_branch.pdi_edges,
"PPI edges": gene_branch.ppi_edges,
"hidden genes": len(gene_branch.hidden_genes),
"pathways": len(gene_branch.pathways),
"mask density": gene_branch.biological_mask.astype(bool).mean(),
},
{
"branch": "DNA methylation",
"samples": len(dna_frame),
"input genes": len(dna_branch.input_genes),
"PDI edges": dna_branch.pdi_edges,
"PPI edges": dna_branch.ppi_edges,
"hidden genes": len(dna_branch.hidden_genes),
"pathways": len(dna_branch.pathways),
"mask density": dna_branch.biological_mask.astype(bool).mean(),
},
]
)
enrichments = pd.concat(
[gene_enrichment.assign(branch="Gene expression"), dna_enrichment.assign(branch="DNA methylation")],
ignore_index=True,
)
final_len = len(gene_frame)
resolved_rows = []
for key in ("gene", "dna", "labels"):
raw = frames[key]
aligned = len(raw) == final_len
resolved_rows.append(
{
"File": file_labels[key],
"Shape": f"{len(raw):,} × {raw.shape[1]}",
"Rows matched": f"{final_len:,}",
"Status": "Aligned" if aligned else "Trimmed",
}
)
for key, branch in (("gene_pathways", gene_branch), ("dna_pathways", dna_branch)):
raw = frames[key]
resolved_rows.append(
{
"File": file_labels[key],
"Shape": f"{len(raw):,} × {raw.shape[1]}",
"Rows matched": f"{len(branch.hidden_genes):,}",
"Status": "Enriched",
}
)
resolved_files = pd.DataFrame(resolved_rows)
warning_html = (
"<br>" + " · ".join(html.escape(item) for item in warnings) if warnings else ""
)
summary = ui.simple_status_html(
f"<strong>Architecture ready.</strong> {len(gene_frame):,} paired samples · "
f"{len(label_names)} classes · {len(gene_branch.pathways) + len(dna_branch.pathways):,} "
f"branch-specific pathways · GenePT: {html.escape(embedding_file)}{warning_html}"
)
note = (
'<div class="actbar-note">'
+ ui.esc(
f"{len(resolved_rows)} of {len(resolved_rows)} inputs resolved"
+ (f" · {len(warnings)} warning(s)" if warnings else "")
+ f" · architecture built {datetime.now().strftime('%H:%M:%S')}"
)
+ "</div>"
)
run_meta = {
"source_name": source_name,
"embedding_file": embedding_file,
"unmatched_pathway_symbols": unmatched_symbols,
}
mean_density = float(
np.mean(
[
gene_branch.biological_mask.astype(bool).mean(),
dna_branch.biological_mask.astype(bool).mean(),
]
)
)
genept_dim = gene_branch.embeddings.shape[1] if gene_branch.embeddings is not None else 0
graph_preview_html = _graph_preview_html(
f"{mean_density * 100:.1f}%",
f"{len(gene_branch.hidden_genes) + len(dna_branch.hidden_genes):,}",
f"{len(gene_branch.pathways) + len(dna_branch.pathways):,}",
f"{genept_dim:,}",
)
progress(1.0, desc="Ready to train")
return workspace, summary, architecture, enrichments.head(100), resolved_files, run_meta, note, graph_preview_html
except Exception as exc:
fallback_note = '<div class="actbar-note">' + ui.esc(
"Choose a source, then build the biological architecture."
) + "</div>"
empty_graph_preview = _graph_preview_html("—", "—", "—", "—")
return None, _status(str(exc), error=True), empty, empty, empty, {}, fallback_note, empty_graph_preview
def _reveal_symbols(run_meta: dict):
symbols = (run_meta or {}).get("unmatched_pathway_symbols") or []
if not symbols:
return gr.update(visible=True, value=ui.empty_note_html("No unmatched pathway symbols were recorded."))
shown = ", ".join(symbols[:40])
more = f" … and {len(symbols) - 40:,} more" if len(symbols) > 40 else ""
return gr.update(visible=True, value=f'<div class="num" style="word-break:break-word">{html.escape(shown)}{more}</div>')
# ═══════════════════════════════════════════════════════════════════════
# Train Model
# ═══════════════════════════════════════════════════════════════════════
def _loss_plot(history: list[dict[str, float]]) -> go.Figure:
frame = pd.DataFrame(history)
figure = go.Figure()
figure.add_trace(
go.Scatter(
x=frame["epoch"], y=frame["training_loss"], mode="lines", name="Training",
line={"color": "#5980a6", "width": 2.5},
)
)
figure.add_trace(
go.Scatter(
x=frame["epoch"], y=frame["validation_loss"], mode="lines", name="Validation",
line={"color": "#2c455d", "width": 2.5},
)
)
figure.update_layout(
template="plotly_white",
paper_bgcolor="#f2f2f3",
plot_bgcolor="#f2f2f3",
font={"family": "Barlow, sans-serif", "color": "#1d1f20"},
margin={"l": 40, "r": 15, "t": 15, "b": 35},
legend={"orientation": "h", "y": 1.12},
xaxis={"title": "Epoch", "gridcolor": "rgba(29,31,32,.1)"},
yaxis={"title": "Cross-entropy loss", "gridcolor": "rgba(29,31,32,.1)"},
)
return figure
def estimate_training_duration(
workspace: PreparedWorkspace | None,
epochs: int,
batch_size: int,
learning_rate: float,
weight_decay: float,
dropout: float,
projection_dim: int,
fusion_dim: int,
validation_fraction: float,
optimizer: str,
class_weighting: bool,
progress=None,
) -> int:
"""Estimate a conservative ZeroGPU reservation from the prepared graph.
ZeroGPU checks the declared duration against each visitor's remaining quota
before the call starts. Keep small jobs short for better queue priority and
cap a single free-tier training request at five minutes.
"""
if workspace is None:
return 10
samples = max(int(len(workspace.labels)), 1)
biological_parameters = int(workspace.gene_branch.biological_mask.size) + int(
workspace.dna_branch.biological_mask.size
)
sample_factor = max(samples / 875.0, 0.25)
graph_factor = max(math.sqrt(biological_parameters / 1_850_000.0), 0.3)
batch_factor = max((16.0 / max(int(batch_size), 1)) ** 0.35, 0.55)
seconds = 25 + int(epochs) * 0.8 * sample_factor * graph_factor * batch_factor
return int(min(300, max(30, math.ceil(seconds))))
def _architecture_audit(bundle: ModelBundle) -> pd.DataFrame:
"""Read-only introspection of the trained model's own tensors — no change
to biolmnet.model, just reporting the shapes/sparsity it already has."""
model = bundle.model
rows = []
def add(name: str, tensor: torch.Tensor, nonzero: int | None = None):
total = tensor.numel()
nnz = int((tensor != 0).sum().item()) if nonzero is None else nonzero
rows.append(
{
"Layer": name,
"Shape": " × ".join(str(d) for d in tensor.shape),
"Nonzero": nnz,
"Density": nnz / total if total else 0.0,
}
)
for branch_name, branch in (("gene", model.gene_branch), ("dna", model.dna_branch)):
add(f"{branch_name} · biological mask", branch.biological.mask)
pw_mask = branch.pathway_attention.pathway_mask
add(f"{branch_name} · pathway attention", pw_mask, nonzero=int(pw_mask.sum().item()))
add(f"{branch_name} · branch projection", branch.projection.weight, nonzero=branch.projection.weight.numel())
add("fusion", model.fusion.weight, nonzero=model.fusion.weight.numel())
add("classifier", model.output.weight, nonzero=model.output.weight.numel())
return pd.DataFrame(rows)
def _macro_auc(validation_predictions: pd.DataFrame, label_names: list[str]) -> float | None:
"""Macro-averaged one-vs-rest AUC, computed here from the stored validation
probabilities — presentation-layer only; does not change what train()
computes or returns."""
try:
probability_columns = [f"P({label})" for label in label_names]
y_true = validation_predictions["observed"]
y_score = validation_predictions[probability_columns].to_numpy()
y_true_indices = y_true.map({label: index for index, label in enumerate(label_names)}).to_numpy()
if len(label_names) < 2 or len(set(y_true_indices)) < 2:
return None
return float(
roc_auc_score(y_true_indices, y_score, multi_class="ovr", average="macro", labels=list(range(len(label_names))))
)
except Exception:
return None
@spaces.GPU(duration=estimate_training_duration)
def train_workspace(
workspace: PreparedWorkspace | None,
epochs: int,
batch_size: int,
learning_rate: float,
weight_decay: float,
dropout: float,
projection_dim: int,
fusion_dim: int,
validation_fraction: float,
optimizer: str,
class_weighting: bool,
progress=gr.Progress(track_tqdm=False),
):
"""Signature intentionally mirrors ``estimate_training_duration`` exactly
(both take the same positional hyperparameters) — ``@spaces.GPU`` calls
the duration estimator with the same *args it received here, so the two
argument lists must stay aligned. Presentation-only metadata (source
dataset, elapsed time, …) is threaded through a separate ``gr.State``
merge step in the Blocks wiring instead of a new parameter here."""
empty = pd.DataFrame()
run_meta_update: dict = {}
if workspace is None:
return (
None, _status("Prepare data and priors in Phase 1 before training.", True), None, empty, None,
empty, empty, run_meta_update,
)
try:
parameters = Hyperparameters(
epochs=int(epochs), batch_size=int(batch_size), learning_rate=float(learning_rate),
weight_decay=float(weight_decay), dropout=float(dropout), projection_dim=int(projection_dim),
fusion_dim=int(fusion_dim), validation_fraction=float(validation_fraction), optimizer=optimizer,
class_weighting=bool(class_weighting),
)
def report(fraction: float, description: str) -> None:
progress(fraction, desc=description)
started = time.monotonic()
result = train(workspace, parameters, progress=report)
elapsed_seconds = time.monotonic() - started
bundle = result.bundle
artifact_path = save_bundle(bundle)
metrics = bundle.metrics
macro_auc = _macro_auc(result.validation_predictions, bundle.label_names)
minutes, seconds = divmod(int(elapsed_seconds), 60)
metrics_html = ui.simple_status_html(
f"<strong>Training complete.</strong> Validation accuracy {metrics['accuracy']:.3f} · "
f"elapsed {minutes:02d}:{seconds:02d}. Best validation checkpoint restored; the "
"downloadable artifact includes architecture, preprocessing, weights and metrics."
)
loss_figure = _loss_plot(bundle.history)
importance = pathway_importance(bundle).head(100)
architecture_audit = _architecture_audit(bundle)
confusion_html = ui.blueprint_div(
f'<div class="mono" style="margin-bottom:12px">Confusion matrix · validation split</div>'
+ ui.confusion_matrix_html(result.confusion.tolist(), bundle.label_names)
)
run_meta_update = {
"elapsed": f"{minutes:02d}:{seconds:02d}",
"finished_at": datetime.now().strftime("%H:%M:%S"),
"macro_auc": macro_auc,
"session_id": secrets.token_hex(3),
"train_loss": bundle.history[-1]["training_loss"],
}
return (
bundle, metrics_html, loss_figure, result.validation_predictions, artifact_path,
importance, architecture_audit, run_meta_update, confusion_html,
)
except Exception as exc:
return (
None, _status(str(exc), error=True), None, empty, None, empty, empty, run_meta_update,
ui.empty_note_html("No confusion matrix yet."),
)
# ═══════════════════════════════════════════════════════════════════════
# Export Artifacts
# ═══════════════════════════════════════════════════════════════════════
def _export_panel(bundle: ModelBundle | None, artifact_path: str | None, run_meta: dict):
run_meta = run_meta or {}
if not bundle or not artifact_path or not Path(artifact_path).exists():
empty_manifest = ui.table_html(["Entry", "Format", "Size", "Reproduces"], [])
empty_note = ui.empty_note_html("No trained artifact yet. Train a model to populate this page.")
stats = "".join(ui.stat_plate("—", label) for label in ("Gene features", "DNA features", "Pathways", "Classes"))
return empty_manifest, empty_note, "", stats, gr.update(value=None)
archive = Path(artifact_path)
with zipfile.ZipFile(archive) as zf:
infos = {info.filename: info for info in zf.infolist()}
reproduces = {
"config.json": "Hyperparameters &amp; architecture",
"arrays.npz": "Biological masks, embeddings &amp; scalers",
"model.safetensors": "Trained weights",
"metrics.json": "Metrics &amp; training history",
}
rows = []
for name in sorted(EXPECTED_FILES):
info = infos.get(name)
size = info.file_size if info else 0
rows.append(
[
html.escape(name),
Path(name).suffix.lstrip(".") or "—",
ui.human_bytes(size),
reproduces.get(name, "—"),
]
)
manifest = ui.table_html(
["Entry", "Format", "Size", "Reproduces"], rows, aligns=["left", "left", "left", "right"],
numeric_cols=[False, True, True, False],
)
sha256 = hashlib.sha256(archive.read_bytes()).hexdigest()
total_size = archive.stat().st_size
provenance = "".join(
[
f'<div class="srow-row" style="display:grid;grid-template-columns:196px minmax(0,1fr);gap:20px;padding:10px 0" >'
f'<div class="srow-label" style="padding:0">Checksum · sha256</div>'
f'<div class="num" style="word-break:break-all">{sha256}</div></div>',
f'<div class="srow-row srow-row-rule" style="display:grid;grid-template-columns:196px minmax(0,1fr);gap:20px;padding:10px 0">'
f'<div class="srow-label" style="padding:0">Torch / Python</div>'
f'<div class="num">{html.escape(torch.__version__)} / {html.escape(platform.python_version())}</div></div>',
f'<div class="srow-row srow-row-rule" style="display:grid;grid-template-columns:196px minmax(0,1fr);gap:20px;padding:10px 0">'
f'<div class="srow-label" style="padding:0">Upstream repository</div>'
f'<div class="num">{html.escape(UPSTREAM_REPOSITORY)}</div></div>',
f'<div class="srow-row srow-row-rule" style="display:grid;grid-template-columns:196px minmax(0,1fr);gap:20px;padding:10px 0;border-bottom:1px solid var(--color-row-rule)">'
f'<div class="srow-label" style="padding:0">GenePT embedding</div>'
f'<div class="num">{html.escape(run_meta.get("embedding_file", "—"))}</div></div>',
]
)
stats = "".join(
[
ui.stat_plate(f"{len(bundle.gene_features):,}", "Gene features"),
ui.stat_plate(f"{len(bundle.dna_features):,}", "DNA features"),
ui.stat_plate(
f"{len(bundle.config['gene_pathways']) + len(bundle.config['dna_pathways']):,}", "Pathways"
),
ui.stat_plate(f"{len(bundle.label_names):,}", "Classes"),
]
)
artifact_panel = ui.blueprint_div(
'<div class="mono" style="font-size:9px">Artifact</div>'
f'<div style="font:600 26px/1.1 var(--font-heading);margin-top:6px">Trained bundle ready</div>'
f'<div style="margin-top:8px;font-size:12.5px;color:var(--color-muted)">'
f'Produced by the run that finished at {html.escape(run_meta.get("finished_at", "—"))} with '
f'validation accuracy {bundle.metrics["accuracy"]:.3f} over {len(bundle.label_names)} classes. '
f'Bundle size {ui.human_bytes(total_size)}.</div>',
extra_class="",
)
return manifest, provenance, artifact_panel, stats, gr.update(value=artifact_path)
# ═══════════════════════════════════════════════════════════════════════
# Predict
# ═══════════════════════════════════════════════════════════════════════
def _resolve_predict_bundle(model_state, artifact_path, use_upload: bool):
if use_upload and artifact_path:
return load_bundle(artifact_path)
return model_state
def _workspace_prediction_frames(workspace: PreparedWorkspace) -> tuple[pd.DataFrame, pd.DataFrame]:
return (
pd.DataFrame(workspace.gene_expression, columns=workspace.gene_branch.input_genes),
pd.DataFrame(workspace.dna_methylation, columns=workspace.dna_branch.input_genes),
)
def refresh_alignment(model_state, workspace, artifact_path, use_upload, use_prepared, gene_path, dna_path):
"""Presentation-only pre-flight check: mirrors the two structural checks
``validate_prediction_frames`` performs (row-count match, required
columns present) purely to render the alignment table live. The actual
gating decision on submit still goes through the real, untouched
``predict()`` call."""
align = {"ok": None, "required": None, "matched": None, "source_label": "Prepared dataset" if use_prepared else "Session model"}
try:
bundle = _resolve_predict_bundle(model_state, artifact_path, use_upload)
except Exception as exc:
rows = [[html.escape("Artifact"), "—", "—", f'<span class="error">{html.escape(str(exc))}</span>']]
table = ui.table_html(["Check", "Artifact", "Uploaded", "Result"], rows, aligns=["left", "left", "left", "right"])
return table, gr.update(visible=False), gr.update(value="Run inference", interactive=False, elem_classes=["actbar-primary-btn"]), align
if use_upload:
align["source_label"] = "Uploaded artifact" if artifact_path else "Upload artifact"
if bundle is None:
table = ui.table_html(
["Check", "Artifact", "Uploaded", "Result"],
[["Trained model available", "—", "—", '<span class="error">Fail</span>']],
aligns=["left", "left", "left", "right"],
)
return table, gr.update(visible=False), gr.update(value="Run inference", interactive=False, elem_classes=["actbar-primary-btn"]), align
align["required"] = len(bundle.gene_features) + len(bundle.dna_features)
rows = []
ok = True
gene_frame = dna_frame = None
if use_prepared:
if workspace is None:
rows.append(["Prepared dataset available", "—", "—", '<span class="error">Fail</span>'])
ok = False
else:
gene_frame, dna_frame = _workspace_prediction_frames(workspace)
rows.append(["Prepared dataset available", "—", f"{len(gene_frame):,} samples", _pass_fail(True)])
elif gene_path and dna_path:
try:
gene_frame = read_csv(gene_path)
dna_frame = read_csv(dna_path)
except Exception as exc:
rows.append(["Reading uploaded files", "—", "—", f'<span class="error">{html.escape(str(exc))}</span>'])
ok = False
if gene_frame is not None and dna_frame is not None:
rows_match = len(gene_frame) == len(dna_frame)
ok = ok and rows_match
rows.append(
[
"Samples paired across matrices", "—", f"{len(gene_frame):,} / {len(dna_frame):,}",
_pass_fail(rows_match),
]
)
missing_gene = sorted(set(bundle.gene_features) - set(gene_frame.columns))
missing_dna = sorted(set(bundle.dna_features) - set(dna_frame.columns))
gene_ok = not missing_gene
dna_ok = not missing_dna
ok = ok and gene_ok and dna_ok
rows.append(["Expression columns", f"{len(bundle.gene_features):,}", f"{len(gene_frame.columns):,}", _pass_fail(gene_ok)])
rows.append(["Methylation columns", f"{len(bundle.dna_features):,}", f"{len(dna_frame.columns):,}", _pass_fail(dna_ok)])
matched = align["required"] - len(missing_gene) - len(missing_dna)
align["matched"] = matched
align["missing_gene"] = missing_gene
align["missing_dna"] = missing_dna
else:
rows.append(["Gene expression uploaded", "—", "—", _pass_fail(gene_path is not None)])
rows.append(["DNA methylation uploaded", "—", "—", _pass_fail(dna_path is not None)])
ok = False
align["ok"] = ok
table = ui.table_html(["Check", "Artifact", "Uploaded", "Result"], rows, aligns=["left", "left", "left", "right"])
missing_total = len(align.get("missing_gene", [])) + len(align.get("missing_dna", []))
strip_visible = not ok and (gene_frame is not None and dna_frame is not None)
return table, gr.update(visible=strip_visible, value=(
ui.strip_text_html("Blocking · features missing", f"{missing_total:,} required column(s) are missing from the uploaded files.")
if strip_visible else ""
)), gr.update(value="Run inference", interactive=bool(ok), elem_classes=["actbar-primary-btn"]), align
def _pass_fail(ok: bool) -> str:
return '<span class="accent">Pass</span>' if ok else '<span class="error">Fail</span>'
def _list_missing(align: dict):
missing = list((align or {}).get("missing_gene", [])) + list((align or {}).get("missing_dna", []))
if not missing:
return gr.update(visible=True, value=ui.empty_note_html("Nothing missing."))
shown = ", ".join(missing[:40])
more = f" … and {len(missing) - 40:,} more" if len(missing) > 40 else ""
return gr.update(visible=True, value=f'<div class="num" style="word-break:break-word">{html.escape(shown)}{more}</div>')
def run_prediction(
bundle: ModelBundle | None,
workspace: PreparedWorkspace | None,
uploaded_artifact: str | None,
use_upload: bool,
use_prepared: bool,
gene_file: str | None,
dna_file: str | None,
predict_meta: dict,
):
predict_meta = dict(predict_meta or {})
try:
active_bundle = load_bundle(uploaded_artifact) if (use_upload and uploaded_artifact) else bundle
if active_bundle is None:
raise ValueError("Train a model in Phase 2 or upload a BioLM-NET model artifact.")
if use_prepared:
if workspace is None:
raise ValueError("Prepare data and priors before scoring the prepared dataset.")
gene_frame, dna_frame = _workspace_prediction_frames(workspace)
else:
gene_frame = _read_required_upload(gene_file, "a prediction gene-expression CSV")
dna_frame = _read_required_upload(dna_file, "a prediction DNA-methylation CSV")
output = predict(gene_frame, dna_frame, active_bundle)
destination = Path(tempfile.mkdtemp(prefix="biolmnet-prediction-")) / "biolm-net-predictions.csv"
output.to_csv(destination, index=False)
counts = output["predicted_class"].value_counts()
total = int(counts.sum())
bars = "".join(
ui.labeled_bar_row(str(label), (count / total) * 100 if total else 0, f"{count:,}")
for label, count in counts.items()
)
distribution_panel = ui.blueprint_div(
f'<div class="mono" style="font-size:9px;margin-bottom:12px">Class distribution · {total:,} samples</div>'
f'<div style="display:flex;flex-direction:column;gap:9px">{bars}</div>'
)
status = _status(f"Predicted {len(output):,} samples. Mean confidence: {output['confidence'].mean():.3f}.")
predict_meta.update(
{
"n_samples": total,
"mean_confidence": float(output["confidence"].mean()),
"predicted_at": datetime.now().strftime("%H:%M:%S"),
}
)
return active_bundle, status, output, distribution_panel, str(destination), predict_meta
except Exception as exc:
return bundle, _status(str(exc), True), pd.DataFrame(), ui.empty_note_html("No predictions yet."), None, predict_meta
# ═══════════════════════════════════════════════════════════════════════
# Results
# ═══════════════════════════════════════════════════════════════════════
def _pathway_attention_table(bundle: ModelBundle) -> pd.DataFrame:
importance = pathway_importance(bundle).head(20)
model = bundle.model
gene_counts = model.gene_branch.pathway_attention.pathway_mask.sum(axis=0).cpu().numpy()
dna_counts = model.dna_branch.pathway_attention.pathway_mask.sum(axis=0).cpu().numpy()
gene_index = {pathway: int(count) for pathway, count in zip(bundle.config["gene_pathways"], gene_counts)}
dna_index = {pathway: int(count) for pathway, count in zip(bundle.config["dna_pathways"], dna_counts)}
def gene_count(row):
table = gene_index if row["branch"] == "Gene expression" else dna_index
return table.get(row["pathway"], 0)
importance = importance.copy()
importance["genes"] = importance.apply(gene_count, axis=1)
return importance[["branch", "pathway", "genes", "peak_gene_attention", "attention_entropy"]]
def refresh_results(bundle: ModelBundle | None, run_meta: dict):
run_meta = run_meta or {}
if bundle is None:
stats = (
'<div class="metric-strip">'
+ "".join(ui.stat_plate("—", label) for label in ("Validation accuracy", "Macro F1", "Macro AUC", "Retained pathways"))
+ "</div>"
)
return (
stats,
gr.update(visible=False),
ui.empty_note_html(
"Build the biological architecture, train a model, or run prediction to populate this page."
),
gr.update(visible=False),
)
macro_auc = run_meta.get("macro_auc")
retained = len(bundle.config["gene_pathways"]) + len(bundle.config["dna_pathways"])
stats = (
'<div class="metric-strip">'
+ "".join(
[
ui.stat_plate(f"{bundle.metrics['accuracy']:.3f}", "Validation accuracy"),
ui.stat_plate(f"{bundle.metrics['f1_macro']:.3f}", "Macro F1"),
ui.stat_plate(f"{macro_auc:.3f}" if macro_auc is not None else "—", "Macro AUC"),
ui.stat_plate(f"{retained:,}", "Retained pathways"),
]
)
+ "</div>"
)
return stats, gr.update(visible=True), "", gr.update(visible=True)
THEME = gr.themes.Base(primary_hue="gray", neutral_hue="gray")
# ═══════════════════════════════════════════════════════════════════════
# Blocks
# ═══════════════════════════════════════════════════════════════════════
with gr.Blocks(title="BioLM-NET Workbench", fill_width=True) as demo:
workspace_state = gr.State(None)
model_state = gr.State(None)
artifact_path_state = gr.State(None)
run_meta_state = gr.State({})
align_state = gr.State({})
predict_meta_state = gr.State({})
active_page_state = gr.State("data")
session_id_state = gr.State("")
intro_step_state = gr.State(0)
intro_dismissed_state = gr.BrowserState(False, storage_key="biolmnet_intro_dismissed")
with gr.Row(elem_classes=["app-shell"]):
with gr.Column(elem_classes=["rail-col"]):
with gr.Column(elem_classes=["rail"]):
gr.HTML(ui.brand_block_html())
rail_row_htmls = []
rail_row_buttons = []
with gr.Column(elem_classes=["rail-nav"]):
for index, label in enumerate(STAGE_LABELS):
with gr.Column(elem_classes=["rail-row-wrap"]):
row_html = gr.HTML(ui.rail_row_html(index + 1, label, "on" if index == 0 else "off", "Active" if index == 0 else "Locked"))
row_button = gr.Button("", elem_classes=["rail-row-click"])
rail_row_htmls.append(row_html)
rail_row_buttons.append(row_button)
with gr.Column(elem_classes=["rail-spacer", "rail-plate-wrap"]):
gr.HTML('<div class="mono" style="margin-bottom:9px">Run state</div>')
run_state_plate_html = gr.HTML(ui.run_state_plate(_run_state_rows("data", None, None, {}, {}, {})))
gr.HTML(ui.footnote_html("Research use only", "ZeroGPU on demand · py 3.12"))
with gr.Column(elem_classes=["workspace-col"]):
with gr.Row(elem_classes=["topbar-host"]):
topbar_html_component = gr.HTML(ui.topbar_html(1, STAGE_TOTAL, STAGE_LABELS[0], "——————", "ZeroGPU idle"))
intro_open_button = gr.Button("Introduction", size="sm", variant="secondary", elem_classes=["btn-ghost", "intro-open-btn"])
# ── Data & Priors ────────────────────────────────────────
with gr.Column(visible=True) as data_page:
with gr.Row(elem_classes=["pghd"]):
gr.HTML(
ui.title_block_html(
"Data & Priors",
"Resolve paired omics and pathway inputs, then assemble the biologically masked "
"graph that training and prediction both reuse.",
)
)
with gr.Row(elem_classes=["stage-grid"], equal_height=False):
with gr.Column(scale=7):
with gr.Row(elem_classes=["section-heading-row"]):
gr.HTML('<h4 style="font-size:16px">Dataset source</h4>')
gr.HTML('<div class="mono" style="font-size:9.5px;text-align:right">Samples in rows · HGNC symbols in columns</div>')
source_mode = gr.Radio(
["BioLM-NET examples", "GitHub folder", "Upload files"],
value="BioLM-NET examples", show_label=False, container=False,
elem_classes=["seg-radio"],
)
with gr.Row(elem_classes=["srow-row", "first"]):
with gr.Column(scale=1, min_width=0):
with gr.Column(visible=True) as example_group:
with gr.Row(elem_classes=["field-pair"]):
with gr.Column(min_width=0):
gr.HTML('<div class="field-label">Repository dataset</div>')
example_dataset = gr.Dropdown(
["BRCA", "COAD", "GBM", "scTrioseq2"], value="BRCA",
show_label=False, container=False,
)
with gr.Column(visible=False) as github_group:
github_folder = gr.Textbox(
label="GitHub dataset folder",
placeholder="https://github.com/owner/repo/tree/main/Dataset/BRCA",
info="The folder must contain the five standard BioLM-NET CSV filenames.",
)
with gr.Column(visible=False) as upload_group:
with gr.Row():
uploaded_gene = gr.File(label="Gene expression", file_types=[".csv"], type="filepath", elem_classes=["file-slot"])
uploaded_dna = gr.File(label="DNA methylation", file_types=[".csv"], type="filepath", elem_classes=["file-slot"])
uploaded_labels = gr.File(label="Labels", file_types=[".csv"], type="filepath", elem_classes=["file-slot"])
with gr.Row():
uploaded_gene_pathway = gr.File(label="Gene → pathway mapping", file_types=[".csv"], type="filepath", elem_classes=["file-slot"])
uploaded_dna_pathway = gr.File(label="DNA → pathway mapping", file_types=[".csv"], type="filepath", elem_classes=["file-slot"])
gr.HTML('<div class="mono sect sect-spaced">Resolved files</div>')
resolved_files_table = gr.Dataframe(
headers=["File", "Shape", "Rows matched", "Status"], interactive=False, wrap=False,
value=pd.DataFrame(),
)
with gr.Row(elem_classes=["strip", "error"], visible=False) as unmatched_strip:
unmatched_note = gr.HTML("")
show_symbols_btn = gr.Button("Show symbols", size="sm", elem_classes=["btn-ghost"])
unmatched_reveal = gr.HTML(visible=False)
with gr.Column(scale=5, min_width=0):
gr.HTML('<div class="mono" style="margin-bottom:2px">Priors</div>')
with gr.Row(elem_classes=["srow-row", "first"]):
with gr.Column(scale=196, min_width=140):
gr.HTML(ui.esc("Pathway files already significant"))
with gr.Column(scale=300, min_width=0):
pathway_files_are_significant = gr.Checkbox(
value=True, label="Yes — skip enrichment", container=False,
)
gr.HTML(
'<div class="mono" style="font-size:9px;text-transform:none;'
'letter-spacing:0;margin-top:2px">Turn off for a full SYMBOL/PathwayID '
'annotation catalog; enrichment will use BH-adjusted p &lt; 0.05.</div>'
)
with gr.Row(elem_classes=["srow-row"]):
with gr.Column(scale=196, min_width=140):
gr.HTML(ui.esc("Enrichment cutoff"))
with gr.Column(scale=300, min_width=0):
gr.HTML(
'<span class="num" style="opacity:.55">BH-adjusted p</span> '
'<span class="num">&lt; 0.05</span>'
)
with gr.Row(elem_classes=["srow-row"]):
with gr.Column(scale=196, min_width=140):
gr.HTML(ui.esc("GenePT context"))
with gr.Column(scale=300, min_width=0):
embedding_option = gr.Dropdown(
list(GENEPT_OPTIONS), value=list(GENEPT_OPTIONS)[0], show_label=False, container=False,
)
with gr.Row(elem_classes=["srow-row"]):
with gr.Column(scale=196, min_width=140):
gr.HTML(ui.esc("PDI source"))
with gr.Column(scale=300, min_width=0):
with gr.Row(elem_classes=["srow-value-row"]):
gr.HTML('<span class="num">DoRothEA</span>')
pdi_override_btn = gr.Button("Override", size="sm", elem_classes=["btn-ghost"])
uploaded_pdi = gr.File(
label="Custom PDI.csv (needs TF, Target columns)", file_types=[".csv"],
type="filepath", visible=False,
)
with gr.Row(elem_classes=["srow-row"]):
with gr.Column(scale=196, min_width=140):
gr.HTML(ui.esc("PPI source"))
with gr.Column(scale=300, min_width=0):
with gr.Row(elem_classes=["srow-value-row"]):
gr.HTML('<span class="num">STRING &gt; 0.7</span>')
ppi_override_btn = gr.Button("Override", size="sm", elem_classes=["btn-ghost"])
uploaded_ppi = gr.File(
label="Custom PPI.csv (needs protein1, protein2, combined_score)",
file_types=[".csv"], type="filepath", visible=False,
)
with gr.Row(elem_classes=["srow-row", "last"]):
with gr.Column(scale=196, min_width=140):
gr.HTML(ui.esc("PPI retention"))
with gr.Column(scale=300, min_width=0):
gr.HTML(
'<span class="num">Top decile</span> '
'<span class="mono" style="font-size:9px">as in paper</span>'
)
gr.HTML(
'<div class="mono" style="font-size:9px;text-transform:none;letter-spacing:0;'
'margin-top:8px">Upload both PDI and PPI to override — the repository priors are '
'used unless both files are present.</div>'
)
gr.HTML('<div class="mono" style="margin:24px 0 9px">Graph preview</div>')
graph_preview = gr.HTML(_graph_preview_html("—", "—", "—", "—"))
with gr.Row(visible=False):
architecture_table = gr.Dataframe(interactive=False, wrap=False, value=pd.DataFrame(), headers=[""], column_count=(1, "dynamic"))
enrichment_table = gr.Dataframe(interactive=False, wrap=False, value=pd.DataFrame(), headers=[""], column_count=(1, "dynamic"))
with gr.Row(elem_classes=["actbar-row"]):
data_actionbar_note = gr.HTML(f'<div class="actbar-note">{ui.esc("Choose a source, then build the biological architecture.")}</div>')
with gr.Row(elem_classes=["actbar-buttons"]):
revalidate_button = gr.Button(
"Re-validate inputs", size="sm", variant="secondary", elem_classes=["actbar-secondary-btn"]
)
prepare_button = gr.Button(
"Build biological architecture", size="sm", variant="primary", elem_classes=["actbar-primary-btn"]
)
preparation_status = gr.HTML(visible=False)
# ── Train Model ──────────────────────────────────────────
with gr.Column(visible=False) as train_page:
with gr.Row(elem_classes=["pghd"]):
gr.HTML(
ui.title_block_html(
"Train Model",
"Fit the prepared architecture with a stratified validation split and balanced "
"loss options.",
)
)
with gr.Row(elem_classes=["stage-grid"], equal_height=False):
with gr.Column(scale=7, min_width=0):
gr.HTML('<div class="mono sect">A · Optimisation</div>')
with gr.Row(elem_classes=["srow-row", "first"]):
with gr.Column(min_width=140):
gr.HTML(ui.esc("Epochs"))
with gr.Column():
epochs = gr.Slider(5, 200, value=50, step=5, show_label=False, container=False)
with gr.Row(elem_classes=["srow-row"]):
with gr.Column(min_width=140):
gr.HTML(ui.esc("Batch size"))
with gr.Column():
batch_size = gr.Dropdown([8, 16, 32, 64, 128], value=16, show_label=False, container=False)
with gr.Row(elem_classes=["srow-row"]):
with gr.Column(min_width=140):
gr.HTML(ui.esc("Learning rate"))
with gr.Column():
learning_rate = gr.Number(value=0.001, minimum=0.000001, show_label=False, container=False)
with gr.Row(elem_classes=["srow-row"]):
with gr.Column(min_width=140):
gr.HTML(ui.esc("L2 weight decay"))
with gr.Column():
weight_decay = gr.Number(value=0.01, minimum=0, show_label=False, container=False)
with gr.Row(elem_classes=["srow-row", "last"]):
with gr.Column(min_width=140):
gr.HTML(ui.esc("Dropout"))
with gr.Column():
dropout = gr.Slider(0, 0.8, value=0.3, step=0.05, show_label=False, container=False)
gr.HTML('<div class="mono sect sect-spaced">B · Validation &amp; loss</div>')
with gr.Row(elem_classes=["srow-row", "first"]):
with gr.Column(min_width=140):
gr.HTML(ui.esc("Branch projection"))
with gr.Column():
projection_dim = gr.Dropdown([16, 32, 64, 128], value=64, show_label=False, container=False)
with gr.Row(elem_classes=["srow-row"]):
with gr.Column(min_width=140):
gr.HTML(ui.esc("Fusion layer"))
with gr.Column():
fusion_dim = gr.Dropdown([8, 12, 16, 32], value=12, show_label=False, container=False)
with gr.Row(elem_classes=["srow-row"]):
with gr.Column(min_width=140):
gr.HTML(ui.esc("Validation fraction"))
with gr.Column():
validation_fraction = gr.Slider(0.1, 0.4, value=0.2, step=0.05, show_label=False, container=False)
with gr.Row(elem_classes=["srow-row"]):
with gr.Column(min_width=140):
gr.HTML(ui.esc("Optimizer"))
with gr.Column():
optimizer = gr.Radio(["Adam", "SGD"], value="Adam", show_label=False, container=False, elem_classes=["seg-radio"])
with gr.Row(elem_classes=["srow-row", "last"]):
with gr.Column(min_width=140):
gr.HTML(ui.esc("Class weighting"))
with gr.Column():
class_weighting = gr.Checkbox(
value=True, label="Balance classes in the loss",
info="Uses N / (classes × samples in class), as in the paper.",
)
with gr.Column(scale=5, min_width=0):
gr.HTML('<div class="mono" style="margin-bottom:9px">Training result</div>')
training_status = gr.HTML(_status("Prepare the architecture before training."))
loss_plot = gr.Plot(label="Loss by epoch", show_label=False)
gr.HTML('<div class="mono" style="margin:20px 0 9px">Paper-faithful defaults</div>')
gr.HTML(
"".join(
[
ui.kv_plain_row("First layer", "trainable W ⊙ M"),
ui.kv_plain_row("PDI weights", "binary"),
ui.kv_plain_row("PPI weights", "STRING norm."),
ui.kv_plain_row("Attention", "GenePT query"),
ui.kv_plain_row("Fusion", "dual → dense → softmax", last=True),
]
)
)
with gr.Row(elem_classes=["actbar-row"]):
train_actionbar_note = gr.HTML('<div class="actbar-note">Set hyperparameters, then train.</div>')
with gr.Row(elem_classes=["actbar-buttons"]):
train_button = gr.Button(
"Train BioLM-NET on ZeroGPU", size="sm", variant="primary", elem_classes=["actbar-primary-btn"]
)
# ── Export Artifacts ─────────────────────────────────────
with gr.Column(visible=False) as export_page:
with gr.Row(elem_classes=["pghd"]):
gr.HTML(
ui.title_block_html(
"Export Artifacts",
"One bundle carries the trained weights, the fitted preprocessing and the "
"architecture — enough to reproduce prediction without rebuilding the graph.",
)
)
with gr.Row(elem_classes=["stage-grid"], equal_height=False):
with gr.Column(scale=7, min_width=0):
gr.HTML('<div class="mono sect">Bundle contents</div>')
bundle_manifest = gr.HTML(ui.table_html(["Entry", "Format", "Size", "Reproduces"], []))
gr.HTML('<div class="mono sect sect-spaced">Provenance</div>')
provenance_html = gr.HTML("")
with gr.Column(scale=5, min_width=0):
artifact_panel = gr.HTML(ui.empty_note_html("No trained artifact yet. Train a model to generate the download."))
model_download = gr.File(label="Trained model artifact", interactive=False, elem_classes=["file-slot"])
gr.HTML('<div class="mono" style="margin:20px 0 9px">Graph carried in the bundle</div>')
export_stats = gr.HTML(
'<div class="stat-grid" style="display:grid;grid-template-columns:1fr 1fr;gap:10px">'
+ "".join(ui.stat_plate("—", label) for label in ("Gene features", "DNA features", "Pathways", "Classes"))
+ "</div>"
)
with gr.Row(elem_classes=["actbar-row"]):
gr.HTML('<div class="actbar-note">Bundle written to a temp file on export · not persisted between Space restarts</div>')
# ── Predict ──────────────────────────────────────────────
with gr.Column(visible=False) as predict_page:
with gr.Row(elem_classes=["pghd"]):
gr.HTML(
ui.title_block_html(
"Predict",
"Score new paired omics with the session model, or upload a previous bundle. "
"Features are aligned to the artifact before inference runs.",
)
)
predict_source_mode = gr.Radio(
["Session model", "Upload artifact"], value="Session model", show_label=False,
container=False, elem_classes=["seg-radio"],
)
with gr.Row(elem_classes=["stage-grid"], equal_height=False):
with gr.Column(scale=6, min_width=0):
gr.HTML('<div class="mono sect">Inputs</div>')
prediction_input_mode = gr.Radio(
["Prepared dataset", "Upload files"], value="Prepared dataset", show_label=False,
container=False, elem_classes=["seg-radio"],
)
with gr.Row(elem_classes=["srow-row", "first"], visible=True) as session_model_row:
with gr.Column(min_width=100):
gr.HTML(ui.esc("Trained artifact"))
with gr.Column():
session_model_note = gr.HTML('<div class="num" style="opacity:.55">No session model yet — train one, or switch to Upload artifact.</div>')
with gr.Row(elem_classes=["srow-row", "first"], visible=False) as artifact_upload_row:
with gr.Column(min_width=100):
gr.HTML(ui.esc("Trained artifact"))
with gr.Column():
prediction_artifact = gr.File(label="", show_label=False, file_types=[".zip"], type="filepath", elem_classes=["file-slot"])
with gr.Row(elem_classes=["srow-row"]) as prepared_dataset_row:
with gr.Column(min_width=100):
gr.HTML(ui.esc("Prediction cohort"))
with gr.Column():
gr.HTML('<div class="num" style="opacity:.65">Use the dataset prepared in Data &amp; Priors.</div>')
with gr.Row(elem_classes=["srow-row"], visible=False) as prediction_gene_row:
with gr.Column(min_width=100):
gr.HTML(ui.esc("Gene expression"))
with gr.Column():
prediction_gene = gr.File(label="", show_label=False, file_types=[".csv"], type="filepath", elem_classes=["file-slot"])
with gr.Row(elem_classes=["srow-row", "last"], visible=False) as prediction_dna_row:
with gr.Column(min_width=100):
gr.HTML(ui.esc("DNA methylation"))
with gr.Column():
prediction_dna = gr.File(label="", show_label=False, file_types=[".csv"], type="filepath", elem_classes=["file-slot"])
gr.HTML('<div class="mono sect sect-spaced">Feature alignment</div>')
alignment_table = gr.HTML(ui.table_html(["Check", "Artifact", "Uploaded", "Result"], []))
with gr.Row(elem_classes=["strip", "error"], visible=False) as alignment_strip:
alignment_note = gr.HTML("")
list_missing_btn = gr.Button("List missing", size="sm", elem_classes=["btn-ghost"])
missing_reveal = gr.HTML(visible=False)
with gr.Column(scale=6, min_width=0):
gr.HTML('<div class="mono" style="margin-bottom:9px">Predictions · last successful run</div>')
distribution_panel = gr.HTML(ui.empty_note_html("No predictions yet."))
prediction_table = gr.Dataframe(interactive=False, wrap=False, value=pd.DataFrame(), headers=[""], column_count=(1, "dynamic"))
prediction_download = gr.File(label="Prediction CSV", interactive=False, elem_classes=["file-slot", "btn-block"])
with gr.Row(elem_classes=["actbar-row"]):
predict_actionbar_note = gr.HTML('<div class="actbar-note">Use the current trained model or upload an artifact.</div>')
with gr.Row(elem_classes=["actbar-buttons"]):
predict_button = gr.Button(
"Run inference", size="sm", variant="primary", interactive=False,
elem_classes=["actbar-primary-btn"],
)
prediction_status = gr.HTML(visible=False)
# ── Results ──────────────────────────────────────────────
with gr.Column(visible=False) as results_page:
with gr.Row(elem_classes=["pghd"]):
gr.HTML(
ui.title_block_html(
"Results",
"Everything the run produced, read in one place: architecture audit, "
"validation outputs, pathway attention and the scored cohort.",
)
)
results_empty_note = gr.HTML(ui.empty_note_html("Build the biological architecture, train a model, or run prediction to populate this page."))
with gr.Row(elem_classes=["results-stat-row"], visible=False) as results_stats_row:
results_stats = gr.HTML("")
with gr.Row(elem_classes=["results-grid"], equal_height=False, visible=False) as results_content:
with gr.Column(scale=6, min_width=0):
gr.HTML('<div class="mono sect">Training history</div>')
results_loss_plot = gr.Plot(show_label=False)
gr.HTML('<div class="mono sect sect-spaced">Sparse architecture audit</div>')
results_architecture_table = gr.Dataframe(interactive=False, wrap=False, value=pd.DataFrame(), headers=[""], column_count=(1, "dynamic"))
with gr.Column(scale=6, min_width=0):
results_confusion = gr.HTML("")
gr.HTML('<div class="mono sect sect-spaced">Pathway attention · top retained</div>')
results_pathway_table = gr.Dataframe(interactive=False, wrap=False, value=pd.DataFrame(), headers=[""], column_count=(1, "dynamic"))
with gr.Column(visible=False) as results_tables_bottom:
gr.HTML('<div class="mono sect sect-spaced">Validation predictions</div>')
results_validation_table = gr.Dataframe(interactive=False, wrap=False, value=pd.DataFrame(), headers=[""], column_count=(1, "dynamic"))
gr.HTML('<div class="mono sect sect-spaced">Predictions and class probabilities</div>')
results_prediction_table = gr.Dataframe(interactive=False, wrap=False, value=pd.DataFrame(), headers=[""], column_count=(1, "dynamic"))
with gr.Row(elem_classes=["actbar-row"]):
gr.HTML(
'<div class="actbar-note">Validate cohorts, preprocessing and performance before '
"drawing biological or clinical conclusions.</div>"
)
with gr.Group(elem_id="intro", visible=True) as intro_group:
intro_cards = []
intro_next_buttons = []
intro_back_buttons = []
intro_skip_buttons = []
intro_start_buttons = []
with gr.Column(elem_classes=["intro-card"]):
for index, (_, _, _, _, note) in enumerate(ui.INTRO_CARDS):
with gr.Column(visible=index == 0, elem_classes=["intro-card-page"]) as intro_card:
gr.HTML(ui.intro_card_html(index))
with gr.Row(elem_classes=["intro-card-foot"]):
if index == 0:
intro_dont_show = gr.Checkbox(value=False, label="Don't show again", container=False)
else:
gr.HTML(f'<span class="mono" style="font-size:9.5px">{ui.esc(note)}</span>')
with gr.Row(elem_classes=["intro-card-actions"]):
gr.HTML(ui.intro_dots_html(index))
if index == 0:
skip_btn = gr.Button("Skip", size="sm", variant="secondary", elem_classes=["btn-ghost"])
next_btn = gr.Button("Next", size="sm", variant="primary")
intro_skip_buttons.append(skip_btn)
intro_next_buttons.append(next_btn)
elif index < len(ui.INTRO_CARDS) - 1:
back_btn = gr.Button("Back", size="sm", variant="secondary", elem_classes=["btn-ghost"])
next_btn = gr.Button("Next", size="sm", variant="primary")
intro_back_buttons.append(back_btn)
intro_next_buttons.append(next_btn)
else:
back_btn = gr.Button("Back", size="sm", variant="secondary", elem_classes=["btn-ghost"])
start_btn = gr.Button("Start with the BRCA example", size="sm", variant="primary")
intro_back_buttons.append(back_btn)
intro_start_buttons.append(start_btn)
intro_cards.append(intro_card)
# ═════════════════════════════════════════════════════════════════
# Wiring
# ═════════════════════════════════════════════════════════════════
chrome_outputs = [*rail_row_htmls, run_state_plate_html, topbar_html_component, rail_row_buttons[1], rail_row_buttons[2]]
page_columns = [data_page, train_page, export_page, predict_page, results_page]
chrome_inputs_tail = [workspace_state, model_state, run_meta_state, align_state, predict_meta_state, session_id_state]
def _nav(page_key):
def _handler(workspace, bundle, run_meta, align, predicted, session_id):
return _enter_page(page_key, workspace, bundle, run_meta, align, predicted, session_id)
return _handler
for button, key in zip(rail_row_buttons, STAGE_KEYS):
button.click(
_nav(key), inputs=chrome_inputs_tail, outputs=[active_page_state, *page_columns, *chrome_outputs]
)
def _intro_updates(index: int, visible: bool = True):
index = max(0, min(len(ui.INTRO_CARDS) - 1, int(index)))
cards = [gr.update(visible=i == index) for i in range(len(ui.INTRO_CARDS))]
return index, gr.update(visible=visible), *cards
def _intro_load(dismissed):
return secrets.token_hex(3), *_intro_updates(0, visible=not bool(dismissed))
def _intro_close(dont_show):
return 0, gr.update(visible=False), bool(dont_show)
def _intro_start_brca(dont_show):
return (
0,
gr.update(visible=False),
bool(dont_show),
"BioLM-NET examples",
"BRCA",
*_source_visibility("BioLM-NET examples"),
)
intro_outputs = [intro_step_state, intro_group, *intro_cards]
demo.load(
_intro_load, inputs=[intro_dismissed_state], outputs=[session_id_state, *intro_outputs]
).then(
lambda sid, *rest: _refresh_chrome("data", *rest, sid), inputs=[session_id_state, *chrome_inputs_tail[:-1]], outputs=chrome_outputs
)
intro_open_button.click(
lambda: _intro_updates(0, True), outputs=intro_outputs, show_progress="hidden"
)
for index, button in enumerate(intro_next_buttons):
button.click(
lambda i=index: _intro_updates(i + 1, True), outputs=intro_outputs, show_progress="hidden"
)
for index, button in enumerate(intro_back_buttons, start=1):
button.click(
lambda i=index: _intro_updates(i - 1, True), outputs=intro_outputs, show_progress="hidden"
)
for button in intro_skip_buttons:
button.click(
_intro_close, inputs=[intro_dont_show],
outputs=[intro_step_state, intro_group, intro_dismissed_state],
show_progress="hidden",
)
for button in intro_start_buttons:
button.click(
_intro_start_brca, inputs=[intro_dont_show],
outputs=[
intro_step_state, intro_group, intro_dismissed_state,
source_mode, example_dataset, example_group, github_group, upload_group,
],
show_progress="hidden",
)
source_mode.change(
_source_visibility, inputs=[source_mode], outputs=[example_group, github_group, upload_group]
)
# The one deliberate loading indicator (a thin animated bar on the busy
# button itself — see `.is-busy` in styles.py) is toggled only through
# these two helpers, explicitly, on exactly the button that was clicked.
# It is never left to Gradio's own per-component pending state, which is
# what caused several of these to appear at once for one click.
def _busy(label, base_classes):
return gr.update(value=label, interactive=False, elem_classes=[*base_classes, "is-busy"])
def _idle(label, base_classes, interactive=True):
return gr.update(value=label, interactive=interactive, elem_classes=list(base_classes))
prepare_inputs = [
source_mode, example_dataset, github_folder, uploaded_gene, uploaded_dna, uploaded_labels,
uploaded_gene_pathway, uploaded_dna_pathway, uploaded_pdi, uploaded_ppi,
pathway_files_are_significant, embedding_option,
]
def _after_prepare(workspace, run_meta, session_id):
chrome = _refresh_chrome("data", workspace, None, run_meta, {}, {}, session_id)
return (
_idle("Build biological architecture", ["actbar-primary-btn"]),
_idle("Re-validate inputs", ["actbar-secondary-btn"]),
*chrome,
)
# Only the button actually clicked gets the sweep; its sibling just goes
# inert (disabled, unchanged label) — two buttons both sweeping for one
# action would itself be the "more than one indicator" problem.
def _busy_pair(which: str):
if which == "primary":
return (
_busy("Building…", ["actbar-primary-btn"]),
gr.update(interactive=False),
)
return (
gr.update(interactive=False),
_busy("Building…", ["actbar-secondary-btn"]),
)
prepare_button.click(
lambda: _busy_pair("primary"), outputs=[prepare_button, revalidate_button],
show_progress="hidden",
).then(
prepare_workspace, inputs=prepare_inputs,
outputs=[
workspace_state, preparation_status, architecture_table, enrichment_table,
resolved_files_table, run_meta_state, data_actionbar_note, graph_preview,
],
show_progress="hidden",
).then(
_after_prepare, inputs=[workspace_state, run_meta_state, session_id_state],
outputs=[prepare_button, revalidate_button, *chrome_outputs],
show_progress="hidden",
)
revalidate_button.click(
lambda: _busy_pair("secondary"), outputs=[prepare_button, revalidate_button],
show_progress="hidden",
).then(
prepare_workspace, inputs=prepare_inputs,
outputs=[
workspace_state, preparation_status, architecture_table, enrichment_table,
resolved_files_table, run_meta_state, data_actionbar_note, graph_preview,
],
show_progress="hidden",
).then(
_after_prepare, inputs=[workspace_state, run_meta_state, session_id_state],
outputs=[prepare_button, revalidate_button, *chrome_outputs],
show_progress="hidden",
)
show_symbols_btn.click(_reveal_symbols, inputs=[run_meta_state], outputs=[unmatched_reveal])
pdi_override_btn.click(lambda: gr.update(visible=True), outputs=[uploaded_pdi])
ppi_override_btn.click(lambda: gr.update(visible=True), outputs=[uploaded_ppi])
def _predict_source_toggle(mode):
return gr.update(visible=mode == "Session model"), gr.update(visible=mode == "Upload artifact")
predict_source_mode.change(
_predict_source_toggle, inputs=[predict_source_mode], outputs=[session_model_row, artifact_upload_row]
)
def _predict_input_toggle(mode):
use_uploads = mode == "Upload files"
return (
gr.update(visible=not use_uploads),
gr.update(visible=use_uploads),
gr.update(visible=use_uploads),
)
prediction_input_mode.change(
_predict_input_toggle,
inputs=[prediction_input_mode],
outputs=[prepared_dataset_row, prediction_gene_row, prediction_dna_row],
)
align_inputs = [
model_state, workspace_state, prediction_artifact, predict_source_mode,
prediction_input_mode, prediction_gene, prediction_dna,
]
def _align_wrapper(bundle, workspace, artifact_path, model_mode, input_mode, gene_path, dna_path):
return refresh_alignment(
bundle, workspace, artifact_path, model_mode == "Upload artifact",
input_mode == "Prepared dataset", gene_path, dna_path,
)
# `train_workspace`'s positional signature intentionally mirrors
# `estimate_training_duration` exactly for `@spaces.GPU` — session
# metadata (dataset name, elapsed time, …) is threaded through
# separately via `run_meta_update_state` and merged below, rather than
# passed into/out of `train_workspace` itself.
train_inputs = [
workspace_state, epochs, batch_size, learning_rate, weight_decay, dropout, projection_dim,
fusion_dim, validation_fraction, optimizer, class_weighting,
]
def _after_train(bundle, run_meta, session_id):
chrome = _refresh_chrome("train", bundle is not None and True, bundle, run_meta, {}, {}, session_id)
return _idle("Train BioLM-NET on ZeroGPU", ["actbar-primary-btn"]), *chrome
validation_predictions_state = gr.State(pd.DataFrame())
importance_state = gr.State(pd.DataFrame())
architecture_audit_state = gr.State(pd.DataFrame())
confusion_html_state = gr.State("")
run_meta_update_state = gr.State({})
train_button.click(
lambda: _busy("Training on ZeroGPU…", ["actbar-primary-btn"]), outputs=[train_button],
show_progress="hidden",
).then(
train_workspace, inputs=train_inputs,
outputs=[
model_state, training_status, loss_plot, validation_predictions_state, artifact_path_state,
importance_state, architecture_audit_state, run_meta_update_state, confusion_html_state,
],
show_progress="hidden",
).then(
lambda old, new: {**(old or {}), **(new or {})},
inputs=[run_meta_state, run_meta_update_state], outputs=[run_meta_state],
show_progress="hidden",
).then(
_after_train, inputs=[model_state, run_meta_state, session_id_state],
outputs=[train_button, *chrome_outputs],
show_progress="hidden",
).then(
_align_wrapper, inputs=align_inputs,
outputs=[alignment_table, alignment_strip, predict_button, align_state],
show_progress="hidden",
).then(
lambda bundle, artifact_path, run_meta: _export_panel(bundle, artifact_path, run_meta),
inputs=[model_state, artifact_path_state, run_meta_state],
outputs=[bundle_manifest, provenance_html, artifact_panel, export_stats, model_download],
show_progress="hidden",
).then(
refresh_results,
inputs=[model_state, run_meta_state],
outputs=[results_stats, results_stats_row, results_empty_note, results_content],
show_progress="hidden",
).then(
lambda df: gr.update(value=df), inputs=[validation_predictions_state], outputs=[results_validation_table],
show_progress="hidden",
).then(
lambda df: gr.update(value=df), inputs=[architecture_audit_state], outputs=[results_architecture_table],
show_progress="hidden",
).then(
lambda bundle: gr.update(value=_pathway_attention_table(bundle)) if bundle else gr.update(value=pd.DataFrame()),
inputs=[model_state], outputs=[results_pathway_table],
show_progress="hidden",
).then(
lambda html_value: gr.update(value=html_value), inputs=[confusion_html_state], outputs=[results_confusion],
show_progress="hidden",
).then(
lambda bundle: gr.update(visible=bundle is not None), inputs=[model_state], outputs=[results_tables_bottom],
show_progress="hidden",
)
# training_status is NOT re-set here: train_workspace's own return
# already carries the final "Training complete · accuracy · elapsed"
# message (it has the elapsed time on hand already), so writing it
# again from run_meta_state afterwards was a second, redundant paint
# of the same component for one click.
# Deliberately NOT including `model_state` here: it only ever changes as
# part of the train/predict chains below, and both of those already call
# `_align_wrapper` explicitly as their own finalize step. Adding it here
# too would fire alignment twice per click (train_workspace/_predict_
# wrapper set model_state -> this listener fires -> the chain's own
# explicit call also fires) — the exact "loading bar twice" symptom.
# `workspace_state` stays: nothing else re-checks alignment when the
# user rebuilds the architecture while already on the Predict page.
for control in (predict_source_mode, prediction_input_mode, prediction_artifact, prediction_gene, prediction_dna, workspace_state):
control.change(
_align_wrapper, inputs=align_inputs,
outputs=[alignment_table, alignment_strip, predict_button, align_state],
)
list_missing_btn.click(_list_missing, inputs=[align_state], outputs=[missing_reveal])
predict_inputs = [
model_state, workspace_state, prediction_artifact, predict_source_mode,
prediction_input_mode, prediction_gene, prediction_dna, predict_meta_state,
]
def _predict_wrapper(bundle, workspace, artifact_path, model_mode, input_mode, gene_path, dna_path, predict_meta):
return run_prediction(
bundle, workspace, artifact_path, model_mode == "Upload artifact",
input_mode == "Prepared dataset", gene_path, dna_path, predict_meta,
)
predict_button.click(
lambda: _busy("Running inference…", ["actbar-primary-btn"]), outputs=[predict_button],
show_progress="hidden",
).then(
_predict_wrapper, inputs=predict_inputs,
outputs=[model_state, prediction_status, prediction_table, distribution_panel, prediction_download, predict_meta_state],
show_progress="hidden",
).then(
# This one call both resets the button's label back to "Run
# inference" (out of its transient "Running inference…" busy state)
# and sets the correct interactive flag — no separate blind
# re-enable step first, which used to write predict_button twice.
_align_wrapper, inputs=align_inputs, outputs=[alignment_table, alignment_strip, predict_button, align_state],
show_progress="hidden",
).then(
lambda page, workspace, bundle, run_meta, align, predicted, session_id: _refresh_chrome(page, workspace, bundle, run_meta, align, predicted, session_id),
inputs=[active_page_state, workspace_state, model_state, run_meta_state, align_state, predict_meta_state, session_id_state],
outputs=chrome_outputs,
show_progress="hidden",
).then(
lambda df: gr.update(value=df), inputs=[prediction_table], outputs=[results_prediction_table],
show_progress="hidden",
).then(
refresh_results, inputs=[model_state, run_meta_state],
outputs=[results_stats, results_stats_row, results_empty_note, results_content],
show_progress="hidden",
)
if __name__ == "__main__":
demo.queue(default_concurrency_limit=1).launch(theme=THEME, css=CSS)