Models / mistral-7b-instruct-v0.3

Mistral 7B Instruct v0.3

A 7B-parameter Llama-architecture decoder with grouped-query attention, a 32768-token vocabulary and context, and an untied output head, instruction-tuned.

7.2B parametersllamaApache-2.0text-generationdecoder-onlygrouped-query-attentionchat
models/mistral-7b-instruct-v0.312 files · 88.9 kB in the registry · GitHub
README.md5.8 kB
md
# Mistral 7B Instruct v0.3

A 7.25B-parameter decoder with the Llama architecture, instruction-tuned:
grouped-query attention (32 query heads, 8 key/value heads, head dimension
128), rotary positions with base 1000000, RMS normalization, a SwiGLU MLP,
and an untied output head, with a KV cache for token-by-token decoding. The
context is 32768 tokens; v0.3 has no sliding window (`config.json` sets
`sliding_window` to `null`), so attention is full rather than windowed, and
its vocabulary of 32768 tokens is larger than v0.2's.

The Linnet source is the same Llama package the other decoders in this zoo
use (`tinyllama-1.1b-chat`, `smollm2-1.7b-instruct`), with `THETA` set to
this model's rotary base and the generics set to its width, head counts, and
depth. Unlike those two, the output head is not tied to the embedding:
`bindings.json` maps `lm_head.weight` to the checkpoint's own separate
`lm_head.weight` tensor.

## Loading

```python
from linnet import nest

model = nest.load("mistral-7b-instruct-v0.3", backend="torch", numerics="fast")
logits = model(tokens)                      # Tensor[1, S; i32] -> Tensor[1, S, 32768; bf16]
```

The weights are the published, sharded `model-0000N-of-00003.safetensors`
(bf16); `bindings.json` maps Linnet parameter paths to their tensor names.
Every name, shape, and dtype is checked before anything runs.

## Entries

| Entry | |
| --- | --- |
| `forward<B, S>(tokens)` | logits for a whole sequence |
| `next_token<B, S>(tokens)` | logits for the last position |
| `decode(token, pos)` | one token through the KV caches (`Batch = 1`, `MaxSeq = 32768`) |
| `prefill<S>(tokens, pos)` | a whole prompt through the KV caches in one pass; logits after its last token, `decode` continues at `pos + S` |
| `prefill_slots<M, S>(tokens, slots, lengths)` | `M` requests' prompts into rows `slots` of the caches in one pass (each padded to `S`, its first `lengths[m]` tokens real), as they join a batch being served |
| `prefill_packed<P>(tokens, rows, positions, segments, last)` | several requests' prompts packed end to end into one pass of `P` tokens, each token given its cache row, its position, and its prompt: no padding between prompts, and each prompt sees only itself |
| `decode_rows(tokens, positions)` | one token for every row of the caches, each at its own position: the step `linnet.serve` takes for continuous batching |
| `prefill_paged<P, Rows, Pages>(tokens, positions, rows, slots, last, table)` | for serving from pages: prompts packed end to end, each token written at its place in the pool and attending over its row's pages up to its position (chunks, shared prefixes) |
| `decode_paged<Rows, Pages>(tokens, positions, table)` | `decode_rows` for serving from pages: loaded with `Batch = 1`, the caches' one row is a pool of `MaxSeq` positions in pages of `PageSize` (64), and row `b`'s positions lie in the pages `table[b]` lists |
| `step_paged<P, Rows, Pages>(...)` | `prefill_paged` and `decode_paged` in one pass |
| `generate<Steps>(token, pos)` | greedy decoding in the graph |
| `sample<Steps>(token, pos, key, temperature)` | sampling with `std.random` |
| `generate_until<MaxNew>(token, pos, eos)` | decoding until an end token |

## Shards

`Shards` (1 unless bound) splits the model across devices: one shard holds
`Heads / Shards` query heads, `KvHeads / Shards` key and value heads and
`Inner / Shards` hidden units, with the KV caches to match, and the
attention's and the MLP's output projections end in
`std.nn.parallel::all_reduce`, the sum over the shards. `lm_head` holds
`Vocab / Shards` rows of the vocabulary, and `std.nn.parallel::all_gather`
sets the shards' slices of the logits side by side. On one device the sum
and the gather are the value itself.
`linnet.torch.load(..., tensor_parallel=mesh)` binds `Shards` to the mesh
size and gives each process its part of the checkpoint.
The training entries train split as well: each split computation reads its
input through `std.nn.parallel::shared`, and `loss_packed` and
`log_probs_packed` run over the vocabulary's parts
(`std.nn.loss::split_cross_entropy`, `split_token_log_probs`).


## Provenance

