Models / bert-base-uncased

BERT Base (uncased)

Google's original BERT: 12 post-norm encoder layers of width 768, word/position/token-type embeddings, and the next-sentence-prediction pooler.

109.5M parametersbertApache-2.0fill-masktext-encoderencoder-only
models/bert-base-uncased8 files · 34.9 kB in the registry · GitHub
README.md2.6 kB
md
# BERT Base (uncased)

The original BERT: word, learned-position, and token-type embeddings summed
and normalized once, then 12 post-norm encoder layers of width 768 with 12
heads. Post-norm is the block order BERT and RoBERTa share and the
pre-norm ViT and GPT-2 already in Nest do not: the residual add happens
first and LayerNorm reads the sum, `h = LayerNorm(x + attention(x))`, then
`out = LayerNorm(h + ffn(h))`, rather than normalizing each branch's input.
There is no final normalization after the last layer, since every block
already ends with one.

The checkpoint's `LayerNorm` weight and bias are stored as `gamma` and
`beta` (an older naming convention this repository predates); `bindings.json`
maps Linnet's `weight`/`bias` names onto them.

Upstream's `hidden_act` is plain GELU, the error-function form, so the
source uses `std.nn.activations::gelu_erf` rather than the tanh
approximation `gelu`. The distinction is not cosmetic: over twelve post-norm
layers the tanh form moves some tokens' hidden states by 1.8e-1.

## Loading

```python
from linnet import nest

model = nest.load("bert-base-uncased", backend="torch")
hidden = model(tokens, token_types)        # Tensor[B, S; i32], Tensor[B, S; i32] -> Tensor[B, S, 768; f32]
pooled = model.pool(tokens, token_types)   # -> Tensor[B, 768; f32]
```

`token_types` is the segment-id tensor (0 for a single sentence); pass a
zero tensor of the same shape as `tokens` unless doing sentence-pair tasks.

Sequences longer than 512 positions are rejected at compile time
(`where S <= MaxPositions`). The encoder runs full attention over every
position: `std.nn.attention::attention` takes an unbatched `[Q, K]` mask, so
a per-sequence padding mask cannot be expressed here. Pad-free batches (or a
batch of one) give exact results; padded batches attend onto the padding.

## Entries

| Entry | |
| --- | --- |
| `forward<B, S>(tokens, token_types)` | last hidden state, `Tensor[B, S, 768; f32]` |
| `pool<B, S>(tokens, token_types)` | `tanh(dense(hidden[:, 0, :]))`, BERT's pooled output, `Tensor[B, 768; f32]` |

## Numerics

Validated against `transformers.BertModel` (`AutoModel.from_pretrained`,
`numerics="fast"`) on CPU in f32, comparing `last_hidden_state` for
`forward` and `pooler_output` for `pool` on pad-free sentences. The maximum
absolute difference is 3.7e-5 for `forward` and 1.1e-6 for `pool`.

## Provenance

