models/qwen3-8bREADME.md7.0 kB
md
# Qwen3 8B
An 8.2B-parameter decoder with the Qwen3 architecture: 36 layers of width
4096, grouped-query attention (32 query heads, 8 key/value heads, head
dimension 128 -- not `H / Heads`, which is 128 here only by coincidence;
Qwen3-4B's head dimension is the same 128 while its width is 2560), rotary
positions with base 1,000,000, RMS normalization with epsilon 1e-6, and a
SwiGLU MLP of width 12288. The query and key projections carry no bias
(`attention_bias` is false in config.json, unlike Qwen2.5); the change from
Qwen2 that does matter here is QK-norm, an RMS normalization of every
head's `head_dim`-wide slice of the query and key projections, applied
after the projection is split into heads but before rope turns it. The
output head is untied from the token embedding for this checkpoint.
The Linnet source is a self-contained Qwen3 package (`src/lib.linnet`,
`src/attention.linnet`, `src/norm.linnet`, `src/rope.linnet`), adapted from
the Qwen2 decoder `models/qwen2.5-0.5b-instruct` uses. Three changes from
that source:
- `HeadDim` is its own generic rather than `H / Heads`: Qwen3 sizes the
query, key, and value projections by `Heads * HeadDim` and
`KvHeads * HeadDim`, and the two need not agree with `H` (they don't, in
the 4B checkpoint of this family). `q_proj`/`k_proj`/`v_proj` therefore
map `H` to `Heads * HeadDim` / `KvHeads * HeadDim`, and `o_proj` maps
`Heads * HeadDim` back to `H`, not `H` to `H`.
- `q_norm` and `k_norm`, in `src/norm.linnet`'s `RmsNorm` block instantiated
at width `HeadDim`, called on the `[B, Heads, S, HeadDim]` tensor
`split_heads` produces (and the one-position slice `decode` uses), before
rope. `RmsNorm::forward` normalizes only its last axis, so this
normalizes each head's vector independently rather than the flattened
`Heads * HeadDim` projection -- putting it before the `split_heads`
reshape would normalize over all heads at once, which is a different
(and wrong) computation.
- `attention_bias` is false, so `q_proj`, `k_proj`, and `v_proj` carry no
bias -- `std.nn.linear::Linear`'s optional bias parameter already
defaults to `none`, so this needs no structural change, only that
`bindings.json` has no `*.bias` entries for them (nor did Qwen2's
`o_proj` and MLP projections).
`bindings.json` maps `lm_head.weight` to the checkpoint's own
`lm_head.weight` tensor (`tie_word_embeddings` is false in config.json).
The epsilon Qwen3 uses for `RMSNorm` (1e-6) is tighter than the stdlib
`RmsNorm` block's default (1e-5), so `src/norm.linnet` defines its own
`RmsNorm` block that calls `std.nn.norm::rms_norm` directly with the
checkpoint's epsilon; the same block, instantiated at `H` and at `HeadDim`,
serves both the per-layer norms and QK-norm.
## Loading
```python
from linnet import nest
model = nest.load("qwen3-8b", backend="torch", numerics="fast", compile="inductor")
logits = model(tokens) # Tensor[1, S; i32] -> Tensor[1, S, 151936; bf16]
```
The weights are the published, sharded `model-0000N-of-00005.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 = 40960`) |
| `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/Qwen3-8B](https://huggingface.co/Qwen/Qwen3-8B), Apache-2.0.
- Code: [QwenLM/Qwen3](https://github.com/QwenLM/Qwen3).
- Paper: [Qwen3 Technical Report](https://arxiv.org/abs/2505.09388).
The model is large enough that comparing the full checkpoint against
`transformers` in one pass on shared CPU hardware is impractical, so this
was validated on CPU with both sides truncated to the first 2 layers
(`num_hidden_layers = 2` on the `transformers` side, `generics =
{"Layers": 2}` on this one, both loaded in f32, real weights read from the
published checkpoint): on a chat-templated prompt, the Linnet forward pass
matches `transformers` logits to a maximum absolute difference of 5.2e-6
on the final position (well inside the usual 2e-3), with the top-1 next
token agreeing exactly; peak RSS for the comparison was 10.8 GiB. This
proves the architecture -- every weight this checkpoint has, read
correctly, through the right shapes -- but not the whole depth. A second
pass, comparing the full model in bf16 (top-1 agreement and the logits'
overall magnitude of difference), is deferred to a batched GPU run across
every Nest model and **has not been done here**; treat that pass as
outstanding, not as having passed.bench.json16.0 kB
json
{
"date": "2026-09-28T15:53:27+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": 23.708746768534184,
"decode_tok_s": 48.62133772028787,
"load_s": 3.8542761355638504,
"peak_vram_mib": 16774.0
},
"max_abs_diff": 0.0,
"notes": "",
"error": null,
"key": "transformers-eager"
},
{
"method": "transformers (torch.compile, static cache)",
"kind": "reference",
"metrics": {
"ttft_ms": 16.62532053887844,
"decode_tok_s": 107.35717706128929,
"load_s": 3.6620775032788515,
"peak_vram_mib": 16832.0
},
"max_abs_diff": 0.09375,
"notes": "",
"error": null,
"key": "transformers-compile"
},
{
"method": "vLLM",
"kind": "reference",
"metrics": {
"ttft_ms": 16.308249905705452,
"decode_tok_s": 146.7107114149884,
"load_s": 88.96389130875468,
"first_token": 198.0,
"peak_vram_mib": 69534.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": 19.221420399844646,
"decode_tok_s": 68.7898915856546,
"load_s": 3.060630140826106,
"peak_vram_mib": 16564.0
},
"max_abs_diff": 0.1171875,
"notes": "KV cache compiled for 768 positions",
"error": null,
"key": "linnet-torch"
},
{
"method": "Linnet torch (CUDA graphs)",
"kind": "linnet",
"metrics": {
"ttft_ms": 12.723592109978199,
"decode_tok_s": 153.6268574727338,
"load_s": 3.4564142283052206,
"peak_vram_mib": 23378.9
},
"max_abs_diff": 0.125,
"notes": "KV cache compiled for 768 positions",
"error": null,
"key": "linnet-cudagraphs"
},
{
"method": "Linnet torch (torch.compile, inductor)",
"kind": "linnet",
"metrics": {
"ttft_ms": 13.182359747588634,
"decode_tok_s": 131.61832210217798,
"load_s": 3.4364029075950384,
"peak_vram_mib": 23378.9
},
"max_abs_diff": 0.125,
"notes": "KV cache compiled for 768 positions",
"error": null,
"key": "linnet-inductor"
},
{
"method": "Linnet JAX (XLA, StableHLO)",
"kind": "linnet",
"metrics": {
"ttft_ms": 17.256347462534904,
"decode_tok_s": 158.44350948479033,
"load_s": 10.569810584187508,
"peak_vram_mib": 17205.2,
"driver_vram_mib": 33406.9
},
"max_abs_diff": 0.125,
"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": 16.228516586124897,
"decode_tok_s": 158.79566791138257,
"load_s": 10.384865283966064,
"peak_vram_mib": 17205.2,
"driver_vram_mib": 33404.9
},
"max_abs_diff": 0.09375,
"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": 29.06346507370472,
"decode_tok_s": 149.22733142466805,
"load_s": 30.462697291746736,
"peak_vram_mib": 18054.6,
"driver_vram_mib": 33408.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": 49.841469153761864,
"decode_tok_s": 46.70835463241617,
"load_s": 0.5813951920717955,
"peak_vram_mib": 37192.0
},
"max_abs_diff": 0.08458328247070312,
"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-trt",
"kind": "linnet",
"metrics": {
"peak_vram_mib": 68692.0
},
"max_abs_diff": null,
"notes": "",
"error": "Fail: [ONNXRuntimeError] : 1 : FAIL : TensorRT EP failed to create engine from network for fused node: TensorrtExecutionProvider_TRTKernel_graph_main_17650212983882034366_0_0",
"key": "linnet-onnx-f32-trt"
},
{
"method": "Linnet ONNX f16 -> ONNX Runtime (CUDA)",
"kind": "linnet",
"metrics": {
"ttft_ms": 36.007994785904884,
"decode_tok_s": 59.285335038904975,
"load_s": 0.5848892014473677,
"peak_vram_mib": 19064.0
},
"max_abs_diff": 1.3671875,
"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": 28.889787383377552,
"decode_tok_s": 46.78718922503968,
"load_s": 0.556227320805192,
"peak_vram_mib": 42760.0
},
"max_abs_diff": 11.3125,
"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": 37.05498669296503,
"decode_tok_s": 58.39534015321274,
"load_s": 0.6108948606997728,
"peak_vram_mib": 17350.0
},
"max_abs_diff": 2.5,
"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": 27.859910391271114,
"decode_tok_s": 47.32616385271398,
"load_s": 0.5565030835568905,
"peak_vram_mib": 41426.0
},
"max_abs_diff": 1.59375,
"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": 5295.3415519460095,
"serve_s": 6.188080538064241,
"load_s": 78.51854255795479,
"peak_vram_mib": 70217.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 gpu_memory_utilization 0.85",
"error": null,
"key": "serve-vllm"
},
{
"method": "transformers (generate_batch, continuous batching)",
"kind": "reference",
"metrics": {
"serve_tok_s": 815.8552334031966,
"serve_s": 40.16398824006319,
"load_s": 3.512190042063594,
"peak_vram_mib": 80992.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": 353.267722808459,
"serve_s": 92.75684667564929,
"load_s": 29.681373229250312,
"peak_vram_mib": 28195.4,
"driver_vram_mib": 33398.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": 1970.248196305766,
"serve_s": 16.631407180801034,
"load_s": 72.07779429107904,
"serve_ttft_ms": 6765.235970728099,
"peak_vram_mib": 68762.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": 5399.199116469079,
"serve_s": 6.069048259407282,
"load_s": 8.730941414833069,
"serve_ttft_ms": 2609.773548319936,
"peak_vram_mib": 29190.2
},
"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": 4485.43598625187,
"serve_s": 7.3054213905707,
"load_s": 13.346043360419571,
"serve_ttft_ms": 3270.068143494427,
"peak_vram_mib": 23622.1,
"driver_vram_mib": 33488.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": 3245.0828236796765,
"serve_s": 10.097739188931882,
"load_s": 0.2737073050811887,
"serve_ttft_ms": 4745.791389606893,
"peak_vram_mib": 38498.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": 5118.189548509118,
"serve_s": 6.40226386487484,
"load_s": 852.549588598311,
"serve_ttft_ms": 2685.1379703730345,
"peak_vram_mib": 30234.0
},
"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 torch (offloaded)",
"kind": "linnet",
"metrics": {
"ttft_ms": 191.44272059202194,
"decode_tok_s": 5.4433087391937365,
"load_s": 19.87111039273441,
"peak_vram_mib": 9174.0
},
"max_abs_diff": 0.1171875,
"notes": "KV cache compiled for 768 positions; cuda:0: embedding, layers.0-14; host, streamed in: layers.15-35, norm, lm_head",
"error": null,
"key": "linnet-offload"
},
{
"method": "Linnet torch (all GPUs)",
"kind": "linnet",
"metrics": {
"ttft_ms": 19.59872990846634,
"decode_tok_s": 68.17684474426542,
"load_s": 7.166691355407238,
"peak_vram_mib": 17082.0
},
"max_abs_diff": 0.1171875,
"notes": "KV cache compiled for 768 positions",
"error": null,
"key": "linnet-gpus"
},
{
"method": "Linnet torch (tensor parallel on 2 GPUs, shards)",
"kind": "linnet",
"metrics": {
"ttft_ms": 11.704625561833382,
"decode_tok_s": 211.75611549749615,
"load_s": 4.9715527929365635,
"peak_vram_mib": 27039.2
},
"max_abs_diff": null,
"notes": "one process per GPU under torchrun, NCCL, CUDA graphs; KV cache compiled for 768 positions",
"error": null,
"key": "linnet-tp-torch"
},
{
"method": "Linnet JAX (XLA, tensor parallel on 2 GPUs)",
"kind": "linnet",
"metrics": {
"ttft_ms": 18.369032070040703,
"decode_tok_s": 130.79896289786538,
"load_s": 16.26223200932145,
"peak_vram_mib": 19832.4,
"driver_vram_mib": 39209.8
},
"max_abs_diff": 0.125,
"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-tp-jax"
},
{
"method": "vLLM (tensor parallel on 2 GPUs)",
"kind": "reference",
"metrics": {
"ttft_ms": 14.14920762181282,
"decode_tok_s": 220.33461552536124,
"load_s": 78.78161194548011,
"first_token": 198.0,
"peak_vram_mib": 143659.8
},
"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-tp"
},
{
"method": "Linnet -> vLLM",
"kind": "linnet",
"metrics": {
"ttft_ms": 13.281981460750103,
"decode_tok_s": 153.15888238702155,
"load_s": 67.15713514667004,
"first_token": 198.0,
"peak_vram_mib": 69534.8,
"driver_vram_mib": 69534.8
},
"max_abs_diff": null,
"notes": "linnet.hf.export (17 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": 5524.171693121377,
"serve_s": 5.9317490151152015,
"load_s": 61.386786988936365,
"peak_vram_mib": 70220.8,
"driver_vram_mib": 70220.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"
},
{
"method": "Linnet -> llama.cpp (GGUF, f16)",
"kind": "linnet",
"metrics": {
"ttft_ms": 25.14982503549167,
"decode_tok_s": 165.366973,
"load_s": 60.42819887213409,
"peak_vram_mib": 15526.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"
}
],
"reference": "transformers-eager"
}bindings.json30.2 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.attention.q_proj.weight": "model.layers.0.self_attn.q_proj.weight",
"layers.0.attention.q_norm.weight": "model.layers.0.self_attn.q_norm.weight",
"layers.0.attention.k_proj.weight": "model.layers.0.self_attn.k_proj.weight",
"layers.0.attention.k_norm.weight": "model.layers.0.self_attn.k_norm.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_norm.weight": "model.layers.0.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.1.self_attn.q_proj.weight",
"layers.1.attention.q_norm.weight": "model.layers.1.self_attn.q_norm.weight",
"layers.1.attention.k_proj.weight": "model.layers.1.self_attn.k_proj.weight",
"layers.1.attention.k_norm.weight": "model.layers.1.self_attn.k_norm.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_norm.weight": "model.layers.1.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.2.self_attn.q_proj.weight",
"layers.2.attention.q_norm.weight": "model.layers.2.self_attn.q_norm.weight",
"layers.2.attention.k_proj.weight": "model.layers.2.self_attn.k_proj.weight",
"layers.2.attention.k_norm.weight": "model.layers.2.self_attn.k_norm.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_norm.weight": "model.layers.2.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.3.self_attn.q_proj.weight",
"layers.3.attention.q_norm.weight": "model.layers.3.self_attn.q_norm.weight",
"layers.3.attention.k_proj.weight": "model.layers.3.self_attn.k_proj.weight",
"layers.3.attention.k_norm.weight": "model.layers.3.self_attn.k_norm.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_norm.weight": "model.layers.3.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.4.self_attn.q_proj.weight",
"layers.4.attention.q_norm.weight": "model.layers.4.self_attn.q_norm.weight",
"layers.4.attention.k_proj.weight": "model.layers.4.self_attn.k_proj.weight",
"layers.4.attention.k_norm.weight": "model.layers.4.self_attn.k_norm.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_norm.weight": "model.layers.4.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.5.self_attn.q_proj.weight",
"layers.5.attention.q_norm.weight": "model.layers.5.self_attn.q_norm.weight",
"layers.5.attention.k_proj.weight": "model.layers.5.self_attn.k_proj.weight",
"layers.5.attention.k_norm.weight": "model.layers.5.self_attn.k_norm.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_norm.weight": "model.layers.5.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.6.self_attn.q_proj.weight",
"layers.6.attention.q_norm.weight": "model.layers.6.self_attn.q_norm.weight",
"layers.6.attention.k_proj.weight": "model.layers.6.self_attn.k_proj.weight",
"layers.6.attention.k_norm.weight": "model.layers.6.self_attn.k_norm.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_norm.weight": "model.layers.6.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.7.self_attn.q_proj.weight",
"layers.7.attention.q_norm.weight": "model.layers.7.self_attn.q_norm.weight",
"layers.7.attention.k_proj.weight": "model.layers.7.self_attn.k_proj.weight",
"layers.7.attention.k_norm.weight": "model.layers.7.self_attn.k_norm.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_norm.weight": "model.layers.7.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.8.self_attn.q_proj.weight",
"layers.8.attention.q_norm.weight": "model.layers.8.self_attn.q_norm.weight",
"layers.8.attention.k_proj.weight": "model.layers.8.self_attn.k_proj.weight",
"layers.8.attention.k_norm.weight": "model.layers.8.self_attn.k_norm.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_norm.weight": "model.layers.8.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.9.self_attn.q_proj.weight",
"layers.9.attention.q_norm.weight": "model.layers.9.self_attn.q_norm.weight",
"layers.9.attention.k_proj.weight": "model.layers.9.self_attn.k_proj.weight",
"layers.9.attention.k_norm.weight": "model.layers.9.self_attn.k_norm.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_norm.weight": "model.layers.9.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.10.self_attn.q_proj.weight",
"layers.10.attention.q_norm.weight": "model.layers.10.self_attn.q_norm.weight",
"layers.10.attention.k_proj.weight": "model.layers.10.self_attn.k_proj.weight",
"layers.10.attention.k_norm.weight": "model.layers.10.self_attn.k_norm.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_norm.weight": "model.layers.10.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.11.self_attn.q_proj.weight",
"layers.11.attention.q_norm.weight": "model.layers.11.self_attn.q_norm.weight",
"layers.11.attention.k_proj.weight": "model.layers.11.self_attn.k_proj.weight",
"layers.11.attention.k_norm.weight": "model.layers.11.self_attn.k_norm.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_norm.weight": "model.layers.11.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.12.self_attn.q_proj.weight",
"layers.12.attention.q_norm.weight": "model.layers.12.self_attn.q_norm.weight",
"layers.12.attention.k_proj.weight": "model.layers.12.self_attn.k_proj.weight",
"layers.12.attention.k_norm.weight": "model.layers.12.self_attn.k_norm.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_norm.weight": "model.layers.12.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.13.self_attn.q_proj.weight",
"layers.13.attention.q_norm.weight": "model.layers.13.self_attn.q_norm.weight",
"layers.13.attention.k_proj.weight": "model.layers.13.self_attn.k_proj.weight",
"layers.13.attention.k_norm.weight": "model.layers.13.self_attn.k_norm.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_norm.weight": "model.layers.13.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.14.self_attn.q_proj.weight",
"layers.14.attention.q_norm.weight": "model.layers.14.self_attn.q_norm.weight",
"layers.14.attention.k_proj.weight": "model.layers.14.self_attn.k_proj.weight",
"layers.14.attention.k_norm.weight": "model.layers.14.self_attn.k_norm.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_norm.weight": "model.layers.14.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.15.self_attn.q_proj.weight",
"layers.15.attention.q_norm.weight": "model.layers.15.self_attn.q_norm.weight",
"layers.15.attention.k_proj.weight": "model.layers.15.self_attn.k_proj.weight",
"layers.15.attention.k_norm.weight": "model.layers.15.self_attn.k_norm.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_norm.weight": "model.layers.15.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.16.self_attn.q_proj.weight",
"layers.16.attention.q_norm.weight": "model.layers.16.self_attn.q_norm.weight",
"layers.16.attention.k_proj.weight": "model.layers.16.self_attn.k_proj.weight",
"layers.16.attention.k_norm.weight": "model.layers.16.self_attn.k_norm.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_norm.weight": "model.layers.16.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.17.self_attn.q_proj.weight",
"layers.17.attention.q_norm.weight": "model.layers.17.self_attn.q_norm.weight",
"layers.17.attention.k_proj.weight": "model.layers.17.self_attn.k_proj.weight",
"layers.17.attention.k_norm.weight": "model.layers.17.self_attn.k_norm.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_norm.weight": "model.layers.17.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.18.self_attn.q_proj.weight",
"layers.18.attention.q_norm.weight": "model.layers.18.self_attn.q_norm.weight",
"layers.18.attention.k_proj.weight": "model.layers.18.self_attn.k_proj.weight",
"layers.18.attention.k_norm.weight": "model.layers.18.self_attn.k_norm.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_norm.weight": "model.layers.18.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.19.self_attn.q_proj.weight",
"layers.19.attention.q_norm.weight": "model.layers.19.self_attn.q_norm.weight",
"layers.19.attention.k_proj.weight": "model.layers.19.self_attn.k_proj.weight",
"layers.19.attention.k_norm.weight": "model.layers.19.self_attn.k_norm.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_norm.weight": "model.layers.19.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.20.self_attn.q_proj.weight",
"layers.20.attention.q_norm.weight": "model.layers.20.self_attn.q_norm.weight",
"layers.20.attention.k_proj.weight": "model.layers.20.self_attn.k_proj.weight",
"layers.20.attention.k_norm.weight": "model.layers.20.self_attn.k_norm.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_norm.weight": "model.layers.20.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.21.self_attn.q_proj.weight",
"layers.21.attention.q_norm.weight": "model.layers.21.self_attn.q_norm.weight",
"layers.21.attention.k_proj.weight": "model.layers.21.self_attn.k_proj.weight",
"layers.21.attention.k_norm.weight": "model.layers.21.self_attn.k_norm.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_norm.weight": "model.layers.21.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.22.self_attn.q_proj.weight",
"layers.22.attention.q_norm.weight": "model.layers.22.self_attn.q_norm.weight",
"layers.22.attention.k_proj.weight": "model.layers.22.self_attn.k_proj.weight",
"layers.22.attention.k_norm.weight": "model.layers.22.self_attn.k_norm.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_norm.weight": "model.layers.22.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.23.self_attn.q_proj.weight",
"layers.23.attention.q_norm.weight": "model.layers.23.self_attn.q_norm.weight",
"layers.23.attention.k_proj.weight": "model.layers.23.self_attn.k_proj.weight",
"layers.23.attention.k_norm.weight": "model.layers.23.self_attn.k_norm.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_norm.weight": "model.layers.23.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.24.self_attn.q_proj.weight",
"layers.24.attention.q_norm.weight": "model.layers.24.self_attn.q_norm.weight",
"layers.24.attention.k_proj.weight": "model.layers.24.self_attn.k_proj.weight",
"layers.24.attention.k_norm.weight": "model.layers.24.self_attn.k_norm.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_norm.weight": "model.layers.24.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.25.self_attn.q_proj.weight",
"layers.25.attention.q_norm.weight": "model.layers.25.self_attn.q_norm.weight",
"layers.25.attention.k_proj.weight": "model.layers.25.self_attn.k_proj.weight",
"layers.25.attention.k_norm.weight": "model.layers.25.self_attn.k_norm.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_norm.weight": "model.layers.25.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.26.self_attn.q_proj.weight",
"layers.26.attention.q_norm.weight": "model.layers.26.self_attn.q_norm.weight",
"layers.26.attention.k_proj.weight": "model.layers.26.self_attn.k_proj.weight",
"layers.26.attention.k_norm.weight": "model.layers.26.self_attn.k_norm.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_norm.weight": "model.layers.26.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.27.self_attn.q_proj.weight",
"layers.27.attention.q_norm.weight": "model.layers.27.self_attn.q_norm.weight",
"layers.27.attention.k_proj.weight": "model.layers.27.self_attn.k_proj.weight",
"layers.27.attention.k_norm.weight": "model.layers.27.self_attn.k_norm.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_norm.weight": "model.layers.27.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.28.self_attn.q_proj.weight",
"layers.28.attention.q_norm.weight": "model.layers.28.self_attn.q_norm.weight",
"layers.28.attention.k_proj.weight": "model.layers.28.self_attn.k_proj.weight",
"layers.28.attention.k_norm.weight": "model.layers.28.self_attn.k_norm.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_norm.weight": "model.layers.28.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.29.self_attn.q_proj.weight",
"layers.29.attention.q_norm.weight": "model.layers.29.self_attn.q_norm.weight",
"layers.29.attention.k_proj.weight": "model.layers.29.self_attn.k_proj.weight",
"layers.29.attention.k_norm.weight": "model.layers.29.self_attn.k_norm.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_norm.weight": "model.layers.29.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.30.self_attn.q_proj.weight",
"layers.30.attention.q_norm.weight": "model.layers.30.self_attn.q_norm.weight",
"layers.30.attention.k_proj.weight": "model.layers.30.self_attn.k_proj.weight",
"layers.30.attention.k_norm.weight": "model.layers.30.self_attn.k_norm.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_norm.weight": "model.layers.30.post_attention_layernorm.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.attention.q_proj.weight": "model.layers.31.self_attn.q_proj.weight",
"layers.31.attention.q_norm.weight": "model.layers.31.self_attn.q_norm.weight",
"layers.31.attention.k_proj.weight": "model.layers.31.self_attn.k_proj.weight",
"layers.31.attention.k_norm.weight": "model.layers.31.self_attn.k_norm.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_norm.weight": "model.layers.31.post_attention_layernorm.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",
"layers.32.attention_norm.weight": "model.layers.32.input_layernorm.weight",
"layers.32.attention.q_proj.weight": "model.layers.32.self_attn.q_proj.weight",
"layers.32.attention.q_norm.weight": "model.layers.32.self_attn.q_norm.weight",
"layers.32.attention.k_proj.weight": "model.layers.32.self_attn.k_proj.weight",
"layers.32.attention.k_norm.weight": "model.layers.32.self_attn.k_norm.weight",
"layers.32.attention.v_proj.weight": "model.layers.32.self_attn.v_proj.weight",
"layers.32.attention.o_proj.weight": "model.layers.32.self_attn.o_proj.weight",
"layers.32.mlp_norm.weight": "model.layers.32.post_attention_layernorm.weight",
"layers.32.mlp.gate.weight": "model.layers.32.mlp.gate_proj.weight",
"layers.32.mlp.up.weight": "model.layers.32.mlp.up_proj.weight",
"layers.32.mlp.down.weight": "model.layers.32.mlp.down_proj.weight",
"layers.33.attention_norm.weight": "model.layers.33.input_layernorm.weight",
"layers.33.attention.q_proj.weight": "model.layers.33.self_attn.q_proj.weight",
"layers.33.attention.q_norm.weight": "model.layers.33.self_attn.q_norm.weight",
"layers.33.attention.k_proj.weight": "model.layers.33.self_attn.k_proj.weight",
"layers.33.attention.k_norm.weight": "model.layers.33.self_attn.k_norm.weight",
"layers.33.attention.v_proj.weight": "model.layers.33.self_attn.v_proj.weight",
"layers.33.attention.o_proj.weight": "model.layers.33.self_attn.o_proj.weight",
"layers.33.mlp_norm.weight": "model.layers.33.post_attention_layernorm.weight",
"layers.33.mlp.gate.weight": "model.layers.33.mlp.gate_proj.weight",
"layers.33.mlp.up.weight": "model.layers.33.mlp.up_proj.weight",
"layers.33.mlp.down.weight": "model.layers.33.mlp.down_proj.weight",
"layers.34.attention_norm.weight": "model.layers.34.input_layernorm.weight",
"layers.34.attention.q_proj.weight": "model.layers.34.self_attn.q_proj.weight",
"layers.34.attention.q_norm.weight": "model.layers.34.self_attn.q_norm.weight",
"layers.34.attention.k_proj.weight": "model.layers.34.self_attn.k_proj.weight",
"layers.34.attention.k_norm.weight": "model.layers.34.self_attn.k_norm.weight",
"layers.34.attention.v_proj.weight": "model.layers.34.self_attn.v_proj.weight",
"layers.34.attention.o_proj.weight": "model.layers.34.self_attn.o_proj.weight",
"layers.34.mlp_norm.weight": "model.layers.34.post_attention_layernorm.weight",
"layers.34.mlp.gate.weight": "model.layers.34.mlp.gate_proj.weight",
"layers.34.mlp.up.weight": "model.layers.34.mlp.up_proj.weight",
"layers.34.mlp.down.weight": "model.layers.34.mlp.down_proj.weight",
"layers.35.attention_norm.weight": "model.layers.35.input_layernorm.weight",
"layers.35.attention.q_proj.weight": "model.layers.35.self_attn.q_proj.weight",
"layers.35.attention.q_norm.weight": "model.layers.35.self_attn.q_norm.weight",
"layers.35.attention.k_proj.weight": "model.layers.35.self_attn.k_proj.weight",
"layers.35.attention.k_norm.weight": "model.layers.35.self_attn.k_norm.weight",
"layers.35.attention.v_proj.weight": "model.layers.35.self_attn.v_proj.weight",
"layers.35.attention.o_proj.weight": "model.layers.35.self_attn.o_proj.weight",
"layers.35.mlp_norm.weight": "model.layers.35.post_attention_layernorm.weight",
"layers.35.mlp.gate.weight": "model.layers.35.mlp.gate_proj.weight",
"layers.35.mlp.up.weight": "model.layers.35.mlp.up_proj.weight",
"layers.35.mlp.down.weight": "model.layers.35.mlp.down_proj.weight"
}linnet.toml60 B
toml
[package]
name = "qwen3"
version = "0.1.0"
language = "0.1"nest.toml1.0 kB
toml
[model]
name = "qwen3-8b"
title = "Qwen3 8B"
summary = "An 8B-parameter Qwen3 decoder with grouped-query attention, per-head RMS normalization of the query and key projections, unbiased projections, and an untied output head."
license = "Apache-2.0"
family = "qwen3"
tags = ["text-generation", "decoder-only", "grouped-query-attention", "qk-norm"]
[links]
huggingface = "https://huggingface.co/Qwen/Qwen3-8B"
github = "https://github.com/QwenLM/Qwen3"
arxiv = "https://arxiv.org/abs/2505.09388"
[source]
path = "src/lib.linnet"
root = "Model"
entry = "forward"
[generics]
Vocab = 151936
H = 4096
Heads = 32
KvHeads = 8
HeadDim = 128
Inner = 12288
Layers = 36
Batch = 1
MaxSeq = 40960
T = "bf16"
[check]
B = 1
S = 8
[weights]
repo = "Qwen/Qwen3-8B"
revision = "b968826d9c46dd6066d109eabc6255188de91218"
files = [
"model-00001-of-00005.safetensors",
"model-00002-of-00005.safetensors",
"model-00003-of-00005.safetensors",
"model-00004-of-00005.safetensors",
"model-00005-of-00005.safetensors",
]
bindings = "bindings.json"samples.json1.1 kB
json
{
"kind": "text",
"prompt": "Explain in two sentences why the sky is blue.",
"chat": true,
"output": "<think>\nOkay, the user wants to know why the sky is blue in two sentences. Let me start by recalling the basic explanation. The sky appears blue because of Rayleigh scattering. Sunlight, which is white, is made up of different colors, each with different wavelengths. When sunlight enters Earth's atmosphere, the shorter wavelengths (like blue and violet) scatter more than the longer ones (like red and yellow). \n\nWait, but why blue and not violet? Oh right",
"reference": "<think>\nOkay, the user wants to know why the sky is blue in two sentences. Let me start by recalling the basic explanation. The sky appears blue because of Rayleigh scattering. Sunlight, which is white, is made up of different colors, each with different wavelengths. When sunlight enters Earth's atmosphere, the shorter wavelengths (like blue and violet) scatter more than the longer ones (like red and yellow). \n\nWait, but why blue and not violet? Oh right",
"reference_stack": "transformers (KV cache, argmax, bf16)",
"tokens": 96,
"agreeing_prefix": 96
}src
attention.linnet21.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.
//
// Two changes from Qwen2's attention, both from config.json:
//
// - `head_dim` (128 here) is its own generic, not `H / Heads` -- Qwen3
// sizes its projections by `Heads * HeadDim` and `KvHeads * HeadDim`,
// which need not equal `H` (Qwen3-4B's `H` is 2560 but `Heads * HeadDim`
// is 4096). `o_proj` therefore maps `Heads * HeadDim` back to `H`, not
// `H` to `H`.
// - `q_norm` and `k_norm` (`attention_bias` is false, so there is no
// projection bias to speak of either) apply RMS normalization to every
// head's `HeadDim`-wide vector, after the projection is split into heads
// but before rope turns it. `RmsNorm::forward` normalizes only its last
// axis, so calling it on `[B, Heads, S, HeadDim]` (or the one-position
// slice `decode` uses) normalizes each head independently, not the
// flattened `Heads * HeadDim` projection.
//
// `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 qwen3.attention
use crate.norm::{RmsNorm}
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,
HeadDim: 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,
Heads % KvHeads == 0,
HeadDim % 2 == 0,
MaxSeq > 0,
Batch > 0,
PageSize > 0
{
sub q_proj: Linear<H, Heads / Shards * HeadDim, T>
sub q_norm: RmsNorm<HeadDim, T>
sub k_proj: Linear<H, KvHeads / Shards * HeadDim, T>
sub k_norm: RmsNorm<HeadDim, T>
sub v_proj: Linear<H, KvHeads / Shards * HeadDim, T>
sub o_proj: Linear<Heads / Shards * HeadDim, H, T>
state cache_k: Tensor[Batch, KvHeads / Shards, MaxSeq, HeadDim; T]
state cache_v: Tensor[Batch, KvHeads / Shards, MaxSeq, HeadDim; 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, HeadDim, T>(THETA)
let q = rope(
q_norm.forward(split_heads<B, S, Heads / Shards, HeadDim, T>(q_proj.forward(x))),
cos_table,
sin_table,
)
let k = rope(
k_norm.forward(split_heads<B, S, KvHeads / Shards, HeadDim, T>(k_proj.forward(x))),
cos_table,
sin_table,
)
let v = split_heads<B, S, KvHeads / Shards, HeadDim, T>(v_proj.forward(x))
let mixed = grouped_attention(
q,
k,
v,
rsqrt(cast<f32>(HeadDim)),
some(causal_mask<S, S>()),
)
return all_reduce(o_proj.forward(merge_heads<B, S, Heads / Shards, HeadDim, 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, HeadDim, 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, HeadDim])
let sin_row = reshape(sin_at, [1, HeadDim])
let q = rope(
q_norm.forward(split_heads<Batch, 1, Heads / Shards, HeadDim, T>(q_proj.forward(x))),
cos_row,
sin_row,
)
let k = rope(
k_norm.forward(split_heads<Batch, 1, KvHeads / Shards, HeadDim, T>(k_proj.forward(x))),
cos_row,
sin_row,
)
let v = split_heads<Batch, 1, KvHeads / Shards, HeadDim, 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>(HeadDim)),
some(reshape(seen, [1, MaxSeq])),
)
return all_reduce(o_proj.forward(merge_heads<Batch, 1, Heads / Shards, HeadDim, 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. QK-norm applies in
// the same place `decode` applies it: after `split_heads`, before rope.
// 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, HeadDim, 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(
q_norm.forward(split_heads<Batch, S, Heads / Shards, HeadDim, T>(q_proj.forward(x))),
cos_rows,
sin_rows,
)
let k = rope(
k_norm.forward(split_heads<Batch, S, KvHeads / Shards, HeadDim, T>(k_proj.forward(x))),
cos_rows,
sin_rows,
)
let v = split_heads<Batch, S, KvHeads / Shards, HeadDim, 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>(HeadDim)), some(seen))
return all_reduce(o_proj.forward(merge_heads<Batch, S, Heads / Shards, HeadDim, 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, HeadDim, 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, HeadDim])
let sin_rows = reshape(sin_at, [Batch, 1, HeadDim])
let q = rope_rows(
q_norm.forward(split_heads<Batch, 1, Heads / Shards, HeadDim, T>(q_proj.forward(x))),
cos_rows,
sin_rows,
)
let k = rope_rows(
k_norm.forward(split_heads<Batch, 1, KvHeads / Shards, HeadDim, T>(k_proj.forward(x))),
cos_rows,
sin_rows,
)
let v = split_heads<Batch, 1, KvHeads / Shards, HeadDim, 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>(HeadDim)),
reshape(seen, [Batch, 1, MaxSeq]),
)
return all_reduce(o_proj.forward(merge_heads<Batch, 1, Heads / Shards, HeadDim, 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, HeadDim, T>(THETA)
let q = rope(
q_norm.forward(split_heads<M, S, Heads / Shards, HeadDim, T>(q_proj.forward(x))),
cos_table,
sin_table,
)
let k = rope(
k_norm.forward(split_heads<M, S, KvHeads / Shards, HeadDim, T>(k_proj.forward(x))),
cos_table,
sin_table,
)
let v = split_heads<M, S, KvHeads / Shards, HeadDim, 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>(HeadDim)),
some(causal_mask<S, S>()),
)
return all_reduce(o_proj.forward(merge_heads<M, S, Heads / Shards, HeadDim, 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, HeadDim, 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(
q_norm.forward(split_heads<1, P, Heads / Shards, HeadDim, T>(q_proj.forward(x))),
cos_at,
sin_at,
)
let k = rope(
k_norm.forward(split_heads<1, P, KvHeads / Shards, HeadDim, T>(k_proj.forward(x))),
cos_at,
sin_at,
)
let v = split_heads<1, P, KvHeads / Shards, HeadDim, 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>(HeadDim)),
some(own),
)
return all_reduce(o_proj.forward(merge_heads<1, P, Heads / Shards, HeadDim, 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, HeadDim, 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(
q_norm.forward(split_heads<1, P, Heads / Shards, HeadDim, T>(q_proj.forward(x))),
cos_at,
sin_at,
)
let k = rope(
k_norm.forward(split_heads<1, P, KvHeads / Shards, HeadDim, T>(k_proj.forward(x))),
cos_at,
sin_at,
)
let v = split_heads<1, P, KvHeads / Shards, HeadDim, 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>(HeadDim)),
some(own),
)
return all_reduce(o_proj.forward(merge_heads<1, P, Heads / Shards, HeadDim, 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, HeadDim, 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(
q_norm.forward(
split_heads<1, P + Batch, Heads / Shards, HeadDim, T>(q_proj.forward(x)),
),
cos_at,
sin_at,
)
let k = rope(
k_norm.forward(
split_heads<1, P + Batch, KvHeads / Shards, HeadDim, T>(k_proj.forward(x)),
),
cos_at,
sin_at,
)
let v = split_heads<1, P + Batch, KvHeads / Shards, HeadDim, 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>(HeadDim)),
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>(HeadDim)),
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, HeadDim, 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, HeadDim, 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, HeadDim])
let sin_rows = reshape(sin_at, [Rows, 1, HeadDim])
let q = rope_rows(
q_norm.forward(split_heads<Rows, 1, Heads / Shards, HeadDim, T>(q_proj.forward(x))),
cos_rows,
sin_rows,
)
let k = rope_rows(
k_norm.forward(split_heads<Rows, 1, KvHeads / Shards, HeadDim, T>(k_proj.forward(x))),
cos_rows,
sin_rows,
)
let v = split_heads<Rows, 1, KvHeads / Shards, HeadDim, 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, HeadDim, T>(
q,
cache_k[0:1, :, :, :],
cache_v[0:1, :, :, :],
table,
positions,
rsqrt(cast<f32>(HeadDim)),
)
return all_reduce(o_proj.forward(merge_heads<Rows, 1, Heads / Shards, HeadDim, 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, HeadDim, 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(
q_norm.forward(split_heads<1, P, Heads / Shards, HeadDim, T>(q_proj.forward(x))),
cos_at,
sin_at,
)
let k = rope(
k_norm.forward(split_heads<1, P, KvHeads / Shards, HeadDim, T>(k_proj.forward(x))),
cos_at,
sin_at,
)
let v = split_heads<1, P, KvHeads / Shards, HeadDim, 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, HeadDim, T>(
q,
cache_k[0:1, :, :, :],
cache_v[0:1, :, :, :],
table,
rows,
positions,
rsqrt(cast<f32>(HeadDim)),
)
return all_reduce(o_proj.forward(merge_heads<1, P, Heads / Shards, HeadDim, 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, HeadDim, 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(
q_norm.forward(
split_heads<1, P + Rows, Heads / Shards, HeadDim, T>(q_proj.forward(x)),
),
cos_at,
sin_at,
)
let k = rope(
k_norm.forward(
split_heads<1, P + Rows, KvHeads / Shards, HeadDim, T>(k_proj.forward(x)),
),
cos_at,
sin_at,
)
let v = split_heads<1, P + Rows, KvHeads / Shards, HeadDim, 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, HeadDim, T>(
q[:, :, 0:P, :],
cache_k[0:1, :, :, :],
cache_v[0:1, :, :, :],
table,
rows,
positions,
rsqrt(cast<f32>(HeadDim)),
)
let steps = paged_attention<Rows, Heads / Shards, KvHeads /
Shards, MaxSeq, Pages, PageSize, HeadDim, T>(
permute(q[:, :, P:P + Rows, :], [2, 1, 0, 3]),
cache_k[0:1, :, :, :],
cache_v[0:1, :, :, :],
table,
step_positions,
rsqrt(cast<f32>(HeadDim)),
)
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, HeadDim, 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.5 kB
linnet
// A Qwen3-style decoder: RMSNorm, grouped-query attention with rotary
// positions and per-head RMS normalization of the query and key
// projections (QK-norm), SwiGLU, and an output head that is untied from
// the token embedding for this checkpoint (`bindings.json` binds `lm_head`
// to its own tensor; the 4B checkpoint of the same family ties it instead).
// 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 qwen3
use crate.attention::{GroupedQueryAttention}
use crate.norm::{RmsNorm}
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}
pub block DecoderLayer<
H: Dim,
Heads: Dim,
KvHeads: Dim,
HeadDim: 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,
Heads % KvHeads == 0,
HeadDim % 2 == 0,
MaxSeq > 0,
Batch > 0,
PageSize > 0
{
sub attention_norm: RmsNorm<H, T>
sub attention: GroupedQueryAttention<H, Heads, KvHeads, HeadDim, 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,
HeadDim: 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,
Heads % KvHeads == 0,
HeadDim % 2 == 0,
MaxSeq > 0,
Batch > 0,
PageSize > 0
{
sub embedding: Embedding<Vocab, H, T>
sub layers: [DecoderLayer<H, Heads, KvHeads, HeadDim, Inner, Batch, MaxSeq, T, Shards, PageSize>; Layers]
sub norm: RmsNorm<H, T>
// This checkpoint does not tie the output head to the token embedding
// (`tie_word_embeddings` is false in config.json), so `bindings.json`
// maps this to the checkpoint's own `lm_head.weight`.
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
}
}norm.linnet680 B
linnet
// RMS normalization with Qwen3's epsilon, shared by the per-layer norms
// (width `H`) and the query/key norms (width `HeadDim`): `rms_norm`
// normalizes only its last axis, so the same block, instantiated at two
// different widths, serves both.
module qwen3.norm
use std.nn.norm::{rms_norm}
// Qwen3'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)
}
}rope.linnet909 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 qwen3.rope
// Qwen3 uses a base frequency of 1,000,000 (config.json's `rope_theta`),
// the same as Qwen2.5, for the same long-context resolution.
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)))
}model-00001-of-00005.safetensorsHugging Face ↗
model-00002-of-00005.safetensorsHugging Face ↗
model-00003-of-00005.safetensorsHugging Face ↗