- Weights: [mistralai/Mistral-7B-Instruct-v0.3](https://huggingface.co/mistralai/Mistral-7B-Instruct-v0.3), Apache-2.0.
- Code: [mistralai/mistral-inference](https://github.com/mistralai/mistral-inference).
- Paper: [Mistral 7B](https://arxiv.org/abs/2310.06825) (the base architecture; v0.3 extends its tokenizer and drops the sliding window).

## Validation

Checked against `transformers` (`MistralForCausalLM`) on CPU, `torch.no_grad()`,
under the project's memory discipline for multi-billion-parameter models.

- **Truncated, f32, tight tolerance -- done.** Both sides built from only
  the first 2 of 32 layers' real weights (`generics={"Layers": 2, "T":
  "f32"}` on the Linnet side, `num_hidden_layers=2` on the reference, both
  reading the same 21 tensors, cast to f32), prompt `"The capital of
  France is"`. Max absolute difference on the logits: **4.3e-6**; the
  last-position argmax agrees. (A first attempt left the Linnet side at
  its native bf16 against the f32 reference and saw a 6e-2 max absolute
  difference with the same argmax match -- bf16 storage-rounding noise
  consistent with this zoo's other bf16-only checkpoints, not an
  architecture bug; the all-f32 comparison above is the one that isolates
  the architecture.)
- **Whole model, bf16, loose tolerance -- deferred, not run here.** Loading
  all 32 layers in bf16 on both sides needs roughly 14 GB of RSS per side,
  which this pass would hold on the shared CPU machine this card was
  written on. By instruction, that measurement is deferred to a batched
  GPU validation run instead of being run here; treat it as **not yet
  checked** until that run reports a top-1 token agreement and a bf16-scale
  (order 1e-2) logits max absolute difference for this model.

Open on GitHub

bench.json18.5 kB
json
{
  "date": "2026-09-29T15:07:15+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": {
        "ttft_ms": 21.291871555149555,
        "decode_tok_s": 48.701821492335746,
        "load_s": 5.0697523560374975,
        "peak_vram_mib": 14844.9
      },
      "max_abs_diff": 0.0,
      "notes": "",
      "error": null,
      "key": "transformers-eager"
    },
    {
      "method": "transformers (torch.compile, static cache)",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 17.62433908879757,
        "decode_tok_s": 118.21297799149646,
        "load_s": 3.807374821975827,
        "peak_vram_mib": 14858.0
      },
      "max_abs_diff": 0.078125,
      "notes": "",
      "error": null,
      "key": "transformers-compile"
    },
    {
      "method": "vLLM",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 14.84755054116249,
        "decode_tok_s": 158.4807568922746,
        "load_s": 73.46515459939837,
        "first_token": 1054.0,
        "peak_vram_mib": 67804.9
      },
      "max_abs_diff": null,
      "notes": "reserves a KV-cache pool up front (gpu_memory_utilization 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "vllm"
    },
    {
      "method": "Linnet torch (generated source)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 17.19226036220789,
        "decode_tok_s": 75.2503109833677,
        "load_s": 3.1651483811438084,
        "peak_vram_mib": 14752.0
      },
      "max_abs_diff": 0.078125,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-torch"
    },
    {
      "method": "Linnet torch (CUDA graphs)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 12.724582105875015,
        "decode_tok_s": 169.2229410858576,
        "load_s": 2.940107148140669,
        "peak_vram_mib": 21816.9
      },
      "max_abs_diff": 0.0859375,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-cudagraphs"
    },
    {
      "method": "Linnet torch (torch.compile, inductor)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 12.86735013127327,
        "decode_tok_s": 143.29303738252253,
        "load_s": 3.179628014564514,
        "peak_vram_mib": 21816.9
      },
      "max_abs_diff": 0.0859375,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-inductor"
    },
    {
      "method": "Linnet JAX (XLA, StableHLO)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 16.004810109734535,
        "decode_tok_s": 174.16092750222566,
        "load_s": 10.189460884779692,
        "peak_vram_mib": 14528.6,
        "driver_vram_mib": 17020.9
      },
      "max_abs_diff": 0.13037109375,
      "notes": "KV cache compiled for 768 positions; 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": {
        "ttft_ms": 15.299675986170769,
        "decode_tok_s": 174.07943131337353,
        "load_s": 10.907724453136325,
        "peak_vram_mib": 14528.6,
        "driver_vram_mib": 17018.9
      },
      "max_abs_diff": 0.1279296875,
      "notes": "KV cache compiled for 768 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "linnet-jax-source"
    },
    {
      "method": "KerasHub (JAX)",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 27.385671623051167,
        "decode_tok_s": 166.14281268483572,
        "load_s": 35.09686706773937,
        "peak_vram_mib": 14904.7,
        "driver_vram_mib": 17020.0
      },
      "max_abs_diff": 0.09375,
      "notes": "peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "keras-hub"
    },
    {
      "method": "Linnet ONNX f32 -> ONNX Runtime (CUDA)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 41.49291478097439,
        "decode_tok_s": 55.866589468949925,
        "load_s": 0.438414016738534,
        "peak_vram_mib": 32940.0
      },
      "max_abs_diff": 0.07344776391983032,
      "notes": "f32; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-f32"
    },
    {
      "method": "Linnet ONNX f32 -> ONNX Runtime (TensorRT)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 50.22773239761591,
        "decode_tok_s": 29.660580071467624,
        "load_s": 0.4473560955375433,
        "peak_vram_mib": 75228.0
      },
      "max_abs_diff": 0.06902849674224854,
      "notes": "f32; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-f32-trt"
    },
    {
      "method": "Linnet ONNX f16 -> ONNX Runtime (CUDA)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 29.916991479694843,
        "decode_tok_s": 75.14534621440954,
        "load_s": 0.4814409166574478,
        "peak_vram_mib": 16906.0
      },
      "max_abs_diff": 0.07275390625,
      "notes": "f16; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-f16"
    },
    {
      "method": "Linnet ONNX f16 -> ONNX Runtime (TensorRT)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 25.960758328437805,
        "decode_tok_s": 51.87752159087435,
        "load_s": 0.42362662218511105,
        "peak_vram_mib": 38750.0
      },
      "max_abs_diff": 1.81640625,
      "notes": "f16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-f16-trt"
    },
    {
      "method": "Linnet ONNX bf16 -> ONNX Runtime (CUDA)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 30.519812367856503,
        "decode_tok_s": 73.98039909877826,
        "load_s": 0.44579144194722176,
        "peak_vram_mib": 15426.0
      },
      "max_abs_diff": 0.15625,
      "notes": "bf16; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-bf16"
    },
    {
      "method": "Linnet ONNX bf16 -> ONNX Runtime (TensorRT)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 26.95709653198719,
        "decode_tok_s": 50.293274726396305,
        "load_s": 0.4525479879230261,
        "peak_vram_mib": 37580.0
      },
      "max_abs_diff": 0.1484375,
      "notes": "bf16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
      "error": null,
      "key": "linnet-onnx-bf16-trt"
    },
    {
      "method": "vLLM (offline, continuous batching)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 5545.832737676782,
        "serve_s": 5.908580649644136,
        "load_s": 78.80427030473948,
        "peak_vram_mib": 70410.2
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at gpu_memory_utilization 0.85",
      "error": null,
      "key": "serve-vllm"
    },
    {
      "method": "transformers (generate_batch, continuous batching)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 914.5561570152817,
        "serve_s": 35.82940177991986,
        "load_s": 3.7513412758708,
        "peak_vram_mib": 81042.0
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; paged|sdpa attention; reserves a paged KV-cache pool up front, so its memory is a setting, not a need; no per-request timestamps",
      "error": null,
      "key": "serve-transformers"
    },
    {
      "method": "KerasHub (JAX, static batches)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 393.68417986785676,
        "serve_s": 83.2342310808599,
        "load_s": 37.684305580332875,
        "peak_vram_mib": 24669.0,
        "driver_vram_mib": 33394.0
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; every batch runs to the longest possible prompt plus the new tokens; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "serve-keras-hub"
    },
    {
      "method": "Triton Inference Server (vLLM backend)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 2070.5276965460675,
        "serve_s": 15.82591725513339,
        "load_s": 58.069482401013374,
        "serve_ttft_ms": 5190.573105588555,
        "peak_vram_mib": 70840.4
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed, first tokens timed at the client",
      "error": null,
      "key": "serve-triton-vllm"
    },
    {
      "method": "Linnet torch (linnet.serve, CUDA graphs)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 5764.448327956281,
        "serve_s": 5.684498868882656,
        "load_s": 8.005591470748186,
        "serve_ttft_ms": 2447.2957085818052,
        "peak_vram_mib": 27929.3
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions",
      "error": null,
      "key": "serve-linnet-torch"
    },
    {
      "method": "Linnet jax (linnet.serve, XLA)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 4796.311015290441,
        "serve_s": 6.831917257979512,
        "load_s": 10.897419090382755,
        "serve_ttft_ms": 3064.247688744217,
        "peak_vram_mib": 21097.2,
        "driver_vram_mib": 33466.0
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "serve-linnet-jax"
    },
    {
      "method": "Linnet onnx (linnet.serve, ONNX Runtime, f16)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 3616.6256272209207,
        "serve_s": 9.060379308648407,
        "load_s": 0.25492167938500643,
        "serve_ttft_ms": 4283.054237719625,
        "peak_vram_mib": 36408.0
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions",
      "error": null,
      "key": "serve-linnet-onnx"
    },
    {
      "method": "Triton Inference Server (Python backend, linnet.serve)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 5487.142701813233,
        "serve_s": 5.971778351813555,
        "load_s": 652.4048110358417,
        "serve_ttft_ms": 2513.456841930747,
        "peak_vram_mib": 27538.4
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed, first tokens timed at the client",
      "error": null,
      "key": "serve-triton-linnet"
    },
    {
      "method": "Linnet -> vLLM",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 16.070714220404625,
        "decode_tok_s": 165.57741345792672,
        "load_s": 34.13225184008479,
        "first_token": 1054.0,
        "peak_vram_mib": 67804.8,
        "driver_vram_mib": 67804.8
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (14 s), then vLLM on the exported checkpoint; reserves a KV-cache pool up front (gpu_memory_utilization 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "linnet-vllm"
    },
    {
      "method": "Linnet -> vLLM (offline, continuous batching)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 5848.769994291768,
        "serve_s": 5.602545497938991,
        "load_s": 29.347902312874794,
        "peak_vram_mib": 70338.8,
        "driver_vram_mib": 70338.8
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (0 s), then vLLM on the exported checkpoint; 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at gpu_memory_utilization 0.85",
      "error": null,
      "key": "serve-linnet-vllm"
    },
    {
      "method": "Linnet -> llama.cpp (GGUF, f16)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 21.10782248534182,
        "decode_tok_s": 179.083426,
        "load_s": 108.03528477065265,
        "peak_vram_mib": 14456.8
      },
      "max_abs_diff": null,
      "notes": "linnet.gguf.export, then llama-bench with every layer on the GPU; the first token is the prompt at llama-bench's prompt rate",
      "error": null,
      "key": "linnet-llamacpp"
    },
    {
      "method": "SGLang",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 19.374131225049496,
        "decode_tok_s": 165.49173198449475,
        "load_s": 59.06863775849342,
        "first_token": 1054.0,
        "peak_vram_mib": 69635.6
      },
      "max_abs_diff": null,
      "notes": "reserves a KV-cache pool up front (mem_fraction_static 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "sglang"
    },
    {
      "method": "Linnet -> SGLang",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 19.409666769206524,
        "decode_tok_s": 165.28226222907185,
        "load_s": 31.485597802326083,
        "first_token": 1054.0,
        "peak_vram_mib": 69635.6,
        "driver_vram_mib": 69635.6
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (15 s), then SGLang on the exported checkpoint; reserves a KV-cache pool up front (mem_fraction_static 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "linnet-sglang"
    },
    {
      "method": "SGLang (offline, continuous batching)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 4905.692160738505,
        "serve_s": 6.679587492719293,
        "load_s": 49.675609743222594,
        "peak_vram_mib": 71673.6
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at mem_fraction_static 0.85",
      "error": null,
      "key": "serve-sglang"
    },
    {
      "method": "Linnet -> SGLang (offline, continuous batching)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 4926.704639670798,
        "serve_s": 6.6510989386588335,
        "load_s": 36.524502135813236,
        "peak_vram_mib": 71673.6,
        "driver_vram_mib": 71673.6
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (0 s), then SGLang (offline, continuous batching) on the exported checkpoint; 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at mem_fraction_static 0.85",
      "error": null,
      "key": "serve-linnet-sglang"
    },
    {
      "method": "Text Generation Inference",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 23.002566769719124,
        "decode_tok_s": 118.05026112411038,
        "load_s": 52.11064511537552,
        "peak_vram_mib": 60642.9
      },
      "max_abs_diff": null,
      "notes": "over HTTP, streamed; the prompt is the token ids decoded and tokenized again; reserves a KV-cache pool up front (cuda-memory-fraction 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "tgi"
    },
    {
      "method": "Linnet -> Text Generation Inference",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 22.629539482295513,
        "decode_tok_s": 121.30967088917365,
        "load_s": 36.07101030834019,
        "peak_vram_mib": 60642.9,
        "driver_vram_mib": 60642.9
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (17 s), then Text Generation Inference on the exported checkpoint; over HTTP, streamed; the prompt is the token ids decoded and tokenized again; reserves a KV-cache pool up front (cuda-memory-fraction 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "linnet-tgi"
    },
    {
      "method": "Text Generation Inference (continuous batching)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 1664.2477352960555,
        "serve_s": 6.843332130461931,
        "load_s": 42.088825173676014,
        "serve_ttft_ms": 1649.5383251458406,
        "peak_vram_mib": 62520.9
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed; stops at end-of-sequence (TGI cannot ignore it), so throughput counts the tokens produced; KV-cache pool at cuda-memory-fraction 0.85",
      "error": null,
      "key": "serve-tgi"
    },
    {
      "method": "Linnet -> Text Generation Inference (continuous batching)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 1825.622420567375,
        "serve_s": 6.318393042311072,
        "load_s": 36.08311059884727,
        "serve_ttft_ms": 1408.4483832120895,
        "peak_vram_mib": 62646.9,
        "driver_vram_mib": 62646.9
      },
      "max_abs_diff": null,
      "notes": "linnet.hf.export (0 s), then Text Generation Inference (continuous batching) on the exported checkpoint; 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed; stops at end-of-sequence (TGI cannot ignore it), so throughput counts the tokens produced; KV-cache pool at cuda-memory-fraction 0.85",
      "error": null,
      "key": "serve-linnet-tgi"
    }
  ],
  "reference": "transformers-eager"
}

Open on GitHub

