Reference: Python API

Reference

Python API

Signatures and parameter notes for the public objects used in the guides.

This reference is generated from the package source during the website build. The workflow guides show how these objects fit together; use this page when you need an exact signature or argument name.

39 entries

Load and configure

Top-level functions used before running an evaluation.

mrna_benchload_datasetLoad a registered dataset.
load_dataset(dataset_name: str, force_redownload_hf: bool = False, force_rebuild_raw: bool = False) -> BenchmarkDataset

Parameters

dataset_name
Name of the dataset.
force_redownload_hf
Forces redownload from HuggingFace.
force_rebuild_raw
Forces rebuild from raw data source.

Returns

Initialized benchmark dataset.

Source: mrna_bench/loader/loader.py

mrna_benchload_modelLoad a registered model and set it to inference mode.
load_model(model_name: str, model_version: str | None = None, device: 'torch.device | None' = None, attn_implementation: str | None = None) -> 'EmbeddingModel'

Parameters

model_name
Name of model class.
model_version
Specific model version to load. Defaults to model's default_version if not specified.
device
PyTorch device to load model to. Defaults to CUDA if available.
attn_implementation
Attention implementation override. One of "eager", "sdpa", or "flash_attention_2". Pass None to use the model's default.

Returns

Initialized EmbeddingModel in inference mode.

Source: mrna_bench/loader/loader.py

mrna_benchupdate_data_pathSet the directory used for processed data, embeddings, and results.
update_data_path(path_to_data: str)

Parameters

path_to_data
New path to directory where data is stored.

Source: mrna_bench/utils.py

mrna_benchget_data_pathReturn the configured data directory.
get_data_path() -> str

Returns

Directory where benchmark data is stored.

Source: mrna_bench/utils.py

mrna_benchupdate_model_weights_pathSet the directory used for downloaded model weights.
update_model_weights_path(path_to_weights: str)

Parameters

path_to_weights
New path to directory where model weights are stored.

Source: mrna_bench/utils.py

mrna_benchget_model_weights_pathReturn the configured model-weights directory.
get_model_weights_path() -> str

Returns

Directory where model weights are stored.

Source: mrna_bench/utils.py

Datasets and splits

Dataset metadata, paired variants, and train-test splits.

mrna_bench.datasetsDatasetMetadataMetadata used to select targets, splits, and evaluation routes.
DatasetMetadata(dataset_name: str, species: str, task: list[str], target_col: list[str], default_split_type: str, benchmark_set: str, evaluations: tuple[EvaluationMethod | str, ...], variant_region: str | None = None, vep_target_col: str | None = None)

Source: mrna_bench/datasets/benchmark_dataset.py

mrna_bench.datasetsBenchmarkDataset.get_splitsGet data splits for the dataset.
BenchmarkDataset.get_splits(split_ratios: tuple[float, float, float], random_seed: int = 2541, split_type: str | None = None, split_kwargs: dict | None = None) -> dict[str, pd.DataFrame]

Parameters

split_ratios
Ratios for train, val, test splits.
random_seed
Random seed for reproducibility.
split_type
Type of split to use. Defaults to metadata default.
split_kwargs
Additional arguments for the split type.

Returns

Dictionary of dataframes containing splits.

Source: mrna_bench/datasets/benchmark_dataset.py

mrna_bench.datasetsBenchmarkDataset.get_vep_pairsReturn the dataset's variants in canonical ref/alt columns.
BenchmarkDataset.get_vep_pairs(dataframe: pd.DataFrame, value_columns: tuple[str, ...] = ('sequence', 'cds', 'splice')) -> pd.DataFrame

Parameters

dataframe
Dataset rows containing reference and alternate variants.
value_columns
Columns to copy into ref- and alt-prefixed output columns.

Returns

One row per variant with aligned reference and alternate values.

Source: mrna_bench/datasets/benchmark_dataset.py

Models and embeddings

Model outputs and dataset-wide embedding generation.

mrna_bench.modelsEmbeddingModel.embedEmbed sequences, optionally using cds/splice tracks.
EmbeddingModel.embed(sequences: list[str], cds: list[np.ndarray] | None = None, splice: list[np.ndarray] | None = None, agg_fn: Callable = mean_pool) -> list[torch.Tensor]

Parameters

sequences
List of nucleotide sequences to embed (uses DNA bases).
cds
List of binary encodings of first nucleotide of each codon.
splice
List of binary encodings of splice site locations.
agg_fn
Method used to aggregate across sequence dimension.

Returns

