Instructions to use pantoniadis/RAGenome with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use pantoniadis/RAGenome with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="pantoniadis/RAGenome", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("pantoniadis/RAGenome", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
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