bindings.json21.8 kB
json
{
  "embedding.weight": "model.embed_tokens.weight",
  "norm.weight": "model.norm.weight",
  "lm_head.weight": "lm_head.weight",
  "layers.0.attention_norm.weight": "model.layers.0.input_layernorm.weight",
  "layers.0.mlp_norm.weight": "model.layers.0.post_attention_layernorm.weight",
  "layers.0.attention.q_proj.weight": "model.layers.0.self_attn.q_proj.weight",
  "layers.0.attention.k_proj.weight": "model.layers.0.self_attn.k_proj.weight",
  "layers.0.attention.v_proj.weight": "model.layers.0.self_attn.v_proj.weight",
  "layers.0.attention.o_proj.weight": "model.layers.0.self_attn.o_proj.weight",
  "layers.0.mlp.gate.weight": "model.layers.0.mlp.gate_proj.weight",
  "layers.0.mlp.up.weight": "model.layers.0.mlp.up_proj.weight",
  "layers.0.mlp.down.weight": "model.layers.0.mlp.down_proj.weight",
  "layers.1.attention_norm.weight": "model.layers.1.input_layernorm.weight",
  "layers.1.mlp_norm.weight": "model.layers.1.post_attention_layernorm.weight",
  "layers.1.attention.q_proj.weight": "model.layers.1.self_attn.q_proj.weight",
  "layers.1.attention.k_proj.weight": "model.layers.1.self_attn.k_proj.weight",
  "layers.1.attention.v_proj.weight": "model.layers.1.self_attn.v_proj.weight",
  "layers.1.attention.o_proj.weight": "model.layers.1.self_attn.o_proj.weight",
  "layers.1.mlp.gate.weight": "model.layers.1.mlp.gate_proj.weight",
  "layers.1.mlp.up.weight": "model.layers.1.mlp.up_proj.weight",
  "layers.1.mlp.down.weight": "model.layers.1.mlp.down_proj.weight",
  "layers.2.attention_norm.weight": "model.layers.2.input_layernorm.weight",
  "layers.2.mlp_norm.weight": "model.layers.2.post_attention_layernorm.weight",
  "layers.2.attention.q_proj.weight": "model.layers.2.self_attn.q_proj.weight",
  "layers.2.attention.k_proj.weight": "model.layers.2.self_attn.k_proj.weight",
  "layers.2.attention.v_proj.weight": "model.layers.2.self_attn.v_proj.weight",
  "layers.2.attention.o_proj.weight": "model.layers.2.self_attn.o_proj.weight",
  "layers.2.mlp.gate.weight": "model.layers.2.mlp.gate_proj.weight",
  "layers.2.mlp.up.weight": "model.layers.2.mlp.up_proj.weight",
  "layers.2.mlp.down.weight": "model.layers.2.mlp.down_proj.weight",
  "layers.3.attention_norm.weight": "model.layers.3.input_layernorm.weight",
  "layers.3.mlp_norm.weight": "model.layers.3.post_attention_layernorm.weight",
  "layers.3.attention.q_proj.weight": "model.layers.3.self_attn.q_proj.weight",
  "layers.3.attention.k_proj.weight": "model.layers.3.self_attn.k_proj.weight",
  "layers.3.attention.v_proj.weight": "model.layers.3.self_attn.v_proj.weight",
  "layers.3.attention.o_proj.weight": "model.layers.3.self_attn.o_proj.weight",
  "layers.3.mlp.gate.weight": "model.layers.3.mlp.gate_proj.weight",
  "layers.3.mlp.up.weight": "model.layers.3.mlp.up_proj.weight",
  "layers.3.mlp.down.weight": "model.layers.3.mlp.down_proj.weight",
  "layers.4.attention_norm.weight": "model.layers.4.input_layernorm.weight",
  "layers.4.mlp_norm.weight": "model.layers.4.post_attention_layernorm.weight",
  "layers.4.attention.q_proj.weight": "model.layers.4.self_attn.q_proj.weight",
  "layers.4.attention.k_proj.weight": "model.layers.4.self_attn.k_proj.weight",
  "layers.4.attention.v_proj.weight": "model.layers.4.self_attn.v_proj.weight",
  "layers.4.attention.o_proj.weight": "model.layers.4.self_attn.o_proj.weight",
  "layers.4.mlp.gate.weight": "model.layers.4.mlp.gate_proj.weight",
  "layers.4.mlp.up.weight": "model.layers.4.mlp.up_proj.weight",
  "layers.4.mlp.down.weight": "model.layers.4.mlp.down_proj.weight",
  "layers.5.attention_norm.weight": "model.layers.5.input_layernorm.weight",
  "layers.5.mlp_norm.weight": "model.layers.5.post_attention_layernorm.weight",
  "layers.5.attention.q_proj.weight": "model.layers.5.self_attn.q_proj.weight",
  "layers.5.attention.k_proj.weight": "model.layers.5.self_attn.k_proj.weight",
  "layers.5.attention.v_proj.weight": "model.layers.5.self_attn.v_proj.weight",
  "layers.5.attention.o_proj.weight": "model.layers.5.self_attn.o_proj.weight",
  "layers.5.mlp.gate.weight": "model.layers.5.mlp.gate_proj.weight",
  "layers.5.mlp.up.weight": "model.layers.5.mlp.up_proj.weight",
  "layers.5.mlp.down.weight": "model.layers.5.mlp.down_proj.weight",
  "layers.6.attention_norm.weight": "model.layers.6.input_layernorm.weight",
  "layers.6.mlp_norm.weight": "model.layers.6.post_attention_layernorm.weight",
  "layers.6.attention.q_proj.weight": "model.layers.6.self_attn.q_proj.weight",
  "layers.6.attention.k_proj.weight": "model.layers.6.self_attn.k_proj.weight",
  "layers.6.attention.v_proj.weight": "model.layers.6.self_attn.v_proj.weight",
  "layers.6.attention.o_proj.weight": "model.layers.6.self_attn.o_proj.weight",
  "layers.6.mlp.gate.weight": "model.layers.6.mlp.gate_proj.weight",
  "layers.6.mlp.up.weight": "model.layers.6.mlp.up_proj.weight",
  "layers.6.mlp.down.weight": "model.layers.6.mlp.down_proj.weight",
  "layers.7.attention_norm.weight": "model.layers.7.input_layernorm.weight",
  "layers.7.mlp_norm.weight": "model.layers.7.post_attention_layernorm.weight",
  "layers.7.attention.q_proj.weight": "model.layers.7.self_attn.q_proj.weight",
  "layers.7.attention.k_proj.weight": "model.layers.7.self_attn.k_proj.weight",
  "layers.7.attention.v_proj.weight": "model.layers.7.self_attn.v_proj.weight",
  "layers.7.attention.o_proj.weight": "model.layers.7.self_attn.o_proj.weight",
  "layers.7.mlp.gate.weight": "model.layers.7.mlp.gate_proj.weight",
  "layers.7.mlp.up.weight": "model.layers.7.mlp.up_proj.weight",
  "layers.7.mlp.down.weight": "model.layers.7.mlp.down_proj.weight",
  "layers.8.attention_norm.weight": "model.layers.8.input_layernorm.weight",
  "layers.8.mlp_norm.weight": "model.layers.8.post_attention_layernorm.weight",
  "layers.8.attention.q_proj.weight": "model.layers.8.self_attn.q_proj.weight",
  "layers.8.attention.k_proj.weight": "model.layers.8.self_attn.k_proj.weight",
  "layers.8.attention.v_proj.weight": "model.layers.8.self_attn.v_proj.weight",
  "layers.8.attention.o_proj.weight": "model.layers.8.self_attn.o_proj.weight",
  "layers.8.mlp.gate.weight": "model.layers.8.mlp.gate_proj.weight",
  "layers.8.mlp.up.weight": "model.layers.8.mlp.up_proj.weight",
  "layers.8.mlp.down.weight": "model.layers.8.mlp.down_proj.weight",
  "layers.9.attention_norm.weight": "model.layers.9.input_layernorm.weight",
  "layers.9.mlp_norm.weight": "model.layers.9.post_attention_layernorm.weight",
  "layers.9.attention.q_proj.weight": "model.layers.9.self_attn.q_proj.weight",
  "layers.9.attention.k_proj.weight": "model.layers.9.self_attn.k_proj.weight",
  "layers.9.attention.v_proj.weight": "model.layers.9.self_attn.v_proj.weight",
  "layers.9.attention.o_proj.weight": "model.layers.9.self_attn.o_proj.weight",
  "layers.9.mlp.gate.weight": "model.layers.9.mlp.gate_proj.weight",
  "layers.9.mlp.up.weight": "model.layers.9.mlp.up_proj.weight",
  "layers.9.mlp.down.weight": "model.layers.9.mlp.down_proj.weight",
  "layers.10.attention_norm.weight": "model.layers.10.input_layernorm.weight",
  "layers.10.mlp_norm.weight": "model.layers.10.post_attention_layernorm.weight",
  "layers.10.attention.q_proj.weight": "model.layers.10.self_attn.q_proj.weight",
  "layers.10.attention.k_proj.weight": "model.layers.10.self_attn.k_proj.weight",
  "layers.10.attention.v_proj.weight": "model.layers.10.self_attn.v_proj.weight",
  "layers.10.attention.o_proj.weight": "model.layers.10.self_attn.o_proj.weight",
  "layers.10.mlp.gate.weight": "model.layers.10.mlp.gate_proj.weight",
  "layers.10.mlp.up.weight": "model.layers.10.mlp.up_proj.weight",
  "layers.10.mlp.down.weight": "model.layers.10.mlp.down_proj.weight",
  "layers.11.attention_norm.weight": "model.layers.11.input_layernorm.weight",
  "layers.11.mlp_norm.weight": "model.layers.11.post_attention_layernorm.weight",
  "layers.11.attention.q_proj.weight": "model.layers.11.self_attn.q_proj.weight",
  "layers.11.attention.k_proj.weight": "model.layers.11.self_attn.k_proj.weight",
  "layers.11.attention.v_proj.weight": "model.layers.11.self_attn.v_proj.weight",
  "layers.11.attention.o_proj.weight": "model.layers.11.self_attn.o_proj.weight",
  "layers.11.mlp.gate.weight": "model.layers.11.mlp.gate_proj.weight",
  "layers.11.mlp.up.weight": "model.layers.11.mlp.up_proj.weight",
  "layers.11.mlp.down.weight": "model.layers.11.mlp.down_proj.weight",
  "layers.12.attention_norm.weight": "model.layers.12.input_layernorm.weight",
  "layers.12.mlp_norm.weight": "model.layers.12.post_attention_layernorm.weight",
  "layers.12.attention.q_proj.weight": "model.layers.12.self_attn.q_proj.weight",
  "layers.12.attention.k_proj.weight": "model.layers.12.self_attn.k_proj.weight",
  "layers.12.attention.v_proj.weight": "model.layers.12.self_attn.v_proj.weight",
  "layers.12.attention.o_proj.weight": "model.layers.12.self_attn.o_proj.weight",
  "layers.12.mlp.gate.weight": "model.layers.12.mlp.gate_proj.weight",
  "layers.12.mlp.up.weight": "model.layers.12.mlp.up_proj.weight",
  "layers.12.mlp.down.weight": "model.layers.12.mlp.down_proj.weight",
  "layers.13.attention_norm.weight": "model.layers.13.input_layernorm.weight",
  "layers.13.mlp_norm.weight": "model.layers.13.post_attention_layernorm.weight",
  "layers.13.attention.q_proj.weight": "model.layers.13.self_attn.q_proj.weight",
  "layers.13.attention.k_proj.weight": "model.layers.13.self_attn.k_proj.weight",
  "layers.13.attention.v_proj.weight": "model.layers.13.self_attn.v_proj.weight",
  "layers.13.attention.o_proj.weight": "model.layers.13.self_attn.o_proj.weight",
  "layers.13.mlp.gate.weight": "model.layers.13.mlp.gate_proj.weight",
  "layers.13.mlp.up.weight": "model.layers.13.mlp.up_proj.weight",
  "layers.13.mlp.down.weight": "model.layers.13.mlp.down_proj.weight",
  "layers.14.attention_norm.weight": "model.layers.14.input_layernorm.weight",
  "layers.14.mlp_norm.weight": "model.layers.14.post_attention_layernorm.weight",
  "layers.14.attention.q_proj.weight": "model.layers.14.self_attn.q_proj.weight",
  "layers.14.attention.k_proj.weight": "model.layers.14.self_attn.k_proj.weight",
  "layers.14.attention.v_proj.weight": "model.layers.14.self_attn.v_proj.weight",
  "layers.14.attention.o_proj.weight": "model.layers.14.self_attn.o_proj.weight",
  "layers.14.mlp.gate.weight": "model.layers.14.mlp.gate_proj.weight",
  "layers.14.mlp.up.weight": "model.layers.14.mlp.up_proj.weight",
  "layers.14.mlp.down.weight": "model.layers.14.mlp.down_proj.weight",
  "layers.15.attention_norm.weight": "model.layers.15.input_layernorm.weight",
  "layers.15.mlp_norm.weight": "model.layers.15.post_attention_layernorm.weight",
  "layers.15.attention.q_proj.weight": "model.layers.15.self_attn.q_proj.weight",
  "layers.15.attention.k_proj.weight": "model.layers.15.self_attn.k_proj.weight",
  "layers.15.attention.v_proj.weight": "model.layers.15.self_attn.v_proj.weight",
  "layers.15.attention.o_proj.weight": "model.layers.15.self_attn.o_proj.weight",
  "layers.15.mlp.gate.weight": "model.layers.15.mlp.gate_proj.weight",
  "layers.15.mlp.up.weight": "model.layers.15.mlp.up_proj.weight",
  "layers.15.mlp.down.weight": "model.layers.15.mlp.down_proj.weight",
  "layers.16.attention_norm.weight": "model.layers.16.input_layernorm.weight",
  "layers.16.mlp_norm.weight": "model.layers.16.post_attention_layernorm.weight",
  "layers.16.attention.q_proj.weight": "model.layers.16.self_attn.q_proj.weight",
  "layers.16.attention.k_proj.weight": "model.layers.16.self_attn.k_proj.weight",
  "layers.16.attention.v_proj.weight": "model.layers.16.self_attn.v_proj.weight",
  "layers.16.attention.o_proj.weight": "model.layers.16.self_attn.o_proj.weight",
  "layers.16.mlp.gate.weight": "model.layers.16.mlp.gate_proj.weight",
  "layers.16.mlp.up.weight": "model.layers.16.mlp.up_proj.weight",
  "layers.16.mlp.down.weight": "model.layers.16.mlp.down_proj.weight",
  "layers.17.attention_norm.weight": "model.layers.17.input_layernorm.weight",
  "layers.17.mlp_norm.weight": "model.layers.17.post_attention_layernorm.weight",
  "layers.17.attention.q_proj.weight": "model.layers.17.self_attn.q_proj.weight",
  "layers.17.attention.k_proj.weight": "model.layers.17.self_attn.k_proj.weight",
  "layers.17.attention.v_proj.weight": "model.layers.17.self_attn.v_proj.weight",
  "layers.17.attention.o_proj.weight": "model.layers.17.self_attn.o_proj.weight",
  "layers.17.mlp.gate.weight": "model.layers.17.mlp.gate_proj.weight",
  "layers.17.mlp.up.weight": "model.layers.17.mlp.up_proj.weight",
  "layers.17.mlp.down.weight": "model.layers.17.mlp.down_proj.weight",
  "layers.18.attention_norm.weight": "model.layers.18.input_layernorm.weight",
  "layers.18.mlp_norm.weight": "model.layers.18.post_attention_layernorm.weight",
  "layers.18.attention.q_proj.weight": "model.layers.18.self_attn.q_proj.weight",
  "layers.18.attention.k_proj.weight": "model.layers.18.self_attn.k_proj.weight",
  "layers.18.attention.v_proj.weight": "model.layers.18.self_attn.v_proj.weight",
  "layers.18.attention.o_proj.weight": "model.layers.18.self_attn.o_proj.weight",
  "layers.18.mlp.gate.weight": "model.layers.18.mlp.gate_proj.weight",
  "layers.18.mlp.up.weight": "model.layers.18.mlp.up_proj.weight",
  "layers.18.mlp.down.weight": "model.layers.18.mlp.down_proj.weight",
  "layers.19.attention_norm.weight": "model.layers.19.input_layernorm.weight",
  "layers.19.mlp_norm.weight": "model.layers.19.post_attention_layernorm.weight",
  "layers.19.attention.q_proj.weight": "model.layers.19.self_attn.q_proj.weight",
  "layers.19.attention.k_proj.weight": "model.layers.19.self_attn.k_proj.weight",
  "layers.19.attention.v_proj.weight": "model.layers.19.self_attn.v_proj.weight",
  "layers.19.attention.o_proj.weight": "model.layers.19.self_attn.o_proj.weight",
  "layers.19.mlp.gate.weight": "model.layers.19.mlp.gate_proj.weight",
  "layers.19.mlp.up.weight": "model.layers.19.mlp.up_proj.weight",
  "layers.19.mlp.down.weight": "model.layers.19.mlp.down_proj.weight",
  "layers.20.attention_norm.weight": "model.layers.20.input_layernorm.weight",
  "layers.20.mlp_norm.weight": "model.layers.20.post_attention_layernorm.weight",
  "layers.20.attention.q_proj.weight": "model.layers.20.self_attn.q_proj.weight",
  "layers.20.attention.k_proj.weight": "model.layers.20.self_attn.k_proj.weight",
  "layers.20.attention.v_proj.weight": "model.layers.20.self_attn.v_proj.weight",
  "layers.20.attention.o_proj.weight": "model.layers.20.self_attn.o_proj.weight",
  "layers.20.mlp.gate.weight": "model.layers.20.mlp.gate_proj.weight",
  "layers.20.mlp.up.weight": "model.layers.20.mlp.up_proj.weight",
  "layers.20.mlp.down.weight": "model.layers.20.mlp.down_proj.weight",
  "layers.21.attention_norm.weight": "model.layers.21.input_layernorm.weight",
  "layers.21.mlp_norm.weight": "model.layers.21.post_attention_layernorm.weight",
  "layers.21.attention.q_proj.weight": "model.layers.21.self_attn.q_proj.weight",
  "layers.21.attention.k_proj.weight": "model.layers.21.self_attn.k_proj.weight",
  "layers.21.attention.v_proj.weight": "model.layers.21.self_attn.v_proj.weight",
  "layers.21.attention.o_proj.weight": "model.layers.21.self_attn.o_proj.weight",
  "layers.21.mlp.gate.weight": "model.layers.21.mlp.gate_proj.weight",
  "layers.21.mlp.up.weight": "model.layers.21.mlp.up_proj.weight",
  "layers.21.mlp.down.weight": "model.layers.21.mlp.down_proj.weight",
  "layers.22.attention_norm.weight": "model.layers.22.input_layernorm.weight",
  "layers.22.mlp_norm.weight": "model.layers.22.post_attention_layernorm.weight",
  "layers.22.attention.q_proj.weight": "model.layers.22.self_attn.q_proj.weight",
  "layers.22.attention.k_proj.weight": "model.layers.22.self_attn.k_proj.weight",
  "layers.22.attention.v_proj.weight": "model.layers.22.self_attn.v_proj.weight",
  "layers.22.attention.o_proj.weight": "model.layers.22.self_attn.o_proj.weight",
  "layers.22.mlp.gate.weight": "model.layers.22.mlp.gate_proj.weight",
  "layers.22.mlp.up.weight": "model.layers.22.mlp.up_proj.weight",
  "layers.22.mlp.down.weight": "model.layers.22.mlp.down_proj.weight",
  "layers.23.attention_norm.weight": "model.layers.23.input_layernorm.weight",
  "layers.23.mlp_norm.weight": "model.layers.23.post_attention_layernorm.weight",
  "layers.23.attention.q_proj.weight": "model.layers.23.self_attn.q_proj.weight",
  "layers.23.attention.k_proj.weight": "model.layers.23.self_attn.k_proj.weight",
  "layers.23.attention.v_proj.weight": "model.layers.23.self_attn.v_proj.weight",
  "layers.23.attention.o_proj.weight": "model.layers.23.self_attn.o_proj.weight",
  "layers.23.mlp.gate.weight": "model.layers.23.mlp.gate_proj.weight",
  "layers.23.mlp.up.weight": "model.layers.23.mlp.up_proj.weight",
  "layers.23.mlp.down.weight": "model.layers.23.mlp.down_proj.weight",
  "layers.24.attention_norm.weight": "model.layers.24.input_layernorm.weight",
  "layers.24.mlp_norm.weight": "model.layers.24.post_attention_layernorm.weight",
  "layers.24.attention.q_proj.weight": "model.layers.24.self_attn.q_proj.weight",
  "layers.24.attention.k_proj.weight": "model.layers.24.self_attn.k_proj.weight",
  "layers.24.attention.v_proj.weight": "model.layers.24.self_attn.v_proj.weight",
  "layers.24.attention.o_proj.weight": "model.layers.24.self_attn.o_proj.weight",
  "layers.24.mlp.gate.weight": "model.layers.24.mlp.gate_proj.weight",
  "layers.24.mlp.up.weight": "model.layers.24.mlp.up_proj.weight",
  "layers.24.mlp.down.weight": "model.layers.24.mlp.down_proj.weight",
  "layers.25.attention_norm.weight": "model.layers.25.input_layernorm.weight",
  "layers.25.mlp_norm.weight": "model.layers.25.post_attention_layernorm.weight",
  "layers.25.attention.q_proj.weight": "model.layers.25.self_attn.q_proj.weight",
  "layers.25.attention.k_proj.weight": "model.layers.25.self_attn.k_proj.weight",
  "layers.25.attention.v_proj.weight": "model.layers.25.self_attn.v_proj.weight",
  "layers.25.attention.o_proj.weight": "model.layers.25.self_attn.o_proj.weight",
  "layers.25.mlp.gate.weight": "model.layers.25.mlp.gate_proj.weight",
  "layers.25.mlp.up.weight": "model.layers.25.mlp.up_proj.weight",
  "layers.25.mlp.down.weight": "model.layers.25.mlp.down_proj.weight",
  "layers.26.attention_norm.weight": "model.layers.26.input_layernorm.weight",
  "layers.26.mlp_norm.weight": "model.layers.26.post_attention_layernorm.weight",
  "layers.26.attention.q_proj.weight": "model.layers.26.self_attn.q_proj.weight",
  "layers.26.attention.k_proj.weight": "model.layers.26.self_attn.k_proj.weight",
  "layers.26.attention.v_proj.weight": "model.layers.26.self_attn.v_proj.weight",
  "layers.26.attention.o_proj.weight": "model.layers.26.self_attn.o_proj.weight",
  "layers.26.mlp.gate.weight": "model.layers.26.mlp.gate_proj.weight",
  "layers.26.mlp.up.weight": "model.layers.26.mlp.up_proj.weight",
  "layers.26.mlp.down.weight": "model.layers.26.mlp.down_proj.weight",
  "layers.27.attention_norm.weight": "model.layers.27.input_layernorm.weight",
  "layers.27.mlp_norm.weight": "model.layers.27.post_attention_layernorm.weight",
  "layers.27.attention.q_proj.weight": "model.layers.27.self_attn.q_proj.weight",
  "layers.27.attention.k_proj.weight": "model.layers.27.self_attn.k_proj.weight",
  "layers.27.attention.v_proj.weight": "model.layers.27.self_attn.v_proj.weight",
  "layers.27.attention.o_proj.weight": "model.layers.27.self_attn.o_proj.weight",
  "layers.27.mlp.gate.weight": "model.layers.27.mlp.gate_proj.weight",
  "layers.27.mlp.up.weight": "model.layers.27.mlp.up_proj.weight",
  "layers.27.mlp.down.weight": "model.layers.27.mlp.down_proj.weight",
  "layers.28.attention_norm.weight": "model.layers.28.input_layernorm.weight",
  "layers.28.mlp_norm.weight": "model.layers.28.post_attention_layernorm.weight",
  "layers.28.attention.q_proj.weight": "model.layers.28.self_attn.q_proj.weight",
  "layers.28.attention.k_proj.weight": "model.layers.28.self_attn.k_proj.weight",
  "layers.28.attention.v_proj.weight": "model.layers.28.self_attn.v_proj.weight",
  "layers.28.attention.o_proj.weight": "model.layers.28.self_attn.o_proj.weight",
  "layers.28.mlp.gate.weight": "model.layers.28.mlp.gate_proj.weight",
  "layers.28.mlp.up.weight": "model.layers.28.mlp.up_proj.weight",
  "layers.28.mlp.down.weight": "model.layers.28.mlp.down_proj.weight",
  "layers.29.attention_norm.weight": "model.layers.29.input_layernorm.weight",
  "layers.29.mlp_norm.weight": "model.layers.29.post_attention_layernorm.weight",
  "layers.29.attention.q_proj.weight": "model.layers.29.self_attn.q_proj.weight",
  "layers.29.attention.k_proj.weight": "model.layers.29.self_attn.k_proj.weight",
  "layers.29.attention.v_proj.weight": "model.layers.29.self_attn.v_proj.weight",
  "layers.29.attention.o_proj.weight": "model.layers.29.self_attn.o_proj.weight",
  "layers.29.mlp.gate.weight": "model.layers.29.mlp.gate_proj.weight",
  "layers.29.mlp.up.weight": "model.layers.29.mlp.up_proj.weight",
  "layers.29.mlp.down.weight": "model.layers.29.mlp.down_proj.weight",
  "layers.30.attention_norm.weight": "model.layers.30.input_layernorm.weight",
  "layers.30.mlp_norm.weight": "model.layers.30.post_attention_layernorm.weight",
  "layers.30.attention.q_proj.weight": "model.layers.30.self_attn.q_proj.weight",
  "layers.30.attention.k_proj.weight": "model.layers.30.self_attn.k_proj.weight",
  "layers.30.attention.v_proj.weight": "model.layers.30.self_attn.v_proj.weight",
  "layers.30.attention.o_proj.weight": "model.layers.30.self_attn.o_proj.weight",
  "layers.30.mlp.gate.weight": "model.layers.30.mlp.gate_proj.weight",
  "layers.30.mlp.up.weight": "model.layers.30.mlp.up_proj.weight",
  "layers.30.mlp.down.weight": "model.layers.30.mlp.down_proj.weight",
  "layers.31.attention_norm.weight": "model.layers.31.input_layernorm.weight",
  "layers.31.mlp_norm.weight": "model.layers.31.post_attention_layernorm.weight",
  "layers.31.attention.q_proj.weight": "model.layers.31.self_attn.q_proj.weight",
  "layers.31.attention.k_proj.weight": "model.layers.31.self_attn.k_proj.weight",
  "layers.31.attention.v_proj.weight": "model.layers.31.self_attn.v_proj.weight",
  "layers.31.attention.o_proj.weight": "model.layers.31.self_attn.o_proj.weight",
  "layers.31.mlp.gate.weight": "model.layers.31.mlp.gate_proj.weight",
  "layers.31.mlp.up.weight": "model.layers.31.mlp.up_proj.weight",
  "layers.31.mlp.down.weight": "model.layers.31.mlp.down_proj.weight"
}