Embedded sequences with shape (batch_size x H).

Source: mrna_bench/models/embedding_model.py

mrna_bench.modelsEmbeddingModel.sequence_scoreScore complete sequences with a model-native likelihood method.
EmbeddingModel.sequence_score(sequences: list[str], method: ModelBehavior | str | None = None, normalization: str = 'mean', cds: list[np.ndarray] | None = None, splice: list[np.ndarray] | None = None) -> list[float]

Parameters

sequences
Nucleotide sequences to score.
method
Causal or pseudo-likelihood behavior to use. May be omitted when the model supports exactly one likelihood.
normalization
Return the mean or sum of token log-probabilities.
cds
Optional CDS tracks aligned to sequences.
splice
Optional splice tracks aligned to sequences.

Returns

One log-likelihood score per sequence.

Source: mrna_bench/models/embedding_model.py

mrna_bench.modelsEmbeddingModel.masked_marginal_llrCompare alleles by masking only tokenizer positions they change.
EmbeddingModel.masked_marginal_llr(reference_sequences: list[str], alternate_sequences: list[str], normalization: str = 'mean', cds: list[np.ndarray] | None = None, splice: list[np.ndarray] | None = None) -> list[float]

Parameters

reference_sequences
Reference allele sequences.
alternate_sequences
Alternate allele sequences aligned to the reference sequences.
normalization
Return the mean or sum over changed token scores.
cds
Optional CDS tracks aligned to reference sequences.
splice
Optional splice tracks aligned to reference sequences.

Returns

Reference-minus-alternate log-likelihood ratios.

Source: mrna_bench/models/embedding_model.py

mrna_bench.embedderDatasetEmbedderGenerate and save embeddings for a dataset or sequence dataframe.
DatasetEmbedder(model: EmbeddingModel, dataset: BenchmarkDataset, d_chunk_ind: int = 0, d_num_chunks: int = 0, agg_fn: Callable = mean_pool, ragged_out: bool = False, batch_size: int = 1)

Parameters

model
Model used to embed sequences.
dataset
Dataset to embed.
d_chunk_ind
Current dataset chunk to be processed.
d_num_chunks
Total number of chunks to divide dataset into.
agg_fn
Aggregation function to apply to sequence embeddings.
ragged_out
Whether the model produces ragged output under agg_fn.
batch_size
Number of dataset rows passed to model.embed at once.

Source: mrna_bench/embedder/dataset_embedder.py

mrna_bench.embedderDatasetEmbedder.from_dataframeCreate a DatasetEmbedder from a custom sequence dataframe.
DatasetEmbedder.from_dataframe(model: EmbeddingModel, data_df: pd.DataFrame) -> 'DatasetEmbedder'

Source: mrna_bench/embedder/dataset_embedder.py

mrna_bench.embedderDatasetEmbedder.embed_datasetCompute embeddings for current dataset chunk.
DatasetEmbedder.embed_dataset() -> list[torch.Tensor]

Returns

Embeddings for current dataset chunk in original order. - pooled: (1, H) - unpooled: (1, L_i, H)

Source: mrna_bench/embedder/dataset_embedder.py

mrna_bench.embedderDatasetEmbedder.persist_embeddingsPersist embeddings at global data storage location.
DatasetEmbedder.persist_embeddings(embeddings: list[torch.Tensor])

Parameters

embeddings
Embeddings to persist.

Source: mrna_bench/embedder/dataset_embedder.py

mrna_bench.embedderDatasetEmbedder.merge_embeddingsMerge persisted processed dataset chunks into single file.
DatasetEmbedder.merge_embeddings()

Source: mrna_bench/embedder/dataset_embedder.py

Linear probes

Builder configuration, fitted probes, and seeded runs.

mrna_bench.linear_probeLinearProbeBuilderConfigure embeddings, splits, targets, estimators, and persistence.
LinearProbeBuilder(dataset: BenchmarkDataset | None = None, dataset_name: str | None = None)

Parameters

dataset
BenchmarkDataset to linearly probe.
dataset_name
Name of dataset to linearly probe.

Source: mrna_bench/linear_probe/linear_probe_builder.py

mrna_bench.linear_probeLinearProbeBuilder.fetch_embedding_by_embedding_instanceStore embeddings for LinearProbe using an embedding instance.
LinearProbeBuilder.fetch_embedding_by_embedding_instance(model_short_name: str, embedding: np.ndarray) -> 'LinearProbeBuilder'

Parameters

model_short_name
Short name of model used to generate embeddings.
embedding
Locally generated embedding for dataset.

