Models / all-MiniLM-L6-v2

all-MiniLM-L6-v2

The most downloaded model on the Hub: a 6-layer, width-384 BERT encoder distilled for sentence embeddings, mean-pooled and L2-normalized.

22.6M parametersbertApache-2.0sentence-similaritytext-encoderencoder-onlyfeature-extraction
models/all-MiniLM-L6-v28 files · 28.2 kB in the registry · GitHub
README.md2.9 kB
md
# all-MiniLM-L6-v2

The most downloaded model on the Hugging Face Hub: a BERT encoder distilled
down to 6 post-norm layers of width 384 with 12 heads, wrapped by
sentence-transformers into a general-purpose sentence embedding model. The
`Transformer` stage is exactly a small BERT -- word, learned-position, and
token-type embeddings summed and normalized, then post-norm blocks,
`h = LayerNorm(x + attention(x))` then `out = LayerNorm(h + ffn(h))`, the
same order `bert-base-uncased` and `roberta-base` use and the pre-norm ViT
and GPT-2 already in Nest do not.

The Hub repository's `modules.json` chains three stages: `Transformer`,
`1_Pooling` (`1_Pooling/config.json` sets `pooling_mode_mean_tokens: true`,
i.e. mean pooling over tokens, no CLS or max pooling), and `2_Normalize`
(unit L2 norm). Neither pooling stage has learned weights -- the
checkpoint's `pooler.dense` tensors belong to the underlying `BertModel`
architecture and go unused by the sentence-transformers pipeline, so this
card does not bind them. The `embed` entry below computes the two stages
directly.

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`.

## Loading

```python
from linnet import nest