Open on GitHub

linnet.toml60 B
toml
[package]
name = "llama"
version = "0.1.0"
language = "0.1"

Open on GitHub

nest.toml1023 B
toml
[model]
name = "mistral-7b-instruct-v0.3"
title = "Mistral 7B Instruct v0.3"
summary = "A 7B-parameter Llama-architecture decoder with grouped-query attention, a 32768-token vocabulary and context, and an untied output head, instruction-tuned."
license = "Apache-2.0"
family = "llama"
tags = ["text-generation", "decoder-only", "grouped-query-attention", "chat"]

[links]
huggingface = "https://huggingface.co/mistralai/Mistral-7B-Instruct-v0.3"
github = "https://github.com/mistralai/mistral-inference"
arxiv = "https://arxiv.org/abs/2310.06825"

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

[generics]
Vocab = 32768
H = 4096
Heads = 32
KvHeads = 8
Inner = 14336
Layers = 32
Batch = 1
MaxSeq = 32768
T = "bf16"

[check]
B = 1
S = 8

[weights]
repo = "mistralai/Mistral-7B-Instruct-v0.3"
revision = "c170c708c41dac9275d15a8fff4eca08d52bab71"
files = [
    "model-00001-of-00003.safetensors",
    "model-00002-of-00003.safetensors",
    "model-00003-of-00003.safetensors",
]
bindings = "bindings.json"