Returns

LinearProbeBuilder with set embeddings.

Source: mrna_bench/linear_probe/linear_probe_builder.py

mrna_bench.linear_probeLinearProbeBuilder.build_splitterSet data splitter for LinearProbe.
LinearProbeBuilder.build_splitter(split_type: str, split_ratios: tuple[float, float, float] = (0.7, 0.15, 0.15), eval_all_splits: bool = False, **split_args) -> 'LinearProbeBuilder'

Parameters

split_type
Method used for data split generation.
split_ratios
Ratio of data split sizes as a fraction of dataset.
eval_all_splits
Evaluate metrics on all splits. Only evaluates validation split otherwise.
**split_args
Additional arguments for data splitter.

Returns

LinearProbeBuilder with set data splitter.

Source: mrna_bench/linear_probe/linear_probe_builder.py

mrna_bench.linear_probeLinearProbeBuilder.set_targetSet linear probing target column.
LinearProbeBuilder.set_target(target_col: str) -> 'LinearProbeBuilder'

Parameters

target_col
Column from dataframe to use as labels.

Returns

LinearProbeBuilder with set task and target column.

Source: mrna_bench/linear_probe/linear_probe_builder.py

mrna_bench.linear_probeLinearProbeBuilder.build_evaluatorSet evaluator for LinearProbe.
LinearProbeBuilder.build_evaluator(task: str) -> 'LinearProbeBuilder'

Parameters

task
Linear probing task.

Returns

LinearProbeBuilder with set evaluator.

Source: mrna_bench/linear_probe/linear_probe_builder.py

mrna_bench.linear_probeLinearProbeBuilder.set_regressorSelect the estimator used for regression tasks.
LinearProbeBuilder.set_regressor(regressor: str) -> 'LinearProbeBuilder'

Parameters

regressor
Regression estimator, either ``ols`` or ``ridge``.

Returns

LinearProbeBuilder with the selected regression estimator.

Source: mrna_bench/linear_probe/linear_probe_builder.py

mrna_bench.linear_probeLinearProbeBuilder.use_persisterIndicate that persister for LinearProbe should be built.
LinearProbeBuilder.use_persister() -> 'LinearProbeBuilder'

Returns

LinearProbeBuilder with persister flag set.

Source: mrna_bench/linear_probe/linear_probe_builder.py

mrna_bench.linear_probeLinearProbeBuilder.buildBuild a LinearProbe instance.
LinearProbeBuilder.build() -> LinearProbe

Returns

Configured linear probe ready to run.

Source: mrna_bench/linear_probe/linear_probe_builder.py

mrna_bench.linear_probeLinearProbe.run_linear_probePerform data split and run linear probe.
LinearProbe.run_linear_probe(random_seed: int = 2541, persist: bool = False, dropna: bool = True) -> dict[str, float]

Parameters

random_seed
Random seed used for data split.
persist
Save results to disk.
dropna
Drop rows with NaN values in target column

Returns

Dictionary of linear probing metrics per split.

Source: mrna_bench/linear_probe/linear_probe.py

mrna_bench.linear_probeLinearProbe.linear_probe_multirunRun one linear probe per random seed.
LinearProbe.linear_probe_multirun(random_seeds: list[int], persist: bool = False) -> dict[int, dict[str, float]]

Parameters

random_seeds
Random seeds used to generate dataset splits.
persist
Save each run to the configured persister.

Returns

Metrics for each random seed.

Source: mrna_bench/linear_probe/linear_probe.py

mrna_bench.linear_probeLinearProbe.get_fit_modelReturn the model trained for a specific split seed.
LinearProbe.get_fit_model(random_seed: int) -> BaseEstimator

Parameters

random_seed
Random seed identifying the fitted split.

Returns

Estimator fitted for the requested seed.

Source: mrna_bench/linear_probe/linear_probe.py

Variant effects

Embedding-difference and likelihood-based VEP.

mrna_bench.zeroshotZeroShotVEP.from_embeddingsBuild embedding VEP directly from a dataset and embeddings.
ZeroShotVEP.from_embeddings(dataset: 'BenchmarkDataset', embeddings: np.ndarray, target_col: str | None = None, task: str | None = None, **kwargs) -> 'ZeroShotVEP'

Source: mrna_bench/zeroshot/vep.py

mrna_bench.zeroshotZeroShotVEP.from_modelBuild likelihood VEP directly from a dataset and model.
ZeroShotVEP.from_model(dataset: 'BenchmarkDataset', model: 'EmbeddingModel', score_method: 'ModelBehavior | str | None' = None, target_col: str | None = None, task: str | None = None, normalization: str = 'sum', likelihood_direction: str | None = None, **kwargs) -> 'ZeroShotVEP'