- Weights: [google-bert/bert-base-uncased](https://huggingface.co/google-bert/bert-base-uncased), Apache-2.0.
- Code: [google-research/bert](https://github.com/google-research/bert).
- Paper: [BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding](https://arxiv.org/abs/1810.04805).

Open on GitHub

bench.json9.4 kB
json
{
  "date": "2026-09-28T18:07:07+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": 3.6013582721352577,
        "throughput_per_s": 11043.048408516197,
        "load_s": 1.4411308970302343,
        "peak_vram_mib": 1162.0
      },
      "max_abs_diff": 0.0,
      "notes": "unpadded batches of exactly 128 tokens (this card takes no padding mask)",
      "error": null,
      "key": "transformers-eager"
    },
    {
      "method": "transformers (torch.compile)",
      "kind": "reference",
      "metrics": {
        "latency_ms": 1.6716653481125832,
        "throughput_per_s": 14662.857694943665,
        "load_s": 1.6261053308844566,
        "peak_vram_mib": 1142.0
      },
      "max_abs_diff": 0.109375,
      "notes": "unpadded batches of exactly 128 tokens (this card takes no padding mask)",
      "error": null,
      "key": "transformers-compile"
    },
    {
      "method": "Linnet torch (generated source)",
      "kind": "linnet",
      "metrics": {
        "latency_ms": 2.79293954372406,
        "throughput_per_s": 13867.08498335816,
        "load_s": 1.0936322510242462,
        "peak_vram_mib": 1108.0
      },
      "max_abs_diff": 0.1171875,
      "notes": "unpadded batches of exactly 128 tokens; bf16 on cuda",
      "error": null,
      "key": "linnet-torch"
    },
    {
      "method": "Linnet torch (CUDA graphs)",
      "kind": "linnet",
      "metrics": {
        "latency_ms": 0.7357271388173103,
        "throughput_per_s": 17770.208849641735,
        "load_s": 1.1175704188644886,
        "peak_vram_mib": 1148.0
      },
      "max_abs_diff": 0.140625,
      "notes": "unpadded batches of exactly 128 tokens; bf16 on cuda",
      "error": null,
      "key": "linnet-cudagraphs"
    },
    {
      "method": "Linnet torch (torch.compile, inductor)",
      "kind": "linnet",
      "metrics": {
        "latency_ms": 1.7542997375130653,
        "throughput_per_s": 17002.309812055253,
        "load_s": 1.1143667306751013,
        "peak_vram_mib": 1060.0
      },
      "max_abs_diff": 0.140625,
      "notes": "unpadded batches of exactly 128 tokens; bf16 on cuda",
      "error": null,
      "key": "linnet-inductor"
    },
    {
      "method": "Linnet JAX (XLA, StableHLO)",
      "kind": "linnet",
      "metrics": {
        "latency_ms": 1.555175520479679,
        "throughput_per_s": 13223.314608371777,
        "load_s": 0.5223176646977663,
        "peak_vram_mib": 734.4,
        "driver_vram_mib": 1646.0
      },
      "max_abs_diff": 0.109375,
      "notes": "unpadded batches of exactly 128 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": 0.9256266057491302,
        "throughput_per_s": 11749.313148251746,
        "load_s": 0.5791833251714706,
        "peak_vram_mib": 726.6,
        "driver_vram_mib": 1646.0
      },
      "max_abs_diff": 0.25,
      "notes": "unpadded batches of exactly 128 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": 1.5163514763116837,
        "throughput_per_s": 5273.784703467455,
        "load_s": 0.33368977159261703,
        "peak_vram_mib": 2586.0
      },
      "max_abs_diff": 0.1145319938659668,
      "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": 0.738825649023056,
        "throughput_per_s": 9036.098744588096,
        "load_s": 0.2823925744742155,
        "peak_vram_mib": 3282.0
      },
      "max_abs_diff": 0.11189699172973633,
      "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": 1.2513808906078339,
        "throughput_per_s": 9557.05211873476,
        "load_s": 0.5559203401207924,
        "peak_vram_mib": 1748.0
      },
      "max_abs_diff": 0.11328125,
      "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": 0.5308650434017181,
        "throughput_per_s": 19915.629406496973,
        "load_s": 0.2835025265812874,
        "peak_vram_mib": 3042.0
      },
      "max_abs_diff": 0.11328125,
      "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": 1.250801607966423,
        "throughput_per_s": 9728.454393629683,
        "load_s": 0.2895093970000744,
        "peak_vram_mib": 1748.0
      },
      "max_abs_diff": 0.125,
      "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": 0.5469843745231628,
        "throughput_per_s": 18309.78428181509,
        "load_s": 0.3062436003237963,
        "peak_vram_mib": 3038.0
      },
      "max_abs_diff": 0.109375,
      "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": "KerasHub (JAX)",
      "kind": "reference",
      "metrics": {
        "latency_ms": 4.673472605645657,
        "throughput_per_s": 5158.8344350616635,
        "load_s": 6.311374740675092,
        "peak_vram_mib": 525.0,
        "driver_vram_mib": 1650.0
      },
      "max_abs_diff": 0.2265625,
      "notes": "bf16, under jax.jit; unpadded batches of exactly 128 tokens; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "keras-hub"
    },
    {
      "method": "transformers -> torch.onnx -> ONNX Runtime",
      "kind": "reference",
      "metrics": {
        "latency_ms": 2.057228237390518,
        "throughput_per_s": 4942.8163677230195,
        "load_s": 17.537603374570608,
        "peak_vram_mib": 3382.0
      },
      "max_abs_diff": 0.11434412002563477,
      "notes": "the reference model's own ONNX export, f32 like Linnet's; unpadded batches of exactly 128 tokens",
      "error": null,
      "key": "onnx-reference"
    },
    {
      "method": "Triton Inference Server (ONNX Runtime backend, torch.onnx export)",
      "kind": "reference",
      "metrics": {
        "latency_ms": 3.6560427397489548,
        "throughput_per_s": 578.6878095224938,
        "load_s": 9.10310436040163,
        "peak_vram_mib": 2428.9
      },
      "max_abs_diff": 0.11434412002563477,
      "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": 3.2435525208711624,
        "throughput_per_s": 594.4761908508923,
        "load_s": 16.062782548367977,
        "peak_vram_mib": 2424.9
      },
      "max_abs_diff": 0.1145319938659668,
      "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": 2.7097295969724655,
        "throughput_per_s": 557.9240686320418,
        "load_s": 6.024971269071102,
        "peak_vram_mib": 2880.1
      },
      "max_abs_diff": 0.140625,
      "notes": "over HTTP; bf16, CUDA graphs; throughput with 4 requests of batch 64 in flight",
      "error": null,
      "key": "triton-linnet-python"
    }
  ],
  "reference": "transformers-eager"
}