Open on GitHub

samples.json914 B
json
{
  "kind": "text",
  "prompt": "Explain in two sentences why the sky is blue.",
  "chat": true,
  "output": "The sky appears blue because molecules in the Earth's atmosphere, such as nitrogen and oxygen, scatter sunlight in all directions. However, blue light is scattered more because it has a shorter wavelength and is less energetic, which allows it to be scattered more easily. This scattering of blue light is what we perceive as the blue sky.",
  "reference": "The sky appears blue because molecules in the Earth's atmosphere, such as nitrogen and oxygen, scatter sunlight in all directions. However, blue light is scattered more because it has a shorter wavelength and is less energetic, which allows it to be scattered more easily. This scattering of blue light is what we perceive as the blue sky.",
  "reference_stack": "transformers (KV cache, argmax, bf16)",
  "tokens": 72,
  "agreeing_prefix": 72
}

Open on GitHub

src
attention.linnet20.0 kB
linnet
// Grouped-query attention: `KvHeads` key/value heads shared by `Heads`
// query heads, through `std.nn.attention::grouped_attention`. The `where`
// clause states the divisibility the reshapes depend on; the checker proves
// every shape from it.
//
// `forward` attends over a whole sequence. `decode` takes one position and
// keeps the keys and values seen so far in two `state` members, sized by
// the block's `Batch` and `MaxSeq`: the block assigns them, the runtime
// keeps them between calls, and graph exports thread them in and out.
// `prefill` is `decode` for a whole prompt: every position in one pass.
// `decode_rows` and `prefill_slots` serve many requests at once: each row of
// the caches belongs to one request, at its own length.
module llama.attention

use crate.rope::{THETA, tables}
use std.nn.attention::{
    causal_mask,
    grouped_attention,
    grouped_attention_rows,
    paged_attention,
    paged_prefill_attention,
}
use std.nn.cache::{page_slots, write_at, write_rows, write_slots, write_span, write_tokens}
use std.nn.parallel::{all_reduce}
use std.nn.linear::{Linear}
use std.nn.rope::{rope, rope_rows}

pub block GroupedQueryAttention<
    H: Dim,
    Heads: Dim,
    KvHeads: Dim,
    Batch: Dim,
    MaxSeq: Dim,
    T: Float,
    Shards: Dim = 1,
    PageSize: Dim = 64,