Source: mrna_bench/zeroshot/vep.py

mrna_bench.zeroshotZeroShotVEP.runCompute zero-shot VEP scores and evaluate.
ZeroShotVEP.run(persist: bool = False) -> dict[str, float]

Parameters

persist
Write results to results.db when True.

Returns

Classification or regression metrics.

Source: mrna_bench/zeroshot/vep.py

Fine-tuning

Optional task heads, LoRA adapters, and training loops.

mrna_bench.fine_tuneTaskHeadPrediction head for regression, classification, or multilabel tasks.
TaskHead(input_dim: int, output_dim: int, task_type: str, hidden_dims: list[int] | None = None, dropout: float = 0.1)

Parameters

input_dim
Embedding dimension from backbone model.
output_dim
Number of output units (classes or targets).
task_type
One of regression, classification, multilabel.
hidden_dims
Hidden layer dimensions. None for linear head.
dropout
Dropout probability between layers.

Source: mrna_bench/fine_tune/task_heads.py

mrna_bench.fine_tuneFineTuneWrapperCombine an EmbeddingModel with a task head and optional LoRA adapters.
FineTuneWrapper(model: EmbeddingModel, task_head: nn.Module)

Parameters

model
EmbeddingModel instance to use as backbone.
task_head
Task head module. Must have task_type and get_loss_fn() attributes.

Source: mrna_bench/fine_tune/fine_tune_wrapper.py

mrna_bench.fine_tuneFineTuneWrapper.apply_loraApply LoRA adapters to the backbone model.
FineTuneWrapper.apply_lora(rank: int = 8, alpha: int = 16, dropout: float = 0.0, target_modules: list[str] | None = None)

Parameters

rank
Rank of LoRA decomposition.
alpha
LoRA scaling factor.
dropout
Dropout probability for LoRA layers.
target_modules
Module names to apply LoRA to.

Source: mrna_bench/fine_tune/fine_tune_wrapper.py

mrna_bench.fine_tuneTrainerConfigOptions used by FineTuneTrainer.
TrainerConfig(learning_rate: float = 0.0001, epochs: int = 10, warmup_steps: int = 100, early_stopping_patience: int = 3, gradient_accumulation_steps: int = 1, max_grad_norm: float = 1.0, lr_schedule: str = 'none', total_steps: int | None = None)

Source: mrna_bench/fine_tune/trainer.py

mrna_bench.fine_tuneFineTuneTrainerTrain and evaluate a FineTuneWrapper.
FineTuneTrainer(wrapper: FineTuneWrapper, config: TrainerConfig | None = None)

Parameters

wrapper
FineTuneWrapper with backbone and task head.
config
Training configuration. Uses defaults if not provided.

Source: mrna_bench/fine_tune/trainer.py

mrna_bench.fine_tuneFineTuneTrainer.fitFull training loop with early stopping and best-model restore.
FineTuneTrainer.fit(train_dataloader: DataLoader, val_dataloader: DataLoader | None = None) -> dict[str, list[float]]

Parameters

train_dataloader
Training data loader.
val_dataloader
Validation data loader (optional).

Returns

Training history dictionary.

Source: mrna_bench/fine_tune/trainer.py

mrna_bench.fine_tuneFineTuneTrainer.evaluateEvaluate model on validation data.
FineTuneTrainer.evaluate(dataloader: DataLoader) -> dict[str, float]

Parameters

dataloader
Validation data loader.

Returns

Dictionary of evaluation metrics.

Source: mrna_bench/fine_tune/trainer.py

mrna_bench.fine_tunecreate_dataloadersCreate train/val/test DataLoaders from a BenchmarkDataset.
create_dataloaders(dataset: BenchmarkDataset, target_col: str, split_type: str, random_seed: int, batch_size: int, split_ratios: tuple[float, float, float] = (0.7, 0.15, 0.15)) -> tuple[DataLoader, DataLoader, DataLoader]

Parameters

dataset
Benchmark dataset instance.
target_col
Target column name in the dataset dataframe.
split_type
Type of data split (e.g. "default", "homology").
random_seed
Random seed for reproducible splits.
batch_size
Batch size for DataLoaders.
split_ratios
Train/val/test split ratios.

Returns

Tuple of (train_loader, val_loader, test_loader).

Source: mrna_bench/fine_tune/dataloader.py