Open on GitHub

bindings.json15.8 kB
json
{
  "embeddings.word_embeddings": "bert.embeddings.word_embeddings.weight",
  "embeddings.position_embeddings": "bert.embeddings.position_embeddings.weight",
  "embeddings.token_type_embeddings": "bert.embeddings.token_type_embeddings.weight",
  "embeddings.norm.weight": "bert.embeddings.LayerNorm.gamma",
  "embeddings.norm.bias": "bert.embeddings.LayerNorm.beta",
  "layers.0.attention.query.weight": "bert.encoder.layer.0.attention.self.query.weight",
  "layers.0.attention.query.bias": "bert.encoder.layer.0.attention.self.query.bias",
  "layers.0.attention.key.weight": "bert.encoder.layer.0.attention.self.key.weight",
  "layers.0.attention.key.bias": "bert.encoder.layer.0.attention.self.key.bias",
  "layers.0.attention.value.weight": "bert.encoder.layer.0.attention.self.value.weight",
  "layers.0.attention.value.bias": "bert.encoder.layer.0.attention.self.value.bias",
  "layers.0.attention.out.weight": "bert.encoder.layer.0.attention.output.dense.weight",
  "layers.0.attention.out.bias": "bert.encoder.layer.0.attention.output.dense.bias",
  "layers.0.attention_norm.weight": "bert.encoder.layer.0.attention.output.LayerNorm.gamma",
  "layers.0.attention_norm.bias": "bert.encoder.layer.0.attention.output.LayerNorm.beta",
  "layers.0.up.weight": "bert.encoder.layer.0.intermediate.dense.weight",
  "layers.0.up.bias": "bert.encoder.layer.0.intermediate.dense.bias",
  "layers.0.down.weight": "bert.encoder.layer.0.output.dense.weight",
  "layers.0.down.bias": "bert.encoder.layer.0.output.dense.bias",
  "layers.0.output_norm.weight": "bert.encoder.layer.0.output.LayerNorm.gamma",
  "layers.0.output_norm.bias": "bert.encoder.layer.0.output.LayerNorm.beta",
  "layers.1.attention.query.weight": "bert.encoder.layer.1.attention.self.query.weight",
  "layers.1.attention.query.bias": "bert.encoder.layer.1.attention.self.query.bias",
  "layers.1.attention.key.weight": "bert.encoder.layer.1.attention.self.key.weight",
  "layers.1.attention.key.bias": "bert.encoder.layer.1.attention.self.key.bias",
  "layers.1.attention.value.weight": "bert.encoder.layer.1.attention.self.value.weight",
  "layers.1.attention.value.bias": "bert.encoder.layer.1.attention.self.value.bias",
  "layers.1.attention.out.weight": "bert.encoder.layer.1.attention.output.dense.weight",
  "layers.1.attention.out.bias": "bert.encoder.layer.1.attention.output.dense.bias",
  "layers.1.attention_norm.weight": "bert.encoder.layer.1.attention.output.LayerNorm.gamma",
  "layers.1.attention_norm.bias": "bert.encoder.layer.1.attention.output.LayerNorm.beta",
  "layers.1.up.weight": "bert.encoder.layer.1.intermediate.dense.weight",
  "layers.1.up.bias": "bert.encoder.layer.1.intermediate.dense.bias",
  "layers.1.down.weight": "bert.encoder.layer.1.output.dense.weight",
  "layers.1.down.bias": "bert.encoder.layer.1.output.dense.bias",
  "layers.1.output_norm.weight": "bert.encoder.layer.1.output.LayerNorm.gamma",
  "layers.1.output_norm.bias": "bert.encoder.layer.1.output.LayerNorm.beta",
  "layers.2.attention.query.weight": "bert.encoder.layer.2.attention.self.query.weight",
  "layers.2.attention.query.bias": "bert.encoder.layer.2.attention.self.query.bias",
  "layers.2.attention.key.weight": "bert.encoder.layer.2.attention.self.key.weight",
  "layers.2.attention.key.bias": "bert.encoder.layer.2.attention.self.key.bias",
  "layers.2.attention.value.weight": "bert.encoder.layer.2.attention.self.value.weight",
  "layers.2.attention.value.bias": "bert.encoder.layer.2.attention.self.value.bias",
  "layers.2.attention.out.weight": "bert.encoder.layer.2.attention.output.dense.weight",
  "layers.2.attention.out.bias": "bert.encoder.layer.2.attention.output.dense.bias",
  "layers.2.attention_norm.weight": "bert.encoder.layer.2.attention.output.LayerNorm.gamma",
  "layers.2.attention_norm.bias": "bert.encoder.layer.2.attention.output.LayerNorm.beta",
  "layers.2.up.weight": "bert.encoder.layer.2.intermediate.dense.weight",
  "layers.2.up.bias": "bert.encoder.layer.2.intermediate.dense.bias",
  "layers.2.down.weight": "bert.encoder.layer.2.output.dense.weight",
  "layers.2.down.bias": "bert.encoder.layer.2.output.dense.bias",
  "layers.2.output_norm.weight": "bert.encoder.layer.2.output.LayerNorm.gamma",
  "layers.2.output_norm.bias": "bert.encoder.layer.2.output.LayerNorm.beta",
  "layers.3.attention.query.weight": "bert.encoder.layer.3.attention.self.query.weight",
  "layers.3.attention.query.bias": "bert.encoder.layer.3.attention.self.query.bias",
  "layers.3.attention.key.weight": "bert.encoder.layer.3.attention.self.key.weight",
  "layers.3.attention.key.bias": "bert.encoder.layer.3.attention.self.key.bias",
  "layers.3.attention.value.weight": "bert.encoder.layer.3.attention.self.value.weight",
  "layers.3.attention.value.bias": "bert.encoder.layer.3.attention.self.value.bias",
  "layers.3.attention.out.weight": "bert.encoder.layer.3.attention.output.dense.weight",
  "layers.3.attention.out.bias": "bert.encoder.layer.3.attention.output.dense.bias",
  "layers.3.attention_norm.weight": "bert.encoder.layer.3.attention.output.LayerNorm.gamma",
  "layers.3.attention_norm.bias": "bert.encoder.layer.3.attention.output.LayerNorm.beta",
  "layers.3.up.weight": "bert.encoder.layer.3.intermediate.dense.weight",
  "layers.3.up.bias": "bert.encoder.layer.3.intermediate.dense.bias",
  "layers.3.down.weight": "bert.encoder.layer.3.output.dense.weight",
  "layers.3.down.bias": "bert.encoder.layer.3.output.dense.bias",
  "layers.3.output_norm.weight": "bert.encoder.layer.3.output.LayerNorm.gamma",
  "layers.3.output_norm.bias": "bert.encoder.layer.3.output.LayerNorm.beta",
  "layers.4.attention.query.weight": "bert.encoder.layer.4.attention.self.query.weight",
  "layers.4.attention.query.bias": "bert.encoder.layer.4.attention.self.query.bias",
  "layers.4.attention.key.weight": "bert.encoder.layer.4.attention.self.key.weight",
  "layers.4.attention.key.bias": "bert.encoder.layer.4.attention.self.key.bias",
  "layers.4.attention.value.weight": "bert.encoder.layer.4.attention.self.value.weight",
  "layers.4.attention.value.bias": "bert.encoder.layer.4.attention.self.value.bias",
  "layers.4.attention.out.weight": "bert.encoder.layer.4.attention.output.dense.weight",
  "layers.4.attention.out.bias": "bert.encoder.layer.4.attention.output.dense.bias",
  "layers.4.attention_norm.weight": "bert.encoder.layer.4.attention.output.LayerNorm.gamma",
  "layers.4.attention_norm.bias": "bert.encoder.layer.4.attention.output.LayerNorm.beta",
  "layers.4.up.weight": "bert.encoder.layer.4.intermediate.dense.weight",
  "layers.4.up.bias": "bert.encoder.layer.4.intermediate.dense.bias",
  "layers.4.down.weight": "bert.encoder.layer.4.output.dense.weight",
  "layers.4.down.bias": "bert.encoder.layer.4.output.dense.bias",
  "layers.4.output_norm.weight": "bert.encoder.layer.4.output.LayerNorm.gamma",
  "layers.4.output_norm.bias": "bert.encoder.layer.4.output.LayerNorm.beta",
  "layers.5.attention.query.weight": "bert.encoder.layer.5.attention.self.query.weight",
  "layers.5.attention.query.bias": "bert.encoder.layer.5.attention.self.query.bias",
  "layers.5.attention.key.weight": "bert.encoder.layer.5.attention.self.key.weight",
  "layers.5.attention.key.bias": "bert.encoder.layer.5.attention.self.key.bias",
  "layers.5.attention.value.weight": "bert.encoder.layer.5.attention.self.value.weight",
  "layers.5.attention.value.bias": "bert.encoder.layer.5.attention.self.value.bias",
  "layers.5.attention.out.weight": "bert.encoder.layer.5.attention.output.dense.weight",
  "layers.5.attention.out.bias": "bert.encoder.layer.5.attention.output.dense.bias",
  "layers.5.attention_norm.weight": "bert.encoder.layer.5.attention.output.LayerNorm.gamma",
  "layers.5.attention_norm.bias": "bert.encoder.layer.5.attention.output.LayerNorm.beta",
  "layers.5.up.weight": "bert.encoder.layer.5.intermediate.dense.weight",
  "layers.5.up.bias": "bert.encoder.layer.5.intermediate.dense.bias",
  "layers.5.down.weight": "bert.encoder.layer.5.output.dense.weight",
  "layers.5.down.bias": "bert.encoder.layer.5.output.dense.bias",
  "layers.5.output_norm.weight": "bert.encoder.layer.5.output.LayerNorm.gamma",
  "layers.5.output_norm.bias": "bert.encoder.layer.5.output.LayerNorm.beta",
  "layers.6.attention.query.weight": "bert.encoder.layer.6.attention.self.query.weight",
  "layers.6.attention.query.bias": "bert.encoder.layer.6.attention.self.query.bias",
  "layers.6.attention.key.weight": "bert.encoder.layer.6.attention.self.key.weight",
  "layers.6.attention.key.bias": "bert.encoder.layer.6.attention.self.key.bias",
  "layers.6.attention.value.weight": "bert.encoder.layer.6.attention.self.value.weight",
  "layers.6.attention.value.bias": "bert.encoder.layer.6.attention.self.value.bias",
  "layers.6.attention.out.weight": "bert.encoder.layer.6.attention.output.dense.weight",
  "layers.6.attention.out.bias": "bert.encoder.layer.6.attention.output.dense.bias",
  "layers.6.attention_norm.weight": "bert.encoder.layer.6.attention.output.LayerNorm.gamma",
  "layers.6.attention_norm.bias": "bert.encoder.layer.6.attention.output.LayerNorm.beta",
  "layers.6.up.weight": "bert.encoder.layer.6.intermediate.dense.weight",
  "layers.6.up.bias": "bert.encoder.layer.6.intermediate.dense.bias",
  "layers.6.down.weight": "bert.encoder.layer.6.output.dense.weight",
  "layers.6.down.bias": "bert.encoder.layer.6.output.dense.bias",
  "layers.6.output_norm.weight": "bert.encoder.layer.6.output.LayerNorm.gamma",
  "layers.6.output_norm.bias": "bert.encoder.layer.6.output.LayerNorm.beta",
  "layers.7.attention.query.weight": "bert.encoder.layer.7.attention.self.query.weight",
  "layers.7.attention.query.bias": "bert.encoder.layer.7.attention.self.query.bias",
  "layers.7.attention.key.weight": "bert.encoder.layer.7.attention.self.key.weight",
  "layers.7.attention.key.bias": "bert.encoder.layer.7.attention.self.key.bias",
  "layers.7.attention.value.weight": "bert.encoder.layer.7.attention.self.value.weight",
  "layers.7.attention.value.bias": "bert.encoder.layer.7.attention.self.value.bias",
  "layers.7.attention.out.weight": "bert.encoder.layer.7.attention.output.dense.weight",
  "layers.7.attention.out.bias": "bert.encoder.layer.7.attention.output.dense.bias",
  "layers.7.attention_norm.weight": "bert.encoder.layer.7.attention.output.LayerNorm.gamma",
  "layers.7.attention_norm.bias": "bert.encoder.layer.7.attention.output.LayerNorm.beta",
  "layers.7.up.weight": "bert.encoder.layer.7.intermediate.dense.weight",
  "layers.7.up.bias": "bert.encoder.layer.7.intermediate.dense.bias",
  "layers.7.down.weight": "bert.encoder.layer.7.output.dense.weight",
  "layers.7.down.bias": "bert.encoder.layer.7.output.dense.bias",
  "layers.7.output_norm.weight": "bert.encoder.layer.7.output.LayerNorm.gamma",
  "layers.7.output_norm.bias": "bert.encoder.layer.7.output.LayerNorm.beta",
  "layers.8.attention.query.weight": "bert.encoder.layer.8.attention.self.query.weight",
  "layers.8.attention.query.bias": "bert.encoder.layer.8.attention.self.query.bias",
  "layers.8.attention.key.weight": "bert.encoder.layer.8.attention.self.key.weight",
  "layers.8.attention.key.bias": "bert.encoder.layer.8.attention.self.key.bias",
  "layers.8.attention.value.weight": "bert.encoder.layer.8.attention.self.value.weight",
  "layers.8.attention.value.bias": "bert.encoder.layer.8.attention.self.value.bias",
  "layers.8.attention.out.weight": "bert.encoder.layer.8.attention.output.dense.weight",
  "layers.8.attention.out.bias": "bert.encoder.layer.8.attention.output.dense.bias",
  "layers.8.attention_norm.weight": "bert.encoder.layer.8.attention.output.LayerNorm.gamma",
  "layers.8.attention_norm.bias": "bert.encoder.layer.8.attention.output.LayerNorm.beta",
  "layers.8.up.weight": "bert.encoder.layer.8.intermediate.dense.weight",
  "layers.8.up.bias": "bert.encoder.layer.8.intermediate.dense.bias",
  "layers.8.down.weight": "bert.encoder.layer.8.output.dense.weight",
  "layers.8.down.bias": "bert.encoder.layer.8.output.dense.bias",
  "layers.8.output_norm.weight": "bert.encoder.layer.8.output.LayerNorm.gamma",
  "layers.8.output_norm.bias": "bert.encoder.layer.8.output.LayerNorm.beta",
  "layers.9.attention.query.weight": "bert.encoder.layer.9.attention.self.query.weight",
  "layers.9.attention.query.bias": "bert.encoder.layer.9.attention.self.query.bias",
  "layers.9.attention.key.weight": "bert.encoder.layer.9.attention.self.key.weight",
  "layers.9.attention.key.bias": "bert.encoder.layer.9.attention.self.key.bias",
  "layers.9.attention.value.weight": "bert.encoder.layer.9.attention.self.value.weight",
  "layers.9.attention.value.bias": "bert.encoder.layer.9.attention.self.value.bias",
  "layers.9.attention.out.weight": "bert.encoder.layer.9.attention.output.dense.weight",
  "layers.9.attention.out.bias": "bert.encoder.layer.9.attention.output.dense.bias",
  "layers.9.attention_norm.weight": "bert.encoder.layer.9.attention.output.LayerNorm.gamma",
  "layers.9.attention_norm.bias": "bert.encoder.layer.9.attention.output.LayerNorm.beta",
  "layers.9.up.weight": "bert.encoder.layer.9.intermediate.dense.weight",
  "layers.9.up.bias": "bert.encoder.layer.9.intermediate.dense.bias",
  "layers.9.down.weight": "bert.encoder.layer.9.output.dense.weight",
  "layers.9.down.bias": "bert.encoder.layer.9.output.dense.bias",
  "layers.9.output_norm.weight": "bert.encoder.layer.9.output.LayerNorm.gamma",
  "layers.9.output_norm.bias": "bert.encoder.layer.9.output.LayerNorm.beta",
  "layers.10.attention.query.weight": "bert.encoder.layer.10.attention.self.query.weight",
  "layers.10.attention.query.bias": "bert.encoder.layer.10.attention.self.query.bias",
  "layers.10.attention.key.weight": "bert.encoder.layer.10.attention.self.key.weight",
  "layers.10.attention.key.bias": "bert.encoder.layer.10.attention.self.key.bias",
  "layers.10.attention.value.weight": "bert.encoder.layer.10.attention.self.value.weight",
  "layers.10.attention.value.bias": "bert.encoder.layer.10.attention.self.value.bias",
  "layers.10.attention.out.weight": "bert.encoder.layer.10.attention.output.dense.weight",
  "layers.10.attention.out.bias": "bert.encoder.layer.10.attention.output.dense.bias",
  "layers.10.attention_norm.weight": "bert.encoder.layer.10.attention.output.LayerNorm.gamma",
  "layers.10.attention_norm.bias": "bert.encoder.layer.10.attention.output.LayerNorm.beta",
  "layers.10.up.weight": "bert.encoder.layer.10.intermediate.dense.weight",
  "layers.10.up.bias": "bert.encoder.layer.10.intermediate.dense.bias",
  "layers.10.down.weight": "bert.encoder.layer.10.output.dense.weight",
  "layers.10.down.bias": "bert.encoder.layer.10.output.dense.bias",
  "layers.10.output_norm.weight": "bert.encoder.layer.10.output.LayerNorm.gamma",
  "layers.10.output_norm.bias": "bert.encoder.layer.10.output.LayerNorm.beta",
  "layers.11.attention.query.weight": "bert.encoder.layer.11.attention.self.query.weight",
  "layers.11.attention.query.bias": "bert.encoder.layer.11.attention.self.query.bias",
  "layers.11.attention.key.weight": "bert.encoder.layer.11.attention.self.key.weight",
  "layers.11.attention.key.bias": "bert.encoder.layer.11.attention.self.key.bias",
  "layers.11.attention.value.weight": "bert.encoder.layer.11.attention.self.value.weight",
  "layers.11.attention.value.bias": "bert.encoder.layer.11.attention.self.value.bias",
  "layers.11.attention.out.weight": "bert.encoder.layer.11.attention.output.dense.weight",
  "layers.11.attention.out.bias": "bert.encoder.layer.11.attention.output.dense.bias",
  "layers.11.attention_norm.weight": "bert.encoder.layer.11.attention.output.LayerNorm.gamma",
  "layers.11.attention_norm.bias": "bert.encoder.layer.11.attention.output.LayerNorm.beta",
  "layers.11.up.weight": "bert.encoder.layer.11.intermediate.dense.weight",
  "layers.11.up.bias": "bert.encoder.layer.11.intermediate.dense.bias",
  "layers.11.down.weight": "bert.encoder.layer.11.output.dense.weight",
  "layers.11.down.bias": "bert.encoder.layer.11.output.dense.bias",
  "layers.11.output_norm.weight": "bert.encoder.layer.11.output.LayerNorm.gamma",
  "layers.11.output_norm.bias": "bert.encoder.layer.11.output.LayerNorm.beta",
  "pooler.dense.weight": "bert.pooler.dense.weight",
  "pooler.dense.bias": "bert.pooler.dense.bias"
}