>
where
    Shards > 0,
    Heads % Shards == 0,
    KvHeads % Shards == 0,
    KvHeads / Shards > 0,
    (Heads / Shards) % (KvHeads / Shards) == 0,
    Heads > 0,
    KvHeads > 0,
    H % Heads == 0,
    Heads % KvHeads == 0,
    (H / Heads) % 2 == 0,
    MaxSeq > 0,
    Batch > 0,
    PageSize > 0
{
    sub q_proj: Linear<H, Heads / Shards * (H / Heads), T>
    sub k_proj: Linear<H, KvHeads / Shards * (H / Heads), T>
    sub v_proj: Linear<H, KvHeads / Shards * (H / Heads), T>
    sub o_proj: Linear<Heads / Shards * (H / Heads), H, T>

    state cache_k: Tensor[Batch, KvHeads / Shards, MaxSeq, H / Heads; T]
    state cache_v: Tensor[Batch, KvHeads / Shards, MaxSeq, H / Heads; T]

    pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T]) -> Tensor[B, S, H; T] {
        let (cos_table, sin_table) = tables<S, H / Heads, T>(THETA)
        let q = rope(
            split_heads<B, S, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_table,
            sin_table,
        )
        let k = rope(
            split_heads<B, S, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_table,
            sin_table,
        )
        let v = split_heads<B, S, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        let mixed = grouped_attention(
            q,
            k,
            v,
            rsqrt(cast<f32>(H / Heads)),
            some(causal_mask<S, S>()),
        )
        return all_reduce(o_proj.forward(merge_heads<B, S, Heads / Shards, H / Heads, T>(mixed)))
    }

    // One new position `pos`: its key and value are written into the caches
    // at that position (a masked write, no scatter), and the query attends
    // over every cached position up to and including it.
    pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32) -> Tensor[Batch, 1, H; T] {
        let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
        let cos_at[d] = cos_table[cast<i64>(pos), d]
        let sin_at[d] = sin_table[cast<i64>(pos), d]
        let cos_row = reshape(cos_at, [1, H / Heads])
        let sin_row = reshape(sin_at, [1, H / Heads])
        let q = rope(
            split_heads<Batch, 1, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_row,
            sin_row,
        )
        let k = rope(
            split_heads<Batch, 1, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_row,
            sin_row,
        )
        let v = split_heads<Batch, 1, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        cache_k = write_at(cache_k, k, pos)
        cache_v = write_at(cache_v, v, pos)
        let positions = iota<i32>(MaxSeq)
        let seen[s] = positions[s] <= pos
        let mixed = grouped_attention(
            q,
            cache_k,
            cache_v,
            rsqrt(cast<f32>(H / Heads)),
            some(reshape(seen, [1, MaxSeq])),
        )
        return all_reduce(
            o_proj.forward(merge_heads<Batch, 1, Heads / Shards, H / Heads, T>(mixed)),
        )
    }

    // `S` new positions starting at `pos`, as a prompt arrives: their keys
    // and values go into the caches as one span, and each query attends over
    // every cached position up to and including its own. A prompt costs one
    // pass instead of one per token, which is most of the time to the first
    // generated token.
    pub fn prefill<S: Dim>(x: Tensor[Batch, S, H; T], pos: i32) -> Tensor[Batch, S, H; T]
    where S > 0 {
        let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
        let rows[i] = cast<i64>(pos) + iota<i64>(S)[i]
        let cos_rows[i, d] = cos_table[rows[i], d]
        let sin_rows[i, d] = sin_table[rows[i], d]
        let q = rope(
            split_heads<Batch, S, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_rows,
            sin_rows,
        )
        let k = rope(
            split_heads<Batch, S, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_rows,
            sin_rows,
        )
        let v = split_heads<Batch, S, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        cache_k = write_span(cache_k, k, pos)
        cache_v = write_span(cache_v, v, pos)
        let positions = iota<i32>(MaxSeq)
        let queries = iota<i32>(S)
        let seen[i, s] = positions[s] <= pos + queries[i]
        let mixed = grouped_attention(q, cache_k, cache_v, rsqrt(cast<f32>(H / Heads)), some(seen))
        return all_reduce(
            o_proj.forward(merge_heads<Batch, S, Heads / Shards, H / Heads, T>(mixed)),
        )
    }

    // One new position per sequence, each at its own `positions[b]`: the
    // step a server takes for every request it is decoding together. Row `b`
    // of the caches is written at `positions[b]`, and its query attends over
    // that row up to there.
    pub fn decode_rows(
        x: Tensor[Batch, 1, H; T],
        positions: Tensor[Batch; i32],
    ) -> Tensor[Batch, 1, H; T] {
        let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
        let cos_at[b, d] = cos_table[cast<i64>(positions[b]), d]
        let sin_at[b, d] = sin_table[cast<i64>(positions[b]), d]
        let cos_rows = reshape(cos_at, [Batch, 1, H / Heads])
        let sin_rows = reshape(sin_at, [Batch, 1, H / Heads])
        let q = rope_rows(
            split_heads<Batch, 1, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_rows,
            sin_rows,
        )
        let k = rope_rows(
            split_heads<Batch, 1, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_rows,
            sin_rows,
        )
        let v = split_heads<Batch, 1, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        cache_k = write_rows(cache_k, k, positions)
        cache_v = write_rows(cache_v, v, positions)
        let slots = iota<i32>(MaxSeq)
        let seen[b, s] = slots[s] <= positions[b]
        let mixed = grouped_attention_rows(
            q,
            cache_k,
            cache_v,
            rsqrt(cast<f32>(H / Heads)),
            reshape(seen, [Batch, 1, MaxSeq]),
        )
        return all_reduce(
            o_proj.forward(merge_heads<Batch, 1, Heads / Shards, H / Heads, T>(mixed)),
        )
    }

    // `M` requests' prompts of `S` tokens, from their first positions, into
    // rows `slots` of the caches: each prompt's keys and values fill its row,
    // and its queries attend causally among themselves. The rest of a row
    // keeps what an earlier request left there; decoding never reads past its
    // own position, so nothing needs clearing.
    pub fn prefill_slots<M: Dim, S: Dim>(
        x: Tensor[M, S, H; T],
        slots: Tensor[M; i32],
    ) -> Tensor[M, S, H; T]
    where
        M > 0,
        S > 0
    {
        let (cos_table, sin_table) = tables<S, H / Heads, T>(THETA)
        let q = rope(
            split_heads<M, S, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_table,
            sin_table,
        )
        let k = rope(
            split_heads<M, S, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_table,
            sin_table,
        )
        let v = split_heads<M, S, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        cache_k = write_slots(cache_k, k, slots, 0)
        cache_v = write_slots(cache_v, v, slots, 0)
        let mixed = grouped_attention(
            q,
            k,
            v,
            rsqrt(cast<f32>(H / Heads)),
            some(causal_mask<S, S>()),
        )
        return all_reduce(o_proj.forward(merge_heads<M, S, Heads / Shards, H / Heads, T>(mixed)))
    }

    // Prompts packed end to end, `P` tokens in all: token `p` is position
    // `positions[p]` of prompt `segments[p]`, whose row of the caches is
    // `rows[p]`. Each token's key and value go to its row at its position, and
    // its query attends to the tokens of its own prompt up to itself -- no
    // padding between prompts, and none of the pass's other prompts seen.
    pub fn prefill_packed<P: Dim>(
        x: Tensor[1, P, H; T],
        rows: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
    ) -> Tensor[1, P, H; T]
    where P > 0 {
        let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
        let cos_at[p, d] = cos_table[cast<i64>(positions[p]), d]
        let sin_at[p, d] = sin_table[cast<i64>(positions[p]), d]
        let q = rope(
            split_heads<1, P, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_at,
            sin_at,
        )
        let k = rope(
            split_heads<1, P, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_at,
            sin_at,
        )
        let v = split_heads<1, P, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        cache_k = write_tokens(cache_k, k, rows, positions)
        cache_v = write_tokens(cache_v, v, rows, positions)
        let own[i, j] = segments[i] == segments[j] && positions[j] <= positions[i]
        let mixed = grouped_attention(
            q,
            k,
            v,
            rsqrt(cast<f32>(H / Heads)),
            some(own),
        )
        return all_reduce(o_proj.forward(merge_heads<1, P, Heads / Shards, H / Heads, T>(mixed)))
    }

    // `prefill_packed` for training: every packed token's output, and no
    // cache written.
    pub fn packed<P: Dim>(
        x: Tensor[1, P, H; T],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
    ) -> Tensor[1, P, H; T]
    where P > 0 {
        let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
        let cos_at[p, d] = cos_table[cast<i64>(positions[p]), d]
        let sin_at[p, d] = sin_table[cast<i64>(positions[p]), d]
        let q = rope(
            split_heads<1, P, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_at,
            sin_at,
        )
        let k = rope(
            split_heads<1, P, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_at,
            sin_at,
        )
        let v = split_heads<1, P, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        let own[i, j] = segments[i] == segments[j] && positions[j] <= positions[i]
        let mixed = grouped_attention(
            q,
            k,
            v,
            rsqrt(cast<f32>(H / Heads)),
            some(own),
        )
        return all_reduce(o_proj.forward(merge_heads<1, P, Heads / Shards, H / Heads, T>(mixed)))
    }

    // `prefill_packed`'s prompts, `P` tokens, then one step for each of the
    // `Batch` rows, at `step_positions[b]`: every token's key and value go to
    // its row at its position, the prompts attend among themselves as in
    // `prefill_packed`, and each step over its row as in `decode_rows`.
    pub fn step_packed<P: Dim>(
        x: Tensor[1, P + Batch, H; T],
        rows: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
        step_positions: Tensor[Batch; i32],
    ) -> Tensor[1, P + Batch, H; T]
    where P > 0 {
        let every_at = concat(positions, step_positions, axis = 0)
        let every_row = concat(rows, iota<i32>(Batch), axis = 0)
        let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
        let cos_at[p, d] = cos_table[cast<i64>(every_at[p]), d]
        let sin_at[p, d] = sin_table[cast<i64>(every_at[p]), d]
        let q = rope(
            split_heads<1, P + Batch, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_at,
            sin_at,
        )
        let k = rope(
            split_heads<1, P + Batch, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_at,
            sin_at,
        )
        let v = split_heads<1, P + Batch, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        cache_k = write_tokens(cache_k, k, every_row, every_at)
        cache_v = write_tokens(cache_v, v, every_row, every_at)
        let own[i, j] = segments[i] == segments[j] && positions[j] <= positions[i]
        let prompts = grouped_attention(
            q[:, :, 0:P, :],
            k[:, :, 0:P, :],
            v[:, :, 0:P, :],
            rsqrt(cast<f32>(H / Heads)),
            some(own),
        )
        let slots = iota<i32>(MaxSeq)
        let seen[b, s] = slots[s] <= step_positions[b]
        let steps = grouped_attention_rows(
            permute(q[:, :, P:P + Batch, :], [2, 1, 0, 3]),
            cache_k,
            cache_v,
            rsqrt(cast<f32>(H / Heads)),
            reshape(seen, [Batch, 1, MaxSeq]),
        )
        let mixed = concat(prompts, permute(steps, [2, 1, 0, 3]), axis = 2)
        return all_reduce(
            o_proj.forward(merge_heads<1, P + Batch, Heads / Shards, H / Heads, T>(mixed)),
        )
    }

    // Paged serving: with `Batch` 1 the caches' one row is a pool of
    // pages, `PageSize` positions each. `decode_paged` is `decode_rows` for
    // `Rows` rows whose positions lie in the pages `table` lists: each row's
    // key and value go to its position's place in the pool, and its query
    // reads its own pages.
    pub fn decode_paged<Rows: Dim, Pages: Dim>(
        x: Tensor[Rows, 1, H; T],
        positions: Tensor[Rows; i32],
        table: Tensor[Rows, Pages; i32],
    ) -> Tensor[Rows, 1, H; T]
    where Rows > 0 {
        let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
        let cos_at[b, d] = cos_table[cast<i64>(positions[b]), d]
        let sin_at[b, d] = sin_table[cast<i64>(positions[b]), d]
        let cos_rows = reshape(cos_at, [Rows, 1, H / Heads])
        let sin_rows = reshape(sin_at, [Rows, 1, H / Heads])
        let q = rope_rows(
            split_heads<Rows, 1, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_rows,
            sin_rows,
        )
        let k = rope_rows(
            split_heads<Rows, 1, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_rows,
            sin_rows,
        )
        let v = split_heads<Rows, 1, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        let slots = page_slots<Rows, Pages, PageSize>(table, positions)
        let pool = fill<i32>([Rows], 0)
        cache_k = write_tokens(cache_k, permute(k, [2, 1, 0, 3]), pool, slots)
        cache_v = write_tokens(cache_v, permute(v, [2, 1, 0, 3]), pool, slots)
        let mixed = paged_attention<Rows, Heads / Shards, KvHeads /
            Shards, MaxSeq, Pages, PageSize, H / Heads, T>(
            q,
            cache_k[0:1, :, :, :],
            cache_v[0:1, :, :, :],
            table,
            positions,
            rsqrt(cast<f32>(H / Heads)),
        )
        return all_reduce(
            o_proj.forward(merge_heads<Rows, 1, Heads / Shards, H / Heads, T>(mixed)),
        )
    }

    // Prompts packed end to end over the pool: each token, written at its
    // place (`slots`), attends over its row's pages (`table[rows[p]]`) up
    // to its position, so a prompt can pass in chunks or start after pages
    // it shares.
    pub fn prefill_paged<P: Dim, Rows: Dim, Pages: Dim>(
        x: Tensor[1, P, H; T],
        positions: Tensor[P; i32],
        rows: Tensor[P; i32],
        slots: Tensor[P; i32],
        table: Tensor[Rows, Pages; i32],
    ) -> Tensor[1, P, H; T]
    where P > 0 {
        let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
        let cos_at[p, d] = cos_table[cast<i64>(positions[p]), d]
        let sin_at[p, d] = sin_table[cast<i64>(positions[p]), d]
        let q = rope(
            split_heads<1, P, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_at,
            sin_at,
        )
        let k = rope(
            split_heads<1, P, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_at,
            sin_at,
        )
        let v = split_heads<1, P, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        let pool = fill<i32>([P], 0)
        cache_k = write_tokens(cache_k, k, pool, slots)
        cache_v = write_tokens(cache_v, v, pool, slots)
        let mixed = paged_prefill_attention<P, Heads / Shards, KvHeads /
            Shards, MaxSeq, Rows, Pages, PageSize, H / Heads, T>(
            q,
            cache_k[0:1, :, :, :],
            cache_v[0:1, :, :, :],
            table,
            rows,
            positions,
            rsqrt(cast<f32>(H / Heads)),
        )
        return all_reduce(o_proj.forward(merge_heads<1, P, Heads / Shards, H / Heads, T>(mixed)))
    }

    // `prefill_paged`'s prompts and `decode_paged`'s steps in one pass.
    pub fn step_paged<P: Dim, Rows: Dim, Pages: Dim>(
        x: Tensor[1, P + Rows, H; T],
        positions: Tensor[P; i32],
        rows: Tensor[P; i32],
        slots: Tensor[P; i32],
        step_positions: Tensor[Rows; i32],
        table: Tensor[Rows, Pages; i32],
    ) -> Tensor[1, P + Rows, H; T]
    where
        P > 0,
        Rows > 0
    {
        let every_at = concat(positions, step_positions, axis = 0)
        let every_slot = concat(
            slots,
            page_slots<Rows, Pages, PageSize>(table, step_positions),
            axis = 0,
        )
        let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
        let cos_at[p, d] = cos_table[cast<i64>(every_at[p]), d]
        let sin_at[p, d] = sin_table[cast<i64>(every_at[p]), d]
        let q = rope(
            split_heads<1, P + Rows, Heads / Shards, H / Heads, T>(q_proj.forward(x)),
            cos_at,
            sin_at,
        )
        let k = rope(
            split_heads<1, P + Rows, KvHeads / Shards, H / Heads, T>(k_proj.forward(x)),
            cos_at,
            sin_at,
        )
        let v = split_heads<1, P + Rows, KvHeads / Shards, H / Heads, T>(v_proj.forward(x))
        let pool = fill<i32>([P + Rows], 0)
        cache_k = write_tokens(cache_k, k, pool, every_slot)
        cache_v = write_tokens(cache_v, v, pool, every_slot)
        let prompts = paged_prefill_attention<P, Heads / Shards, KvHeads /
            Shards, MaxSeq, Rows, Pages, PageSize, H / Heads, T>(
            q[:, :, 0:P, :],
            cache_k[0:1, :, :, :],
            cache_v[0:1, :, :, :],
            table,
            rows,
            positions,
            rsqrt(cast<f32>(H / Heads)),
        )
        let steps = paged_attention<Rows, Heads / Shards, KvHeads /
            Shards, MaxSeq, Pages, PageSize, H / Heads, T>(
            permute(q[:, :, P:P + Rows, :], [2, 1, 0, 3]),
            cache_k[0:1, :, :, :],
            cache_v[0:1, :, :, :],
            table,
            step_positions,
            rsqrt(cast<f32>(H / Heads)),
        )
        let mixed = concat(prompts, permute(steps, [2, 1, 0, 3]), axis = 2)
        return all_reduce(
            o_proj.forward(merge_heads<1, P + Rows, Heads / Shards, H / Heads, T>(mixed)),
        )
    }
}

