models/qwen3-8b-fp8README.md1.6 kB
md
# Qwen3 8B FP8
Qwen3 8B with its decoder projections in FP8: Qwen's own
[FP8 checkpoint](https://huggingface.co/Qwen/Qwen3-8B-FP8). Each query,
key, value, output, gate, up and down projection holds E4M3 weights (one
byte each) with one bf16 scale per block of 128 by 128 weights, the scheme
DeepSeek-V3 uses; the embedding, the norms and the output head stay bf16.
The source is the `qwen3-8b` package with those projections as
`std.quant::Fp8BlockLinear`. Its blocks need every projection's widths to
be multiples of 128, which the `where` clauses state. `bindings.json` binds
each `scale` to the checkpoint's `weight_scale_inv` (a multiplier, despite
its name) and each FP8 `weight` to its bytes.
## Loading
```python
from linnet import nest
model = nest.load("qwen3-8b-fp8", backend="torch", numerics="fast")
```
On a GPU with compute capability 9 or later, the projections multiply the
FP8 weights as they are: one row in a kernel that reads each byte once, up
to 32 rows on tensor cores, and longer inputs rounded to FP8 128 inputs at
a time (`torch._scaled_mm` with block scales). Elsewhere the weights are
decoded to bf16 for each product.
## Against bf16
One H100, the same host; CUDA graphs for decoding, a 512-token prompt run
eagerly:
| | Perplexity, WikiText-2 | Time to first token | Decoding | Peak memory |
| --- | ---: | ---: | ---: | ---: |
| bf16 (`qwen3-8b`) | 12.71 | 21.6 ms | 157 tokens/s | 16.9 GiB |
| FP8 | 12.84 | 23.2 ms | 238 tokens/s | 10.4 GiB |
Under `numerics="equivalent"`, which never rounds the input, the perplexity
is 12.77. The prompt runs slower than bf16's: cuBLAS's block-scaled FP8
product is no faster than bf16 at these sizes.bindings.json50.5 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",
"layers.0.attention.q_proj.scale": "model.layers.0.self_attn.q_proj.weight_scale_inv",
"layers.0.attention.k_proj.scale": "model.layers.0.self_attn.k_proj.weight_scale_inv",
"layers.0.attention.v_proj.scale": "model.layers.0.self_attn.v_proj.weight_scale_inv",
"layers.0.attention.o_proj.scale": "model.layers.0.self_attn.o_proj.weight_scale_inv",
"layers.0.mlp.gate.scale": "model.layers.0.mlp.gate_proj.weight_scale_inv",
"layers.0.mlp.up.scale": "model.layers.0.mlp.up_proj.weight_scale_inv",
"layers.0.mlp.down.scale": "model.layers.0.mlp.down_proj.weight_scale_inv",
"layers.1.attention.q_proj.scale": "model.layers.1.self_attn.q_proj.weight_scale_inv",
"layers.1.attention.k_proj.scale": "model.layers.1.self_attn.k_proj.weight_scale_inv",
"layers.1.attention.v_proj.scale": "model.layers.1.self_attn.v_proj.weight_scale_inv",
"layers.1.attention.o_proj.scale": "model.layers.1.self_attn.o_proj.weight_scale_inv",
"layers.1.mlp.gate.scale": "model.layers.1.mlp.gate_proj.weight_scale_inv",
"layers.1.mlp.up.scale": "model.layers.1.mlp.up_proj.weight_scale_inv",
"layers.1.mlp.down.scale": "model.layers.1.mlp.down_proj.weight_scale_inv",
"layers.2.attention.q_proj.scale": "model.layers.2.self_attn.q_proj.weight_scale_inv",
"layers.2.attention.k_proj.scale": "model.layers.2.self_attn.k_proj.weight_scale_inv",
"layers.2.attention.v_proj.scale": "model.layers.2.self_attn.v_proj.weight_scale_inv",
"layers.2.attention.o_proj.scale": "model.layers.2.self_attn.o_proj.weight_scale_inv",
"layers.2.mlp.gate.scale": "model.layers.2.mlp.gate_proj.weight_scale_inv",
"layers.2.mlp.up.scale": "model.layers.2.mlp.up_proj.weight_scale_inv",
"layers.2.mlp.down.scale": "model.layers.2.mlp.down_proj.weight_scale_inv",
"layers.3.attention.q_proj.scale": "model.layers.3.self_attn.q_proj.weight_scale_inv",
"layers.3.attention.k_proj.scale": "model.layers.3.self_attn.k_proj.weight_scale_inv",
"layers.3.attention.v_proj.scale": "model.layers.3.self_attn.v_proj.weight_scale_inv",
"layers.3.attention.o_proj.scale": "model.layers.3.self_attn.o_proj.weight_scale_inv",
"layers.3.mlp.gate.scale": "model.layers.3.mlp.gate_proj.weight_scale_inv",
"layers.3.mlp.up.scale": "model.layers.3.mlp.up_proj.weight_scale_inv",
"layers.3.mlp.down.scale": "model.layers.3.mlp.down_proj.weight_scale_inv",
"layers.4.attention.q_proj.scale": "model.layers.4.self_attn.q_proj.weight_scale_inv",
"layers.4.attention.k_proj.scale": "model.layers.4.self_attn.k_proj.weight_scale_inv",
"layers.4.attention.v_proj.scale": "model.layers.4.self_attn.v_proj.weight_scale_inv",
"layers.4.attention.o_proj.scale": "model.layers.4.self_attn.o_proj.weight_scale_inv",
"layers.4.mlp.gate.scale": "model.layers.4.mlp.gate_proj.weight_scale_inv",
"layers.4.mlp.up.scale": "model.layers.4.mlp.up_proj.weight_scale_inv",
"layers.4.mlp.down.scale": "model.layers.4.mlp.down_proj.weight_scale_inv",
"layers.5.attention.q_proj.scale": "model.layers.5.self_attn.q_proj.weight_scale_inv",
"layers.5.attention.k_proj.scale": "model.layers.5.self_attn.k_proj.weight_scale_inv",
"layers.5.attention.v_proj.scale": "model.layers.5.self_attn.v_proj.weight_scale_inv",
"layers.5.attention.o_proj.scale": "model.layers.5.self_attn.o_proj.weight_scale_inv",
"layers.5.mlp.gate.scale": "model.layers.5.mlp.gate_proj.weight_scale_inv",
"layers.5.mlp.up.scale": "model.layers.5.mlp.up_proj.weight_scale_inv",
"layers.5.mlp.down.scale": "model.layers.5.mlp.down_proj.weight_scale_inv",
"layers.6.attention.q_proj.scale": "model.layers.6.self_attn.q_proj.weight_scale_inv",
"layers.6.attention.k_proj.scale": "model.layers.6.self_attn.k_proj.weight_scale_inv",
"layers.6.attention.v_proj.scale": "model.layers.6.self_attn.v_proj.weight_scale_inv",
"layers.6.attention.o_proj.scale": "model.layers.6.self_attn.o_proj.weight_scale_inv",
"layers.6.mlp.gate.scale": "model.layers.6.mlp.gate_proj.weight_scale_inv",
"layers.6.mlp.up.scale": "model.layers.6.mlp.up_proj.weight_scale_inv",
"layers.6.mlp.down.scale": "model.layers.6.mlp.down_proj.weight_scale_inv",
"layers.7.attention.q_proj.scale": "model.layers.7.self_attn.q_proj.weight_scale_inv",
"layers.7.attention.k_proj.scale": "model.layers.7.self_attn.k_proj.weight_scale_inv",
"layers.7.attention.v_proj.scale": "model.layers.7.self_attn.v_proj.weight_scale_inv",
"layers.7.attention.o_proj.scale": "model.layers.7.self_attn.o_proj.weight_scale_inv",
"layers.7.mlp.gate.scale": "model.layers.7.mlp.gate_proj.weight_scale_inv",
"layers.7.mlp.up.scale": "model.layers.7.mlp.up_proj.weight_scale_inv",
"layers.7.mlp.down.scale": "model.layers.7.mlp.down_proj.weight_scale_inv",
"layers.8.attention.q_proj.scale": "model.layers.8.self_attn.q_proj.weight_scale_inv",
"layers.8.attention.k_proj.scale": "model.layers.8.self_attn.k_proj.weight_scale_inv",
"layers.8.attention.v_proj.scale": "model.layers.8.self_attn.v_proj.weight_scale_inv",
"layers.8.attention.o_proj.scale": "model.layers.8.self_attn.o_proj.weight_scale_inv",
"layers.8.mlp.gate.scale": "model.layers.8.mlp.gate_proj.weight_scale_inv",
"layers.8.mlp.up.scale": "model.layers.8.mlp.up_proj.weight_scale_inv",
"layers.8.mlp.down.scale": "model.layers.8.mlp.down_proj.weight_scale_inv",
"layers.9.attention.q_proj.scale": "model.layers.9.self_attn.q_proj.weight_scale_inv",
"layers.9.attention.k_proj.scale": "model.layers.9.self_attn.k_proj.weight_scale_inv",
"layers.9.attention.v_proj.scale": "model.layers.9.self_attn.v_proj.weight_scale_inv",
"layers.9.attention.o_proj.scale": "model.layers.9.self_attn.o_proj.weight_scale_inv",
"layers.9.mlp.gate.scale": "model.layers.9.mlp.gate_proj.weight_scale_inv",
"layers.9.mlp.up.scale": "model.layers.9.mlp.up_proj.weight_scale_inv",
"layers.9.mlp.down.scale": "model.layers.9.mlp.down_proj.weight_scale_inv",
"layers.10.attention.q_proj.scale": "model.layers.10.self_attn.q_proj.weight_scale_inv",
"layers.10.attention.k_proj.scale": "model.layers.10.self_attn.k_proj.weight_scale_inv",
"layers.10.attention.v_proj.scale": "model.layers.10.self_attn.v_proj.weight_scale_inv",
"layers.10.attention.o_proj.scale": "model.layers.10.self_attn.o_proj.weight_scale_inv",
"layers.10.mlp.gate.scale": "model.layers.10.mlp.gate_proj.weight_scale_inv",
"layers.10.mlp.up.scale": "model.layers.10.mlp.up_proj.weight_scale_inv",
"layers.10.mlp.down.scale": "model.layers.10.mlp.down_proj.weight_scale_inv",
"layers.11.attention.q_proj.scale": "model.layers.11.self_attn.q_proj.weight_scale_inv",
"layers.11.attention.k_proj.scale": "model.layers.11.self_attn.k_proj.weight_scale_inv",
"layers.11.attention.v_proj.scale": "model.layers.11.self_attn.v_proj.weight_scale_inv",
"layers.11.attention.o_proj.scale": "model.layers.11.self_attn.o_proj.weight_scale_inv",
"layers.11.mlp.gate.scale": "model.layers.11.mlp.gate_proj.weight_scale_inv",
"layers.11.mlp.up.scale": "model.layers.11.mlp.up_proj.weight_scale_inv",
"layers.11.mlp.down.scale": "model.layers.11.mlp.down_proj.weight_scale_inv",
"layers.12.attention.q_proj.scale": "model.layers.12.self_attn.q_proj.weight_scale_inv",
"layers.12.attention.k_proj.scale": "model.layers.12.self_attn.k_proj.weight_scale_inv",
"layers.12.attention.v_proj.scale": "model.layers.12.self_attn.v_proj.weight_scale_inv",
"layers.12.attention.o_proj.scale": "model.layers.12.self_attn.o_proj.weight_scale_inv",
"layers.12.mlp.gate.scale": "model.layers.12.mlp.gate_proj.weight_scale_inv",
"layers.12.mlp.up.scale": "model.layers.12.mlp.up_proj.weight_scale_inv",
"layers.12.mlp.down.scale": "model.layers.12.mlp.down_proj.weight_scale_inv",
"layers.13.attention.q_proj.scale": "model.layers.13.self_attn.q_proj.weight_scale_inv",
"layers.13.attention.k_proj.scale": "model.layers.13.self_attn.k_proj.weight_scale_inv",
"layers.13.attention.v_proj.scale": "model.layers.13.self_attn.v_proj.weight_scale_inv",
"layers.13.attention.o_proj.scale": "model.layers.13.self_attn.o_proj.weight_scale_inv",
"layers.13.mlp.gate.scale": "model.layers.13.mlp.gate_proj.weight_scale_inv",
"layers.13.mlp.up.scale": "model.layers.13.mlp.up_proj.weight_scale_inv",
"layers.13.mlp.down.scale": "model.layers.13.mlp.down_proj.weight_scale_inv",
"layers.14.attention.q_proj.scale": "model.layers.14.self_attn.q_proj.weight_scale_inv",
"layers.14.attention.k_proj.scale": "model.layers.14.self_attn.k_proj.weight_scale_inv",
"layers.14.attention.v_proj.scale": "model.layers.14.self_attn.v_proj.weight_scale_inv",
"layers.14.attention.o_proj.scale": "model.layers.14.self_attn.o_proj.weight_scale_inv",
"layers.14.mlp.gate.scale": "model.layers.14.mlp.gate_proj.weight_scale_inv",
"layers.14.mlp.up.scale": "model.layers.14.mlp.up_proj.weight_scale_inv",
"layers.14.mlp.down.scale": "model.layers.14.mlp.down_proj.weight_scale_inv",
"layers.15.attention.q_proj.scale": "model.layers.15.self_attn.q_proj.weight_scale_inv",
"layers.15.attention.k_proj.scale": "model.layers.15.self_attn.k_proj.weight_scale_inv",
"layers.15.attention.v_proj.scale": "model.layers.15.self_attn.v_proj.weight_scale_inv",
"layers.15.attention.o_proj.scale": "model.layers.15.self_attn.o_proj.weight_scale_inv",
"layers.15.mlp.gate.scale": "model.layers.15.mlp.gate_proj.weight_scale_inv",
"layers.15.mlp.up.scale": "model.layers.15.mlp.up_proj.weight_scale_inv",
"layers.15.mlp.down.scale": "model.layers.15.mlp.down_proj.weight_scale_inv",
"layers.16.attention.q_proj.scale": "model.layers.16.self_attn.q_proj.weight_scale_inv",
"layers.16.attention.k_proj.scale": "model.layers.16.self_attn.k_proj.weight_scale_inv",
"layers.16.attention.v_proj.scale": "model.layers.16.self_attn.v_proj.weight_scale_inv",
"layers.16.attention.o_proj.scale": "model.layers.16.self_attn.o_proj.weight_scale_inv",
"layers.16.mlp.gate.scale": "model.layers.16.mlp.gate_proj.weight_scale_inv",
"layers.16.mlp.up.scale": "model.layers.16.mlp.up_proj.weight_scale_inv",
"layers.16.mlp.down.scale": "model.layers.16.mlp.down_proj.weight_scale_inv",
"layers.17.attention.q_proj.scale": "model.layers.17.self_attn.q_proj.weight_scale_inv",
"layers.17.attention.k_proj.scale": "model.layers.17.self_attn.k_proj.weight_scale_inv",
"layers.17.attention.v_proj.scale": "model.layers.17.self_attn.v_proj.weight_scale_inv",
"layers.17.attention.o_proj.scale": "model.layers.17.self_attn.o_proj.weight_scale_inv",
"layers.17.mlp.gate.scale": "model.layers.17.mlp.gate_proj.weight_scale_inv",
"layers.17.mlp.up.scale": "model.layers.17.mlp.up_proj.weight_scale_inv",
"layers.17.mlp.down.scale": "model.layers.17.mlp.down_proj.weight_scale_inv",
"layers.18.attention.q_proj.scale": "model.layers.18.self_attn.q_proj.weight_scale_inv",
"layers.18.attention.k_proj.scale": "model.layers.18.self_attn.k_proj.weight_scale_inv",
"layers.18.attention.v_proj.scale": "model.layers.18.self_attn.v_proj.weight_scale_inv",
"layers.18.attention.o_proj.scale": "model.layers.18.self_attn.o_proj.weight_scale_inv",
"layers.18.mlp.gate.scale": "model.layers.18.mlp.gate_proj.weight_scale_inv",
"layers.18.mlp.up.scale": "model.layers.18.mlp.up_proj.weight_scale_inv",
"layers.18.mlp.down.scale": "model.layers.18.mlp.down_proj.weight_scale_inv",
"layers.19.attention.q_proj.scale": "model.layers.19.self_attn.q_proj.weight_scale_inv",
"layers.19.attention.k_proj.scale": "model.layers.19.self_attn.k_proj.weight_scale_inv",
"layers.19.attention.v_proj.scale": "model.layers.19.self_attn.v_proj.weight_scale_inv",
"layers.19.attention.o_proj.scale": "model.layers.19.self_attn.o_proj.weight_scale_inv",
"layers.19.mlp.gate.scale": "model.layers.19.mlp.gate_proj.weight_scale_inv",
"layers.19.mlp.up.scale": "model.layers.19.mlp.up_proj.weight_scale_inv",
"layers.19.mlp.down.scale": "model.layers.19.mlp.down_proj.weight_scale_inv",
"layers.20.attention.q_proj.scale": "model.layers.20.self_attn.q_proj.weight_scale_inv",
"layers.20.attention.k_proj.scale": "model.layers.20.self_attn.k_proj.weight_scale_inv",
"layers.20.attention.v_proj.scale": "model.layers.20.self_attn.v_proj.weight_scale_inv",
"layers.20.attention.o_proj.scale": "model.layers.20.self_attn.o_proj.weight_scale_inv",
"layers.20.mlp.gate.scale": "model.layers.20.mlp.gate_proj.weight_scale_inv",
"layers.20.mlp.up.scale": "model.layers.20.mlp.up_proj.weight_scale_inv",
"layers.20.mlp.down.scale": "model.layers.20.mlp.down_proj.weight_scale_inv",
"layers.21.attention.q_proj.scale": "model.layers.21.self_attn.q_proj.weight_scale_inv",
"layers.21.attention.k_proj.scale": "model.layers.21.self_attn.k_proj.weight_scale_inv",
"layers.21.attention.v_proj.scale": "model.layers.21.self_attn.v_proj.weight_scale_inv",
"layers.21.attention.o_proj.scale": "model.layers.21.self_attn.o_proj.weight_scale_inv",
"layers.21.mlp.gate.scale": "model.layers.21.mlp.gate_proj.weight_scale_inv",
"layers.21.mlp.up.scale": "model.layers.21.mlp.up_proj.weight_scale_inv",
"layers.21.mlp.down.scale": "model.layers.21.mlp.down_proj.weight_scale_inv",
"layers.22.attention.q_proj.scale": "model.layers.22.self_attn.q_proj.weight_scale_inv",
"layers.22.attention.k_proj.scale": "model.layers.22.self_attn.k_proj.weight_scale_inv",
"layers.22.attention.v_proj.scale": "model.layers.22.self_attn.v_proj.weight_scale_inv",
"layers.22.attention.o_proj.scale": "model.layers.22.self_attn.o_proj.weight_scale_inv",
"layers.22.mlp.gate.scale": "model.layers.22.mlp.gate_proj.weight_scale_inv",
"layers.22.mlp.up.scale": "model.layers.22.mlp.up_proj.weight_scale_inv",
"layers.22.mlp.down.scale": "model.layers.22.mlp.down_proj.weight_scale_inv",
"layers.23.attention.q_proj.scale": "model.layers.23.self_attn.q_proj.weight_scale_inv",
"layers.23.attention.k_proj.scale": "model.layers.23.self_attn.k_proj.weight_scale_inv",
"layers.23.attention.v_proj.scale": "model.layers.23.self_attn.v_proj.weight_scale_inv",
"layers.23.attention.o_proj.scale": "model.layers.23.self_attn.o_proj.weight_scale_inv",
"layers.23.mlp.gate.scale": "model.layers.23.mlp.gate_proj.weight_scale_inv",
"layers.23.mlp.up.scale": "model.layers.23.mlp.up_proj.weight_scale_inv",
"layers.23.mlp.down.scale": "model.layers.23.mlp.down_proj.weight_scale_inv",
"layers.24.attention.q_proj.scale": "model.layers.24.self_attn.q_proj.weight_scale_inv",
"layers.24.attention.k_proj.scale": "model.layers.24.self_attn.k_proj.weight_scale_inv",
"layers.24.attention.v_proj.scale": "model.layers.24.self_attn.v_proj.weight_scale_inv",
"layers.24.attention.o_proj.scale": "model.layers.24.self_attn.o_proj.weight_scale_inv",
"layers.24.mlp.gate.scale": "model.layers.24.mlp.gate_proj.weight_scale_inv",
"layers.24.mlp.up.scale": "model.layers.24.mlp.up_proj.weight_scale_inv",
"layers.24.mlp.down.scale": "model.layers.24.mlp.down_proj.weight_scale_inv",
"layers.25.attention.q_proj.scale": "model.layers.25.self_attn.q_proj.weight_scale_inv",
"layers.25.attention.k_proj.scale": "model.layers.25.self_attn.k_proj.weight_scale_inv",
"layers.25.attention.v_proj.scale": "model.layers.25.self_attn.v_proj.weight_scale_inv",
"layers.25.attention.o_proj.scale": "model.layers.25.self_attn.o_proj.weight_scale_inv",
"layers.25.mlp.gate.scale": "model.layers.25.mlp.gate_proj.weight_scale_inv",
"layers.25.mlp.up.scale": "model.layers.25.mlp.up_proj.weight_scale_inv",
"layers.25.mlp.down.scale": "model.layers.25.mlp.down_proj.weight_scale_inv",
"layers.26.attention.q_proj.scale": "model.layers.26.self_attn.q_proj.weight_scale_inv",
"layers.26.attention.k_proj.scale": "model.layers.26.self_attn.k_proj.weight_scale_inv",
"layers.26.attention.v_proj.scale": "model.layers.26.self_attn.v_proj.weight_scale_inv",
"layers.26.attention.o_proj.scale": "model.layers.26.self_attn.o_proj.weight_scale_inv",
"layers.26.mlp.gate.scale": "model.layers.26.mlp.gate_proj.weight_scale_inv",
"layers.26.mlp.up.scale": "model.layers.26.mlp.up_proj.weight_scale_inv",
"layers.26.mlp.down.scale": "model.layers.26.mlp.down_proj.weight_scale_inv",
"layers.27.attention.q_proj.scale": "model.layers.27.self_attn.q_proj.weight_scale_inv",
"layers.27.attention.k_proj.scale": "model.layers.27.self_attn.k_proj.weight_scale_inv",
"layers.27.attention.v_proj.scale": "model.layers.27.self_attn.v_proj.weight_scale_inv",
"layers.27.attention.o_proj.scale": "model.layers.27.self_attn.o_proj.weight_scale_inv",
"layers.27.mlp.gate.scale": "model.layers.27.mlp.gate_proj.weight_scale_inv",
"layers.27.mlp.up.scale": "model.layers.27.mlp.up_proj.weight_scale_inv",
"layers.27.mlp.down.scale": "model.layers.27.mlp.down_proj.weight_scale_inv",
"layers.28.attention.q_proj.scale": "model.layers.28.self_attn.q_proj.weight_scale_inv",
"layers.28.attention.k_proj.scale": "model.layers.28.self_attn.k_proj.weight_scale_inv",
"layers.28.attention.v_proj.scale": "model.layers.28.self_attn.v_proj.weight_scale_inv",
"layers.28.attention.o_proj.scale": "model.layers.28.self_attn.o_proj.weight_scale_inv",
"layers.28.mlp.gate.scale": "model.layers.28.mlp.gate_proj.weight_scale_inv",
"layers.28.mlp.up.scale": "model.layers.28.mlp.up_proj.weight_scale_inv",
"layers.28.mlp.down.scale": "model.layers.28.mlp.down_proj.weight_scale_inv",
"layers.29.attention.q_proj.scale": "model.layers.29.self_attn.q_proj.weight_scale_inv",
"layers.29.attention.k_proj.scale": "model.layers.29.self_attn.k_proj.weight_scale_inv",
"layers.29.attention.v_proj.scale": "model.layers.29.self_attn.v_proj.weight_scale_inv",
"layers.29.attention.o_proj.scale": "model.layers.29.self_attn.o_proj.weight_scale_inv",
"layers.29.mlp.gate.scale": "model.layers.29.mlp.gate_proj.weight_scale_inv",
"layers.29.mlp.up.scale": "model.layers.29.mlp.up_proj.weight_scale_inv",
"layers.29.mlp.down.scale": "model.layers.29.mlp.down_proj.weight_scale_inv",
"layers.30.attention.q_proj.scale": "model.layers.30.self_attn.q_proj.weight_scale_inv",
"layers.30.attention.k_proj.scale": "model.layers.30.self_attn.k_proj.weight_scale_inv",
"layers.30.attention.v_proj.scale": "model.layers.30.self_attn.v_proj.weight_scale_inv",
"layers.30.attention.o_proj.scale": "model.layers.30.self_attn.o_proj.weight_scale_inv",
"layers.30.mlp.gate.scale": "model.layers.30.mlp.gate_proj.weight_scale_inv",
"layers.30.mlp.up.scale": "model.layers.30.mlp.up_proj.weight_scale_inv",
"layers.30.mlp.down.scale": "model.layers.30.mlp.down_proj.weight_scale_inv",
"layers.31.attention.q_proj.scale": "model.layers.31.self_attn.q_proj.weight_scale_inv",
"layers.31.attention.k_proj.scale": "model.layers.31.self_attn.k_proj.weight_scale_inv",
"layers.31.attention.v_proj.scale": "model.layers.31.self_attn.v_proj.weight_scale_inv",
"layers.31.attention.o_proj.scale": "model.layers.31.self_attn.o_proj.weight_scale_inv",
"layers.31.mlp.gate.scale": "model.layers.31.mlp.gate_proj.weight_scale_inv",
"layers.31.mlp.up.scale": "model.layers.31.mlp.up_proj.weight_scale_inv",
"layers.31.mlp.down.scale": "model.layers.31.mlp.down_proj.weight_scale_inv",
"layers.32.attention.q_proj.scale": "model.layers.32.self_attn.q_proj.weight_scale_inv",
"layers.32.attention.k_proj.scale": "model.layers.32.self_attn.k_proj.weight_scale_inv",
"layers.32.attention.v_proj.scale": "model.layers.32.self_attn.v_proj.weight_scale_inv",
"layers.32.attention.o_proj.scale": "model.layers.32.self_attn.o_proj.weight_scale_inv",
"layers.32.mlp.gate.scale": "model.layers.32.mlp.gate_proj.weight_scale_inv",
"layers.32.mlp.up.scale": "model.layers.32.mlp.up_proj.weight_scale_inv",
"layers.32.mlp.down.scale": "model.layers.32.mlp.down_proj.weight_scale_inv",
"layers.33.attention.q_proj.scale": "model.layers.33.self_attn.q_proj.weight_scale_inv",
"layers.33.attention.k_proj.scale": "model.layers.33.self_attn.k_proj.weight_scale_inv",
"layers.33.attention.v_proj.scale": "model.layers.33.self_attn.v_proj.weight_scale_inv",
"layers.33.attention.o_proj.scale": "model.layers.33.self_attn.o_proj.weight_scale_inv",
"layers.33.mlp.gate.scale": "model.layers.33.mlp.gate_proj.weight_scale_inv",
"layers.33.mlp.up.scale": "model.layers.33.mlp.up_proj.weight_scale_inv",
"layers.33.mlp.down.scale": "model.layers.33.mlp.down_proj.weight_scale_inv",
"layers.34.attention.q_proj.scale": "model.layers.34.self_attn.q_proj.weight_scale_inv",
"layers.34.attention.k_proj.scale": "model.layers.34.self_attn.k_proj.weight_scale_inv",
"layers.34.attention.v_proj.scale": "model.layers.34.self_attn.v_proj.weight_scale_inv",
"layers.34.attention.o_proj.scale": "model.layers.34.self_attn.o_proj.weight_scale_inv",
"layers.34.mlp.gate.scale": "model.layers.34.mlp.gate_proj.weight_scale_inv",
"layers.34.mlp.up.scale": "model.layers.34.mlp.up_proj.weight_scale_inv",
"layers.34.mlp.down.scale": "model.layers.34.mlp.down_proj.weight_scale_inv",
"layers.35.attention.q_proj.scale": "model.layers.35.self_attn.q_proj.weight_scale_inv",
"layers.35.attention.k_proj.scale": "model.layers.35.self_attn.k_proj.weight_scale_inv",
"layers.35.attention.v_proj.scale": "model.layers.35.self_attn.v_proj.weight_scale_inv",
"layers.35.attention.o_proj.scale": "model.layers.35.self_attn.o_proj.weight_scale_inv",
"layers.35.mlp.gate.scale": "model.layers.35.mlp.gate_proj.weight_scale_inv",
"layers.35.mlp.up.scale": "model.layers.35.mlp.up_proj.weight_scale_inv",
"layers.35.mlp.down.scale": "model.layers.35.mlp.down_proj.weight_scale_inv"
}linnet.toml60 B
toml
[package]
name = "qwen3"
version = "0.1.0"
language = "0.1"nest.toml940 B
toml
[model]
name = "qwen3-8b-fp8"
title = "Qwen3 8B FP8"
summary = "Qwen3 8B with every decoder projection in FP8 E4M3, one scale per block of 128 by 128 weights: Qwen's own FP8 checkpoint. 60% of bf16's memory, faster decoding."
license = "Apache-2.0"
family = "qwen3"
tags = ["text-generation", "decoder-only", "grouped-query-attention", "qk-norm", "fp8"]
[links]
huggingface = "https://huggingface.co/Qwen/Qwen3-8B-FP8"
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-FP8"
revision = "220b46e3b2180893580a4454f21f22d3ebb187d3"
files = [
"model-00001-of-00002.safetensors",
"model-00002-of-00002.safetensors",
]
bindings = "bindings.json"src
attention.linnet21.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.
//
// 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.quant::{Fp8BlockLinear}
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
H % 128 == 0,
(Heads / Shards * HeadDim) % 128 == 0,
(KvHeads / Shards * HeadDim) % 128 == 0,
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: Fp8BlockLinear<H, Heads / Shards * HeadDim, 128, T>
sub q_norm: RmsNorm<HeadDim, T>
sub k_proj: Fp8BlockLinear<H, KvHeads / Shards * HeadDim, 128, T>
sub k_norm: RmsNorm<HeadDim, T>
sub v_proj: Fp8BlockLinear<H, KvHeads / Shards * HeadDim, 128, T>
sub o_proj: Fp8BlockLinear<Heads / Shards * HeadDim, H, 128, 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.linnet21.4 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.activations::{silu}
use std.quant::{Fp8BlockLinear}
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
H % 128 == 0,
(Heads / Shards * HeadDim) % 128 == 0,
(KvHeads / Shards * HeadDim) % 128 == 0,
(Inner / Shards) % 128 == 0,
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: Fp8BlockSwiGlu<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
H % 128 == 0,
(Heads / Shards * HeadDim) % 128 == 0,
(KvHeads / Shards * HeadDim) % 128 == 0,
(Inner / Shards) % 128 == 0,
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
}
}
// `std.nn.mlp::SwiGlu` over FP8 projections in blocks of 128 by 128: the
// checkpoint quantizes every linear layer of the decoder but the output head.
pub block Fp8BlockSwiGlu<H: Dim, Inner: Dim, T: Float = bf16>
where
H % 128 == 0,
Inner % 128 == 0
{
sub gate: Fp8BlockLinear<H, Inner, 128, T>
sub up: Fp8BlockLinear<H, Inner, 128, T>
sub down: Fp8BlockLinear<Inner, H, 128, T>
pub fn forward<*S: Shape>(x: Tensor[*S, H; T]) -> Tensor[*S, H; T] {
return down.forward(silu(gate.forward(x)) * up.forward(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)))
}