Spaces:
Running on Zero
Running on Zero
| from __future__ import annotations | |
| import html | |
| import math | |
| import tempfile | |
| from functools import lru_cache | |
| from pathlib import Path | |
| import spaces | |
| import gradio as gr | |
| import pandas as pd | |
| import plotly.express as px | |
| import plotly.graph_objects as go | |
| from biolmnet.artifacts import load_bundle, save_bundle | |
| from biolmnet.data import ( | |
| PreparedWorkspace, | |
| 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, | |
| ) | |
| 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" | |
| ), | |
| } | |
| CSS = """ | |
| :root { | |
| --ink: #102824; | |
| --muted: #60726e; | |
| --line: #d7e2de; | |
| --mint: #e8f4ee; | |
| --leaf: #167c5a; | |
| --leaf-dark: #0e5b43; | |
| --amber: #e5a63c; | |
| } | |
| .gradio-container { | |
| max-width: 1260px !important; | |
| margin: 0 auto !important; | |
| background: | |
| radial-gradient(circle at 92% 4%, rgba(209, 235, 224, .7), transparent 26rem), | |
| #f8fbf9 !important; | |
| color: var(--ink) !important; | |
| } | |
| .biolm-hero { | |
| position: relative; | |
| overflow: hidden; | |
| padding: 34px 36px 30px; | |
| margin: 12px 0 18px; | |
| border: 1px solid #cfe0d9; | |
| border-radius: 24px; | |
| background: linear-gradient(135deg, #0b2f28 0%, #124b3d 70%, #17694f 100%); | |
| box-shadow: 0 18px 50px rgba(19, 63, 52, .12); | |
| } | |
| .biolm-hero::after { | |
| content: ""; | |
| position: absolute; | |
| width: 250px; | |
| height: 250px; | |
| right: -60px; | |
| top: -90px; | |
| border: 1px solid rgba(255,255,255,.2); | |
| border-radius: 50%; | |
| box-shadow: 0 0 0 34px rgba(255,255,255,.035), 0 0 0 68px rgba(255,255,255,.025); | |
| } | |
| .eyebrow { | |
| color: #a9e2ca; | |
| font: 700 12px/1.2 ui-monospace, SFMono-Regular, Menlo, monospace; | |
| letter-spacing: .14em; | |
| text-transform: uppercase; | |
| } | |
| .biolm-hero h1 { | |
| color: white; | |
| font-size: clamp(34px, 5vw, 58px); | |
| line-height: .98; | |
| letter-spacing: -.045em; | |
| margin: 12px 0 14px; | |
| } | |
| .biolm-hero p { | |
| max-width: 780px; | |
| color: #d6e8e1; | |
| font-size: 17px; | |
| line-height: 1.55; | |
| margin: 0; | |
| } | |
| .hero-meta { | |
| display: flex; | |
| gap: 10px; | |
| flex-wrap: wrap; | |
| margin-top: 22px; | |
| } | |
| .hero-chip { | |
| color: #e9f7f1; | |
| border: 1px solid rgba(255,255,255,.23); | |
| background: rgba(255,255,255,.07); | |
| border-radius: 999px; | |
| padding: 7px 11px; | |
| font: 600 12px/1 ui-monospace, SFMono-Regular, Menlo, monospace; | |
| } | |
| .phase-card { | |
| border: 1px solid var(--line) !important; | |
| border-radius: 18px !important; | |
| background: rgba(255,255,255,.86) !important; | |
| box-shadow: 0 9px 26px rgba(27, 68, 57, .05) !important; | |
| } | |
| .phase-intro { | |
| border-left: 3px solid var(--leaf); | |
| padding: 2px 0 2px 15px; | |
| color: var(--muted); | |
| } | |
| .phase-intro strong { color: var(--ink); } | |
| .status-box { | |
| border-radius: 15px; | |
| padding: 14px 16px; | |
| background: var(--mint); | |
| border: 1px solid #cbe3d8; | |
| color: #254d41; | |
| } | |
| .error-box { | |
| border-radius: 15px; | |
| padding: 14px 16px; | |
| background: #fff1ee; | |
| border: 1px solid #f0cac2; | |
| color: #793b31; | |
| } | |
| .metric-grid { | |
| display: grid; | |
| grid-template-columns: repeat(4, minmax(0,1fr)); | |
| gap: 10px; | |
| } | |
| .metric { | |
| padding: 13px 14px; | |
| border: 1px solid #d4e2dd; | |
| border-radius: 14px; | |
| background: white; | |
| } | |
| .metric span { | |
| display:block; | |
| color: var(--muted); | |
| font-size: 11px; | |
| font-weight: 700; | |
| letter-spacing: .07em; | |
| text-transform: uppercase; | |
| } | |
| .metric b { | |
| display:block; | |
| margin-top: 3px; | |
| color: var(--ink); | |
| font-size: 22px; | |
| } | |
| button.primary { | |
| background: var(--leaf) !important; | |
| border-color: var(--leaf) !important; | |
| } | |
| button.primary:hover { background: var(--leaf-dark) !important; } | |
| .footnote { | |
| color: #6d7e79; | |
| font-size: 12px; | |
| line-height: 1.55; | |
| } | |
| @media (max-width: 720px) { | |
| .biolm-hero { padding: 26px 22px; border-radius: 18px; } | |
| .metric-grid { grid-template-columns: repeat(2, minmax(0,1fr)); } | |
| } | |
| """ | |
| def _status(message: str, error: bool = False) -> str: | |
| class_name = "error-box" if error else "status-box" | |
| return f'<div class="{class_name}">{html.escape(message)}</div>' | |
| 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 _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 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), | |
| ): | |
| 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()} | |
| 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()} | |
| source_name = github_folder.strip() | |
| precomputed_significant = pathway_files_are_significant | |
| allow_preset_trim = False | |
| else: | |
| frames = { | |
| "gene": _read_required_upload( | |
| uploaded_gene, "Gene_Expression.csv" | |
| ), | |
| "dna": _read_required_upload( | |
| uploaded_dna, "DNA_Methylation.csv" | |
| ), | |
| "labels": _read_required_upload( | |
| uploaded_labels, "label.csv" | |
| ), | |
| "gene_pathways": _read_required_upload( | |
| uploaded_gene_pathway, | |
| "the gene-expression pathway mapping CSV", | |
| ), | |
| "dna_pathways": _read_required_upload( | |
| uploaded_dna_pathway, | |
| "the DNA-methylation pathway mapping CSV", | |
| ), | |
| } | |
| 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, | |
| ) | |
| dna_enrichment = attach_embeddings_and_pathways( | |
| dna_branch, | |
| embeddings, | |
| frames["dna_pathways"], | |
| precomputed_significant=precomputed_significant, | |
| ) | |
| 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, | |
| ) | |
| warning_text = ( | |
| "<br><small>" + " · ".join(html.escape(item) for item in warnings) + "</small>" | |
| if warnings | |
| else "" | |
| ) | |
| summary = ( | |
| '<div class="status-box"><strong>Architecture ready.</strong> ' | |
| f"{len(gene_frame):,} paired samples · {len(label_names)} classes · " | |
| f"{len(gene_branch.pathways) + len(dna_branch.pathways):,} " | |
| f"branch-specific pathways · GenePT: {html.escape(embedding_file)}" | |
| f"{warning_text}</div>" | |
| ) | |
| progress(1.0, desc="Ready to train") | |
| return workspace, summary, architecture, enrichments.head(100) | |
| except Exception as exc: | |
| return None, _status(str(exc), error=True), pd.DataFrame(), pd.DataFrame() | |
| def _training_plots(history: list[dict[str, float]], confusion, labels): | |
| history_frame = pd.DataFrame(history) | |
| loss_figure = go.Figure() | |
| loss_figure.add_trace( | |
| go.Scatter( | |
| x=history_frame["epoch"], | |
| y=history_frame["training_loss"], | |
| mode="lines", | |
| name="Training", | |
| line={"color": "#167c5a", "width": 3}, | |
| ) | |
| ) | |
| loss_figure.add_trace( | |
| go.Scatter( | |
| x=history_frame["epoch"], | |
| y=history_frame["validation_loss"], | |
| mode="lines", | |
| name="Validation", | |
| line={"color": "#e5a63c", "width": 3}, | |
| ) | |
| ) | |
| loss_figure.update_layout( | |
| title="Loss by epoch", | |
| xaxis_title="Epoch", | |
| yaxis_title="Cross-entropy loss", | |
| template="plotly_white", | |
| margin={"l": 30, "r": 15, "t": 50, "b": 35}, | |
| legend={"orientation": "h", "y": 1.12}, | |
| ) | |
| confusion_figure = px.imshow( | |
| confusion, | |
| x=labels, | |
| y=labels, | |
| text_auto=True, | |
| color_continuous_scale=[[0, "#eef6f2"], [1, "#167c5a"]], | |
| labels={"x": "Predicted", "y": "Observed", "color": "Samples"}, | |
| title="Validation confusion matrix", | |
| ) | |
| confusion_figure.update_layout( | |
| template="plotly_white", margin={"l": 30, "r": 15, "t": 50, "b": 35} | |
| ) | |
| return loss_figure, confusion_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 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), | |
| ): | |
| if workspace is None: | |
| return ( | |
| None, | |
| _status("Prepare data and priors in Phase 1 before training.", True), | |
| None, | |
| None, | |
| pd.DataFrame(), | |
| None, | |
| pd.DataFrame(), | |
| ) | |
| 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) | |
| result = train(workspace, parameters, progress=report) | |
| bundle = result.bundle | |
| artifact = save_bundle(bundle) | |
| metrics = bundle.metrics | |
| metrics_html = f""" | |
| <div class="status-box"><strong>Training complete.</strong> | |
| Best validation checkpoint restored; the downloadable artifact includes | |
| architecture, preprocessing, weights, and metrics.</div> | |
| <div class="metric-grid" style="margin-top:10px"> | |
| <div class="metric"><span>Macro F1</span><b>{metrics['f1_macro']:.3f}</b></div> | |
| <div class="metric"><span>Accuracy</span><b>{metrics['accuracy']:.3f}</b></div> | |
| <div class="metric"><span>Macro precision</span><b>{metrics['precision_macro']:.3f}</b></div> | |
| <div class="metric"><span>Macro recall</span><b>{metrics['recall_macro']:.3f}</b></div> | |
| </div> | |
| """ | |
| loss_plot, confusion_plot = _training_plots( | |
| bundle.history, result.confusion, bundle.label_names | |
| ) | |
| importance = pathway_importance(bundle).head(100) | |
| return ( | |
| bundle, | |
| metrics_html, | |
| loss_plot, | |
| confusion_plot, | |
| result.validation_predictions, | |
| artifact, | |
| importance, | |
| ) | |
| except Exception as exc: | |
| return ( | |
| None, | |
| _status(str(exc), error=True), | |
| None, | |
| None, | |
| pd.DataFrame(), | |
| None, | |
| pd.DataFrame(), | |
| ) | |
| def run_prediction( | |
| bundle: ModelBundle | None, | |
| uploaded_artifact: str | None, | |
| gene_file: str | None, | |
| dna_file: str | None, | |
| ): | |
| try: | |
| active_bundle = ( | |
| load_bundle(uploaded_artifact) if uploaded_artifact else bundle | |
| ) | |
| if active_bundle is None: | |
| raise ValueError( | |
| "Train a model in Phase 2 or upload a BioLM-NET model artifact." | |
| ) | |
| 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() | |
| .rename_axis("class") | |
| .reset_index(name="samples") | |
| ) | |
| figure = px.bar( | |
| counts, | |
| x="class", | |
| y="samples", | |
| color="class", | |
| color_discrete_sequence=[ | |
| "#167c5a", | |
| "#e5a63c", | |
| "#497f93", | |
| "#8d6fa8", | |
| "#be6f55", | |
| ], | |
| title="Predicted class distribution", | |
| ) | |
| figure.update_layout( | |
| showlegend=False, | |
| template="plotly_white", | |
| margin={"l": 30, "r": 15, "t": 50, "b": 35}, | |
| ) | |
| status = _status( | |
| f"Predicted {len(output):,} samples. Mean confidence: " | |
| f"{output['confidence'].mean():.3f}." | |
| ) | |
| return active_bundle, status, output, figure, str(destination) | |
| except Exception as exc: | |
| return bundle, _status(str(exc), True), pd.DataFrame(), None, None | |
| THEME = gr.themes.Base( | |
| primary_hue="emerald", | |
| neutral_hue="slate", | |
| ) | |
| with gr.Blocks(title="BioLM-NET Workbench") as demo: | |
| workspace_state = gr.State(None) | |
| model_state = gr.State(None) | |
| gr.HTML( | |
| """ | |
| <section class="biolm-hero"> | |
| <div class="eyebrow">Interpretable multi-omics modeling</div> | |
| <h1>BioLM-NET<br>Workbench</h1> | |
| <p>Build a biologically masked network from paired gene expression and | |
| DNA methylation, train it with GenePT-guided pathway attention, then | |
| carry the exact preprocessing and architecture into prediction.</p> | |
| <div class="hero-meta"> | |
| <span class="hero-chip">PDI · DoRothEA</span> | |
| <span class="hero-chip">PPI · STRING</span> | |
| <span class="hero-chip">Pathways · KEGG</span> | |
| <span class="hero-chip">Context · GenePT</span> | |
| <span class="hero-chip">Compute · ZeroGPU on demand</span> | |
| </div> | |
| </section> | |
| """ | |
| ) | |
| with gr.Tabs(): | |
| with gr.Tab("1 · Data & priors", id="data"): | |
| gr.HTML( | |
| """ | |
| <p class="phase-intro"><strong>Assemble the model graph.</strong> | |
| Select an upstream example, point to a GitHub folder that follows | |
| the BioLM-NET file convention, or upload paired omics files.</p> | |
| """ | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=7, elem_classes=["phase-card"]): | |
| source_mode = gr.Radio( | |
| [ | |
| "BioLM-NET examples", | |
| "GitHub folder", | |
| "Upload files", | |
| ], | |
| value="BioLM-NET examples", | |
| label="Dataset source", | |
| ) | |
| with gr.Column(visible=True) as example_group: | |
| example_dataset = gr.Dropdown( | |
| ["BRCA", "COAD", "GBM", "scTrioseq2"], | |
| value="BRCA", | |
| label="Repository dataset", | |
| ) | |
| 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", | |
| ) | |
| uploaded_dna = gr.File( | |
| label="DNA methylation", | |
| file_types=[".csv"], | |
| type="filepath", | |
| ) | |
| uploaded_labels = gr.File( | |
| label="Labels", | |
| file_types=[".csv"], | |
| type="filepath", | |
| ) | |
| with gr.Row(): | |
| uploaded_gene_pathway = gr.File( | |
| label="Gene → pathway mapping", | |
| file_types=[".csv"], | |
| type="filepath", | |
| ) | |
| uploaded_dna_pathway = gr.File( | |
| label="DNA → pathway mapping", | |
| file_types=[".csv"], | |
| type="filepath", | |
| ) | |
| pathway_files_are_significant = gr.Checkbox( | |
| value=True, | |
| label="Pathway files already contain significant pathways", | |
| info=( | |
| "Turn off for a full SYMBOL/PathwayID annotation " | |
| "catalog; enrichment will use BH-adjusted p < 0.05." | |
| ), | |
| ) | |
| with gr.Column(scale=5, elem_classes=["phase-card"]): | |
| embedding_option = gr.Dropdown( | |
| list(GENEPT_OPTIONS), | |
| value=list(GENEPT_OPTIONS)[0], | |
| label="GenePT context", | |
| ) | |
| with gr.Accordion("Interaction priors", open=False): | |
| gr.Markdown( | |
| "By default, the app retrieves `PDI.csv` and `PPI.csv` " | |
| "from `bozdaglab/BioLM-NET`. Upload both only to override." | |
| ) | |
| uploaded_pdi = gr.File( | |
| label="Custom PDI.csv", | |
| file_types=[".csv"], | |
| type="filepath", | |
| ) | |
| uploaded_ppi = gr.File( | |
| label="Custom PPI.csv", | |
| file_types=[".csv"], | |
| type="filepath", | |
| ) | |
| prepare_button = gr.Button( | |
| "Build biological architecture", | |
| variant="primary", | |
| size="lg", | |
| ) | |
| preparation_status = gr.HTML( | |
| _status("Choose a source, then build the biological architecture.") | |
| ) | |
| with gr.Row(): | |
| architecture_table = gr.Dataframe( | |
| label="Sparse architecture audit", | |
| interactive=False, | |
| wrap=True, | |
| ) | |
| enrichment_table = gr.Dataframe( | |
| label="Retained enriched pathways (first 100)", | |
| interactive=False, | |
| wrap=True, | |
| ) | |
| gr.HTML( | |
| """ | |
| <p class="footnote">Expected orientation: samples in rows and | |
| HGNC gene symbols in columns. PDI requires <code>TF, Target</code>; | |
| PPI requires <code>protein1, protein2, combined_score</code>; | |
| pathways require <code>SYMBOL, PathwayID</code>. PPI is filtered | |
| to score > 0.7 and the retained top decile, following the paper.</p> | |
| """ | |
| ) | |
| with gr.Tab("2 · Train", id="train"): | |
| gr.HTML( | |
| """ | |
| <p class="phase-intro"><strong>Fit and evaluate.</strong> The | |
| split is stratified; scaling is fit on training samples only; | |
| the best validation checkpoint is exported as a safe, | |
| self-contained model artifact. A shared GPU is requested only | |
| while this training callback is running.</p> | |
| """ | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=4, elem_classes=["phase-card"]): | |
| epochs = gr.Slider(5, 200, value=50, step=5, label="Epochs") | |
| batch_size = gr.Dropdown( | |
| [8, 16, 32, 64, 128], value=16, label="Batch size" | |
| ) | |
| learning_rate = gr.Number( | |
| value=0.001, label="Learning rate", minimum=0.000001 | |
| ) | |
| weight_decay = gr.Number( | |
| value=0.01, label="L2 weight decay", minimum=0 | |
| ) | |
| dropout = gr.Slider( | |
| 0, 0.8, value=0.3, step=0.05, label="Dropout" | |
| ) | |
| with gr.Column(scale=4, elem_classes=["phase-card"]): | |
| projection_dim = gr.Dropdown( | |
| [16, 32, 64, 128], value=64, label="Branch projection" | |
| ) | |
| fusion_dim = gr.Dropdown( | |
| [8, 12, 16, 32], value=12, label="Fusion layer" | |
| ) | |
| validation_fraction = gr.Slider( | |
| 0.1, | |
| 0.4, | |
| value=0.2, | |
| step=0.05, | |
| label="Validation fraction", | |
| ) | |
| optimizer = gr.Radio( | |
| ["Adam", "SGD"], value="Adam", label="Optimizer" | |
| ) | |
| 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=4, elem_classes=["phase-card"]): | |
| gr.Markdown( | |
| """ | |
| **Paper-faithful defaults** | |
| - First layer: trainable `W ⊙ M` | |
| - PDI weights: binary | |
| - PPI weights: normalized STRING score | |
| - Pathway attention: GenePT query attention | |
| - Fusion: dual branch → dense → softmax | |
| """ | |
| ) | |
| train_button = gr.Button( | |
| "Train BioLM-NET on ZeroGPU", variant="primary", size="lg" | |
| ) | |
| model_download = gr.File( | |
| label="Trained model artifact", interactive=False | |
| ) | |
| gr.Markdown( | |
| "ZeroGPU reserves 30–300 seconds according to dataset " | |
| "size and epochs. Visitors use their own daily quota." | |
| ) | |
| training_status = gr.HTML( | |
| _status("Phase 2 unlocks after the architecture is prepared.") | |
| ) | |
| with gr.Row(): | |
| loss_plot = gr.Plot(label="Training history") | |
| confusion_plot = gr.Plot(label="Confusion matrix") | |
| with gr.Row(): | |
| validation_table = gr.Dataframe( | |
| label="Validation predictions", | |
| interactive=False, | |
| wrap=True, | |
| ) | |
| importance_table = gr.Dataframe( | |
| label="Pathway attention audit (first 100)", | |
| interactive=False, | |
| wrap=True, | |
| ) | |
| with gr.Tab("3 · Predict", id="predict"): | |
| gr.HTML( | |
| """ | |
| <p class="phase-intro"><strong>Apply a trained model.</strong> | |
| Continue with the model from this session or upload a previous | |
| artifact. Feature names are validated and reordered exactly as | |
| they were during training.</p> | |
| """ | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=4, elem_classes=["phase-card"]): | |
| prediction_artifact = gr.File( | |
| label="Optional trained model artifact", | |
| file_types=[".zip"], | |
| type="filepath", | |
| ) | |
| prediction_gene = gr.File( | |
| label="Prediction gene expression", | |
| file_types=[".csv"], | |
| type="filepath", | |
| ) | |
| prediction_dna = gr.File( | |
| label="Prediction DNA methylation", | |
| file_types=[".csv"], | |
| type="filepath", | |
| ) | |
| predict_button = gr.Button( | |
| "Make predictions", variant="primary", size="lg" | |
| ) | |
| prediction_download = gr.File( | |
| label="Prediction CSV", interactive=False | |
| ) | |
| with gr.Column(scale=8): | |
| prediction_status = gr.HTML( | |
| _status("Use the current trained model or upload an artifact.") | |
| ) | |
| prediction_plot = gr.Plot(label="Class distribution") | |
| prediction_table = gr.Dataframe( | |
| label="Predictions and class probabilities", | |
| interactive=False, | |
| wrap=True, | |
| ) | |
| gr.Markdown( | |
| """ | |
| <p class="footnote">Research use only. This interface reproduces the | |
| architecture described by Rifat et al. and uses the upstream BioLM-NET | |
| repository and GenePT embeddings. Validate cohorts, preprocessing, and | |
| model performance before drawing biological or clinical conclusions.</p> | |
| """ | |
| ) | |
| source_mode.change( | |
| _source_visibility, | |
| inputs=[source_mode], | |
| outputs=[example_group, github_group, upload_group], | |
| ) | |
| prepare_button.click( | |
| prepare_workspace, | |
| 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, | |
| ], | |
| outputs=[ | |
| workspace_state, | |
| preparation_status, | |
| architecture_table, | |
| enrichment_table, | |
| ], | |
| ) | |
| train_button.click( | |
| train_workspace, | |
| inputs=[ | |
| workspace_state, | |
| epochs, | |
| batch_size, | |
| learning_rate, | |
| weight_decay, | |
| dropout, | |
| projection_dim, | |
| fusion_dim, | |
| validation_fraction, | |
| optimizer, | |
| class_weighting, | |
| ], | |
| outputs=[ | |
| model_state, | |
| training_status, | |
| loss_plot, | |
| confusion_plot, | |
| validation_table, | |
| model_download, | |
| importance_table, | |
| ], | |
| ) | |
| predict_button.click( | |
| run_prediction, | |
| inputs=[ | |
| model_state, | |
| prediction_artifact, | |
| prediction_gene, | |
| prediction_dna, | |
| ], | |
| outputs=[ | |
| model_state, | |
| prediction_status, | |
| prediction_table, | |
| prediction_plot, | |
| prediction_download, | |
| ], | |
| ) | |
| if __name__ == "__main__": | |
| demo.queue(default_concurrency_limit=1).launch(theme=THEME, css=CSS) | |