// `[B, S, N * D]` -> `[B, N, S, D]`.
fn split_heads<B: Dim, S: Dim, N: Dim, D: Dim, T: Float>(
    x: Tensor[B, S, N * D; T],
) -> Tensor[B, N, S, D; T] {
    return permute(reshape(x, [B, S, N, D]), [0, 2, 1, 3])
}

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

Open on GitHub

lib.linnet20.1 kB
linnet
// A Llama-style decoder: RMSNorm, grouped-query attention with rotary
// positions, SwiGLU, and an untied output head. Widths, head counts, and
// depth are generic parameters; the dtype defaults to bf16 and every
// intermediate accumulation is spelled out in the standard library. The
// `decode` entry generates one token at a time from the KV caches the
// attention blocks own, sized by `Batch` and `MaxSeq`.
module llama

use crate.attention::{GroupedQueryAttention}
use std.nn.decoding::{argmax}
use std.random::{categorical, split}
use std.nn.embedding::{Embedding}
use std.nn.parallel::{all_gather, all_reduce, shared}
use std.nn.linear::{Linear}
use std.nn.loss::{split_cross_entropy, split_token_log_probs}
use std.nn.mlp::{SwiGlu}
use std.nn.norm::{RmsNorm}

pub block DecoderLayer<
    H: Dim,
    Heads: Dim,
    KvHeads: Dim,
    Inner: Dim,
    Batch: Dim,
    MaxSeq: Dim,
    T: Float,
    Shards: Dim = 1,
    PageSize: Dim = 64,
>
where
    Shards > 0,
    Heads % Shards == 0,
    KvHeads % Shards == 0,
    KvHeads / Shards > 0,
    (Heads / Shards) % (KvHeads / Shards) == 0,
    Inner % Shards == 0,
    Heads > 0,
    KvHeads > 0,
    H % Heads == 0,
    Heads % KvHeads == 0,
    (H / Heads) % 2 == 0,
    MaxSeq > 0,
    Batch > 0,
    PageSize > 0
{
    sub attention_norm: RmsNorm<H, T>
    sub attention: GroupedQueryAttention<H, Heads, KvHeads, Batch, MaxSeq, T, Shards, PageSize>
    sub mlp_norm: RmsNorm<H, T>
    sub mlp: SwiGlu<H, Inner / Shards, T>

    pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T]) -> Tensor[B, S, H; T] {
        let attended = x + attention.forward(shared(attention_norm.forward(x)))
        return attended + all_reduce(mlp.forward(shared(mlp_norm.forward(attended))))
    }

    pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32) -> Tensor[Batch, 1, H; T] {
        let attended = x + attention.decode(attention_norm.forward(x), pos)
        return attended + all_reduce(mlp.forward(mlp_norm.forward(attended)))
    }

    pub fn prefill<S: Dim>(x: Tensor[Batch, S, H; T], pos: i32) -> Tensor[Batch, S, H; T]
    where S > 0 {
        let attended = x + attention.prefill(attention_norm.forward(x), pos)
        return attended + all_reduce(mlp.forward(mlp_norm.forward(attended)))
    }

    pub fn decode_rows(
        x: Tensor[Batch, 1, H; T],
        positions: Tensor[Batch; i32],
    ) -> Tensor[Batch, 1, H; T] {
        let attended = x + attention.decode_rows(attention_norm.forward(x), positions)
        return attended + all_reduce(mlp.forward(mlp_norm.forward(attended)))
    }

    pub fn prefill_slots<M: Dim, S: Dim>(
        x: Tensor[M, S, H; T],
        slots: Tensor[M; i32],
    ) -> Tensor[M, S, H; T]
    where
        M > 0,
        S > 0
    {
        let attended = x + attention.prefill_slots(attention_norm.forward(x), slots)
        return attended + all_reduce(mlp.forward(mlp_norm.forward(attended)))
    }

    pub fn prefill_packed<P: Dim>(
        x: Tensor[1, P, H; T],
        rows: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
    ) -> Tensor[1, P, H; T]
    where P > 0 {
        let attended =
            x + attention.prefill_packed(attention_norm.forward(x), rows, positions, segments)
        return attended + all_reduce(mlp.forward(mlp_norm.forward(attended)))
    }

    pub fn packed<P: Dim>(
        x: Tensor[1, P, H; T],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
    ) -> Tensor[1, P, H; T]
    where P > 0 {
        let attended = x + attention.packed(shared(attention_norm.forward(x)), positions, segments)
        return attended + all_reduce(mlp.forward(shared(mlp_norm.forward(attended))))
    }

    pub fn step_packed<P: Dim>(
        x: Tensor[1, P + Batch, H; T],
        rows: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
        step_positions: Tensor[Batch; i32],
    ) -> Tensor[1, P + Batch, H; T]
    where P > 0 {
        let attended =
            x +
            attention.step_packed(
                attention_norm.forward(x),
                rows,
                positions,
                segments,
                step_positions,
            )
        return attended + all_reduce(mlp.forward(mlp_norm.forward(attended)))
    }

    pub fn decode_paged<Rows: Dim, Pages: Dim>(
        x: Tensor[Rows, 1, H; T],
        positions: Tensor[Rows; i32],
        table: Tensor[Rows, Pages; i32],
    ) -> Tensor[Rows, 1, H; T]
    where Rows > 0 {
        let attended = x + attention.decode_paged(attention_norm.forward(x), positions, table)
        return attended + all_reduce(mlp.forward(mlp_norm.forward(attended)))
    }

    pub fn prefill_paged<P: Dim, Rows: Dim, Pages: Dim>(
        x: Tensor[1, P, H; T],
        positions: Tensor[P; i32],
        rows: Tensor[P; i32],
        slots: Tensor[P; i32],
        table: Tensor[Rows, Pages; i32],
    ) -> Tensor[1, P, H; T]
    where P > 0 {
        let attended =
            x + attention.prefill_paged(attention_norm.forward(x), positions, rows, slots, table)
        return attended + all_reduce(mlp.forward(mlp_norm.forward(attended)))
    }

    pub fn step_paged<P: Dim, Rows: Dim, Pages: Dim>(
        x: Tensor[1, P + Rows, H; T],
        positions: Tensor[P; i32],
        rows: Tensor[P; i32],
        slots: Tensor[P; i32],
        step_positions: Tensor[Rows; i32],
        table: Tensor[Rows, Pages; i32],
    ) -> Tensor[1, P + Rows, H; T]
    where
        P > 0,
        Rows > 0
    {
        let attended =
            x +
            attention.step_paged(
                attention_norm.forward(x),
                positions,
                rows,
                slots,
                step_positions,
                table,
            )
        return attended + all_reduce(mlp.forward(mlp_norm.forward(attended)))
    }
}

pub block Model<
    Vocab: Dim,
    H: Dim,
    Heads: Dim,
    KvHeads: Dim,
    Inner: Dim,
    Layers: Dim,
    Batch: Dim,
    MaxSeq: Dim,
    T: Float = bf16,
    Shards: Dim = 1,
    PageSize: Dim = 64,
