Spaces:
Running on Zero
Running on Zero
Download app.py from bozdaglab/BioLM-NET-Workbench: direct link, hf CLI and curl.
- Browser
- Download file 90.2 kB
-
https://huggingface.co/spaces/bozdaglab/BioLM-NET-Workbench/resolve/main/app.py
- Command line
-
hf download hf://spaces/bozdaglab/BioLM-NET-Workbench/app.py
-
curl -L -o app.py https://huggingface.co/spaces/bozdaglab/BioLM-NET-Workbench/resolve/main/app.py
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) | |
| 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 | |
| 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 & architecture", | |
| "arrays.npz": "Biological masks, embeddings & scalers", | |
| "model.safetensors": "Trained weights", | |
| "metrics.json": "Metrics & 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 < 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">< 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 > 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 & 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 & 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) | |