RAGenome

Paper: RAGenome: Scaling Retrieval-Based Genomic Language Models to Long Contexts

RAGenome is a retrieval-based genomic language model (168M parameters) trained on a whole-genome alignment of 100 vertebrates. It retrieves homologous sequences from the alignment at pretraining and inference time, letting the model capture evolutionary signals explicitly while scaling to much longer contexts (13,312 nucleotides) than MSA-based genomic LMs.

Usage

1. Import and load the model.

import torch
from transformers import AutoModel

model = AutoModel.from_pretrained(
    "pantoniadis/RAGenome", trust_remote_code=True
).to("cuda").bfloat16().eval()

2. Load the tokenizer and tokenize your query sequence.

from transformers import AutoTokenizer

tokenizer = AutoTokenizer.from_pretrained("pantoniadis/RAGenome")
genome_input_ids = tokenizer("ACGTAC", return_tensors="pt")["input_ids"].to("cuda")

3. Build the retrieved-species inputs. model(...) expects the retrieved sequences already packed into a single array ([species_0 | species_1 | ... | species_{N-1}]), with gap tokens removed and each retrieved token's original alignment-column index kept as its position.

B, L, N = genome_input_ids.shape[0], genome_input_ids.shape[1], 5  # batch size, query length, number of retrieved species
lengths = torch.randint(50, 500, (B, N), device="cuda")   # real length of each retrieved species
max_total = int(lengths.sum(1).max())

packed_retrieved_ids = torch.randint(7, 11, (B, max_total), device="cuda")      # retrieved tokens, packed
packed_aligned_pos = torch.randint(0, 100_000, (B, max_total), device="cuda")   # alignment column of each retrieved token
retrieved_taxonomy = torch.randint(0, 317, (B, N, 8), device="cuda")            # NCBI lineage of each retrieved species
query_taxonomy = torch.randint(0, 317, (B, 8), device="cuda")                   # NCBI lineage of the query species (human)

The snippet above uses random tensors to show the expected shapes. To build the inputs from a real whole-genome alignment, see scripts/inference.py.

4. Run the model.

out = model(
    genome_input_ids=genome_input_ids,
    packed_retrieved_ids=packed_retrieved_ids,
    retrieved_lengths=lengths,
    packed_aligned_pos=packed_aligned_pos,
    retrieved_taxonomy=retrieved_taxonomy,
    query_taxonomy=query_taxonomy,
)
out.logits  # (B, L, 11)
Downloads last month
54
Safetensors
Model size
0.2B params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Space using pantoniadis/RAGenome 1

Paper for pantoniadis/RAGenome