Open on GitHub

linnet.toml59 B
toml
[package]
name = "bert"
version = "0.1.0"
language = "0.1"

Open on GitHub

nest.toml847 B
toml
[model]
name = "bert-base-uncased"
title = "BERT Base (uncased)"
summary = "Google's original BERT: 12 post-norm encoder layers of width 768, word/position/token-type embeddings, and the next-sentence-prediction pooler."
license = "Apache-2.0"
family = "bert"
tags = ["fill-mask", "text-encoder", "encoder-only"]

[links]
huggingface = "https://huggingface.co/google-bert/bert-base-uncased"
github = "https://github.com/google-research/bert"
arxiv = "https://arxiv.org/abs/1810.04805"

[source]
path = "src/lib.linnet"
root = "Model"
entry = "forward"

[generics]
Vocab = 30522
MaxPositions = 512
TypeVocab = 2
D = 768
Heads = 12
Inner = 3072
Layers = 12
T = "f32"

[check]
B = 1
S = 8

[weights]
repo = "google-bert/bert-base-uncased"
revision = "86b5e0934494bd15c9632b12f734a8a67f723594"
files = ["model.safetensors"]
bindings = "bindings.json"

Open on GitHub

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.647789,
      0.583012,
      0.548669,
      0.425511,
      0.474948
    ],
    [
      0.647789,
      1.0,
      0.571008,
      0.570269,
      0.416706,
      0.501394
    ],
    [
      0.583012,
      0.571008,
      1.0,
      0.648335,
      0.447611,
      0.49322
    ],
    [
      0.548669,
      0.570269,
      0.648335,
      1.0,
      0.495576,
      0.544291
    ],
    [
      0.425511,
      0.416706,
      0.447611,
      0.495576,
      1.0,
      0.759655
    ],
    [
      0.474948,
      0.501394,
      0.49322,
      0.544291,
      0.759655,
      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."
}

