models/qwen2.5-0.5b-instructREADME.md5.0 kB
md
# Qwen2.5 0.5B Instruct
A 0.5B-parameter decoder with the Qwen2 architecture, instruction-tuned. 24
layers of width 896, grouped-query attention (14 query heads, 2 key/value
heads, head dimension 64), rotary positions with base 1,000,000, RMS
normalization with epsilon 1e-6, and a SwiGLU MLP of width 4864. The query,
key, and value projections carry a bias; the output projection and the MLP
do not. The output head is tied to the token embedding.
The Linnet source is a self-contained Qwen2 package (`src/lib.linnet`,
`src/attention.linnet`, `src/rope.linnet`), adapted from the Llama-style
decoder the other Nest models use. `std.nn.linear::Linear` already carries
an optional bias, so the query/key/value bias needs no structural change --
only `bindings.json` entries for `q_proj.bias`, `k_proj.bias`, and
`v_proj.bias`, left unbound (and so `none`) for `o_proj` and the MLP
projections, which the checkpoint has none of. `bindings.json` maps both
`embedding.weight` and `lm_head.weight` to the checkpoint's single
`model.embed_tokens.weight` (`tie_word_embeddings` in config.json), which is
how tied embeddings are expressed without a language feature for them. The
epsilon Qwen2.5 uses for `RMSNorm` (1e-6) is tighter than the stdlib
`RmsNorm` block's default (1e-5), so this model defines its own `RmsNorm`
block that calls `std.nn.norm::rms_norm` directly with the checkpoint's
epsilon.
## Loading
```python
from linnet import nest
model = nest.load("qwen2.5-0.5b-instruct", backend="torch", numerics="fast", compile="inductor")
logits = model(tokens) # Tensor[1, S; i32] -> Tensor[1, S, 151936; bf16]
```
The weights are the published `model.safetensors` (bf16); `bindings.json`
maps Linnet parameter paths to its 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: [Qwen/Qwen2.5-0.5B-Instruct](https://huggingface.co/Qwen/Qwen2.5-0.5B-Instruct), Apache-2.0.
- Code: [QwenLM/Qwen2.5](https://github.com/QwenLM/Qwen2.5).
- Paper: [Qwen2.5 Technical Report](https://arxiv.org/abs/2412.15115).
The Linnet forward pass matches `transformers` logits on the final position
to within 2e-3 absolute in f32, with the top-1 next token agreeing, on a
prompt run through the model's chat template. The `decode` entry's cached
logits match a full forward pass over the same extended sequence to the same
tolerance.bench.json13.8 kB
json
{
"date": "2026-09-28T17:31:05+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": 12.397044338285923,
"decode_tok_s": 93.19387496695063,
"load_s": 1.961898036301136,
"peak_vram_mib": 1910.0
},
"max_abs_diff": 0.0,
"notes": "",
"error": null,
"key": "transformers-eager"
},
{
"method": "transformers (torch.compile, static cache)",
"kind": "reference",
"metrics": {
"ttft_ms": 7.031051442027092,
"decode_tok_s": 188.62422217208638,
"load_s": 1.546073006466031,
"peak_vram_mib": 1986.0
},
"max_abs_diff": 0.19140625,
"notes": "",
"error": null,
"key": "transformers-compile"
},
{
"method": "vLLM",
"kind": "reference",
"metrics": {
"ttft_ms": 7.576551288366318,
"decode_tok_s": 707.2256795467996,
"load_s": 57.82546177133918,
"first_token": 13.0,
"peak_vram_mib": 70740.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": 7.758256047964096,
"decode_tok_s": 107.5998207491478,
"load_s": 0.833325732499361,
"peak_vram_mib": 2010.0
},
"max_abs_diff": 0.1875,
"notes": "KV cache compiled for 768 positions",
"error": null,
"key": "linnet-torch"
},
{
"method": "Linnet torch (CUDA graphs)",
"kind": "linnet",
"metrics": {
"ttft_ms": 1.8201842904090881,
"decode_tok_s": 907.8411096446137,
"load_s": 0.6808824874460697,
"peak_vram_mib": 2364.9
},
"max_abs_diff": 0.140625,
"notes": "KV cache compiled for 768 positions",
"error": null,
"key": "linnet-cudagraphs"
},
{
"method": "Linnet torch (torch.compile, inductor)",
"kind": "linnet",
"metrics": {
"ttft_ms": 3.136279061436653,
"decode_tok_s": 322.1860061734868,
"load_s": 0.7987894974648952,
"peak_vram_mib": 2648.4
},
"max_abs_diff": 0.140625,
"notes": "KV cache compiled for 768 positions",
"error": null,
"key": "linnet-inductor"
},
{
"method": "Linnet JAX (XLA, StableHLO)",
"kind": "linnet",
"metrics": {
"ttft_ms": 3.2251980155706406,
"decode_tok_s": 1046.6275446982527,
"load_s": 0.7698775976896286,
"peak_vram_mib": 1572.0,
"driver_vram_mib": 2680.9
},
"max_abs_diff": 0.3125,
"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": 3.0059758573770523,
"decode_tok_s": 1055.1375392106033,
"load_s": 0.7733310349285603,
"peak_vram_mib": 1572.0,
"driver_vram_mib": 2680.9
},
"max_abs_diff": 0.2734375,
"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": 10.043281130492687,
"decode_tok_s": 716.5388961263709,
"load_s": 7.217734096571803,
"peak_vram_mib": 2085.9,
"driver_vram_mib": 4728.0
},
"max_abs_diff": 0.15625,
"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": 11.057344265282154,
"decode_tok_s": 153.08553185512278,
"load_s": 0.3163565918803215,
"peak_vram_mib": 4120.0
},
"max_abs_diff": 0.14293169975280762,
"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": 7.832708768546581,
"decode_tok_s": 209.89735529311824,
"load_s": 0.2909748684614897,
"peak_vram_mib": 8448.0
},
"max_abs_diff": 0.13897037506103516,
"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": 9.940310381352901,
"decode_tok_s": 137.9356529200322,
"load_s": 0.3352184630930424,
"peak_vram_mib": 2998.0
},
"max_abs_diff": 0.142578125,
"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": 5.738941952586174,
"decode_tok_s": 215.26814482139164,
"load_s": 0.2910300400108099,
"peak_vram_mib": 5886.0
},
"max_abs_diff": 5.984375,
"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": 10.005003772675991,
"decode_tok_s": 145.55794584249023,
"load_s": 0.30839347280561924,
"peak_vram_mib": 2286.0
},
"max_abs_diff": 0.38671875,
"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": 5.584213882684708,
"decode_tok_s": 218.28195362796018,
"load_s": 0.29175250232219696,
"peak_vram_mib": 5168.0
},
"max_abs_diff": 0.25,
"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": 18828.150272560866,
"serve_s": 1.7403727676719427,
"load_s": 18.310222055763006,
"peak_vram_mib": 69678.8
},
"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": 1864.6575697951844,
"serve_s": 17.573199782520533,
"load_s": 1.3440527226775885,
"peak_vram_mib": 78632.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": 2482.352001166264,
"serve_s": 13.200384145602584,
"load_s": 6.5146275125443935,
"peak_vram_mib": 2750.9,
"driver_vram_mib": 4722.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": 2157.4837481956456,
"serve_s": 15.188063422217965,
"load_s": 52.05952265486121,
"serve_ttft_ms": 7317.317272536457,
"peak_vram_mib": 69888.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": 34643.12222022949,
"serve_s": 0.94587317481637,
"load_s": 3.9006778970360756,
"serve_ttft_ms": 403.3740255981684,
"peak_vram_mib": 3064.9
},
"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": 29017.85056115599,
"serve_s": 1.1292359484359622,
"load_s": 2.7512129386886954,
"serve_ttft_ms": 570.5343200825155,
"peak_vram_mib": 2149.1,
"driver_vram_mib": 4792.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": 13227.26884516834,
"serve_s": 2.477306568995118,
"load_s": 0.25077642500400543,
"serve_ttft_ms": 1105.1781452260911,
"peak_vram_mib": 10372.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": 24210.569277050294,
"serve_s": 1.3534584678709507,
"load_s": 621.3704823330045,
"serve_ttft_ms": 558.7800936773419,
"peak_vram_mib": 4092.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 -> llama.cpp (GGUF, f16)",
"kind": "linnet",
"metrics": {
"ttft_ms": 5.949413322821142,
"decode_tok_s": 766.983501,
"load_s": 11.983776700682938,
"peak_vram_mib": 1950.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": "Linnet -> vLLM",
"kind": "linnet",
"metrics": {
"ttft_ms": 12.188578955829144,
"decode_tok_s": 510.20422256728494,
"load_s": 39.15125427953899,
"first_token": 13.0,
"peak_vram_mib": 70740.8,
"driver_vram_mib": 70740.8
},
"max_abs_diff": null,
"notes": "linnet.hf.export (2 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": 26247.5853553564,
"serve_s": 1.2484195996075869,
"load_s": 29.64677148591727,
"peak_vram_mib": 69350.8,
"driver_vram_mib": 69350.8
},
"max_abs_diff": null,
"notes": "linnet.hf.export (0 s), then vLLM (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 gpu_memory_utilization 0.85",
"error": null,
"key": "serve-linnet-vllm"
}
],
"reference": "transformers-eager"
}bindings.json21.7 kB
json
{
"embedding.weight": "model.embed_tokens.weight",
"norm.weight": "model.norm.weight",
"lm_head.weight": "model.embed_tokens.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.q_proj.bias": "model.layers.0.self_attn.q_proj.bias",
"layers.0.attention.k_proj.weight": "model.layers.0.self_attn.k_proj.weight",
"layers.0.attention.k_proj.bias": "model.layers.0.self_attn.k_proj.bias",
"layers.0.attention.v_proj.weight": "model.layers.0.self_attn.v_proj.weight",
"layers.0.attention.v_proj.bias": "model.layers.0.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.1.self_attn.q_proj.bias",
"layers.1.attention.k_proj.weight": "model.layers.1.self_attn.k_proj.weight",
"layers.1.attention.k_proj.bias": "model.layers.1.self_attn.k_proj.bias",
"layers.1.attention.v_proj.weight": "model.layers.1.self_attn.v_proj.weight",
"layers.1.attention.v_proj.bias": "model.layers.1.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.2.self_attn.q_proj.bias",
"layers.2.attention.k_proj.weight": "model.layers.2.self_attn.k_proj.weight",
"layers.2.attention.k_proj.bias": "model.layers.2.self_attn.k_proj.bias",
"layers.2.attention.v_proj.weight": "model.layers.2.self_attn.v_proj.weight",
"layers.2.attention.v_proj.bias": "model.layers.2.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.3.self_attn.q_proj.bias",
"layers.3.attention.k_proj.weight": "model.layers.3.self_attn.k_proj.weight",
"layers.3.attention.k_proj.bias": "model.layers.3.self_attn.k_proj.bias",
"layers.3.attention.v_proj.weight": "model.layers.3.self_attn.v_proj.weight",
"layers.3.attention.v_proj.bias": "model.layers.3.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.4.self_attn.q_proj.bias",
"layers.4.attention.k_proj.weight": "model.layers.4.self_attn.k_proj.weight",
"layers.4.attention.k_proj.bias": "model.layers.4.self_attn.k_proj.bias",
"layers.4.attention.v_proj.weight": "model.layers.4.self_attn.v_proj.weight",
"layers.4.attention.v_proj.bias": "model.layers.4.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.5.self_attn.q_proj.bias",
"layers.5.attention.k_proj.weight": "model.layers.5.self_attn.k_proj.weight",
"layers.5.attention.k_proj.bias": "model.layers.5.self_attn.k_proj.bias",
"layers.5.attention.v_proj.weight": "model.layers.5.self_attn.v_proj.weight",
"layers.5.attention.v_proj.bias": "model.layers.5.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.6.self_attn.q_proj.bias",
"layers.6.attention.k_proj.weight": "model.layers.6.self_attn.k_proj.weight",
"layers.6.attention.k_proj.bias": "model.layers.6.self_attn.k_proj.bias",
"layers.6.attention.v_proj.weight": "model.layers.6.self_attn.v_proj.weight",
"layers.6.attention.v_proj.bias": "model.layers.6.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.7.self_attn.q_proj.bias",
"layers.7.attention.k_proj.weight": "model.layers.7.self_attn.k_proj.weight",
"layers.7.attention.k_proj.bias": "model.layers.7.self_attn.k_proj.bias",
"layers.7.attention.v_proj.weight": "model.layers.7.self_attn.v_proj.weight",
"layers.7.attention.v_proj.bias": "model.layers.7.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.8.self_attn.q_proj.bias",
"layers.8.attention.k_proj.weight": "model.layers.8.self_attn.k_proj.weight",
"layers.8.attention.k_proj.bias": "model.layers.8.self_attn.k_proj.bias",
"layers.8.attention.v_proj.weight": "model.layers.8.self_attn.v_proj.weight",
"layers.8.attention.v_proj.bias": "model.layers.8.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.9.self_attn.q_proj.bias",
"layers.9.attention.k_proj.weight": "model.layers.9.self_attn.k_proj.weight",
"layers.9.attention.k_proj.bias": "model.layers.9.self_attn.k_proj.bias",
"layers.9.attention.v_proj.weight": "model.layers.9.self_attn.v_proj.weight",
"layers.9.attention.v_proj.bias": "model.layers.9.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.10.self_attn.q_proj.bias",
"layers.10.attention.k_proj.weight": "model.layers.10.self_attn.k_proj.weight",
"layers.10.attention.k_proj.bias": "model.layers.10.self_attn.k_proj.bias",
"layers.10.attention.v_proj.weight": "model.layers.10.self_attn.v_proj.weight",
"layers.10.attention.v_proj.bias": "model.layers.10.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.11.self_attn.q_proj.bias",
"layers.11.attention.k_proj.weight": "model.layers.11.self_attn.k_proj.weight",
"layers.11.attention.k_proj.bias": "model.layers.11.self_attn.k_proj.bias",
"layers.11.attention.v_proj.weight": "model.layers.11.self_attn.v_proj.weight",
"layers.11.attention.v_proj.bias": "model.layers.11.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.12.self_attn.q_proj.bias",
"layers.12.attention.k_proj.weight": "model.layers.12.self_attn.k_proj.weight",
"layers.12.attention.k_proj.bias": "model.layers.12.self_attn.k_proj.bias",
"layers.12.attention.v_proj.weight": "model.layers.12.self_attn.v_proj.weight",
"layers.12.attention.v_proj.bias": "model.layers.12.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.13.self_attn.q_proj.bias",
"layers.13.attention.k_proj.weight": "model.layers.13.self_attn.k_proj.weight",
"layers.13.attention.k_proj.bias": "model.layers.13.self_attn.k_proj.bias",
"layers.13.attention.v_proj.weight": "model.layers.13.self_attn.v_proj.weight",
"layers.13.attention.v_proj.bias": "model.layers.13.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.14.self_attn.q_proj.bias",
"layers.14.attention.k_proj.weight": "model.layers.14.self_attn.k_proj.weight",
"layers.14.attention.k_proj.bias": "model.layers.14.self_attn.k_proj.bias",
"layers.14.attention.v_proj.weight": "model.layers.14.self_attn.v_proj.weight",
"layers.14.attention.v_proj.bias": "model.layers.14.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.15.self_attn.q_proj.bias",
"layers.15.attention.k_proj.weight": "model.layers.15.self_attn.k_proj.weight",
"layers.15.attention.k_proj.bias": "model.layers.15.self_attn.k_proj.bias",
"layers.15.attention.v_proj.weight": "model.layers.15.self_attn.v_proj.weight",
"layers.15.attention.v_proj.bias": "model.layers.15.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.16.self_attn.q_proj.bias",
"layers.16.attention.k_proj.weight": "model.layers.16.self_attn.k_proj.weight",
"layers.16.attention.k_proj.bias": "model.layers.16.self_attn.k_proj.bias",
"layers.16.attention.v_proj.weight": "model.layers.16.self_attn.v_proj.weight",
"layers.16.attention.v_proj.bias": "model.layers.16.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.17.self_attn.q_proj.bias",
"layers.17.attention.k_proj.weight": "model.layers.17.self_attn.k_proj.weight",
"layers.17.attention.k_proj.bias": "model.layers.17.self_attn.k_proj.bias",
"layers.17.attention.v_proj.weight": "model.layers.17.self_attn.v_proj.weight",
"layers.17.attention.v_proj.bias": "model.layers.17.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.18.self_attn.q_proj.bias",
"layers.18.attention.k_proj.weight": "model.layers.18.self_attn.k_proj.weight",
"layers.18.attention.k_proj.bias": "model.layers.18.self_attn.k_proj.bias",
"layers.18.attention.v_proj.weight": "model.layers.18.self_attn.v_proj.weight",
"layers.18.attention.v_proj.bias": "model.layers.18.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.19.self_attn.q_proj.bias",
"layers.19.attention.k_proj.weight": "model.layers.19.self_attn.k_proj.weight",
"layers.19.attention.k_proj.bias": "model.layers.19.self_attn.k_proj.bias",
"layers.19.attention.v_proj.weight": "model.layers.19.self_attn.v_proj.weight",
"layers.19.attention.v_proj.bias": "model.layers.19.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.20.self_attn.q_proj.bias",
"layers.20.attention.k_proj.weight": "model.layers.20.self_attn.k_proj.weight",
"layers.20.attention.k_proj.bias": "model.layers.20.self_attn.k_proj.bias",
"layers.20.attention.v_proj.weight": "model.layers.20.self_attn.v_proj.weight",
"layers.20.attention.v_proj.bias": "model.layers.20.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.21.self_attn.q_proj.bias",
"layers.21.attention.k_proj.weight": "model.layers.21.self_attn.k_proj.weight",
"layers.21.attention.k_proj.bias": "model.layers.21.self_attn.k_proj.bias",
"layers.21.attention.v_proj.weight": "model.layers.21.self_attn.v_proj.weight",
"layers.21.attention.v_proj.bias": "model.layers.21.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.22.self_attn.q_proj.bias",
"layers.22.attention.k_proj.weight": "model.layers.22.self_attn.k_proj.weight",
"layers.22.attention.k_proj.bias": "model.layers.22.self_attn.k_proj.bias",
"layers.22.attention.v_proj.weight": "model.layers.22.self_attn.v_proj.weight",
"layers.22.attention.v_proj.bias": "model.layers.22.self_attn.v_proj.bias",
"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.q_proj.bias": "model.layers.23.self_attn.q_proj.bias",
"layers.23.attention.k_proj.weight": "model.layers.23.self_attn.k_proj.weight",
"layers.23.attention.k_proj.bias": "model.layers.23.self_attn.k_proj.bias",
"layers.23.attention.v_proj.weight": "model.layers.23.self_attn.v_proj.weight",
"layers.23.attention.v_proj.bias": "model.layers.23.self_attn.v_proj.bias",
"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"
}linnet.toml60 B
toml
[package]
name = "qwen2"
version = "0.1.0"
language = "0.1"nest.toml893 B
toml
[model]
name = "qwen2.5-0.5b-instruct"
title = "Qwen2.5 0.5B Instruct"
summary = "A 0.5B-parameter Qwen2 decoder with grouped-query attention, biased query/key/value projections, and an output head tied to the token embedding, instruction-tuned."
license = "Apache-2.0"
family = "qwen2"
tags = ["text-generation", "decoder-only", "grouped-query-attention", "chat"]
[links]
huggingface = "https://huggingface.co/Qwen/Qwen2.5-0.5B-Instruct"
github = "https://github.com/QwenLM/Qwen2.5"
arxiv = "https://arxiv.org/abs/2412.15115"
[source]
path = "src/lib.linnet"
root = "Model"
entry = "forward"
[generics]
Vocab = 151936
H = 896
Heads = 14
KvHeads = 2
Inner = 4864
Layers = 24
Batch = 1
MaxSeq = 32768
T = "bf16"
[check]
B = 1
S = 8
[weights]
repo = "Qwen/Qwen2.5-0.5B-Instruct"
revision = "7ae557604adf67be50417f59c2c2f167def9a775"
files = ["model.safetensors"]
bindings = "bindings.json"samples.json802 B
json
{
"kind": "text",
"prompt": "Explain in two sentences why the sky is blue.",
"chat": true,
"output": "The sky is blue because of the scattering of sunlight by tiny particles in the Earth's atmosphere, particularly water vapor and nitrogen oxides. This scattering causes the blue light from the sun to be scattered in all directions, creating the characteristic blue color we see in the sky.",
"reference": "The sky is blue because it is a reflection of the sunlight that falls on Earth. The Earth's atmosphere scatters and disperses the sunlight, causing it to be scattered and refracted by the various particles in the atmosphere, resulting in the characteristic blue color of the sky.",
"reference_stack": "transformers (KV cache, argmax, bf16)",
"tokens": 55,
"agreeing_prefix": 5
}src
attention.linnet20.2 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.
//
// Qwen2 puts a bias on the query, key, and value projections (and only
// those -- the output projection has none). `std.nn.linear::Linear` already
// carries an optional bias parameter that defaults to `none`, so the same
// block serves both cases: `bindings.json` binds `q_proj.bias`, `k_proj.bias`,
// and `v_proj.bias` to the checkpoint's tensors and leaves `o_proj.bias`
// unbound.
//
// `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.
module qwen2.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])
}lib.linnet20.8 kB
linnet
// A Qwen2-style decoder: RMSNorm, grouped-query attention with biased query,
// key, and value projections and rotary positions, SwiGLU, and an output
// head tied to the token embedding. 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 qwen2
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::{rms_norm}
// Qwen2.5-0.5B's `rms_norm_eps` (1e-6) is tighter than the stdlib
// `RmsNorm` block's default (1e-5), so this calls the op directly with the
// checkpoint's epsilon instead of reusing that block.
pub block RmsNorm<H: Dim, T: Float> {
param weight: Tensor[H; T]
pub fn forward<*S: Shape>(x: Tensor[*S, H; T]) -> Tensor[*S, H; T] {
return rms_norm(x, weight, 1e-6)
}
}
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>
// Qwen2.5-0.5B ties the output head to the token embedding
// (`tie_word_embeddings` in config.json); the checkpoint has no separate
// `lm_head` tensor, so `bindings.json` points this at the same
// `model.embed_tokens.weight` the embedding uses.
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
}
}rope.linnet925 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 qwen2.rope
// Qwen2.5 uses a base frequency of 1,000,000 (config.json's `rope_theta`),
// much higher than Llama 2's 10000 so the model resolves its long 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)))
}