models/modernbert-baseREADME.md5.7 kB
md
# ModernBERT-base
Answer.AI and LightOn's 2024 rebuild of the BERT encoder. Only the job is the
same: 22 layers of width 768 with 12 heads, reading a whole sequence at once
and producing one vector per token. Everything inside the layer changed.
Normalization moved in front of both residual branches, every bias was
deleted (`attention_bias`, `mlp_bias` and `norm_bias` are all false), the
absolute position table is gone, and the feed-forward network became a GeGLU.
The part with no analogue in BERT is the attention pattern. Only every third
layer attends over the whole sequence; the two layers between each of those
attend inside a sliding window 128 positions wide, so most of the depth costs
work linear in the sequence length instead of quadratic. The two kinds of
layer also use different rotary base frequencies -- 160000 for the global
layers, which have to resolve positions thousands of tokens apart, and 10000
for the local ones, which never see more than half a window. Missing that
second detail changes every attention score in the model.
## The alternation is in the block structure
A `static for` loop variable is an `i64` scalar, not a compile-time integer,
so it cannot index a sub-block array: there is no way to write "if this
layer's index is a multiple of three". The source encodes the cycle in the
types instead. `LayerCycle` holds the `Period - 1` windowed layers and the
global layer that closes the cycle, and `Model` holds an array of cycles:
```linnet
sub local: [EncoderLayer<D, Heads, Inner, T>; Period - 1]
sub global: EncoderLayer<D, Heads, Inner, T>
```
Each member is reached by name, so each gets its own mask and its own rotary
base with no index arithmetic, and `where (Layers - 1) % Period == 0` makes
the compiler check that the depth actually divides into cycles. Layer 0 is a
global layer too, but it is the one layer whose attention normalization
upstream replaces with an identity -- the embedding LayerNorm has just run --
so it is peeled off as `FirstLayer`, a block with no `attn_norm` at all
rather than an optional parameter that is absent.
## The sliding window
The window is a `[Q, K]` band, which is exactly the shape
`std.nn.attention::attention` takes, so the source writes its own mask in the
style of `causal_mask`:
```linnet
op band_mask<Q: Dim, K: Dim, Window: Dim>() -> Tensor[Q, K; bool] {
let queries = iota<i32>(Q)
let keys = iota<i32>(K)
let mask[q, k] = abs(keys[k] - queries[q]) <= cast<i32>(Window / 2)
return mask
}
```
`local_attention` counts the whole window, so 128 means 64 back and 64
forward, 129 positions in all. A custom mask costs nothing in the kernel
selection that matters: `linnet explain --numerics fast` still picks
`torch.nn.functional.scaled_dot_product_attention` for the attention itself,
because that registration only requires a boolean mask whose true values mean
attend. What a private op does give up is the mask's own kernel -- the stdlib
`causal_mask` lowers to `torch.tril`, while `band_mask` is built from two
`iota`s every time.
## The fused projections
Two weights in each layer hold two or three tensors' worth of rows, and both
are kept exactly as the checkpoint stores them:
- `attn.Wqkv` is `[3 * D, D]`, laid out as query, then key, then value, each
already head by head, so the three parts are three slices of the last axis
of its output.
- `mlp.Wi` is `[2 * Inner, D]`. The **first** half is the branch GELU runs
on and the **second** half multiplies it (`input, gate = Wi(x).chunk(2)`
upstream). Reading the halves the other way round is a silent error.
## Loading
```python
from linnet import nest
model = nest.load("modernbert-base", backend="torch")
hidden = model(tokens) # Tensor[B, S; i32] -> Tensor[B, S, 768; f32]
```
The weights are the published `model.safetensors` (f32) of the masked-LM
checkpoint; `bindings.json` maps every Linnet parameter path to its tensor
name, including the shift between this source's cycle structure and the flat
`model.layers.N` of the checkpoint. The masked-LM head itself
(`head.dense`, `head.norm`, `decoder.bias`) is not part of this card.
Every position attends over every other position it is allowed to see. There
is no padding mask: `std.nn.attention::attention` takes an unbatched
`Tensor[Q, K; bool]`, so a per-sequence mask cannot be expressed today, and a
batch here has to be sequences of the same real length.
## Entries
| Entry | |
| --- | --- |
| `forward<B, S>(tokens)` | the last hidden state, after the final norm |
## Why the GELU matters
`hidden_activation` in the config is plain `gelu`, which to `transformers`
means the error-function form, so the source uses
`std.nn.activations::gelu_erf` rather than the tanh approximation `gelu`.
The two differ by at most 4.7e-4 per activation, which is normally an
acceptable rounding. It is not here, and the reason is worth recording: a
GeGLU multiplies the activated half by an unbounded gate, and twenty-two
pre-norm layers add the result straight into the residual stream. With the
tanh form the last hidden state came out 6.9e-1 from the reference, and
substituting the tanh GELU into `transformers` brought its own output back to
within 1.2e-4 of the Linnet one -- which is how the cause was pinned down
rather than guessed, and why the standard library now has `gelu_erf`.
## Provenance
- Weights: [answerdotai/ModernBERT-base](https://huggingface.co/answerdotai/ModernBERT-base), Apache-2.0.
- Code: [AnswerDotAI/ModernBERT](https://github.com/AnswerDotAI/ModernBERT).
- Paper: [Smarter, Better, Faster, Longer](https://arxiv.org/abs/2412.13663).
The last hidden state matches `transformers.AutoModel` to **1.6e-4** absolute
in f32, over a fixed 256-token sequence -- twice the local window, so the
windowed layers really do mask most of the sequence rather than quietly
degenerating to full attention.bench.json8.9 kB
json
{
"date": "2026-09-28T18:30:34+00:00",
"environment": {
"python": "3.12.3",
"machine": "x86_64",
"torch": "2.14.0",
"jax": "0.11.2",
"transformers": "5.17.0",
"diffusers": "0.40.0",
"onnxruntime": "1.30.0",
"linnet": "0.1.0",
"gpu": "NVIDIA H100 80GB HBM3",
"gpus": "1",
"driver": "580.126.09"
},
"workload": {
"device": "cuda",
"prompt_tokens": 512,
"new_tokens": 128,
"batch": 1,
"warmup": 3,
"iters": 10,
"seed": 0,
"offload_gib": 8.0,
"serve_requests": 256,
"serve_concurrency": 64,
"serve_prompt_min": 128,
"serve_prompt_max": 512,
"serve_new": 128
},
"methods": [
{
"method": "transformers (eager)",
"kind": "reference",
"metrics": {
"latency_ms": 9.18653141707182,
"throughput_per_s": 1357.799344892393,
"load_s": 3.183003354817629,
"peak_vram_mib": 1922.0
},
"max_abs_diff": 0.0,
"notes": "unpadded batches of exactly 256 tokens (this card takes no padding mask)",
"error": null,
"key": "transformers-eager"
},
{
"method": "transformers (torch.compile)",
"kind": "reference",
"metrics": {
"latency_ms": 2.311496064066887,
"throughput_per_s": 4488.460918806441,
"load_s": 1.3915417324751616,
"peak_vram_mib": 1374.0
},
"max_abs_diff": 2.127197265625,
"notes": "unpadded batches of exactly 256 tokens (this card takes no padding mask)",
"error": null,
"key": "transformers-compile"
},
{
"method": "Linnet torch (generated source)",
"kind": "linnet",
"metrics": {
"latency_ms": 5.783703178167343,
"throughput_per_s": 2779.841462719057,
"load_s": 1.38893648609519,
"peak_vram_mib": 1304.0
},
"max_abs_diff": 1.625,
"notes": "unpadded batches of exactly 256 tokens; bf16 on cuda",
"error": null,
"key": "linnet-torch"
},
{
"method": "Linnet torch (CUDA graphs)",
"kind": "linnet",
"metrics": {
"latency_ms": 1.0142279788851738,
"throughput_per_s": 6152.875290063452,
"load_s": 1.35607635602355,
"peak_vram_mib": 1442.0
},
"max_abs_diff": 1.25390625,
"notes": "unpadded batches of exactly 256 tokens; bf16 on cuda",
"error": null,
"key": "linnet-cudagraphs"
},
{
"method": "Linnet torch (torch.compile, inductor)",
"kind": "linnet",
"metrics": {
"latency_ms": 2.151189371943474,
"throughput_per_s": 6041.652441098496,
"load_s": 1.0256760194897652,
"peak_vram_mib": 1306.0
},
"max_abs_diff": 1.25390625,
"notes": "unpadded batches of exactly 256 tokens; bf16 on cuda",
"error": null,
"key": "linnet-inductor"
},
{
"method": "Linnet JAX (XLA, StableHLO)",
"kind": "linnet",
"metrics": {
"latency_ms": 1.9907327368855476,
"throughput_per_s": 3251.68707327197,
"load_s": 0.5805029682815075,
"peak_vram_mib": 927.1,
"driver_vram_mib": 1648.0
},
"max_abs_diff": 1.46875,
"notes": "unpadded batches of exactly 256 tokens; bf16 on cuda; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
"error": null,
"key": "linnet-jax"
},
{
"method": "Linnet JAX (XLA, generated source)",
"kind": "linnet",
"metrics": {
"latency_ms": 1.8067574128508568,
"throughput_per_s": 3506.4462175581734,
"load_s": 0.62601538002491,
"peak_vram_mib": 1234.0,
"driver_vram_mib": 2672.0
},
"max_abs_diff": 2.21142578125,
"notes": "unpadded batches of exactly 256 tokens; bf16 on cuda; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
"error": null,
"key": "linnet-jax-source"
},
{
"method": "Linnet ONNX f32 -> ONNX Runtime (CUDA)",
"kind": "linnet",
"metrics": {
"latency_ms": 3.6254888400435448,
"throughput_per_s": 1347.0558057204555,
"load_s": 0.2884049266576767,
"peak_vram_mib": 4822.0
},
"max_abs_diff": 2.393662452697754,
"notes": "f32; ONNX Runtime's CUDA execution provider; first calls build the sessions; unpadded batches of exactly the sequence length",
"error": null,
"key": "linnet-onnx-f32"
},
{
"method": "Linnet ONNX f32 -> ONNX Runtime (TensorRT)",
"kind": "linnet",
"metrics": {
"latency_ms": 1.7077615484595299,
"throughput_per_s": 2035.4607309036148,
"load_s": 0.3005627356469631,
"peak_vram_mib": 4116.0
},
"max_abs_diff": 2.3549375534057617,
"notes": "f32; ONNX Runtime's TensorRT execution provider; first calls build the sessions; unpadded batches of exactly the sequence length",
"error": null,
"key": "linnet-onnx-f32-trt"
},
{
"method": "Linnet ONNX f16 -> ONNX Runtime (CUDA)",
"kind": "linnet",
"metrics": {
"latency_ms": 2.794858068227768,
"throughput_per_s": 2095.59897492959,
"load_s": 0.3014170601963997,
"peak_vram_mib": 2812.0
},
"max_abs_diff": 2.188232421875,
"notes": "f16; ONNX Runtime's CUDA execution provider; first calls build the sessions; unpadded batches of exactly the sequence length",
"error": null,
"key": "linnet-onnx-f16"
},
{
"method": "Linnet ONNX f16 -> ONNX Runtime (TensorRT)",
"kind": "linnet",
"metrics": {
"latency_ms": 1.2534735724329948,
"throughput_per_s": 5758.625181121119,
"load_s": 0.2768788728863001,
"peak_vram_mib": 3768.0
},
"max_abs_diff": 2.162841796875,
"notes": "f16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; unpadded batches of exactly the sequence length",
"error": null,
"key": "linnet-onnx-f16-trt"
},
{
"method": "Linnet ONNX bf16 -> ONNX Runtime (CUDA)",
"kind": "linnet",
"metrics": {
"latency_ms": 3.237830474972725,
"throughput_per_s": 2040.3568779998031,
"load_s": 0.28330664709210396,
"peak_vram_mib": 2812.0
},
"max_abs_diff": 2.25,
"notes": "bf16; ONNX Runtime's CUDA execution provider; first calls build the sessions; unpadded batches of exactly the sequence length",
"error": null,
"key": "linnet-onnx-bf16"
},
{
"method": "Linnet ONNX bf16 -> ONNX Runtime (TensorRT)",
"kind": "linnet",
"metrics": {
"latency_ms": 1.3609202578663826,
"throughput_per_s": 5455.738809755462,
"load_s": 0.2847806606441736,
"peak_vram_mib": 3770.0
},
"max_abs_diff": 1.9375,
"notes": "bf16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; unpadded batches of exactly the sequence length",
"error": null,
"key": "linnet-onnx-bf16-trt"
},
{
"method": "transformers -> torch.onnx -> ONNX Runtime",
"kind": "reference",
"metrics": {
"latency_ms": 3.5911835730075836,
"throughput_per_s": 1301.835779288961,
"load_s": 31.848577592521906,
"peak_vram_mib": 6960.0
},
"max_abs_diff": 2.3369617462158203,
"notes": "the reference model's own ONNX export, f32 like Linnet's; unpadded batches of exactly 256 tokens",
"error": null,
"key": "onnx-reference"
},
{
"method": "Triton Inference Server (ONNX Runtime backend, torch.onnx export)",
"kind": "reference",
"metrics": {
"latency_ms": 6.785014644265175,
"throughput_per_s": 266.711770133521,
"load_s": 12.02940983325243,
"peak_vram_mib": 7104.9
},
"max_abs_diff": 2.3369617462158203,
"notes": "over HTTP; f32; throughput with 4 requests of batch 64 in flight",
"error": null,
"key": "triton-onnx"
},
{
"method": "Triton Inference Server (ONNX Runtime backend, Linnet ONNX)",
"kind": "linnet",
"metrics": {
"latency_ms": 6.600286811590195,
"throughput_per_s": 263.4712855475231,
"load_s": 18.101466223597527,
"peak_vram_mib": 7064.9
},
"max_abs_diff": 2.393662452697754,
"notes": "over HTTP; f32; throughput with 4 requests of batch 64 in flight",
"error": null,
"key": "triton-linnet-onnx"
},
{
"method": "Triton Inference Server (Python backend, Linnet torch)",
"kind": "linnet",
"metrics": {
"latency_ms": 3.9722491055727005,
"throughput_per_s": 269.6263180345148,
"load_s": 6.024784784764051,
"peak_vram_mib": 3202.1
},
"max_abs_diff": 1.25390625,
"notes": "over HTTP; bf16, CUDA graphs; throughput with 4 requests of batch 64 in flight",
"error": null,
"key": "triton-linnet-python"
}
],
"reference": "transformers-eager"
}bindings.json9.3 kB
json
{
"embeddings.tok_embeddings.weight": "model.embeddings.tok_embeddings.weight",
"embeddings.norm.weight": "model.embeddings.norm.weight",
"first.attn.qkv.weight": "model.layers.0.attn.Wqkv.weight",
"first.attn.out.weight": "model.layers.0.attn.Wo.weight",
"first.mlp_norm.weight": "model.layers.0.mlp_norm.weight",
"first.mlp.up.weight": "model.layers.0.mlp.Wi.weight",
"first.mlp.down.weight": "model.layers.0.mlp.Wo.weight",
"cycles.0.local.0.attn_norm.weight": "model.layers.1.attn_norm.weight",
"cycles.0.local.0.attn.qkv.weight": "model.layers.1.attn.Wqkv.weight",
"cycles.0.local.0.attn.out.weight": "model.layers.1.attn.Wo.weight",
"cycles.0.local.0.mlp_norm.weight": "model.layers.1.mlp_norm.weight",
"cycles.0.local.0.mlp.up.weight": "model.layers.1.mlp.Wi.weight",
"cycles.0.local.0.mlp.down.weight": "model.layers.1.mlp.Wo.weight",
"cycles.0.local.1.attn_norm.weight": "model.layers.2.attn_norm.weight",
"cycles.0.local.1.attn.qkv.weight": "model.layers.2.attn.Wqkv.weight",
"cycles.0.local.1.attn.out.weight": "model.layers.2.attn.Wo.weight",
"cycles.0.local.1.mlp_norm.weight": "model.layers.2.mlp_norm.weight",
"cycles.0.local.1.mlp.up.weight": "model.layers.2.mlp.Wi.weight",
"cycles.0.local.1.mlp.down.weight": "model.layers.2.mlp.Wo.weight",
"cycles.0.global.attn_norm.weight": "model.layers.3.attn_norm.weight",
"cycles.0.global.attn.qkv.weight": "model.layers.3.attn.Wqkv.weight",
"cycles.0.global.attn.out.weight": "model.layers.3.attn.Wo.weight",
"cycles.0.global.mlp_norm.weight": "model.layers.3.mlp_norm.weight",
"cycles.0.global.mlp.up.weight": "model.layers.3.mlp.Wi.weight",
"cycles.0.global.mlp.down.weight": "model.layers.3.mlp.Wo.weight",
"cycles.1.local.0.attn_norm.weight": "model.layers.4.attn_norm.weight",
"cycles.1.local.0.attn.qkv.weight": "model.layers.4.attn.Wqkv.weight",
"cycles.1.local.0.attn.out.weight": "model.layers.4.attn.Wo.weight",
"cycles.1.local.0.mlp_norm.weight": "model.layers.4.mlp_norm.weight",
"cycles.1.local.0.mlp.up.weight": "model.layers.4.mlp.Wi.weight",
"cycles.1.local.0.mlp.down.weight": "model.layers.4.mlp.Wo.weight",
"cycles.1.local.1.attn_norm.weight": "model.layers.5.attn_norm.weight",
"cycles.1.local.1.attn.qkv.weight": "model.layers.5.attn.Wqkv.weight",
"cycles.1.local.1.attn.out.weight": "model.layers.5.attn.Wo.weight",
"cycles.1.local.1.mlp_norm.weight": "model.layers.5.mlp_norm.weight",
"cycles.1.local.1.mlp.up.weight": "model.layers.5.mlp.Wi.weight",
"cycles.1.local.1.mlp.down.weight": "model.layers.5.mlp.Wo.weight",
"cycles.1.global.attn_norm.weight": "model.layers.6.attn_norm.weight",
"cycles.1.global.attn.qkv.weight": "model.layers.6.attn.Wqkv.weight",
"cycles.1.global.attn.out.weight": "model.layers.6.attn.Wo.weight",
"cycles.1.global.mlp_norm.weight": "model.layers.6.mlp_norm.weight",
"cycles.1.global.mlp.up.weight": "model.layers.6.mlp.Wi.weight",
"cycles.1.global.mlp.down.weight": "model.layers.6.mlp.Wo.weight",
"cycles.2.local.0.attn_norm.weight": "model.layers.7.attn_norm.weight",
"cycles.2.local.0.attn.qkv.weight": "model.layers.7.attn.Wqkv.weight",
"cycles.2.local.0.attn.out.weight": "model.layers.7.attn.Wo.weight",
"cycles.2.local.0.mlp_norm.weight": "model.layers.7.mlp_norm.weight",
"cycles.2.local.0.mlp.up.weight": "model.layers.7.mlp.Wi.weight",
"cycles.2.local.0.mlp.down.weight": "model.layers.7.mlp.Wo.weight",
"cycles.2.local.1.attn_norm.weight": "model.layers.8.attn_norm.weight",
"cycles.2.local.1.attn.qkv.weight": "model.layers.8.attn.Wqkv.weight",
"cycles.2.local.1.attn.out.weight": "model.layers.8.attn.Wo.weight",
"cycles.2.local.1.mlp_norm.weight": "model.layers.8.mlp_norm.weight",
"cycles.2.local.1.mlp.up.weight": "model.layers.8.mlp.Wi.weight",
"cycles.2.local.1.mlp.down.weight": "model.layers.8.mlp.Wo.weight",
"cycles.2.global.attn_norm.weight": "model.layers.9.attn_norm.weight",
"cycles.2.global.attn.qkv.weight": "model.layers.9.attn.Wqkv.weight",
"cycles.2.global.attn.out.weight": "model.layers.9.attn.Wo.weight",
"cycles.2.global.mlp_norm.weight": "model.layers.9.mlp_norm.weight",
"cycles.2.global.mlp.up.weight": "model.layers.9.mlp.Wi.weight",
"cycles.2.global.mlp.down.weight": "model.layers.9.mlp.Wo.weight",
"cycles.3.local.0.attn_norm.weight": "model.layers.10.attn_norm.weight",
"cycles.3.local.0.attn.qkv.weight": "model.layers.10.attn.Wqkv.weight",
"cycles.3.local.0.attn.out.weight": "model.layers.10.attn.Wo.weight",
"cycles.3.local.0.mlp_norm.weight": "model.layers.10.mlp_norm.weight",
"cycles.3.local.0.mlp.up.weight": "model.layers.10.mlp.Wi.weight",
"cycles.3.local.0.mlp.down.weight": "model.layers.10.mlp.Wo.weight",
"cycles.3.local.1.attn_norm.weight": "model.layers.11.attn_norm.weight",
"cycles.3.local.1.attn.qkv.weight": "model.layers.11.attn.Wqkv.weight",
"cycles.3.local.1.attn.out.weight": "model.layers.11.attn.Wo.weight",
"cycles.3.local.1.mlp_norm.weight": "model.layers.11.mlp_norm.weight",
"cycles.3.local.1.mlp.up.weight": "model.layers.11.mlp.Wi.weight",
"cycles.3.local.1.mlp.down.weight": "model.layers.11.mlp.Wo.weight",
"cycles.3.global.attn_norm.weight": "model.layers.12.attn_norm.weight",
"cycles.3.global.attn.qkv.weight": "model.layers.12.attn.Wqkv.weight",
"cycles.3.global.attn.out.weight": "model.layers.12.attn.Wo.weight",
"cycles.3.global.mlp_norm.weight": "model.layers.12.mlp_norm.weight",
"cycles.3.global.mlp.up.weight": "model.layers.12.mlp.Wi.weight",
"cycles.3.global.mlp.down.weight": "model.layers.12.mlp.Wo.weight",
"cycles.4.local.0.attn_norm.weight": "model.layers.13.attn_norm.weight",
"cycles.4.local.0.attn.qkv.weight": "model.layers.13.attn.Wqkv.weight",
"cycles.4.local.0.attn.out.weight": "model.layers.13.attn.Wo.weight",
"cycles.4.local.0.mlp_norm.weight": "model.layers.13.mlp_norm.weight",
"cycles.4.local.0.mlp.up.weight": "model.layers.13.mlp.Wi.weight",
"cycles.4.local.0.mlp.down.weight": "model.layers.13.mlp.Wo.weight",
"cycles.4.local.1.attn_norm.weight": "model.layers.14.attn_norm.weight",
"cycles.4.local.1.attn.qkv.weight": "model.layers.14.attn.Wqkv.weight",
"cycles.4.local.1.attn.out.weight": "model.layers.14.attn.Wo.weight",
"cycles.4.local.1.mlp_norm.weight": "model.layers.14.mlp_norm.weight",
"cycles.4.local.1.mlp.up.weight": "model.layers.14.mlp.Wi.weight",
"cycles.4.local.1.mlp.down.weight": "model.layers.14.mlp.Wo.weight",
"cycles.4.global.attn_norm.weight": "model.layers.15.attn_norm.weight",
"cycles.4.global.attn.qkv.weight": "model.layers.15.attn.Wqkv.weight",
"cycles.4.global.attn.out.weight": "model.layers.15.attn.Wo.weight",
"cycles.4.global.mlp_norm.weight": "model.layers.15.mlp_norm.weight",
"cycles.4.global.mlp.up.weight": "model.layers.15.mlp.Wi.weight",
"cycles.4.global.mlp.down.weight": "model.layers.15.mlp.Wo.weight",
"cycles.5.local.0.attn_norm.weight": "model.layers.16.attn_norm.weight",
"cycles.5.local.0.attn.qkv.weight": "model.layers.16.attn.Wqkv.weight",
"cycles.5.local.0.attn.out.weight": "model.layers.16.attn.Wo.weight",
"cycles.5.local.0.mlp_norm.weight": "model.layers.16.mlp_norm.weight",
"cycles.5.local.0.mlp.up.weight": "model.layers.16.mlp.Wi.weight",
"cycles.5.local.0.mlp.down.weight": "model.layers.16.mlp.Wo.weight",
"cycles.5.local.1.attn_norm.weight": "model.layers.17.attn_norm.weight",
"cycles.5.local.1.attn.qkv.weight": "model.layers.17.attn.Wqkv.weight",
"cycles.5.local.1.attn.out.weight": "model.layers.17.attn.Wo.weight",
"cycles.5.local.1.mlp_norm.weight": "model.layers.17.mlp_norm.weight",
"cycles.5.local.1.mlp.up.weight": "model.layers.17.mlp.Wi.weight",
"cycles.5.local.1.mlp.down.weight": "model.layers.17.mlp.Wo.weight",
"cycles.5.global.attn_norm.weight": "model.layers.18.attn_norm.weight",
"cycles.5.global.attn.qkv.weight": "model.layers.18.attn.Wqkv.weight",
"cycles.5.global.attn.out.weight": "model.layers.18.attn.Wo.weight",
"cycles.5.global.mlp_norm.weight": "model.layers.18.mlp_norm.weight",
"cycles.5.global.mlp.up.weight": "model.layers.18.mlp.Wi.weight",
"cycles.5.global.mlp.down.weight": "model.layers.18.mlp.Wo.weight",
"cycles.6.local.0.attn_norm.weight": "model.layers.19.attn_norm.weight",
"cycles.6.local.0.attn.qkv.weight": "model.layers.19.attn.Wqkv.weight",
"cycles.6.local.0.attn.out.weight": "model.layers.19.attn.Wo.weight",
"cycles.6.local.0.mlp_norm.weight": "model.layers.19.mlp_norm.weight",
"cycles.6.local.0.mlp.up.weight": "model.layers.19.mlp.Wi.weight",
"cycles.6.local.0.mlp.down.weight": "model.layers.19.mlp.Wo.weight",
"cycles.6.local.1.attn_norm.weight": "model.layers.20.attn_norm.weight",
"cycles.6.local.1.attn.qkv.weight": "model.layers.20.attn.Wqkv.weight",
"cycles.6.local.1.attn.out.weight": "model.layers.20.attn.Wo.weight",
"cycles.6.local.1.mlp_norm.weight": "model.layers.20.mlp_norm.weight",
"cycles.6.local.1.mlp.up.weight": "model.layers.20.mlp.Wi.weight",
"cycles.6.local.1.mlp.down.weight": "model.layers.20.mlp.Wo.weight",
"cycles.6.global.attn_norm.weight": "model.layers.21.attn_norm.weight",
"cycles.6.global.attn.qkv.weight": "model.layers.21.attn.Wqkv.weight",
"cycles.6.global.attn.out.weight": "model.layers.21.attn.Wo.weight",
"cycles.6.global.mlp_norm.weight": "model.layers.21.mlp_norm.weight",
"cycles.6.global.mlp.up.weight": "model.layers.21.mlp.Wi.weight",
"cycles.6.global.mlp.down.weight": "model.layers.21.mlp.Wo.weight",
"final_norm.weight": "model.final_norm.weight"
}linnet.toml65 B
toml
[package]
name = "modernbert"
version = "0.1.0"
language = "0.1"nest.toml929 B
toml
[model]
name = "modernbert-base"
title = "ModernBERT-base"
summary = "Answer.AI's 2024 redesign of the BERT encoder: 22 pre-norm layers of width 768, rotary positions, GeGLU, no biases, and a sliding attention window that opens up every third layer."
license = "Apache-2.0"
family = "modernbert"
tags = ["fill-mask", "feature-extraction", "encoder-only", "rotary-embeddings", "sliding-window-attention"]
[links]
huggingface = "https://huggingface.co/answerdotai/ModernBERT-base"
github = "https://github.com/AnswerDotAI/ModernBERT"
arxiv = "https://arxiv.org/abs/2412.13663"
[source]
path = "src/lib.linnet"
root = "Model"
entry = "forward"
[generics]
Vocab = 50368
D = 768
Heads = 12
Inner = 1152
Layers = 22
Window = 128
Period = 3
T = "f32"
[check]
B = 1
S = 256
[weights]
repo = "answerdotai/ModernBERT-base"
revision = "8949b909ec900327062f0ebf497f51aef5e6f0c8"
files = ["model.safetensors"]
bindings = "bindings.json"samples.json1.3 kB
json
{
"kind": "similarity",
"sentences": [
"The cat curled up on the warm windowsill and fell asleep in the sun.",
"Our new kitten refuses to eat anything except wet food.",
"The central bank raised interest rates again to cool inflation.",
"She rebalanced her portfolio after the market's sharp swings.",
"The function threw an exception when the array index ran out of bounds.",
"He refactored the module to remove three copies of the same logic."
],
"matrix": [
[
1.0,
0.921806,
0.880466,
0.905468,
0.894608,
0.905349
],
[
0.921806,
1.0,
0.875562,
0.888944,
0.893319,
0.894432
],
[
0.880466,
0.875562,
1.0,
0.927401,
0.867407,
0.898476
],
[
0.905468,
0.888944,
0.927401,
1.0,
0.882896,
0.921987
],
[
0.894608,
0.893319,
0.867407,
0.882896,
1.0,
0.915595
],
[
0.905349,
0.894432,
0.898476,
0.921987,
0.915595,
1.0
]
],
"caption": "Not a trained sentence embedder: the last hidden state, mean-pooled over tokens and L2-normalized by hand, since this checkpoint has no embedding head."
}src
lib.linnet10.3 kB
linnet
// ModernBERT: the 2024 redesign of the BERT encoder. Almost nothing of the
// original block survives. Normalization moved before the residual branches,
// every bias was removed, the absolute position table was replaced by rotary
// embeddings, and the feed-forward network became a GeGLU whose two halves
// come out of one fused projection.
//
// The part with no analogue in BERT is the attention pattern. Only every
// `Period`-th layer looks at the whole sequence; the layers between it attend
// inside a sliding window `Window` positions wide, so most of the depth costs
// linear rather than quadratic work. The two kinds of layer also use
// different rotary base frequencies: a global layer has to resolve positions
// thousands of tokens apart, a local layer never more than half a window, so
// the local layers keep the usual base of 10000 while the global layers
// stretch theirs to 160000.
//
// That alternation is written into the block structure rather than tested
// against a layer index, because a layer index is not something a `static
// for` can carry: `LayerCycle` holds the `Period - 1` local layers and the
// global one that closes the cycle. Layer 0 is a cycle-closing global layer
// too, but it is the one layer whose attention normalization upstream
// replaces with an identity -- the embedding LayerNorm has already normalized
// its input -- so it is peeled off as `FirstLayer`.
module modernbert
use std.nn.activations::{gelu_erf}
use std.nn.attention::{attention}
use std.nn.embedding::{Embedding}
use std.nn.linear::{linear}
use std.nn.norm::{layer_norm}
use std.nn.rope::{rope}
pub const EPS: f32 = 1e-5
// The rotary base of a layer that sees the whole sequence, and of a layer
// that sees only its window (`global_rope_theta` and `local_rope_theta`).
pub const GLOBAL_THETA: f32 = 160000.0
pub const LOCAL_THETA: f32 = 10000.0
// Every projection in the model is a bare weight: `attention_bias` and
// `mlp_bias` are both false, so the standard library's optional bias is
// never declared and no parameter of this model is optional.
pub block Dense<In: Dim, Out: Dim, T: Float> {
param weight: Tensor[Out, In; T]
pub fn forward<*S: Shape>(x: Tensor[*S, In; T]) -> Tensor[*S, Out; T] {
return linear(x, weight)
}
}
// LayerNorm without a bias (`norm_bias` is false), which is every
// normalization in the model.
pub block Norm<D: Dim, T: Float> {
param weight: Tensor[D; T]
pub fn forward<*S: Shape>(x: Tensor[*S, D; T]) -> Tensor[*S, D; T] {
return layer_norm(x, weight, none, EPS)
}
}
// The input side: a token table and a LayerNorm, and nothing else. There is
// no position table at all; position reaches the model only through the
// rotary tables inside each attention block.
pub block Embeddings<Vocab: Dim, D: Dim, T: Float> {
sub tok_embeddings: Embedding<Vocab, D, T>
sub norm: Norm<D, T>
pub fn forward<*S: Shape>(ids: Tensor[*S; i32]) -> Tensor[*S, D; T] {
return norm.forward(tok_embeddings.forward(ids))
}
}
// A `[Q, K]` band: each query keeps the keys within `Window / 2` positions of
// it, on either side. This is the sliding window of the local layers, and the
// halving is the checkpoint's convention -- `local_attention` counts the
// whole window, so 128 means 64 back and 64 forward, 129 positions in all.
//
// Written as an `op` for the same reason `std.nn.attention::causal_mask` is:
// the body is the definition, and a backend with a banded attention kernel is
// free to recognize it and drop the mask.
op band_mask<Q: Dim, K: Dim, Window: Dim>() -> Tensor[Q, K; bool] {
let queries = iota<i32>(Q)
let keys = iota<i32>(K)
let mask[q, k] = abs(keys[k] - queries[q]) <= cast<i32>(Window / 2)
return mask
}
// `(cos, sin)` tables of shape `[S, D]`, each frequency repeated across both
// halves of the head as `rope` expects. The tables come from `iota`, so a
// layer needs no precomputed input and the base frequency is just an
// argument -- which is what lets one block serve both the local and the
// global layers.
fn rope_tables<S: Dim, D: Dim, T: Float>(theta: f32) -> (Tensor[S, D; T], Tensor[S, D; T])
where D % 2 == 0 {
let inv_freq[i] = exp(-(cast<f32>(iota<i64>(D / 2)[i]) * 2.0 / cast<f32>(D)) * log(theta))
let angles[s, i] = cast<f32>(iota<i64>(S)[s]) * inv_freq[i]
let full = concat(angles, angles, axis = -1)
return (cast<T>(cos(full)), cast<T>(sin(full)))
}
// Self-attention with one fused input projection. `Wqkv` stores its rows as
// query, then key, then value, each already laid out head by head, so the
// three parts are three slices of the last axis.
//
// The rotary base and the mask are arguments rather than properties of the
// block: the same weights shape serves a local layer (a band mask, base
// 10000) and a global layer (no mask, base 160000).
pub block Attention<D: Dim, Heads: Dim, T: Float>
where
Heads > 0,
D % Heads == 0,
(D / Heads) % 2 == 0
{
sub qkv: Dense<D, 3 * D, T>
sub out: Dense<D, D, T>
pub fn forward<B: Dim, S: Dim>(
x: Tensor[B, S, D; T],
theta: f32,
window: Tensor[S, S; bool]?,
) -> Tensor[B, S, D; T] {
let fused = qkv.forward(x)
let (cos_table, sin_table) = rope_tables<S, D / Heads, T>(theta)
// Rotation is applied to the queries and the keys; values carry no
// position of their own.
let q = rope(
split_heads<B, S, Heads, D / Heads, T>(fused[..., 0:D]),
cos_table,
sin_table,
)
let k = rope(
split_heads<B, S, Heads, D / Heads, T>(fused[..., D:2 * D]),
cos_table,
sin_table,
)
let v = split_heads<B, S, Heads, D / Heads, T>(fused[..., 2 * D:3 * D])
let mixed = attention(q, k, v, rsqrt(cast<f32>(D / Heads)), window)
return out.forward(merge_heads<B, S, Heads, D / Heads, T>(mixed))
}
}
// `[B, S, N * Dh]` -> `[B, N, S, Dh]`.
fn split_heads<B: Dim, S: Dim, N: Dim, Dh: Dim, T: Float>(
x: Tensor[B, S, N * Dh; T],
) -> Tensor[B, N, S, Dh; T] {
return permute(reshape(x, [B, S, N, Dh]), [0, 2, 1, 3])
}
fn merge_heads<B: Dim, S: Dim, N: Dim, Dh: Dim, T: Float>(
x: Tensor[B, N, S, Dh; T],
) -> Tensor[B, S, N * Dh; T] {
return reshape(permute(x, [0, 2, 1, 3]), [B, S, N * Dh])
}
// The gated feed-forward network. One projection `Wi` produces `2 * Inner`
// values per position: the first half is the branch the activation runs on,
// the second half multiplies it. Reading the halves the other way round is a
// silent error, so the order is spelled out here -- upstream names them
// `input, gate = Wi(x).chunk(2, dim=-1)`.
pub block GeGlu<D: Dim, Inner: Dim, T: Float> {
sub up: Dense<D, 2 * Inner, T>
sub down: Dense<Inner, D, T>
pub fn forward<*S: Shape>(x: Tensor[*S, D; T]) -> Tensor[*S, D; T] {
let fused = up.forward(x)
let activated = gelu_erf(fused[..., 0:Inner])
let gate = fused[..., Inner:2 * Inner]
return down.forward(activated * gate)
}
}
// A pre-norm encoder layer: normalize, attend, add; normalize, gate, add.
pub block EncoderLayer<D: Dim, Heads: Dim, Inner: Dim, T: Float>
where
Heads > 0,
D % Heads == 0,
(D / Heads) % 2 == 0
{
sub attn_norm: Norm<D, T>
sub attn: Attention<D, Heads, T>
sub mlp_norm: Norm<D, T>
sub mlp: GeGlu<D, Inner, T>
pub fn forward<B: Dim, S: Dim>(
x: Tensor[B, S, D; T],
theta: f32,
window: Tensor[S, S; bool]?,
) -> Tensor[B, S, D; T] {
let attended = x + attn.forward(attn_norm.forward(x), theta, window)
return attended + mlp.forward(mlp_norm.forward(attended))
}
}
// Layer 0, which is the same layer without the attention normalization: the
// embedding LayerNorm runs immediately before it, so upstream makes this one
// `attn_norm` an identity and the checkpoint carries no weight for it.
pub block FirstLayer<D: Dim, Heads: Dim, Inner: Dim, T: Float>
where
Heads > 0,
D % Heads == 0,
(D / Heads) % 2 == 0
{
sub attn: Attention<D, Heads, T>
sub mlp_norm: Norm<D, T>
sub mlp: GeGlu<D, Inner, T>
pub fn forward<B: Dim, S: Dim>(
x: Tensor[B, S, D; T],
theta: f32,
window: Tensor[S, S; bool]?,
) -> Tensor[B, S, D; T] {
let attended = x + attn.forward(x, theta, window)
return attended + mlp.forward(mlp_norm.forward(attended))
}
}
// One turn of the alternation: `Period - 1` windowed layers, then the global
// layer that closes the cycle. Holding the two kinds as different members is
// what makes the pattern checkable -- each member is reached by name, so each
// gets its own mask and its own rotary base with no index arithmetic.
pub block LayerCycle<D: Dim, Heads: Dim, Inner: Dim, Period: Dim, T: Float>
where
Heads > 0,
D % Heads == 0,
(D / Heads) % 2 == 0,
Period > 1
{
sub local: [EncoderLayer<D, Heads, Inner, T>; Period - 1]
sub global: EncoderLayer<D, Heads, Inner, T>
pub fn forward<B: Dim, S: Dim>(
x: Tensor[B, S, D; T],
window: Tensor[S, S; bool],
) -> Tensor[B, S, D; T] {
var h = x
static for layer in local {
h = layer.forward(h, LOCAL_THETA, some(window))
}
return global.forward(h, GLOBAL_THETA, none)
}
}
pub block Model<
Vocab: Dim,
D: Dim,
Heads: Dim,
Inner: Dim,
Layers: Dim,
Window: Dim,
Period: Dim,
T: Float = f32,
>
where
Heads > 0,
D % Heads == 0,
(D / Heads) % 2 == 0,
Period > 1,
(Layers - 1) % Period == 0
{
sub embeddings: Embeddings<Vocab, D, T>
sub first: FirstLayer<D, Heads, Inner, T>
sub cycles: [LayerCycle<D, Heads, Inner, Period, T>; (Layers - 1) / Period]
sub final_norm: Norm<D, T>
// The last hidden state, normalized: what a downstream classifier or
// masked-language head reads. Every position attends over every other
// position it is allowed to see; there is no padding mask, because
// `std.nn.attention` masks are per-sequence-pair rather than per-batch
// element, so a batch here must be sequences of the same real length.
pub entry forward<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, D; T] {
let window = band_mask<S, S, Window>()
var x = embeddings.forward(tokens)
x = first.forward(x, GLOBAL_THETA, none)
static for cycle in cycles {
x = cycle.forward(x, window)
}
return final_norm.forward(x)
}
}