Open on GitHub

src
lib.linnet4.9 kB
linnet
// BERT: absolute learned positions summed with token-type embeddings,
// normalized once before the encoder, then twelve post-norm blocks. Post-
// norm is the opposite order from the pre-norm ViT and GPT-2 already in the
// zoo: the residual add happens first and LayerNorm reads the sum --
// `h = LayerNorm(x + attention(x))`, then `out = LayerNorm(h + ffn(h))` --
// rather than normalizing the branch's input. There is no final norm after
// the last layer, since each block already ends with one.
module bert

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}

pub const EPS: f32 = 1e-12

pub block LayerNorm<D: Dim, T: Float> {
    param weight: Tensor[D; T]
    param bias: Tensor[D; T]

    pub fn forward<*S: Shape>(x: Tensor[*S, D; T]) -> Tensor[*S, D; T] {
        return layer_norm(x, weight, some(bias), EPS)
    }
}

pub block Embeddings<Vocab: Dim, MaxPositions: Dim, TypeVocab: Dim, D: Dim, T: Float> {
    param word_embeddings: Tensor[Vocab, D; T]
    param position_embeddings: Tensor[MaxPositions, D; T]
    param token_type_embeddings: Tensor[TypeVocab, D; T]
    sub norm: LayerNorm<D, T>

    pub fn forward<B: Dim, S: Dim>(
        tokens: Tensor[B, S; i32],
        token_types: Tensor[B, S; i32],
    ) -> Tensor[B, S, D; T]
    where S <= MaxPositions {
        // The three lookups agree on width D; positions are the sequence's
        // own indices, shared across the batch by right-aligned broadcast.
        let x =
            embedding(tokens, word_embeddings) +
            embedding(iota<i32>(S), position_embeddings) +
            embedding(token_types, token_type_embeddings)
        return norm.forward(x)
    }
}