>
where
    Shards > 0,
    Heads % Shards == 0,
    KvHeads % Shards == 0,
    KvHeads / Shards > 0,
    (Heads / Shards) % (KvHeads / Shards) == 0,
    Inner % Shards == 0,
    Vocab % Shards == 0,
    Heads > 0,
    KvHeads > 0,
    H % Heads == 0,
    Heads % KvHeads == 0,
    (H / Heads) % 2 == 0,
    MaxSeq > 0,
    Batch > 0,
    PageSize > 0
{
    sub embedding: Embedding<Vocab, H, T>
    sub layers: [DecoderLayer<H, Heads, KvHeads, Inner, Batch, MaxSeq, T, Shards, PageSize>; Layers]
    sub norm: RmsNorm<H, T>
    sub lm_head: Linear<H, Vocab / Shards, T>

    // Logits for every position.
    pub entry forward<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, Vocab; T] {
        let rows = reshape(hidden(tokens), [B * S, H])
        let gathered = all_gather<B * S, Vocab / Shards, Shards, T>(lm_head.forward(shared(rows)))
        return reshape(gathered, [B, S, Vocab])
    }

    // Logits for the last position only, as decoding needs: the final norm
    // and projection run on one row per sequence.
    pub entry next_token<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, Vocab; T]
    where S > 0 {
        let last = hidden_states(tokens)[:, S - 1, :]
        return logits(last)
    }

    // Logits for one new token per sequence at position `pos`, attending
    // over the positions decoded before it through the layers' KV caches.
    // Feed a prompt token by token, then the tokens the model produces.
    pub entry decode(token: Tensor[Batch, 1; i32], pos: i32) -> Tensor[Batch, Vocab; T] {
        return step(token, pos)
    }

    // A prompt of `S` tokens at `pos`, written into the KV caches in one
    // pass. Returns the logits after its last token, so `decode` continues
    // at `pos + S`: the time to the first token is this one call.
    pub entry prefill<S: Dim>(tokens: Tensor[Batch, S; i32], pos: i32) -> Tensor[Batch, Vocab; T]
    where S > 0 {
        var x = embedding.forward(tokens)
        static for layer in layers {
            x = layer.prefill(x, pos)
        }
        return logits(x[:, S - 1, :])
    }

    // Logits for one new token per sequence, each at its own position
    // `positions[b]`: the step a server takes for every request it is
    // decoding together (continuous batching). A row with no request in it
    // computes along and is ignored.
    pub entry decode_rows(
        tokens: Tensor[Batch, 1; i32],
        positions: Tensor[Batch; i32],
    ) -> Tensor[Batch, Vocab; T] {
        var x = embedding.forward(tokens)
        static for layer in layers {
            x = layer.decode_rows(x, positions)
        }
        return logits(x[:, 0, :])
    }

    // `M` requests' prompts into rows `slots` of the caches, as they join a
    // batch being decoded: one pass for all of them. Row `m` of `tokens` holds
    // `S` tokens of which the first `lengths[m]` are its prompt; the rest pad
    // it to one of a few compiled lengths and are never read. Returns the
    // logits after each prompt's last token; `decode_rows` continues row
    // `slots[m]` at position `lengths[m]`.
    pub entry prefill_slots<M: Dim, S: Dim>(
        tokens: Tensor[M, S; i32],
        slots: Tensor[M; i32],
        lengths: Tensor[M; i32],
    ) -> Tensor[M, Vocab; T]
    where
        M > 0,
        S > 0
    {
        var x = embedding.forward(tokens)
        static for layer in layers {
            x = layer.prefill_slots(x, slots)
        }
        let last[b, h] = x[b, cast<i64>(lengths[b]) - 1, h]
        return logits(last)
    }

    // Training over sequences packed into one row, as `prefill_packed`
    // packs prompts: `sum_p weights[p] * -log p(targets[p])` for the token
    // after each position (`std.nn.loss::linear_cross_entropy` with the
    // output head). A weight of 0 leaves a position out; `mask / count`
    // gives the mean. With the model
    // split across processes, each holds its rows of the output head.
    pub entry loss_packed<P: Dim>(
        tokens: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
        targets: Tensor[P; i64],
        weights: Tensor[P; f32],
    ) -> f32
    where P > 0 {
        let states = packed_states(tokens, positions, segments)
        return split_cross_entropy<P, H, Vocab / Shards, Shards, T>(
            shared(states),
            lm_head.weight,
            targets,
            weights,
        )
    }

    // Each packed position's log-probability of `targets[p]`, as a
    // policy-gradient update weighs them.
    pub entry log_probs_packed<P: Dim>(
        tokens: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
        targets: Tensor[P; i64],
    ) -> Tensor[P; f32]
    where P > 0 {
        let states = packed_states(tokens, positions, segments)
        return split_token_log_probs<P, H, Vocab / Shards, Shards, T>(
            shared(states),
            lm_head.weight,
            targets,
        )
    }

    // Each packed position's final hidden state, after the last norm: what
    // the output head (or a value or reward head) reads.
    pub entry hidden_packed<P: Dim>(
        tokens: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
    ) -> Tensor[P, H; T]
    where P > 0 {
        return packed_states(tokens, positions, segments)
    }

    fn packed_states<P: Dim>(
        tokens: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
    ) -> Tensor[P, H; T]
    where P > 0 {
        var x = embedding.forward(reshape(tokens, [1, P]))
        static for layer in layers {
            x = layer.packed(x, positions, segments)
        }
        return reshape(norm.forward(x), [P, H])
    }

    // Prompts of several requests packed end to end into one pass of `P`
    // tokens, with no padding between them: token `p` is position
    // `positions[p]` of prompt `segments[p]`, and its key and value go to row
    // `rows[p]` of the caches. A pack may end in padding that writes where no
    // request reads (the engine uses position `MaxSeq - 1`, which decoding
    // never reaches) with a segment of its own. Returns the logits after
    // prompt `m`'s last token, `last[m]`, for up to `Batch` prompts; rows
    // past the pass's prompts repeat whichever token `last` names.
    pub entry prefill_packed<P: Dim>(
        tokens: Tensor[P; i32],
        rows: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
        last: Tensor[Batch; i32],
    ) -> Tensor[Batch, Vocab; T]
    where P > 0 {
        var x = embedding.forward(reshape(tokens, [1, P]))
        static for layer in layers {
            x = layer.prefill_packed(x, rows, positions, segments)
        }
        let ends[m, h] = x[0, cast<i64>(last[m]), h]
        return logits(ends)
    }

    // `prefill_packed`'s prompts and a step for every row in one pass, as
    // `decode_rows` takes it (`step_tokens[b]` at `step_positions[b]`): the
    // weights are read once for both. A row the prompts fill steps along,
    // written where no request reads (the engine uses `MaxSeq - 1`). Returns
    // the logits after each prompt, as `prefill_packed`, then each row's step.
    pub entry step_packed<P: Dim>(
        tokens: Tensor[P; i32],
        rows: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
        last: Tensor[Batch; i32],
        step_tokens: Tensor[Batch, 1; i32],
        step_positions: Tensor[Batch; i32],
    ) -> Tensor[2 * Batch, Vocab; T]
    where P > 0 {
        let every = concat(tokens, reshape(step_tokens, [Batch]), axis = 0)
        var x = embedding.forward(reshape(every, [1, P + Batch]))
        static for layer in layers {
            x = layer.step_packed(x, rows, positions, segments, step_positions)
        }
        let ends[m, h] = x[0, cast<i64>(last[m]), h]
        return logits(concat(ends, x[0, P:P + Batch, :], axis = 0))
    }

    // Paged serving (`linnet.serve`): loaded with `Batch` 1, the caches'
    // one row is a pool of `MaxSeq` positions in pages of `PageSize`, and
    // each of `Rows` requests takes pages as it grows. `decode_paged` is
    // `decode_rows` for those rows, row `b`'s positions lying in the pages
    // `table[b]` lists; a row with no request lists the first page, where
    // nobody reads.
    pub entry decode_paged<Rows: Dim, Pages: Dim>(
        tokens: Tensor[Rows, 1; i32],
        positions: Tensor[Rows; i32],
        table: Tensor[Rows, Pages; i32],
    ) -> Tensor[Rows, Vocab; T]
    where Rows > 0 {
        var x = embedding.forward(tokens)
        static for layer in layers {
            x = layer.decode_paged(x, positions, table)
        }
        return logits(x[:, 0, :])
    }

    // Prompts packed end to end over the pool: token `p`, position
    // `positions[p]` of row `rows[p]`, is written at `slots[p]` and attends
    // over its row's pages (`table[rows[p]]`) up to its position, so a
    // prompt can pass in chunks or start after pages it shares. Padding
    // writes at the first page's start, which nobody reads. Returns the
    // logits after up to `Rows` prompts.
    pub entry prefill_paged<P: Dim, Rows: Dim, Pages: Dim>(
        tokens: Tensor[P; i32],
        positions: Tensor[P; i32],
        rows: Tensor[P; i32],
        slots: Tensor[P; i32],
        last: Tensor[Rows; i32],
        table: Tensor[Rows, Pages; i32],
    ) -> Tensor[Rows, Vocab; T]
    where P > 0 {
        var x = embedding.forward(reshape(tokens, [1, P]))
        static for layer in layers {
            x = layer.prefill_paged(x, positions, rows, slots, table)
        }
        let ends[m, h] = x[0, cast<i64>(last[m]), h]
        return logits(ends)
    }

    // `prefill_paged`'s prompts and a step for every row as `decode_paged`
    // takes it, in one pass.
    pub entry step_paged<P: Dim, Rows: Dim, Pages: Dim>(
        tokens: Tensor[P; i32],
        positions: Tensor[P; i32],
        rows: Tensor[P; i32],
        slots: Tensor[P; i32],
        last: Tensor[Rows; i32],
        step_tokens: Tensor[Rows, 1; i32],
        step_positions: Tensor[Rows; i32],
        table: Tensor[Rows, Pages; i32],
    ) -> Tensor[2 * Rows, Vocab; T]
    where
        P > 0,
        Rows > 0
    {
        let every = concat(tokens, reshape(step_tokens, [Rows]), axis = 0)
        var x = embedding.forward(reshape(every, [1, P + Rows]))
        static for layer in layers {
            x = layer.step_paged(x, positions, rows, slots, step_positions, table)
        }
        let ends[m, h] = x[0, cast<i64>(last[m]), h]
        return logits(concat(ends, x[0, P:P + Rows, :], axis = 0))
    }

    // Greedy generation inside the graph: `Steps` tokens after `token` at
    // `pos`, each decode step feeding the caches and its `argmax` the next
    // step. The range loop is expanded at compile time, so the whole
    // generation exports as one program.
    pub entry generate<Steps: Dim>(
        token: Tensor[Batch, 1; i32],
        pos: i32,
    ) -> Tensor[Batch, Steps; i32]
    where Steps > 0 {
        var current = token
        var produced = fill<i32>([Batch, Steps], 0)
        let slots = iota<i32>(Steps)
        static for i in 0..Steps {
            let next = argmax(step(current, pos + cast<i32>(i)))
            let written[b, s] = select(slots[s] == cast<i32>(i), next[b], produced[b, s])
            produced = written
            current = reshape(next, [Batch, 1])
        }
        return produced
    }

    // Sampling with an explicit key: `Steps` tokens drawn from the softmax
    // of the logits divided by `temperature`, one derived key per step, so
    // the same key gives the same text on every backend.
    pub entry sample<Steps: Dim>(
        token: Tensor[Batch, 1; i32],
        pos: i32,
        key: Tensor[2; i64],
        temperature: f32,
    ) -> Tensor[Batch, Steps; i32]
    where Steps > 0 {
        var current = token
        var produced = fill<i32>([Batch, Steps], 0)
        let slots = iota<i32>(Steps)
        let keys = split<Steps>(key)
        static for i in 0..Steps {
            let logits = cast<f32>(step(current, pos + cast<i32>(i))) / temperature
            let key_i[j] = keys[i, j]
            let next = categorical(key_i, logits)
            let written[b, s] = select(slots[s] == cast<i32>(i), next[b], produced[b, s])
            produced = written
            current = reshape(next, [Batch, 1])
        }
        return produced
    }

    // Greedy generation that stops early: at most `MaxNew` tokens, or fewer
    // once every sequence has produced `eos`. A runtime `while` loop over
    // scalar state; the count of tokens produced is returned with them, and
    // positions past it hold zeros.
    pub entry generate_until<MaxNew: Dim>(
        token: Tensor[Batch, 1; i32],
        pos: i32,
        eos: i32,
    ) -> (Tensor[Batch, MaxNew; i32], i32)
    where MaxNew > 0 {
        var current = token
        var produced = fill<i32>([Batch, MaxNew], 0)
        var count: i32 = 0
        var running = true
        let slots = iota<i32>(MaxNew)
        while running && count < MaxNew {
            let next = argmax(step(current, pos + count))
            let written[b, s] = select(slots[s] == count, next[b], produced[b, s])
            produced = written
            current = reshape(next, [Batch, 1])
            count = count + 1
            let finished = all[b] (next[b] == eos)
            running = !finished
        }
        return (produced, count)
    }

    fn step(token: Tensor[Batch, 1; i32], pos: i32) -> Tensor[Batch, Vocab; T] {
        var x = embedding.forward(token)
        static for layer in layers {
            x = layer.decode(x, pos)
        }
        return logits(x[:, 0, :])
    }

    // Rows' logits from their final hidden states: each shard's slice of the
    // vocabulary, the slices side by side.
    fn logits<R: Dim>(x: Tensor[R, H; T]) -> Tensor[R, Vocab; T] {
        return all_gather<R, Vocab / Shards, Shards, T>(lm_head.forward(shared(norm.forward(x))))
    }

    fn hidden<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, H; T] {
        return norm.forward(hidden_states(tokens))
    }

    fn hidden_states<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, H; T] {
        var x = embedding.forward(tokens)
        static for layer in layers {
            x = layer.forward(x)
        }
        return x
    }
}

Open on GitHub

rope.linnet890 B
linnet
// Rotary position tables computed from the compile-time sequence length:
// the positions and frequencies come from `iota`, so a model needs no
// precomputed inputs and the tables have exactly the length of the sequence.
module llama.rope

// Mistral v0.3 uses base frequency 1000000, well above Llama 2's 10000,
// consistent with its 32768-token context.
pub const THETA: f32 = 1000000.0

// `(cos, sin)` tables of shape `[S, D]`, each frequency repeated for both
// halves of the head dimension as `rope` expects.
pub fn 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)))
}

Open on GitHub

model-00001-of-00003.safetensorsHugging Face ↗

model-00002-of-00003.safetensorsHugging Face ↗

model-00003-of-00003.safetensorsHugging Face ↗