# MultiMolecule # Copyright (C) 2024-Present MultiMolecule # This file is part of MultiMolecule. # MultiMolecule is free software: you can redistribute it and/or modify # it under the terms of the GNU Affero General Public License as published by # the Free Software Foundation, either version 3 of the License, or # any later version. # MultiMolecule is distributed in the hope that it will be useful, # but WITHOUT ANY WARRANTY; without even the implied warranty of # MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the # GNU Affero General Public License for more details. # You should have received a copy of the GNU Affero General Public License # along with this program. If not, see . # For additional terms and clarifications, please refer to our License FAQ at: # . from __future__ import annotations import csv import json import tempfile from functools import lru_cache from pathlib import Path from typing import Any import gradio as gr import matplotlib import numpy as np import pandas as pd import torch from transformers import pipeline matplotlib.use("Agg") import matplotlib.pyplot as plt # noqa: E402 import multimolecule # noqa: E402, F401 - registers MultiMolecule models and pipelines with Transformers import multimolecule.io as mmio # noqa: E402 TRACK_TASK = "regulatory-track" PROFILE_TASK = "regulatory-profile" TASK_OPTIONS = { "Track": TRACK_TASK, "Profile": PROFILE_TASK, } TRACK_MODEL_OPTIONS = { "Enformer": "multimolecule/enformer", "Basenji": "multimolecule/basenji", } PROFILE_MODEL_OPTIONS = { "BPNet": "multimolecule/bpnet", "ChromBPNet": "multimolecule/chrombpnet", "ProCapNet": "multimolecule/procapnet", } FASTA_SUFFIXES = {f".{suffix}" for suffix in mmio.FASTA} DNA_ALPHABET = set("ACGTN") DEFAULT_SEQUENCE = ("ACGT" * 529)[:2114] def _device() -> int: return 0 if torch.cuda.is_available() else -1 @lru_cache(maxsize=4) def load_predictor(task: str, model_id: str): return pipeline(task, model=model_id, device=_device()) def model_visibility(task_label: str): is_track = TASK_OPTIONS[task_label] == TRACK_TASK return gr.update(visible=is_track), gr.update(visible=not is_track) def clean_sequence(sequence: str) -> str: sequence = "".join(str(sequence).split()).upper() if not sequence: raise gr.Error("Sequence is empty.") invalid = sorted(set(sequence) - DNA_ALPHABET) if invalid: raise gr.Error(f"DNA sequence can only contain A, C, G, T, and N. Found: {', '.join(invalid)}.") return sequence def load_input_file(input_file: Any): if input_file is None: return gr.update() path = Path(getattr(input_file, "name", input_file)) if path.suffix.lower() not in FASTA_SUFFIXES: raise gr.Error("Could not parse uploaded file. Supported formats: FASTA, FA, and FNA.") try: records = mmio.read_fasta_records(path) except mmio.InvalidStructureFile as error: raise gr.Error("Could not parse uploaded file. Supported formats: FASTA, FA, and FNA.") from error if not records: raise gr.Error(f"No FASTA records found in {path.name}.") if len(records) > 1: raise gr.Error(f"This demo supports one sequence at a time. Uploaded FASTA contains {len(records)} records.") return clean_sequence(records[0].sequence) def run_prediction( task_label: str, track_model_label: str, profile_model_label: str, sequence: str, max_table_rows: int, max_display_channels: int, ): task = TASK_OPTIONS[task_label] model_options = TRACK_MODEL_OPTIONS if task == TRACK_TASK else PROFILE_MODEL_OPTIONS model_label = track_model_label if task == TRACK_TASK else profile_model_label model_id = model_options[model_label] sequence = clean_sequence(sequence) predictor = load_predictor(task, model_id) result = predictor(sequence) result = _unwrap_result(result) rows_key = "tracks" if task == TRACK_TASK else "profile" axis_name = "bin" if task == TRACK_TASK else "position" signal_rows = result.get(rows_key) if not isinstance(signal_rows, list) or not signal_rows: raise gr.Error(f"The selected model did not return a non-empty `{rows_key}` signal table.") output_sequence = str(result.get("sequence", sequence)) channels = _resolve_channels(result, signal_rows, axis_name) max_table_rows = int(max_table_rows) max_display_channels = int(max_display_channels) table = _rows_to_table(signal_rows, channels, axis_name, max_table_rows, max_display_channels) plot = _plot_signal(signal_rows, channels, axis_name, task_label, model_label, max_display_channels) metadata = { "task": task_label, "pipeline_task": task, "model": model_id, "device": "cuda" if torch.cuda.is_available() else "cpu", "input_length": len(sequence), "output_sequence_length": len(output_sequence), "axis": axis_name, "signals": len(signal_rows), "channels": len(channels), "displayed_rows": min(max_table_rows, len(signal_rows)), "displayed_channels": min(max_display_channels, len(channels)), "coordinate_scope": "sequence-relative output bins/positions only", } csv_path, json_path = _write_result_files( task=task, model_id=model_id, sequence=output_sequence, rows_key=rows_key, rows=signal_rows, channels=channels, axis_name=axis_name, metadata=metadata, ) return table, metadata, plot, csv_path, json_path def _unwrap_result(result: Any) -> dict[str, Any]: if isinstance(result, list): if len(result) != 1: raise gr.Error(f"Expected one prediction result, got {len(result)}.") result = result[0] if not isinstance(result, dict): raise gr.Error(f"Expected a prediction dictionary, got {type(result).__name__}.") return result def _resolve_channels(result: dict[str, Any], rows: list[dict[str, Any]], axis_name: str) -> list[str]: channels = result.get("channels") if isinstance(channels, list) and channels: return [str(channel) for channel in channels] metadata_columns = {axis_name, "nucleotide"} return [key for key in rows[0] if key not in metadata_columns] def _rows_to_table( rows: list[dict[str, Any]], channels: list[str], axis_name: str, max_rows: int, max_channels: int, ) -> pd.DataFrame: selected_channels = channels[:max_channels] include_nucleotide = any("nucleotide" in row for row in rows[:max_rows]) columns = [axis_name] if include_nucleotide: columns.append("nucleotide") columns.extend(selected_channels) table = [{column: row.get(column) for column in columns} for row in rows[:max_rows]] return pd.DataFrame.from_records(table, columns=columns) def _plot_signal( rows: list[dict[str, Any]], channels: list[str], axis_name: str, task_label: str, model_label: str, max_channels: int, ): selected_channels = channels[:max_channels] x = np.asarray([row.get(axis_name, index) for index, row in enumerate(rows)], dtype=float) fig, ax = plt.subplots(figsize=(11, 4.5)) for channel in selected_channels: y = np.asarray([row.get(channel, np.nan) for row in rows], dtype=float) ax.plot(x, y, linewidth=1.1, label=_short_label(channel)) ax.set_title(f"{model_label} {task_label.lower()} signal") ax.set_xlabel("Output bin (0-based)" if axis_name == "bin" else "Sequence position (0-based)") ax.set_ylabel("Predicted signal") ax.grid(alpha=0.25) if selected_channels: ax.legend(loc="upper right", fontsize="x-small", ncol=2 if len(selected_channels) > 4 else 1) fig.tight_layout() return fig def _short_label(label: str, limit: int = 36) -> str: if len(label) <= limit: return label return f"{label[: limit - 1]}..." def _write_result_files( *, task: str, model_id: str, sequence: str, rows_key: str, rows: list[dict[str, Any]], channels: list[str], axis_name: str, metadata: dict[str, Any], ) -> tuple[str, str]: columns = [axis_name] if any("nucleotide" in row for row in rows): columns.append("nucleotide") columns.extend(channels) csv_file = tempfile.NamedTemporaryFile("w", suffix=".csv", delete=False, newline="") writer = csv.DictWriter(csv_file, fieldnames=columns, extrasaction="ignore") writer.writeheader() writer.writerows(rows) csv_file.close() payload = { "task": task, "model": model_id, "sequence": sequence, "channels": channels, rows_key: rows, "metadata": metadata, } json_file = tempfile.NamedTemporaryFile("w", suffix=".json", delete=False) json.dump(payload, json_file, indent=2) json_file.close() return csv_file.name, json_file.name with gr.Blocks(title="Regulatory Signal") as demo: gr.Markdown( "# Regulatory Signal\n" "Run MultiMolecule DNA regulatory track and profile checkpoints and inspect sequence-relative signal outputs." ) with gr.Row(): task = gr.Radio( choices=list(TASK_OPTIONS.keys()), value="Profile", label="Task", ) track_model = gr.Dropdown( choices=list(TRACK_MODEL_OPTIONS.keys()), value="Enformer", label="Track checkpoint", visible=False, ) profile_model = gr.Dropdown( choices=list(PROFILE_MODEL_OPTIONS.keys()), value="BPNet", label="Profile checkpoint", ) sequence = gr.Textbox( label="DNA sequence", value=DEFAULT_SEQUENCE, lines=5, ) input_file = gr.File( label="Upload FASTA", file_types=[".fa", ".fasta", ".fna"], ) with gr.Row(): max_table_rows = gr.Slider(10, 2000, value=200, step=10, label="Rows shown") max_display_channels = gr.Slider(1, 24, value=8, step=1, label="Channels shown") run = gr.Button("Run prediction", variant="primary") with gr.Row(): signal_table = gr.Dataframe(label="Signal table", interactive=False, wrap=True) metadata = gr.JSON(label="Run metadata") signal_plot = gr.Plot(label="Signal plot") with gr.Row(): csv_download = gr.File(label="Download CSV") json_download = gr.File(label="Download JSON") task.change(model_visibility, inputs=task, outputs=[track_model, profile_model]) input_file.change(load_input_file, inputs=input_file, outputs=sequence) run.click( run_prediction, inputs=[task, track_model, profile_model, sequence, max_table_rows, max_display_channels], outputs=[signal_table, metadata, signal_plot, csv_download, json_download], ) if __name__ == "__main__": demo.launch()