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.compileenabled. 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
- Code: https://github.com/NNeuralDynamics/Spliceformer
- Data:
SaumyaGupta-99/spliceformer-data - GTEx tissues:
SaumyaGupta-99/spliceformer-gtex-tissues - HAEC training data:
mrunyan1/haec-training-data
License
MIT