pub block Attention<D: Dim, Heads: Dim, T: Float>
where
    Heads > 0,
    D % Heads == 0
{
    sub query: Linear<D, D, T>
    sub key: Linear<D, D, T>
    sub value: Linear<D, D, T>
    sub out: Linear<D, D, T>

    // Full attention over every position: the language's attention op takes
    // an unbatched [Q, K] mask, so a per-sequence padding mask cannot be
    // expressed here. Callers pad-free or accept attention onto padding.
    pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, D; T]) -> Tensor[B, S, D; T] {
        let mixed = attention(
            heads<B, S, Heads, D / Heads, T>(query.forward(x)),
            heads<B, S, Heads, D / Heads, T>(key.forward(x)),
            heads<B, S, Heads, D / Heads, T>(value.forward(x)),
            rsqrt(cast<f32>(D / Heads)),
            none,
        )
        return out.forward(reshape(permute(mixed, [0, 2, 1, 3]), [B, S, D]))
    }
}

fn heads<B: Dim, S: Dim, Heads: Dim, Dh: Dim, T: Float>(
    x: Tensor[B, S, Heads * Dh; T],
) -> Tensor[B, Heads, S, Dh; T] {
    return permute(reshape(x, [B, S, Heads, Dh]), [0, 2, 1, 3])
}

