Run: Run a benchmark

Run

Run a benchmark

Generate embeddings, run probes across seeds, score variant effects, and optionally fine-tune a model.

These examples use the installed Python API. They work in a notebook or a normal Python file and show the same objects used by the batch utilities.

1. Generate embeddings

python
import torch
import mrna_bench as mb

from mrna_bench.embedder import DatasetEmbedder

device = torch.device("cuda")
dataset = mb.load_dataset("go-mf")
model = mb.load_model(
    "Orthrus",
    "orthrus-large-6-track",
    device=device,
)

embedder = DatasetEmbedder(
    model,
    dataset,
    batch_size=8,
)
embedding_tensors = embedder.embed_dataset()
embedder.persist_embeddings(embedding_tensors)

embeddings = torch.stack(embedding_tensors).cpu().numpy()
print(embeddings.shape)

DatasetEmbedder keeps dataframe row order and supplies CDS and splice tracks when the dataset contains them. The call topersist_embeddings() writes a compressed NPZ below the configured dataset path.

Process a large dataset in chunks

python
from mrna_bench.embedder import DatasetEmbedder

num_chunks = 8
for chunk_index in range(num_chunks):
    chunk_embedder = DatasetEmbedder(
        model,
        dataset,
        d_chunk_ind=chunk_index,
        d_num_chunks=num_chunks,
        batch_size=8,
    )
    chunk = chunk_embedder.embed_dataset()
    chunk_embedder.persist_embeddings(chunk)
    chunk_embedder.merge_embeddings()

# The merged file appears after the final chunk is present.

Each chunk can run as a separate job. Callingmerge_embeddings() is safe before all chunks exist; it creates the merged file only after every zero-based chunk index is present.

2. Run seeded linear probes

python
from mrna_bench.linear_probe import LinearProbeBuilder

probe = (
    LinearProbeBuilder(dataset)
    .fetch_embedding_by_embedding_instance(
        model.short_name,
        embeddings,
    )
    .build_splitter(
        dataset.metadata.default_split_type,
        eval_all_splits=True,
    )
    .use_persister()
    .build()
)

seeds = [2541, 413, 411, 412, 2547]
for seed in seeds:
    metrics = probe.run_linear_probe(seed, persist=True)
    test_metrics = {
        name: value
        for name, value in metrics.items()
        if name.startswith("test_")
    }
    print(seed, test_metrics)

Dataset metadata selects the biological task and target. This example uses the dataset's default split, evaluates train, validation, and test partitions, and stores each seed in results.db.

Reuse saved embeddings

python
from mrna_bench.linear_probe import LinearProbeBuilder

# Reuse embeddings saved by DatasetEmbedder.persist_embeddings().
probe = (
    LinearProbeBuilder(dataset)
    .fetch_embedding_by_model_instance(model)
    .build_splitter(
        dataset.metadata.default_split_type,
        eval_all_splits=True,
    )
    .build()
)

metrics = probe.run_linear_probe(random_seed=2541)
print(metrics)

fetch_embedding_by_model_instance() reads the existing NPZ using the loaded model's short name.

Variant-effect prediction

Zero-shot embedding magnitude

python
import torch
import mrna_bench as mb

from mrna_bench.embedder import DatasetEmbedder
from mrna_bench.zeroshot import ZeroShotVEP

dataset = mb.load_dataset("vep-traitgym-mendelian")
model = mb.load_model(
    "Orthrus",
    "orthrus-large-4-track",
    device=torch.device("cuda"),
)

embedding_tensors = DatasetEmbedder(
    model,
    dataset,
    batch_size=8,
).embed_dataset()
embeddings = torch.stack(embedding_tensors).cpu().numpy()

evaluator = ZeroShotVEP.from_embeddings(
    dataset,
    embeddings,
)
print(evaluator.run())

Classification uses the L2 norm of the pooled alternate-minus-reference embedding difference. This measures how much the embedding changed, not whether the measured phenotype increased or decreased.