model = nest.load("all-MiniLM-L6-v2", backend="torch")
hidden = model(tokens, token_types)          # Tensor[B, S; i32], Tensor[B, S; i32] -> Tensor[B, S, 384; f32]
vectors = model.embed(tokens, token_types)   # -> Tensor[B, 384; f32], unit norm
```

`token_types` is the segment-id tensor; sentence-transformers always passes
zeros for it. Sequences longer than 512 positions are rejected at compile
time. The encoder runs full attention over every position and `embed`'s
mean pooling averages every position uniformly: `std.nn.attention::
attention` takes an unbatched `[Q, K]` mask, so a per-sequence padding mask
cannot be expressed here, and pooling likewise has no mask to exclude
padding from the mean. A pad-free batch (or a batch of one) matches
upstream's masked mean exactly, since every position is then real; a padded
batch does not.

## Entries

| Entry | |
| --- | --- |
| `forward<B, S>(tokens, token_types)` | last hidden state, `Tensor[B, S, 384; f32]` |
| `embed<B, S>(tokens, token_types)` | mean-pooled, L2-normalized sentence embedding, `Tensor[B, 384; f32]` |

## Numerics

Validated against `transformers.AutoModel.from_pretrained` (`numerics=
"fast"`) on CPU in f32 over pad-free sentences: the maximum absolute
difference is 1.8e-6 on `forward` (`last_hidden_state`) and 1.2e-7 on
`embed`, the mean-pooled, L2-normalized sentence embedding.

## Provenance

- Weights: [sentence-transformers/all-MiniLM-L6-v2](https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2), Apache-2.0.
- Code: [UKPLab/sentence-transformers](https://github.com/UKPLab/sentence-transformers).
- Paper: [Sentence-BERT: Sentence Embeddings using Siamese BERT-Networks](https://arxiv.org/abs/1908.10084).

Open on GitHub

bench.json10.0 kB
json
{
  "date": "2026-09-28T17:57:04+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": 1.90711859613657,
        "throughput_per_s": 26734.13406772362,
        "load_s": 2.193337509408593,
        "peak_vram_mib": 880.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": 0.8391356095671654,
        "throughput_per_s": 50394.62898285235,
        "load_s": 1.4945701844990253,
        "peak_vram_mib": 892.0
      },
      "max_abs_diff": 0.046875,
      "notes": "unpadded batches of exactly 128 tokens (this card takes no padding mask)",
      "error": null,
      "key": "transformers-compile"
    },
    {
      "method": "sentence-transformers (SentenceTransformer.encode)",
      "kind": "reference",
      "metrics": {
        "latency_ms": 3.21220513433218,
        "throughput_per_s": 6430.957616933986,
        "load_s": 2.441088568419218,
        "peak_vram_mib": 824.0
      },
      "max_abs_diff": null,
      "notes": "natural text tokenizing to 65 tokens, identical across the batch (no padding wasted); not the fixed-length random-token workload the other rows use, since this is the library's own string-in interface",
      "error": null,
      "key": "sentence-transformers"
    },
    {
      "method": "Linnet torch (generated source)",
      "kind": "linnet",
      "metrics": {
        "latency_ms": 1.0267235338687897,
        "throughput_per_s": 49771.51191968706,
        "load_s": 0.6230692155659199,
        "peak_vram_mib": 864.0
      },
      "max_abs_diff": 0.0625,
      "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.3216182813048363,
        "throughput_per_s": 74093.31681818966,
        "load_s": 0.5717327613383532,
        "peak_vram_mib": 968.0
      },
      "max_abs_diff": 0.0625,
      "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": 0.7865410298109055,
        "throughput_per_s": 70709.01632843312,
        "load_s": 0.6318587809801102,
        "peak_vram_mib": 840.0
      },
      "max_abs_diff": 0.0625,
      "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": 0.5229143425822258,
        "throughput_per_s": 38961.21353084718,
        "load_s": 0.3341801352798939,
        "peak_vram_mib": 346.9,
        "driver_vram_mib": 1134.0
      },
      "max_abs_diff": 0.046875,
      "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.44460780918598175,
        "throughput_per_s": 39720.15156196986,
        "load_s": 0.2924486342817545,
        "peak_vram_mib": 312.6,
        "driver_vram_mib": 1132.0
      },
      "max_abs_diff": 0.046875,
      "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.2142173945903778,
        "throughput_per_s": 20250.133706120938,
        "load_s": 0.2902862671762705,
        "peak_vram_mib": 1492.0
      },
      "max_abs_diff": 0.06287145614624023,
      "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.3458624705672264,
        "throughput_per_s": 34138.385079146814,
        "load_s": 0.2810468841344118,
        "peak_vram_mib": 2556.0
      },
      "max_abs_diff": 0.06383466720581055,
      "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.0503623634576797,
        "throughput_per_s": 34377.9414880179,
        "load_s": 0.286095617339015,
        "peak_vram_mib": 1140.0
      },
      "max_abs_diff": 0.064453125,
      "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.22809021174907684,
        "throughput_per_s": 80631.77298087563,
        "load_s": 0.28231428004801273,
        "peak_vram_mib": 2464.0
      },
      "max_abs_diff": 0.05859375,
      "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.0450975969433784,
        "throughput_per_s": 34780.49541123745,
        "load_s": 0.28188692033290863,
        "peak_vram_mib": 1140.0
      },
      "max_abs_diff": 0.03515625,
      "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.24310406297445297,
        "throughput_per_s": 74958.57910960921,
        "load_s": 0.27464193664491177,
        "peak_vram_mib": 2464.0
      },
      "max_abs_diff": 0.046875,
      "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": 3.9627309888601303,
        "throughput_per_s": 11660.857981648263,
        "load_s": 6.7153689954429865,
        "peak_vram_mib": 254.6,
        "driver_vram_mib": 880.0
      },
      "max_abs_diff": 0.046875,
      "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": 1.0246848687529564,
        "throughput_per_s": 19155.788798572783,
        "load_s": 12.43218300305307,
        "peak_vram_mib": 1486.0
      },
      "max_abs_diff": 0.06286907196044922,
      "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": 2.147151157259941,
        "throughput_per_s": 1106.6183734645347,
        "load_s": 6.18403435498476,
        "peak_vram_mib": 1558.9
      },
      "max_abs_diff": 0.06286907196044922,
      "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": 2.154892310500145,
        "throughput_per_s": 1110.7647068229883,
        "load_s": 10.043424438685179,
        "peak_vram_mib": 1590.9
      },
      "max_abs_diff": 0.06287145614624023,
      "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": 1.3945326209068298,
        "throughput_per_s": 1100.184687322089,
        "load_s": 6.025819133967161,
        "peak_vram_mib": 2434.1
      },
      "max_abs_diff": 0.0625,
      "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.json7.5 kB
json
{
  "embeddings.word_embeddings": "embeddings.word_embeddings.weight",
  "embeddings.position_embeddings": "embeddings.position_embeddings.weight",
  "embeddings.token_type_embeddings": "embeddings.token_type_embeddings.weight",
  "embeddings.norm.weight": "embeddings.LayerNorm.weight",
  "embeddings.norm.bias": "embeddings.LayerNorm.bias",
  "layers.0.attention.query.weight": "encoder.layer.0.attention.self.query.weight",
  "layers.0.attention.query.bias": "encoder.layer.0.attention.self.query.bias",
  "layers.0.attention.key.weight": "encoder.layer.0.attention.self.key.weight",
  "layers.0.attention.key.bias": "encoder.layer.0.attention.self.key.bias",
  "layers.0.attention.value.weight": "encoder.layer.0.attention.self.value.weight",
  "layers.0.attention.value.bias": "encoder.layer.0.attention.self.value.bias",
  "layers.0.attention.out.weight": "encoder.layer.0.attention.output.dense.weight",
  "layers.0.attention.out.bias": "encoder.layer.0.attention.output.dense.bias",
  "layers.0.attention_norm.weight": "encoder.layer.0.attention.output.LayerNorm.weight",
  "layers.0.attention_norm.bias": "encoder.layer.0.attention.output.LayerNorm.bias",
  "layers.0.up.weight": "encoder.layer.0.intermediate.dense.weight",
  "layers.0.up.bias": "encoder.layer.0.intermediate.dense.bias",
  "layers.0.down.weight": "encoder.layer.0.output.dense.weight",
  "layers.0.down.bias": "encoder.layer.0.output.dense.bias",
  "layers.0.output_norm.weight": "encoder.layer.0.output.LayerNorm.weight",
  "layers.0.output_norm.bias": "encoder.layer.0.output.LayerNorm.bias",
  "layers.1.attention.query.weight": "encoder.layer.1.attention.self.query.weight",
  "layers.1.attention.query.bias": "encoder.layer.1.attention.self.query.bias",
  "layers.1.attention.key.weight": "encoder.layer.1.attention.self.key.weight",
  "layers.1.attention.key.bias": "encoder.layer.1.attention.self.key.bias",
  "layers.1.attention.value.weight": "encoder.layer.1.attention.self.value.weight",
  "layers.1.attention.value.bias": "encoder.layer.1.attention.self.value.bias",
  "layers.1.attention.out.weight": "encoder.layer.1.attention.output.dense.weight",
  "layers.1.attention.out.bias": "encoder.layer.1.attention.output.dense.bias",
  "layers.1.attention_norm.weight": "encoder.layer.1.attention.output.LayerNorm.weight",
  "layers.1.attention_norm.bias": "encoder.layer.1.attention.output.LayerNorm.bias",
  "layers.1.up.weight": "encoder.layer.1.intermediate.dense.weight",
  "layers.1.up.bias": "encoder.layer.1.intermediate.dense.bias",
  "layers.1.down.weight": "encoder.layer.1.output.dense.weight",
  "layers.1.down.bias": "encoder.layer.1.output.dense.bias",
  "layers.1.output_norm.weight": "encoder.layer.1.output.LayerNorm.weight",
  "layers.1.output_norm.bias": "encoder.layer.1.output.LayerNorm.bias",
  "layers.2.attention.query.weight": "encoder.layer.2.attention.self.query.weight",
  "layers.2.attention.query.bias": "encoder.layer.2.attention.self.query.bias",
  "layers.2.attention.key.weight": "encoder.layer.2.attention.self.key.weight",
  "layers.2.attention.key.bias": "encoder.layer.2.attention.self.key.bias",
  "layers.2.attention.value.weight": "encoder.layer.2.attention.self.value.weight",
  "layers.2.attention.value.bias": "encoder.layer.2.attention.self.value.bias",
  "layers.2.attention.out.weight": "encoder.layer.2.attention.output.dense.weight",
  "layers.2.attention.out.bias": "encoder.layer.2.attention.output.dense.bias",
  "layers.2.attention_norm.weight": "encoder.layer.2.attention.output.LayerNorm.weight",
  "layers.2.attention_norm.bias": "encoder.layer.2.attention.output.LayerNorm.bias",
  "layers.2.up.weight": "encoder.layer.2.intermediate.dense.weight",
  "layers.2.up.bias": "encoder.layer.2.intermediate.dense.bias",
  "layers.2.down.weight": "encoder.layer.2.output.dense.weight",
  "layers.2.down.bias": "encoder.layer.2.output.dense.bias",
  "layers.2.output_norm.weight": "encoder.layer.2.output.LayerNorm.weight",
  "layers.2.output_norm.bias": "encoder.layer.2.output.LayerNorm.bias",
  "layers.3.attention.query.weight": "encoder.layer.3.attention.self.query.weight",
  "layers.3.attention.query.bias": "encoder.layer.3.attention.self.query.bias",
  "layers.3.attention.key.weight": "encoder.layer.3.attention.self.key.weight",
  "layers.3.attention.key.bias": "encoder.layer.3.attention.self.key.bias",
  "layers.3.attention.value.weight": "encoder.layer.3.attention.self.value.weight",
  "layers.3.attention.value.bias": "encoder.layer.3.attention.self.value.bias",
  "layers.3.attention.out.weight": "encoder.layer.3.attention.output.dense.weight",
  "layers.3.attention.out.bias": "encoder.layer.3.attention.output.dense.bias",
  "layers.3.attention_norm.weight": "encoder.layer.3.attention.output.LayerNorm.weight",
  "layers.3.attention_norm.bias": "encoder.layer.3.attention.output.LayerNorm.bias",
  "layers.3.up.weight": "encoder.layer.3.intermediate.dense.weight",
  "layers.3.up.bias": "encoder.layer.3.intermediate.dense.bias",
  "layers.3.down.weight": "encoder.layer.3.output.dense.weight",
  "layers.3.down.bias": "encoder.layer.3.output.dense.bias",
  "layers.3.output_norm.weight": "encoder.layer.3.output.LayerNorm.weight",
  "layers.3.output_norm.bias": "encoder.layer.3.output.LayerNorm.bias",
  "layers.4.attention.query.weight": "encoder.layer.4.attention.self.query.weight",
  "layers.4.attention.query.bias": "encoder.layer.4.attention.self.query.bias",
  "layers.4.attention.key.weight": "encoder.layer.4.attention.self.key.weight",
  "layers.4.attention.key.bias": "encoder.layer.4.attention.self.key.bias",
  "layers.4.attention.value.weight": "encoder.layer.4.attention.self.value.weight",
  "layers.4.attention.value.bias": "encoder.layer.4.attention.self.value.bias",
  "layers.4.attention.out.weight": "encoder.layer.4.attention.output.dense.weight",
  "layers.4.attention.out.bias": "encoder.layer.4.attention.output.dense.bias",
  "layers.4.attention_norm.weight": "encoder.layer.4.attention.output.LayerNorm.weight",
  "layers.4.attention_norm.bias": "encoder.layer.4.attention.output.LayerNorm.bias",
  "layers.4.up.weight": "encoder.layer.4.intermediate.dense.weight",
  "layers.4.up.bias": "encoder.layer.4.intermediate.dense.bias",
  "layers.4.down.weight": "encoder.layer.4.output.dense.weight",
  "layers.4.down.bias": "encoder.layer.4.output.dense.bias",
  "layers.4.output_norm.weight": "encoder.layer.4.output.LayerNorm.weight",
  "layers.4.output_norm.bias": "encoder.layer.4.output.LayerNorm.bias",
  "layers.5.attention.query.weight": "encoder.layer.5.attention.self.query.weight",
  "layers.5.attention.query.bias": "encoder.layer.5.attention.self.query.bias",
  "layers.5.attention.key.weight": "encoder.layer.5.attention.self.key.weight",
  "layers.5.attention.key.bias": "encoder.layer.5.attention.self.key.bias",
  "layers.5.attention.value.weight": "encoder.layer.5.attention.self.value.weight",
  "layers.5.attention.value.bias": "encoder.layer.5.attention.self.value.bias",
  "layers.5.attention.out.weight": "encoder.layer.5.attention.output.dense.weight",
  "layers.5.attention.out.bias": "encoder.layer.5.attention.output.dense.bias",
  "layers.5.attention_norm.weight": "encoder.layer.5.attention.output.LayerNorm.weight",
  "layers.5.attention_norm.bias": "encoder.layer.5.attention.output.LayerNorm.bias",
  "layers.5.up.weight": "encoder.layer.5.intermediate.dense.weight",
  "layers.5.up.bias": "encoder.layer.5.intermediate.dense.bias",
  "layers.5.down.weight": "encoder.layer.5.output.dense.weight",
  "layers.5.down.bias": "encoder.layer.5.output.dense.bias",
  "layers.5.output_norm.weight": "encoder.layer.5.output.LayerNorm.weight",
  "layers.5.output_norm.bias": "encoder.layer.5.output.LayerNorm.bias"
}

Open on GitHub

linnet.toml61 B
toml
[package]
name = "minilm"
version = "0.1.0"
language = "0.1"

Open on GitHub

nest.toml894 B
toml
[model]
name = "all-MiniLM-L6-v2"
title = "all-MiniLM-L6-v2"
summary = "The most downloaded model on the Hub: a 6-layer, width-384 BERT encoder distilled for sentence embeddings, mean-pooled and L2-normalized."
license = "Apache-2.0"
family = "bert"
tags = ["sentence-similarity", "text-encoder", "encoder-only", "feature-extraction"]

[links]
huggingface = "https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2"
github = "https://github.com/UKPLab/sentence-transformers"
arxiv = "https://arxiv.org/abs/1908.10084"

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

[generics]
Vocab = 30522
MaxPositions = 512
TypeVocab = 2
D = 384
Heads = 12
Inner = 1536
Layers = 6
T = "f32"

[check]
B = 1
S = 8

[weights]
repo = "sentence-transformers/all-MiniLM-L6-v2"
revision = "1110a243fdf4706b3f48f1d95db1a4f5529b4d41"
files = ["model.safetensors"]
bindings = "bindings.json"

Open on GitHub

samples.json1.9 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.275979,
      0.085964,
      0.051419,
      0.112605,
      0.023628
    ],
    [
      0.275979,
      1.0,
      -0.049659,
      0.050696,
      0.10739,
      -0.01439
    ],
    [
      0.085964,
      -0.049659,
      1.0,
      0.307588,
      -0.010309,
      0.186634
    ],
    [
      0.051419,
      0.050696,
      0.307588,
      1.0,
      0.088557,
      0.126785
    ],
    [
      0.112605,
      0.10739,
      -0.010309,
      0.088557,
      1.0,
      0.085097
    ],
    [
      0.023628,
      -0.01439,
      0.186634,
      0.126785,
      0.085097,
      1.0
    ]
  ],
  "caption": "The card's own `embed` entry: mean-pooled, L2-normalized sentence embeddings.",
  "reference_matrix": [
    [
      1.0,
      0.274703,
      0.085527,
      0.051399,
      0.112673,
      0.023784
    ],
    [
      0.274703,
      1.0,
      -0.050362,
      0.04726,
      0.108435,
      -0.01416
    ],
    [
      0.085527,
      -0.050362,
      1.0,
      0.306304,
      -0.011378,
      0.186581
    ],
    [
      0.051399,
      0.04726,
      0.306304,
      1.0,
      0.08871,
      0.125638
    ],
    [
      0.112673,
      0.108435,
      -0.011378,
      0.08871,
      1.0,
      0.084571
    ],
    [
      0.023784,
      -0.01416,
      0.186581,
      0.125638,
      0.084571,
      1.0
    ]
  ],
  "reference_stack": "sentence-transformers (SentenceTransformer.encode)"
}

Open on GitHub

src
lib.linnet4.9 kB
linnet
// all-MiniLM-L6-v2: the same post-norm BERT encoder as `bert-base-uncased`
// and `roberta-base`, six layers deep and 384 wide, wrapped by
// sentence-transformers into a sentence embedding model. The Transformer
// module produces the last hidden state exactly as a plain BERT would; the
// `1_Pooling` module (mode: mean of token embeddings) and `2_Normalize`
// module (unit L2 norm) that the Hub repository's `modules.json` names are
// not learned weights, so `embed` computes them directly.
module minilm

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))))
    }
}

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]

    // 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 sentence embedding sentence-transformers publishes: every token
    // of the last hidden state averaged (mean pooling, no attention-mask
    // weighting since this encoder runs full attention over every
    // position), then scaled to unit length. This is what `1_Pooling`
    // (mean_tokens) and `2_Normalize` do in the Hub pipeline.
    pub entry embed<B: Dim, S: Dim>(
        tokens: Tensor[B, S; i32],
        token_types: Tensor[B, S; i32],
    ) -> Tensor[B, D; f32]
    where S <= MaxPositions {
        let hidden = cast<f32>(forward(tokens, token_types))
        let pooled[b, d] = sum<f32>[s] hidden[b, s, d] / cast<f32>(S)
        let squared[b] = sum<f32>[d] pooled[b, d] * pooled[b, d]
        let unit[b, d] = pooled[b, d] * rsqrt(squared[b] + 1e-12)
        return unit
    }
}

Open on GitHub

model.safetensorsHugging Face ↗