// Post-norm encoder block: `h = LayerNorm(x + attention(x))`, then
// `out = LayerNorm(h + ffn(h))`.
pub block EncoderLayer<D: Dim, Heads: Dim, Inner: Dim, T: Float>
where
    Heads > 0,
    D % Heads == 0
{
    sub attention: Attention<D, Heads, T>
    sub attention_norm: LayerNorm<D, T>
    sub up: Linear<D, Inner, T>
    sub down: Linear<Inner, D, T>
    sub output_norm: LayerNorm<D, T>

    pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, D; T]) -> Tensor[B, S, D; T] {
        let attended = attention_norm.forward(x + attention.forward(x))
        return output_norm.forward(attended + down.forward(gelu_erf(up.forward(attended))))
    }
}

// `tanh(dense(hidden[:, 0, :]))`: BERT's sentence representation, trained
// against next-sentence prediction and read by downstream classifiers.
pub block Pooler<D: Dim, T: Float> {
    sub dense: Linear<D, D, T>

    pub fn forward<B: Dim>(x: Tensor[B, D; T]) -> Tensor[B, D; T] {
        return tanh(dense.forward(x))
    }
}

pub block Model<
    Vocab: Dim,
    MaxPositions: Dim,
    TypeVocab: Dim,
    D: Dim,
    Heads: Dim,
    Inner: Dim,
    Layers: Dim,
    T: Float = f32,
>
where
    Heads > 0,
    D % Heads == 0
{
    sub embeddings: Embeddings<Vocab, MaxPositions, TypeVocab, D, T>
    sub layers: [EncoderLayer<D, Heads, Inner, T>; Layers]
    sub pooler: Pooler<D, T>

    // The last hidden state: every token's contextualized representation.
    pub entry forward<B: Dim, S: Dim>(
        tokens: Tensor[B, S; i32],
        token_types: Tensor[B, S; i32],
    ) -> Tensor[B, S, D; T]
    where S <= MaxPositions {
        var x = embeddings.forward(tokens, token_types)
        static for layer in layers {
            x = layer.forward(x)
        }
        return x
    }

    // The pooled output classifiers read: the first token of the last
    // hidden state, through the pooler.
    pub entry pool<B: Dim, S: Dim>(
        tokens: Tensor[B, S; i32],
        token_types: Tensor[B, S; i32],
    ) -> Tensor[B, D; T]
    where S <= MaxPositions {
        let hidden = forward(tokens, token_types)
        return pooler.forward(hidden[:, 0, :])
    }
}

Open on GitHub

model.safetensorsHugging Face ↗