Signed regression effects

python
import torch
import mrna_bench as mb

from mrna_bench.embedder import DatasetEmbedder
from mrna_bench.linear_probe import LinearProbeBuilder

dataset = mb.load_dataset("rna-stability-siegel-jurkat")
model = mb.load_model(
    "Orthrus",
    "orthrus-large-4-track",
    device=torch.device("cuda"),
)

embedding_tensors = DatasetEmbedder(
    model,
    dataset,
    batch_size=8,
).embed_dataset()
embeddings = torch.stack(embedding_tensors).cpu().numpy()

effect_spec = dataset.metadata.vep_task_spec
assert effect_spec is not None

effect_probe = (
    LinearProbeBuilder(dataset)
    .fetch_embedding_by_embedding_instance(
        model.short_name,
        embeddings,
    )
    .set_target(effect_spec.target_col)
    .build_evaluator(effect_spec.task)
    .set_regressor("ridge")
    .build_splitter(
        dataset.metadata.default_split_type,
        eval_all_splits=True,
    )
    .build()
)

print(effect_probe.run_linear_probe(random_seed=2541))

A pooled embedding has no universal positive direction. For a labeled regression dataset, the builder pairs each alternate with its reference and computes the alternate-minus-reference embedding. Ridge then learns from the training split which directions correspond to higher or lower measured effects. This route is supervised.

For a zero-shot regression score, the sign must come from the model output or from a biological rule fixed before evaluation. It cannot be taken from the held-out effect label.

Likelihood scoring

python
import torch
import mrna_bench as mb

from mrna_bench.zeroshot import ZeroShotVEP

dataset = mb.load_dataset("vep-traitgym-mendelian")
model = mb.load_model(
    "Evo2",
    "Evo2-7B-8K",
    device=torch.device("cuda"),
    attn_implementation="flash_attention_2",
)

evaluator = ZeroShotVEP.from_model(
    dataset,
    model,
    score_method="causal_likelihood",
    normalization="sum",
)
print(evaluator.run())

Causal and pseudo-likelihood support depends on the model.masked_marginal requires pseudo-likelihood support and scores substitutions only.

Optional LoRA fine-tuning

Install
python -m pip install --no-build-isolation \
  'mrna-bench[base_models,fine_tune]'
python
import numpy as np
import torch
import mrna_bench as mb

from mrna_bench.fine_tune import (
    FineTuneTrainer,
    FineTuneWrapper,
    TaskHead,
    TrainerConfig,
    create_dataloaders,
)

dataset = mb.load_dataset("mrl-sugimoto")
model = mb.load_model(
    "Orthrus",
    "orthrus-large-4-track",
    device=torch.device("cuda"),
)

row = dataset.data_df.iloc[0]
sample_embedding = model.embed(
    [row["sequence"]],
    cds=[np.asarray(row["cds"])],
    splice=[np.asarray(row["splice"])],
)[0]

head = TaskHead(
    input_dim=sample_embedding.shape[-1],
    output_dim=1,
    task_type="regression",
)
wrapper = FineTuneWrapper(model, head)
wrapper.apply_lora(rank=8, alpha=16)

train, validation, test = create_dataloaders(
    dataset=dataset,
    target_col="target",
    split_type=dataset.metadata.default_split_type,
    random_seed=2541,
    batch_size=32,
)
trainer = FineTuneTrainer(
    wrapper,
    TrainerConfig(
        learning_rate=1e-4,
        epochs=15,
        lr_schedule="cosine",
    ),
)
trainer.fit(train, validation)
print(trainer.evaluate(test))

The wrapper applies LoRA adapters to selected backbone modules, whileTaskHead and FineTuneTrainer provide the task-specific loss, early stopping, and metrics.

This in-memory example prints metrics. UseFineTunePersister to write a JSON result underft_results/.

Batch and cluster wrappers

A source checkout also includes thin command-line wrappers around these APIs for scheduling larger runs: