mlm-baseline-hg38

A masked nucleotide language model for the Homo sapiens genome (hg38).

What it predicts

The model is a masked language model over DNA: it learns the distribution of each nucleotide given its surrounding sequence context. At every position it outputs a probability distribution over the four nucleotides (A, C, G, T) — what base it expects there from the context alone. It is trained on the Homo sapiens genome, and is useful for scoring how unexpected a variant is (its likelihood under the model) and as a sequence feature extractor.

Inputs and outputs

The model takes a LongTensor of shape (batch, length) that encodes DNA over a six-token alphabet: A, C, G, T, N, and a padding/mask symbol, mapped to indices 0 through 5. It returns a FloatTensor of shape (batch, length - 1548, 4): at each in-bounds position, four logits over [A, C, G, T]. A softmax over them is the predicted nucleotide distribution. Because it is a masked model, to read the prediction for a position you first replace that position's input token with the mask token (index 5).

Status

This revision is a partial run: it was released at step 93,750, roughly half of the configured cosine cycle, and the learning rate had not annealed. It is published so downstream benchmarks can run; a completed checkpoint is expected to replace it in a later revision. Pin a revision if you need reproducibility.

Training

Trained by masked language modeling: about 15% of input tokens are replaced with the mask token (index 5), and the loss is the cross-entropy between the model's per-position distribution and the true nucleotide at the masked positions. The network is reverse-complement equivariant by construction, so a sequence and its reverse complement receive consistent predictions.

How to use

from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
import importlib.util, torch

repo_id = "songlab/mlm-baseline-hg38"
model_path = hf_hub_download(repo_id=repo_id, filename="model.py")
weights_path = hf_hub_download(repo_id=repo_id, filename="model.safetensors")
spec = importlib.util.spec_from_file_location("released_model", model_path)
module = importlib.util.module_from_spec(spec); spec.loader.exec_module(module)

model = module.CNN()  # defaults match the released checkpoint
model.load_state_dict(load_file(weights_path))
model.eval()

# Input: a (batch, length) LongTensor over [A, C, G, T, N, pad], length >= 1549.
x = torch.randint(0, 4, (1, 2000))
with torch.no_grad():
    output = model(x)         # (1, length - 1548, 4): per-position logits
    embeddings = model.encode(x)  # (1, length - 1548, 512): per-position features

This is a masked language model: to read the prediction for a position, replace that position's input token with the mask token (index 5) before the forward pass, then take a softmax over the four output logits. The output window is centered, so input position p maps to output index p - 774:

masked = x.clone()
masked[0, 874] = 5  # mask one position
with torch.no_grad():
    probs = model(masked).softmax(-1)
    prediction = probs[0, 100]  # distribution over [A, C, G, T] at the masked position

Outputs cover only positions whose full 1549 bp context is in bounds. encode returns the per-position features (use it for embeddings).

Files

  • model.py: a self-contained CNN class with no dependencies beyond PyTorch.
  • model.safetensors: the released weights (~1.06 GB).
  • config.json: architecture hyperparameters, tokenization, and training metadata.

License

MIT

Downloads last month
30
Safetensors
Model size
0.3B params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support