Spliceformer / Merlin โ€” model weights

Weights for Spliceformer, a rotary-position transformer over raw DNA. Merlin is the shared 6-block encoder, pretrained on the human genome with a masked-nucleotide objective; every other checkpoint here is a fine-tune of it.

โš ๏ธ Requires an Ampere (SM80) or newer NVIDIA GPU โ€” A100, H100, RTX 30/40/50. Attention is FlashAttention-2, which has no CPU, MPS or pre-Ampere path.

โš ๏ธ Keep torch.compile enabled. All reported metrics were produced with compilation on. Eager execution changes results in the third decimal under bfloat16 autocast.

Contents

Folder Files What
merlin/ 1 Pretrained backbone (merlin_mlm_6blocks_best.pth, 220 MB)
splice_gencode_10k/ 6 Splice sites, 10 kb context, GENCODE labels (5-seed ensemble)
splice_gencode_400/ 6 Splice sites, 400 nt context, GENCODE labels (5-seed ensemble)
haec_ensemble/ 11 HAEC joint classification + usage regression, 4 folds ร— {cls, reg, joint}
adar/ 2 ADAR A-to-I editing
m6a/ 198 m6A methylation, 11 tissues ร— 3 strategies ร— 5 seeds
rbp/ 108 RBP binding, 37 ENCODE eCLIP targets (single architecture)

Quick start

git clone https://github.com/NNeuralDynamics/Spliceformer.git && cd Spliceformer
pip install -e . && pip install flash-attn --no-build-isolation

python scripts/download_assets.py --merlin                      # backbone only
python scripts/download_assets.py --models splice_gencode_10k   # + splice ensemble
from spliceformer import Paths, SpliceClassifier, load_finetuned
from spliceformer.training import compile_model

model = SpliceClassifier(
    transformer_block_depth=6, embedding_length=512,
    dropout_rate=0.1, attn_dropout=0.05, context_length=10000,
).cuda().eval()
model = compile_model(model)                                    # keep this
load_finetuned(model, Paths().splice_checkpoint("10k", seed=42))

load_finetuned normalises checkpoint prefixes, so a file loads whether the model is currently compiled, DDP-wrapped, or plain.

Checkpoint prefixes

Files were saved from different wrappings, and each adds a state-dict key prefix:

Saved from Prefix Which
raw nn.Module (none) ADAR, m6A, RBP
torch.compile(model) _orig_mod. splice
DDP(torch.compile(model)) module._orig_mod. HAEC

Two things to know before you use these

best_model_400_6blocks_gencode.pth has a legacy head. It predates the seeded ensemble and puts an extra LayerNorm at index 0. infer_splice_head() detects it. Prefer the _seed* files.

RBP is now a single architecture. An earlier release mixed two head generations over disjoint protein sets; the 171 legacy-head checkpoints (15/5/3 conv kernels) have been removed so that everything here uses the current 7/7/7 head. All 108 files are 100.0 MB, covering 37 proteins across frozen / partial / full. Results are directly comparable across proteins.

The HAEC ensemble is 5 folds ร— 3 heads = 15 checkpoints. Four are pending upload (ensemble2_4_reg and all of fold 5). Scripts default to all five folds, skip what is absent, and record the contributing folds in their output, so a partial set cannot be mistaken for a full 5-fold run.

Verification

All 502 fine-tuned checkpoints plus the backbone load with strict=True against the model classes in the repo (pytest tests/test_checkpoints.py).

Related

License

MIT

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. ๐Ÿ™‹ Ask for provider support