Models / llama-3.1-8b-instruct-fp8

Llama 3.1 8B Instruct FP8

Llama 3.1 8B Instruct with every decoder projection in FP8 E4M3, one scale per output row: RedHatAI's FP8-dynamic checkpoint. Half the weight memory of bf16, faster single-sequence decoding.

8B parametersllamallama3.1text-generationdecoder-onlygrouped-query-attentionchatfp8
models/llama-3.1-8b-instruct-fp811 files · 90.2 kB in the registry · GitHub
README.md1.3 kB
md
# Llama 3.1 8B Instruct FP8

Llama 3.1 8B Instruct with its decoder projections in FP8: RedHatAI's
[FP8-dynamic](https://huggingface.co/RedHatAI/Meta-Llama-3.1-8B-Instruct-FP8-dynamic)
checkpoint, made with `llm-compressor`. Each query, key, value, output,
gate, up and down projection holds E4M3 weights (one byte each) with one
bf16 scale per output row; the embedding and the output head stay bf16.

The source is the `llama-3.1-8b-instruct` package with those projections as
`std.quant::Fp8Linear`. `bindings.json` binds each `scale` to the
checkpoint's `weight_scale` and each FP8 `weight` to its bytes.

## Loading

```python
from linnet import nest

model = nest.load("llama-3.1-8b-instruct-fp8", backend="torch", numerics="fast")
```

On a GPU with compute capability 9 or later, `numerics="fast"` multiplies
one sequence's decoding step by the FP8 weights directly and rounds longer
inputs to FP8 a row at a time (`torch._scaled_mm`). 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 (`llama-3.1-8b-instruct`) | 9.55 | 19.1 ms | 162 tokens/s | 16.3 GiB |
| FP8 | 9.62 | 20.3 ms | 197 tokens/s | 9.9 GiB |

Open on GitHub

bench.json3.2 kB
json
{
  "date": "2026-10-10T21:54:35+00:00",
  "environment": {
    "python": "3.12.3",
    "machine": "x86_64",
    "torch": "2.14.1",
    "jax": "0.11.2",
    "transformers": "5.19.0",
    "linnet": "0.1.0",
    "gpu": "NVIDIA H100 80GB HBM3",
    "gpus": "2",
    "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": 87.46246621012688,
        "decode_tok_s": 11.705118022040558,
        "load_s": 20.85195330530405,
        "peak_vram_mib": 20552.0
      },
      "max_abs_diff": 0.0,
      "notes": "",
      "error": null,
      "key": "transformers-eager"
    },
    {
      "method": "Linnet torch (generated source)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 18.419036641716957,
        "decode_tok_s": 52.333152198570396,
        "load_s": 28.883636832237244,
        "peak_vram_mib": 13088.0
      },
      "max_abs_diff": 0.625,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-torch"
    },
    {
      "method": "Linnet torch (CUDA graphs)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 9.133616462349892,
        "decode_tok_s": 194.1926501454968,
        "load_s": 30.52232899144292,
        "peak_vram_mib": 13078.0
      },
      "max_abs_diff": 0.625,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-cudagraphs"
    },
    {
      "method": "vLLM",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 15.930453315377235,
        "decode_tok_s": 228.9141837250821,
        "load_s": 116.21744019538164,
        "first_token": 198.0,
        "peak_vram_mib": 68988.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"
    },
    {
      "method": "Linnet torch (linnet.serve, CUDA graphs)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 7024.278532077484,
        "serve_s": 4.664963077753782,
        "load_s": 72.17693518474698,
        "serve_ttft_ms": 1976.406890898943,
        "peak_vram_mib": 18120.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-torch"
    },
    {
      "method": "vLLM (offline, continuous batching)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 7677.109836153942,
        "serve_s": 4.268272917717695,
        "load_s": 44.680315125733614,
        "peak_vram_mib": 69768.8
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at gpu_memory_utilization 0.85",
      "error": null,
      "key": "serve-vllm"
    }
  ],
  "reference": "transformers-eager"
}

Open on GitHub

bindings.json39.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.mlp_norm.weight": "model.layers.0.post_attention_layernorm.weight",
  "layers.0.attention.q_proj.weight": "model.layers.0.self_attn.q_proj.weight",
  "layers.0.attention.q_proj.scale": "model.layers.0.self_attn.q_proj.weight_scale",
  "layers.0.attention.k_proj.weight": "model.layers.0.self_attn.k_proj.weight",
  "layers.0.attention.k_proj.scale": "model.layers.0.self_attn.k_proj.weight_scale",
  "layers.0.attention.v_proj.weight": "model.layers.0.self_attn.v_proj.weight",
  "layers.0.attention.v_proj.scale": "model.layers.0.self_attn.v_proj.weight_scale",
  "layers.0.attention.o_proj.weight": "model.layers.0.self_attn.o_proj.weight",
  "layers.0.attention.o_proj.scale": "model.layers.0.self_attn.o_proj.weight_scale",
  "layers.0.mlp.gate.weight": "model.layers.0.mlp.gate_proj.weight",
  "layers.0.mlp.gate.scale": "model.layers.0.mlp.gate_proj.weight_scale",
  "layers.0.mlp.up.weight": "model.layers.0.mlp.up_proj.weight",
  "layers.0.mlp.up.scale": "model.layers.0.mlp.up_proj.weight_scale",
  "layers.0.mlp.down.weight": "model.layers.0.mlp.down_proj.weight",
  "layers.0.mlp.down.scale": "model.layers.0.mlp.down_proj.weight_scale",
  "layers.1.attention_norm.weight": "model.layers.1.input_layernorm.weight",
  "layers.1.mlp_norm.weight": "model.layers.1.post_attention_layernorm.weight",
  "layers.1.attention.q_proj.weight": "model.layers.1.self_attn.q_proj.weight",
  "layers.1.attention.q_proj.scale": "model.layers.1.self_attn.q_proj.weight_scale",
  "layers.1.attention.k_proj.weight": "model.layers.1.self_attn.k_proj.weight",
  "layers.1.attention.k_proj.scale": "model.layers.1.self_attn.k_proj.weight_scale",
  "layers.1.attention.v_proj.weight": "model.layers.1.self_attn.v_proj.weight",
  "layers.1.attention.v_proj.scale": "model.layers.1.self_attn.v_proj.weight_scale",
  "layers.1.attention.o_proj.weight": "model.layers.1.self_attn.o_proj.weight",
  "layers.1.attention.o_proj.scale": "model.layers.1.self_attn.o_proj.weight_scale",
  "layers.1.mlp.gate.weight": "model.layers.1.mlp.gate_proj.weight",
  "layers.1.mlp.gate.scale": "model.layers.1.mlp.gate_proj.weight_scale",
  "layers.1.mlp.up.weight": "model.layers.1.mlp.up_proj.weight",
  "layers.1.mlp.up.scale": "model.layers.1.mlp.up_proj.weight_scale",
  "layers.1.mlp.down.weight": "model.layers.1.mlp.down_proj.weight",
  "layers.1.mlp.down.scale": "model.layers.1.mlp.down_proj.weight_scale",
  "layers.2.attention_norm.weight": "model.layers.2.input_layernorm.weight",
  "layers.2.mlp_norm.weight": "model.layers.2.post_attention_layernorm.weight",
  "layers.2.attention.q_proj.weight": "model.layers.2.self_attn.q_proj.weight",
  "layers.2.attention.q_proj.scale": "model.layers.2.self_attn.q_proj.weight_scale",
  "layers.2.attention.k_proj.weight": "model.layers.2.self_attn.k_proj.weight",
  "layers.2.attention.k_proj.scale": "model.layers.2.self_attn.k_proj.weight_scale",
  "layers.2.attention.v_proj.weight": "model.layers.2.self_attn.v_proj.weight",
  "layers.2.attention.v_proj.scale": "model.layers.2.self_attn.v_proj.weight_scale",
  "layers.2.attention.o_proj.weight": "model.layers.2.self_attn.o_proj.weight",
  "layers.2.attention.o_proj.scale": "model.layers.2.self_attn.o_proj.weight_scale",
  "layers.2.mlp.gate.weight": "model.layers.2.mlp.gate_proj.weight",
  "layers.2.mlp.gate.scale": "model.layers.2.mlp.gate_proj.weight_scale",
  "layers.2.mlp.up.weight": "model.layers.2.mlp.up_proj.weight",
  "layers.2.mlp.up.scale": "model.layers.2.mlp.up_proj.weight_scale",
  "layers.2.mlp.down.weight": "model.layers.2.mlp.down_proj.weight",
  "layers.2.mlp.down.scale": "model.layers.2.mlp.down_proj.weight_scale",
  "layers.3.attention_norm.weight": "model.layers.3.input_layernorm.weight",
  "layers.3.mlp_norm.weight": "model.layers.3.post_attention_layernorm.weight",
  "layers.3.attention.q_proj.weight": "model.layers.3.self_attn.q_proj.weight",
  "layers.3.attention.q_proj.scale": "model.layers.3.self_attn.q_proj.weight_scale",
  "layers.3.attention.k_proj.weight": "model.layers.3.self_attn.k_proj.weight",
  "layers.3.attention.k_proj.scale": "model.layers.3.self_attn.k_proj.weight_scale",
  "layers.3.attention.v_proj.weight": "model.layers.3.self_attn.v_proj.weight",
  "layers.3.attention.v_proj.scale": "model.layers.3.self_attn.v_proj.weight_scale",
  "layers.3.attention.o_proj.weight": "model.layers.3.self_attn.o_proj.weight",
  "layers.3.attention.o_proj.scale": "model.layers.3.self_attn.o_proj.weight_scale",
  "layers.3.mlp.gate.weight": "model.layers.3.mlp.gate_proj.weight",
  "layers.3.mlp.gate.scale": "model.layers.3.mlp.gate_proj.weight_scale",
  "layers.3.mlp.up.weight": "model.layers.3.mlp.up_proj.weight",
  "layers.3.mlp.up.scale": "model.layers.3.mlp.up_proj.weight_scale",
  "layers.3.mlp.down.weight": "model.layers.3.mlp.down_proj.weight",
  "layers.3.mlp.down.scale": "model.layers.3.mlp.down_proj.weight_scale",
  "layers.4.attention_norm.weight": "model.layers.4.input_layernorm.weight",
  "layers.4.mlp_norm.weight": "model.layers.4.post_attention_layernorm.weight",
  "layers.4.attention.q_proj.weight": "model.layers.4.self_attn.q_proj.weight",
  "layers.4.attention.q_proj.scale": "model.layers.4.self_attn.q_proj.weight_scale",
  "layers.4.attention.k_proj.weight": "model.layers.4.self_attn.k_proj.weight",
  "layers.4.attention.k_proj.scale": "model.layers.4.self_attn.k_proj.weight_scale",
  "layers.4.attention.v_proj.weight": "model.layers.4.self_attn.v_proj.weight",
  "layers.4.attention.v_proj.scale": "model.layers.4.self_attn.v_proj.weight_scale",
  "layers.4.attention.o_proj.weight": "model.layers.4.self_attn.o_proj.weight",
  "layers.4.attention.o_proj.scale": "model.layers.4.self_attn.o_proj.weight_scale",
  "layers.4.mlp.gate.weight": "model.layers.4.mlp.gate_proj.weight",
  "layers.4.mlp.gate.scale": "model.layers.4.mlp.gate_proj.weight_scale",
  "layers.4.mlp.up.weight": "model.layers.4.mlp.up_proj.weight",
  "layers.4.mlp.up.scale": "model.layers.4.mlp.up_proj.weight_scale",
  "layers.4.mlp.down.weight": "model.layers.4.mlp.down_proj.weight",
  "layers.4.mlp.down.scale": "model.layers.4.mlp.down_proj.weight_scale",
  "layers.5.attention_norm.weight": "model.layers.5.input_layernorm.weight",
  "layers.5.mlp_norm.weight": "model.layers.5.post_attention_layernorm.weight",
  "layers.5.attention.q_proj.weight": "model.layers.5.self_attn.q_proj.weight",
  "layers.5.attention.q_proj.scale": "model.layers.5.self_attn.q_proj.weight_scale",
  "layers.5.attention.k_proj.weight": "model.layers.5.self_attn.k_proj.weight",
  "layers.5.attention.k_proj.scale": "model.layers.5.self_attn.k_proj.weight_scale",
  "layers.5.attention.v_proj.weight": "model.layers.5.self_attn.v_proj.weight",
  "layers.5.attention.v_proj.scale": "model.layers.5.self_attn.v_proj.weight_scale",
  "layers.5.attention.o_proj.weight": "model.layers.5.self_attn.o_proj.weight",
  "layers.5.attention.o_proj.scale": "model.layers.5.self_attn.o_proj.weight_scale",
  "layers.5.mlp.gate.weight": "model.layers.5.mlp.gate_proj.weight",
  "layers.5.mlp.gate.scale": "model.layers.5.mlp.gate_proj.weight_scale",
  "layers.5.mlp.up.weight": "model.layers.5.mlp.up_proj.weight",
  "layers.5.mlp.up.scale": "model.layers.5.mlp.up_proj.weight_scale",
  "layers.5.mlp.down.weight": "model.layers.5.mlp.down_proj.weight",
  "layers.5.mlp.down.scale": "model.layers.5.mlp.down_proj.weight_scale",
  "layers.6.attention_norm.weight": "model.layers.6.input_layernorm.weight",
  "layers.6.mlp_norm.weight": "model.layers.6.post_attention_layernorm.weight",
  "layers.6.attention.q_proj.weight": "model.layers.6.self_attn.q_proj.weight",
  "layers.6.attention.q_proj.scale": "model.layers.6.self_attn.q_proj.weight_scale",
  "layers.6.attention.k_proj.weight": "model.layers.6.self_attn.k_proj.weight",
  "layers.6.attention.k_proj.scale": "model.layers.6.self_attn.k_proj.weight_scale",
  "layers.6.attention.v_proj.weight": "model.layers.6.self_attn.v_proj.weight",
  "layers.6.attention.v_proj.scale": "model.layers.6.self_attn.v_proj.weight_scale",
  "layers.6.attention.o_proj.weight": "model.layers.6.self_attn.o_proj.weight",
  "layers.6.attention.o_proj.scale": "model.layers.6.self_attn.o_proj.weight_scale",
  "layers.6.mlp.gate.weight": "model.layers.6.mlp.gate_proj.weight",
  "layers.6.mlp.gate.scale": "model.layers.6.mlp.gate_proj.weight_scale",
  "layers.6.mlp.up.weight": "model.layers.6.mlp.up_proj.weight",
  "layers.6.mlp.up.scale": "model.layers.6.mlp.up_proj.weight_scale",
  "layers.6.mlp.down.weight": "model.layers.6.mlp.down_proj.weight",
  "layers.6.mlp.down.scale": "model.layers.6.mlp.down_proj.weight_scale",
  "layers.7.attention_norm.weight": "model.layers.7.input_layernorm.weight",
  "layers.7.mlp_norm.weight": "model.layers.7.post_attention_layernorm.weight",
  "layers.7.attention.q_proj.weight": "model.layers.7.self_attn.q_proj.weight",
  "layers.7.attention.q_proj.scale": "model.layers.7.self_attn.q_proj.weight_scale",
  "layers.7.attention.k_proj.weight": "model.layers.7.self_attn.k_proj.weight",
  "layers.7.attention.k_proj.scale": "model.layers.7.self_attn.k_proj.weight_scale",
  "layers.7.attention.v_proj.weight": "model.layers.7.self_attn.v_proj.weight",
  "layers.7.attention.v_proj.scale": "model.layers.7.self_attn.v_proj.weight_scale",
  "layers.7.attention.o_proj.weight": "model.layers.7.self_attn.o_proj.weight",
  "layers.7.attention.o_proj.scale": "model.layers.7.self_attn.o_proj.weight_scale",
  "layers.7.mlp.gate.weight": "model.layers.7.mlp.gate_proj.weight",
  "layers.7.mlp.gate.scale": "model.layers.7.mlp.gate_proj.weight_scale",
  "layers.7.mlp.up.weight": "model.layers.7.mlp.up_proj.weight",
  "layers.7.mlp.up.scale": "model.layers.7.mlp.up_proj.weight_scale",
  "layers.7.mlp.down.weight": "model.layers.7.mlp.down_proj.weight",
  "layers.7.mlp.down.scale": "model.layers.7.mlp.down_proj.weight_scale",
  "layers.8.attention_norm.weight": "model.layers.8.input_layernorm.weight",
  "layers.8.mlp_norm.weight": "model.layers.8.post_attention_layernorm.weight",
  "layers.8.attention.q_proj.weight": "model.layers.8.self_attn.q_proj.weight",
  "layers.8.attention.q_proj.scale": "model.layers.8.self_attn.q_proj.weight_scale",
  "layers.8.attention.k_proj.weight": "model.layers.8.self_attn.k_proj.weight",
  "layers.8.attention.k_proj.scale": "model.layers.8.self_attn.k_proj.weight_scale",
  "layers.8.attention.v_proj.weight": "model.layers.8.self_attn.v_proj.weight",
  "layers.8.attention.v_proj.scale": "model.layers.8.self_attn.v_proj.weight_scale",
  "layers.8.attention.o_proj.weight": "model.layers.8.self_attn.o_proj.weight",
  "layers.8.attention.o_proj.scale": "model.layers.8.self_attn.o_proj.weight_scale",
  "layers.8.mlp.gate.weight": "model.layers.8.mlp.gate_proj.weight",
  "layers.8.mlp.gate.scale": "model.layers.8.mlp.gate_proj.weight_scale",
  "layers.8.mlp.up.weight": "model.layers.8.mlp.up_proj.weight",
  "layers.8.mlp.up.scale": "model.layers.8.mlp.up_proj.weight_scale",
  "layers.8.mlp.down.weight": "model.layers.8.mlp.down_proj.weight",
  "layers.8.mlp.down.scale": "model.layers.8.mlp.down_proj.weight_scale",
  "layers.9.attention_norm.weight": "model.layers.9.input_layernorm.weight",
  "layers.9.mlp_norm.weight": "model.layers.9.post_attention_layernorm.weight",
  "layers.9.attention.q_proj.weight": "model.layers.9.self_attn.q_proj.weight",
  "layers.9.attention.q_proj.scale": "model.layers.9.self_attn.q_proj.weight_scale",
  "layers.9.attention.k_proj.weight": "model.layers.9.self_attn.k_proj.weight",
  "layers.9.attention.k_proj.scale": "model.layers.9.self_attn.k_proj.weight_scale",
  "layers.9.attention.v_proj.weight": "model.layers.9.self_attn.v_proj.weight",
  "layers.9.attention.v_proj.scale": "model.layers.9.self_attn.v_proj.weight_scale",
  "layers.9.attention.o_proj.weight": "model.layers.9.self_attn.o_proj.weight",
  "layers.9.attention.o_proj.scale": "model.layers.9.self_attn.o_proj.weight_scale",
  "layers.9.mlp.gate.weight": "model.layers.9.mlp.gate_proj.weight",
  "layers.9.mlp.gate.scale": "model.layers.9.mlp.gate_proj.weight_scale",
  "layers.9.mlp.up.weight": "model.layers.9.mlp.up_proj.weight",
  "layers.9.mlp.up.scale": "model.layers.9.mlp.up_proj.weight_scale",
  "layers.9.mlp.down.weight": "model.layers.9.mlp.down_proj.weight",
  "layers.9.mlp.down.scale": "model.layers.9.mlp.down_proj.weight_scale",
  "layers.10.attention_norm.weight": "model.layers.10.input_layernorm.weight",
  "layers.10.mlp_norm.weight": "model.layers.10.post_attention_layernorm.weight",
  "layers.10.attention.q_proj.weight": "model.layers.10.self_attn.q_proj.weight",
  "layers.10.attention.q_proj.scale": "model.layers.10.self_attn.q_proj.weight_scale",
  "layers.10.attention.k_proj.weight": "model.layers.10.self_attn.k_proj.weight",
  "layers.10.attention.k_proj.scale": "model.layers.10.self_attn.k_proj.weight_scale",
  "layers.10.attention.v_proj.weight": "model.layers.10.self_attn.v_proj.weight",
  "layers.10.attention.v_proj.scale": "model.layers.10.self_attn.v_proj.weight_scale",
  "layers.10.attention.o_proj.weight": "model.layers.10.self_attn.o_proj.weight",
  "layers.10.attention.o_proj.scale": "model.layers.10.self_attn.o_proj.weight_scale",
  "layers.10.mlp.gate.weight": "model.layers.10.mlp.gate_proj.weight",
  "layers.10.mlp.gate.scale": "model.layers.10.mlp.gate_proj.weight_scale",
  "layers.10.mlp.up.weight": "model.layers.10.mlp.up_proj.weight",
  "layers.10.mlp.up.scale": "model.layers.10.mlp.up_proj.weight_scale",
  "layers.10.mlp.down.weight": "model.layers.10.mlp.down_proj.weight",
  "layers.10.mlp.down.scale": "model.layers.10.mlp.down_proj.weight_scale",
  "layers.11.attention_norm.weight": "model.layers.11.input_layernorm.weight",
  "layers.11.mlp_norm.weight": "model.layers.11.post_attention_layernorm.weight",
  "layers.11.attention.q_proj.weight": "model.layers.11.self_attn.q_proj.weight",
  "layers.11.attention.q_proj.scale": "model.layers.11.self_attn.q_proj.weight_scale",
  "layers.11.attention.k_proj.weight": "model.layers.11.self_attn.k_proj.weight",
  "layers.11.attention.k_proj.scale": "model.layers.11.self_attn.k_proj.weight_scale",
  "layers.11.attention.v_proj.weight": "model.layers.11.self_attn.v_proj.weight",
  "layers.11.attention.v_proj.scale": "model.layers.11.self_attn.v_proj.weight_scale",
  "layers.11.attention.o_proj.weight": "model.layers.11.self_attn.o_proj.weight",
  "layers.11.attention.o_proj.scale": "model.layers.11.self_attn.o_proj.weight_scale",
  "layers.11.mlp.gate.weight": "model.layers.11.mlp.gate_proj.weight",
  "layers.11.mlp.gate.scale": "model.layers.11.mlp.gate_proj.weight_scale",
  "layers.11.mlp.up.weight": "model.layers.11.mlp.up_proj.weight",
  "layers.11.mlp.up.scale": "model.layers.11.mlp.up_proj.weight_scale",
  "layers.11.mlp.down.weight": "model.layers.11.mlp.down_proj.weight",
  "layers.11.mlp.down.scale": "model.layers.11.mlp.down_proj.weight_scale",
  "layers.12.attention_norm.weight": "model.layers.12.input_layernorm.weight",
  "layers.12.mlp_norm.weight": "model.layers.12.post_attention_layernorm.weight",
  "layers.12.attention.q_proj.weight": "model.layers.12.self_attn.q_proj.weight",
  "layers.12.attention.q_proj.scale": "model.layers.12.self_attn.q_proj.weight_scale",
  "layers.12.attention.k_proj.weight": "model.layers.12.self_attn.k_proj.weight",
  "layers.12.attention.k_proj.scale": "model.layers.12.self_attn.k_proj.weight_scale",
  "layers.12.attention.v_proj.weight": "model.layers.12.self_attn.v_proj.weight",
  "layers.12.attention.v_proj.scale": "model.layers.12.self_attn.v_proj.weight_scale",
  "layers.12.attention.o_proj.weight": "model.layers.12.self_attn.o_proj.weight",
  "layers.12.attention.o_proj.scale": "model.layers.12.self_attn.o_proj.weight_scale",
  "layers.12.mlp.gate.weight": "model.layers.12.mlp.gate_proj.weight",
  "layers.12.mlp.gate.scale": "model.layers.12.mlp.gate_proj.weight_scale",
  "layers.12.mlp.up.weight": "model.layers.12.mlp.up_proj.weight",
  "layers.12.mlp.up.scale": "model.layers.12.mlp.up_proj.weight_scale",
  "layers.12.mlp.down.weight": "model.layers.12.mlp.down_proj.weight",
  "layers.12.mlp.down.scale": "model.layers.12.mlp.down_proj.weight_scale",
  "layers.13.attention_norm.weight": "model.layers.13.input_layernorm.weight",
  "layers.13.mlp_norm.weight": "model.layers.13.post_attention_layernorm.weight",
  "layers.13.attention.q_proj.weight": "model.layers.13.self_attn.q_proj.weight",
  "layers.13.attention.q_proj.scale": "model.layers.13.self_attn.q_proj.weight_scale",
  "layers.13.attention.k_proj.weight": "model.layers.13.self_attn.k_proj.weight",
  "layers.13.attention.k_proj.scale": "model.layers.13.self_attn.k_proj.weight_scale",
  "layers.13.attention.v_proj.weight": "model.layers.13.self_attn.v_proj.weight",
  "layers.13.attention.v_proj.scale": "model.layers.13.self_attn.v_proj.weight_scale",
  "layers.13.attention.o_proj.weight": "model.layers.13.self_attn.o_proj.weight",
  "layers.13.attention.o_proj.scale": "model.layers.13.self_attn.o_proj.weight_scale",
  "layers.13.mlp.gate.weight": "model.layers.13.mlp.gate_proj.weight",
  "layers.13.mlp.gate.scale": "model.layers.13.mlp.gate_proj.weight_scale",
  "layers.13.mlp.up.weight": "model.layers.13.mlp.up_proj.weight",
  "layers.13.mlp.up.scale": "model.layers.13.mlp.up_proj.weight_scale",
  "layers.13.mlp.down.weight": "model.layers.13.mlp.down_proj.weight",
  "layers.13.mlp.down.scale": "model.layers.13.mlp.down_proj.weight_scale",
  "layers.14.attention_norm.weight": "model.layers.14.input_layernorm.weight",
  "layers.14.mlp_norm.weight": "model.layers.14.post_attention_layernorm.weight",
  "layers.14.attention.q_proj.weight": "model.layers.14.self_attn.q_proj.weight",
  "layers.14.attention.q_proj.scale": "model.layers.14.self_attn.q_proj.weight_scale",
  "layers.14.attention.k_proj.weight": "model.layers.14.self_attn.k_proj.weight",
  "layers.14.attention.k_proj.scale": "model.layers.14.self_attn.k_proj.weight_scale",
  "layers.14.attention.v_proj.weight": "model.layers.14.self_attn.v_proj.weight",
  "layers.14.attention.v_proj.scale": "model.layers.14.self_attn.v_proj.weight_scale",
  "layers.14.attention.o_proj.weight": "model.layers.14.self_attn.o_proj.weight",
  "layers.14.attention.o_proj.scale": "model.layers.14.self_attn.o_proj.weight_scale",
  "layers.14.mlp.gate.weight": "model.layers.14.mlp.gate_proj.weight",
  "layers.14.mlp.gate.scale": "model.layers.14.mlp.gate_proj.weight_scale",
  "layers.14.mlp.up.weight": "model.layers.14.mlp.up_proj.weight",
  "layers.14.mlp.up.scale": "model.layers.14.mlp.up_proj.weight_scale",
  "layers.14.mlp.down.weight": "model.layers.14.mlp.down_proj.weight",
  "layers.14.mlp.down.scale": "model.layers.14.mlp.down_proj.weight_scale",
  "layers.15.attention_norm.weight": "model.layers.15.input_layernorm.weight",
  "layers.15.mlp_norm.weight": "model.layers.15.post_attention_layernorm.weight",
  "layers.15.attention.q_proj.weight": "model.layers.15.self_attn.q_proj.weight",
  "layers.15.attention.q_proj.scale": "model.layers.15.self_attn.q_proj.weight_scale",
  "layers.15.attention.k_proj.weight": "model.layers.15.self_attn.k_proj.weight",
  "layers.15.attention.k_proj.scale": "model.layers.15.self_attn.k_proj.weight_scale",
  "layers.15.attention.v_proj.weight": "model.layers.15.self_attn.v_proj.weight",
  "layers.15.attention.v_proj.scale": "model.layers.15.self_attn.v_proj.weight_scale",
  "layers.15.attention.o_proj.weight": "model.layers.15.self_attn.o_proj.weight",
  "layers.15.attention.o_proj.scale": "model.layers.15.self_attn.o_proj.weight_scale",
  "layers.15.mlp.gate.weight": "model.layers.15.mlp.gate_proj.weight",
  "layers.15.mlp.gate.scale": "model.layers.15.mlp.gate_proj.weight_scale",
  "layers.15.mlp.up.weight": "model.layers.15.mlp.up_proj.weight",
  "layers.15.mlp.up.scale": "model.layers.15.mlp.up_proj.weight_scale",
  "layers.15.mlp.down.weight": "model.layers.15.mlp.down_proj.weight",
  "layers.15.mlp.down.scale": "model.layers.15.mlp.down_proj.weight_scale",
  "layers.16.attention_norm.weight": "model.layers.16.input_layernorm.weight",
  "layers.16.mlp_norm.weight": "model.layers.16.post_attention_layernorm.weight",
  "layers.16.attention.q_proj.weight": "model.layers.16.self_attn.q_proj.weight",
  "layers.16.attention.q_proj.scale": "model.layers.16.self_attn.q_proj.weight_scale",
  "layers.16.attention.k_proj.weight": "model.layers.16.self_attn.k_proj.weight",
  "layers.16.attention.k_proj.scale": "model.layers.16.self_attn.k_proj.weight_scale",
  "layers.16.attention.v_proj.weight": "model.layers.16.self_attn.v_proj.weight",
  "layers.16.attention.v_proj.scale": "model.layers.16.self_attn.v_proj.weight_scale",
  "layers.16.attention.o_proj.weight": "model.layers.16.self_attn.o_proj.weight",
  "layers.16.attention.o_proj.scale": "model.layers.16.self_attn.o_proj.weight_scale",
  "layers.16.mlp.gate.weight": "model.layers.16.mlp.gate_proj.weight",
  "layers.16.mlp.gate.scale": "model.layers.16.mlp.gate_proj.weight_scale",
  "layers.16.mlp.up.weight": "model.layers.16.mlp.up_proj.weight",
  "layers.16.mlp.up.scale": "model.layers.16.mlp.up_proj.weight_scale",
  "layers.16.mlp.down.weight": "model.layers.16.mlp.down_proj.weight",
  "layers.16.mlp.down.scale": "model.layers.16.mlp.down_proj.weight_scale",
  "layers.17.attention_norm.weight": "model.layers.17.input_layernorm.weight",
  "layers.17.mlp_norm.weight": "model.layers.17.post_attention_layernorm.weight",
  "layers.17.attention.q_proj.weight": "model.layers.17.self_attn.q_proj.weight",
  "layers.17.attention.q_proj.scale": "model.layers.17.self_attn.q_proj.weight_scale",
  "layers.17.attention.k_proj.weight": "model.layers.17.self_attn.k_proj.weight",
  "layers.17.attention.k_proj.scale": "model.layers.17.self_attn.k_proj.weight_scale",
  "layers.17.attention.v_proj.weight": "model.layers.17.self_attn.v_proj.weight",
  "layers.17.attention.v_proj.scale": "model.layers.17.self_attn.v_proj.weight_scale",
  "layers.17.attention.o_proj.weight": "model.layers.17.self_attn.o_proj.weight",
  "layers.17.attention.o_proj.scale": "model.layers.17.self_attn.o_proj.weight_scale",
  "layers.17.mlp.gate.weight": "model.layers.17.mlp.gate_proj.weight",
  "layers.17.mlp.gate.scale": "model.layers.17.mlp.gate_proj.weight_scale",
  "layers.17.mlp.up.weight": "model.layers.17.mlp.up_proj.weight",
  "layers.17.mlp.up.scale": "model.layers.17.mlp.up_proj.weight_scale",
  "layers.17.mlp.down.weight": "model.layers.17.mlp.down_proj.weight",
  "layers.17.mlp.down.scale": "model.layers.17.mlp.down_proj.weight_scale",
  "layers.18.attention_norm.weight": "model.layers.18.input_layernorm.weight",
  "layers.18.mlp_norm.weight": "model.layers.18.post_attention_layernorm.weight",
  "layers.18.attention.q_proj.weight": "model.layers.18.self_attn.q_proj.weight",
  "layers.18.attention.q_proj.scale": "model.layers.18.self_attn.q_proj.weight_scale",
  "layers.18.attention.k_proj.weight": "model.layers.18.self_attn.k_proj.weight",
  "layers.18.attention.k_proj.scale": "model.layers.18.self_attn.k_proj.weight_scale",
  "layers.18.attention.v_proj.weight": "model.layers.18.self_attn.v_proj.weight",
  "layers.18.attention.v_proj.scale": "model.layers.18.self_attn.v_proj.weight_scale",
  "layers.18.attention.o_proj.weight": "model.layers.18.self_attn.o_proj.weight",
  "layers.18.attention.o_proj.scale": "model.layers.18.self_attn.o_proj.weight_scale",
  "layers.18.mlp.gate.weight": "model.layers.18.mlp.gate_proj.weight",
  "layers.18.mlp.gate.scale": "model.layers.18.mlp.gate_proj.weight_scale",
  "layers.18.mlp.up.weight": "model.layers.18.mlp.up_proj.weight",
  "layers.18.mlp.up.scale": "model.layers.18.mlp.up_proj.weight_scale",
  "layers.18.mlp.down.weight": "model.layers.18.mlp.down_proj.weight",
  "layers.18.mlp.down.scale": "model.layers.18.mlp.down_proj.weight_scale",
  "layers.19.attention_norm.weight": "model.layers.19.input_layernorm.weight",
  "layers.19.mlp_norm.weight": "model.layers.19.post_attention_layernorm.weight",
  "layers.19.attention.q_proj.weight": "model.layers.19.self_attn.q_proj.weight",
  "layers.19.attention.q_proj.scale": "model.layers.19.self_attn.q_proj.weight_scale",
  "layers.19.attention.k_proj.weight": "model.layers.19.self_attn.k_proj.weight",
  "layers.19.attention.k_proj.scale": "model.layers.19.self_attn.k_proj.weight_scale",
  "layers.19.attention.v_proj.weight": "model.layers.19.self_attn.v_proj.weight",
  "layers.19.attention.v_proj.scale": "model.layers.19.self_attn.v_proj.weight_scale",
  "layers.19.attention.o_proj.weight": "model.layers.19.self_attn.o_proj.weight",
  "layers.19.attention.o_proj.scale": "model.layers.19.self_attn.o_proj.weight_scale",
  "layers.19.mlp.gate.weight": "model.layers.19.mlp.gate_proj.weight",
  "layers.19.mlp.gate.scale": "model.layers.19.mlp.gate_proj.weight_scale",
  "layers.19.mlp.up.weight": "model.layers.19.mlp.up_proj.weight",
  "layers.19.mlp.up.scale": "model.layers.19.mlp.up_proj.weight_scale",
  "layers.19.mlp.down.weight": "model.layers.19.mlp.down_proj.weight",
  "layers.19.mlp.down.scale": "model.layers.19.mlp.down_proj.weight_scale",
  "layers.20.attention_norm.weight": "model.layers.20.input_layernorm.weight",
  "layers.20.mlp_norm.weight": "model.layers.20.post_attention_layernorm.weight",
  "layers.20.attention.q_proj.weight": "model.layers.20.self_attn.q_proj.weight",
  "layers.20.attention.q_proj.scale": "model.layers.20.self_attn.q_proj.weight_scale",
  "layers.20.attention.k_proj.weight": "model.layers.20.self_attn.k_proj.weight",
  "layers.20.attention.k_proj.scale": "model.layers.20.self_attn.k_proj.weight_scale",
  "layers.20.attention.v_proj.weight": "model.layers.20.self_attn.v_proj.weight",
  "layers.20.attention.v_proj.scale": "model.layers.20.self_attn.v_proj.weight_scale",
  "layers.20.attention.o_proj.weight": "model.layers.20.self_attn.o_proj.weight",
  "layers.20.attention.o_proj.scale": "model.layers.20.self_attn.o_proj.weight_scale",
  "layers.20.mlp.gate.weight": "model.layers.20.mlp.gate_proj.weight",
  "layers.20.mlp.gate.scale": "model.layers.20.mlp.gate_proj.weight_scale",
  "layers.20.mlp.up.weight": "model.layers.20.mlp.up_proj.weight",
  "layers.20.mlp.up.scale": "model.layers.20.mlp.up_proj.weight_scale",
  "layers.20.mlp.down.weight": "model.layers.20.mlp.down_proj.weight",
  "layers.20.mlp.down.scale": "model.layers.20.mlp.down_proj.weight_scale",
  "layers.21.attention_norm.weight": "model.layers.21.input_layernorm.weight",
  "layers.21.mlp_norm.weight": "model.layers.21.post_attention_layernorm.weight",
  "layers.21.attention.q_proj.weight": "model.layers.21.self_attn.q_proj.weight",
  "layers.21.attention.q_proj.scale": "model.layers.21.self_attn.q_proj.weight_scale",
  "layers.21.attention.k_proj.weight": "model.layers.21.self_attn.k_proj.weight",
  "layers.21.attention.k_proj.scale": "model.layers.21.self_attn.k_proj.weight_scale",
  "layers.21.attention.v_proj.weight": "model.layers.21.self_attn.v_proj.weight",
  "layers.21.attention.v_proj.scale": "model.layers.21.self_attn.v_proj.weight_scale",
  "layers.21.attention.o_proj.weight": "model.layers.21.self_attn.o_proj.weight",
  "layers.21.attention.o_proj.scale": "model.layers.21.self_attn.o_proj.weight_scale",
  "layers.21.mlp.gate.weight": "model.layers.21.mlp.gate_proj.weight",
  "layers.21.mlp.gate.scale": "model.layers.21.mlp.gate_proj.weight_scale",
  "layers.21.mlp.up.weight": "model.layers.21.mlp.up_proj.weight",
  "layers.21.mlp.up.scale": "model.layers.21.mlp.up_proj.weight_scale",
  "layers.21.mlp.down.weight": "model.layers.21.mlp.down_proj.weight",
  "layers.21.mlp.down.scale": "model.layers.21.mlp.down_proj.weight_scale",
  "layers.22.attention_norm.weight": "model.layers.22.input_layernorm.weight",
  "layers.22.mlp_norm.weight": "model.layers.22.post_attention_layernorm.weight",
  "layers.22.attention.q_proj.weight": "model.layers.22.self_attn.q_proj.weight",
  "layers.22.attention.q_proj.scale": "model.layers.22.self_attn.q_proj.weight_scale",
  "layers.22.attention.k_proj.weight": "model.layers.22.self_attn.k_proj.weight",
  "layers.22.attention.k_proj.scale": "model.layers.22.self_attn.k_proj.weight_scale",
  "layers.22.attention.v_proj.weight": "model.layers.22.self_attn.v_proj.weight",
  "layers.22.attention.v_proj.scale": "model.layers.22.self_attn.v_proj.weight_scale",
  "layers.22.attention.o_proj.weight": "model.layers.22.self_attn.o_proj.weight",
  "layers.22.attention.o_proj.scale": "model.layers.22.self_attn.o_proj.weight_scale",
  "layers.22.mlp.gate.weight": "model.layers.22.mlp.gate_proj.weight",
  "layers.22.mlp.gate.scale": "model.layers.22.mlp.gate_proj.weight_scale",
  "layers.22.mlp.up.weight": "model.layers.22.mlp.up_proj.weight",
  "layers.22.mlp.up.scale": "model.layers.22.mlp.up_proj.weight_scale",
  "layers.22.mlp.down.weight": "model.layers.22.mlp.down_proj.weight",
  "layers.22.mlp.down.scale": "model.layers.22.mlp.down_proj.weight_scale",
  "layers.23.attention_norm.weight": "model.layers.23.input_layernorm.weight",
  "layers.23.mlp_norm.weight": "model.layers.23.post_attention_layernorm.weight",
  "layers.23.attention.q_proj.weight": "model.layers.23.self_attn.q_proj.weight",
  "layers.23.attention.q_proj.scale": "model.layers.23.self_attn.q_proj.weight_scale",
  "layers.23.attention.k_proj.weight": "model.layers.23.self_attn.k_proj.weight",
  "layers.23.attention.k_proj.scale": "model.layers.23.self_attn.k_proj.weight_scale",
  "layers.23.attention.v_proj.weight": "model.layers.23.self_attn.v_proj.weight",
  "layers.23.attention.v_proj.scale": "model.layers.23.self_attn.v_proj.weight_scale",
  "layers.23.attention.o_proj.weight": "model.layers.23.self_attn.o_proj.weight",
  "layers.23.attention.o_proj.scale": "model.layers.23.self_attn.o_proj.weight_scale",
  "layers.23.mlp.gate.weight": "model.layers.23.mlp.gate_proj.weight",
  "layers.23.mlp.gate.scale": "model.layers.23.mlp.gate_proj.weight_scale",
  "layers.23.mlp.up.weight": "model.layers.23.mlp.up_proj.weight",
  "layers.23.mlp.up.scale": "model.layers.23.mlp.up_proj.weight_scale",
  "layers.23.mlp.down.weight": "model.layers.23.mlp.down_proj.weight",
  "layers.23.mlp.down.scale": "model.layers.23.mlp.down_proj.weight_scale",
  "layers.24.attention_norm.weight": "model.layers.24.input_layernorm.weight",
  "layers.24.mlp_norm.weight": "model.layers.24.post_attention_layernorm.weight",
  "layers.24.attention.q_proj.weight": "model.layers.24.self_attn.q_proj.weight",
  "layers.24.attention.q_proj.scale": "model.layers.24.self_attn.q_proj.weight_scale",
  "layers.24.attention.k_proj.weight": "model.layers.24.self_attn.k_proj.weight",
  "layers.24.attention.k_proj.scale": "model.layers.24.self_attn.k_proj.weight_scale",
  "layers.24.attention.v_proj.weight": "model.layers.24.self_attn.v_proj.weight",
  "layers.24.attention.v_proj.scale": "model.layers.24.self_attn.v_proj.weight_scale",
  "layers.24.attention.o_proj.weight": "model.layers.24.self_attn.o_proj.weight",
  "layers.24.attention.o_proj.scale": "model.layers.24.self_attn.o_proj.weight_scale",
  "layers.24.mlp.gate.weight": "model.layers.24.mlp.gate_proj.weight",
  "layers.24.mlp.gate.scale": "model.layers.24.mlp.gate_proj.weight_scale",
  "layers.24.mlp.up.weight": "model.layers.24.mlp.up_proj.weight",
  "layers.24.mlp.up.scale": "model.layers.24.mlp.up_proj.weight_scale",
  "layers.24.mlp.down.weight": "model.layers.24.mlp.down_proj.weight",
  "layers.24.mlp.down.scale": "model.layers.24.mlp.down_proj.weight_scale",
  "layers.25.attention_norm.weight": "model.layers.25.input_layernorm.weight",
  "layers.25.mlp_norm.weight": "model.layers.25.post_attention_layernorm.weight",
  "layers.25.attention.q_proj.weight": "model.layers.25.self_attn.q_proj.weight",
  "layers.25.attention.q_proj.scale": "model.layers.25.self_attn.q_proj.weight_scale",
  "layers.25.attention.k_proj.weight": "model.layers.25.self_attn.k_proj.weight",
  "layers.25.attention.k_proj.scale": "model.layers.25.self_attn.k_proj.weight_scale",
  "layers.25.attention.v_proj.weight": "model.layers.25.self_attn.v_proj.weight",
  "layers.25.attention.v_proj.scale": "model.layers.25.self_attn.v_proj.weight_scale",
  "layers.25.attention.o_proj.weight": "model.layers.25.self_attn.o_proj.weight",
  "layers.25.attention.o_proj.scale": "model.layers.25.self_attn.o_proj.weight_scale",
  "layers.25.mlp.gate.weight": "model.layers.25.mlp.gate_proj.weight",
  "layers.25.mlp.gate.scale": "model.layers.25.mlp.gate_proj.weight_scale",
  "layers.25.mlp.up.weight": "model.layers.25.mlp.up_proj.weight",
  "layers.25.mlp.up.scale": "model.layers.25.mlp.up_proj.weight_scale",
  "layers.25.mlp.down.weight": "model.layers.25.mlp.down_proj.weight",
  "layers.25.mlp.down.scale": "model.layers.25.mlp.down_proj.weight_scale",
  "layers.26.attention_norm.weight": "model.layers.26.input_layernorm.weight",
  "layers.26.mlp_norm.weight": "model.layers.26.post_attention_layernorm.weight",
  "layers.26.attention.q_proj.weight": "model.layers.26.self_attn.q_proj.weight",
  "layers.26.attention.q_proj.scale": "model.layers.26.self_attn.q_proj.weight_scale",
  "layers.26.attention.k_proj.weight": "model.layers.26.self_attn.k_proj.weight",
  "layers.26.attention.k_proj.scale": "model.layers.26.self_attn.k_proj.weight_scale",
  "layers.26.attention.v_proj.weight": "model.layers.26.self_attn.v_proj.weight",
  "layers.26.attention.v_proj.scale": "model.layers.26.self_attn.v_proj.weight_scale",
  "layers.26.attention.o_proj.weight": "model.layers.26.self_attn.o_proj.weight",
  "layers.26.attention.o_proj.scale": "model.layers.26.self_attn.o_proj.weight_scale",
  "layers.26.mlp.gate.weight": "model.layers.26.mlp.gate_proj.weight",
  "layers.26.mlp.gate.scale": "model.layers.26.mlp.gate_proj.weight_scale",
  "layers.26.mlp.up.weight": "model.layers.26.mlp.up_proj.weight",
  "layers.26.mlp.up.scale": "model.layers.26.mlp.up_proj.weight_scale",
  "layers.26.mlp.down.weight": "model.layers.26.mlp.down_proj.weight",
  "layers.26.mlp.down.scale": "model.layers.26.mlp.down_proj.weight_scale",
  "layers.27.attention_norm.weight": "model.layers.27.input_layernorm.weight",
  "layers.27.mlp_norm.weight": "model.layers.27.post_attention_layernorm.weight",
  "layers.27.attention.q_proj.weight": "model.layers.27.self_attn.q_proj.weight",
  "layers.27.attention.q_proj.scale": "model.layers.27.self_attn.q_proj.weight_scale",
  "layers.27.attention.k_proj.weight": "model.layers.27.self_attn.k_proj.weight",
  "layers.27.attention.k_proj.scale": "model.layers.27.self_attn.k_proj.weight_scale",
  "layers.27.attention.v_proj.weight": "model.layers.27.self_attn.v_proj.weight",
  "layers.27.attention.v_proj.scale": "model.layers.27.self_attn.v_proj.weight_scale",
  "layers.27.attention.o_proj.weight": "model.layers.27.self_attn.o_proj.weight",
  "layers.27.attention.o_proj.scale": "model.layers.27.self_attn.o_proj.weight_scale",
  "layers.27.mlp.gate.weight": "model.layers.27.mlp.gate_proj.weight",
  "layers.27.mlp.gate.scale": "model.layers.27.mlp.gate_proj.weight_scale",
  "layers.27.mlp.up.weight": "model.layers.27.mlp.up_proj.weight",
  "layers.27.mlp.up.scale": "model.layers.27.mlp.up_proj.weight_scale",
  "layers.27.mlp.down.weight": "model.layers.27.mlp.down_proj.weight",
  "layers.27.mlp.down.scale": "model.layers.27.mlp.down_proj.weight_scale",
  "layers.28.attention_norm.weight": "model.layers.28.input_layernorm.weight",
  "layers.28.mlp_norm.weight": "model.layers.28.post_attention_layernorm.weight",
  "layers.28.attention.q_proj.weight": "model.layers.28.self_attn.q_proj.weight",
  "layers.28.attention.q_proj.scale": "model.layers.28.self_attn.q_proj.weight_scale",
  "layers.28.attention.k_proj.weight": "model.layers.28.self_attn.k_proj.weight",
  "layers.28.attention.k_proj.scale": "model.layers.28.self_attn.k_proj.weight_scale",
  "layers.28.attention.v_proj.weight": "model.layers.28.self_attn.v_proj.weight",
  "layers.28.attention.v_proj.scale": "model.layers.28.self_attn.v_proj.weight_scale",
  "layers.28.attention.o_proj.weight": "model.layers.28.self_attn.o_proj.weight",
  "layers.28.attention.o_proj.scale": "model.layers.28.self_attn.o_proj.weight_scale",
  "layers.28.mlp.gate.weight": "model.layers.28.mlp.gate_proj.weight",
  "layers.28.mlp.gate.scale": "model.layers.28.mlp.gate_proj.weight_scale",
  "layers.28.mlp.up.weight": "model.layers.28.mlp.up_proj.weight",
  "layers.28.mlp.up.scale": "model.layers.28.mlp.up_proj.weight_scale",
  "layers.28.mlp.down.weight": "model.layers.28.mlp.down_proj.weight",
  "layers.28.mlp.down.scale": "model.layers.28.mlp.down_proj.weight_scale",
  "layers.29.attention_norm.weight": "model.layers.29.input_layernorm.weight",
  "layers.29.mlp_norm.weight": "model.layers.29.post_attention_layernorm.weight",
  "layers.29.attention.q_proj.weight": "model.layers.29.self_attn.q_proj.weight",
  "layers.29.attention.q_proj.scale": "model.layers.29.self_attn.q_proj.weight_scale",
  "layers.29.attention.k_proj.weight": "model.layers.29.self_attn.k_proj.weight",
  "layers.29.attention.k_proj.scale": "model.layers.29.self_attn.k_proj.weight_scale",
  "layers.29.attention.v_proj.weight": "model.layers.29.self_attn.v_proj.weight",
  "layers.29.attention.v_proj.scale": "model.layers.29.self_attn.v_proj.weight_scale",
  "layers.29.attention.o_proj.weight": "model.layers.29.self_attn.o_proj.weight",
  "layers.29.attention.o_proj.scale": "model.layers.29.self_attn.o_proj.weight_scale",
  "layers.29.mlp.gate.weight": "model.layers.29.mlp.gate_proj.weight",
  "layers.29.mlp.gate.scale": "model.layers.29.mlp.gate_proj.weight_scale",
  "layers.29.mlp.up.weight": "model.layers.29.mlp.up_proj.weight",
  "layers.29.mlp.up.scale": "model.layers.29.mlp.up_proj.weight_scale",
  "layers.29.mlp.down.weight": "model.layers.29.mlp.down_proj.weight",
  "layers.29.mlp.down.scale": "model.layers.29.mlp.down_proj.weight_scale",
  "layers.30.attention_norm.weight": "model.layers.30.input_layernorm.weight",
  "layers.30.mlp_norm.weight": "model.layers.30.post_attention_layernorm.weight",
  "layers.30.attention.q_proj.weight": "model.layers.30.self_attn.q_proj.weight",
  "layers.30.attention.q_proj.scale": "model.layers.30.self_attn.q_proj.weight_scale",
  "layers.30.attention.k_proj.weight": "model.layers.30.self_attn.k_proj.weight",
  "layers.30.attention.k_proj.scale": "model.layers.30.self_attn.k_proj.weight_scale",
  "layers.30.attention.v_proj.weight": "model.layers.30.self_attn.v_proj.weight",
  "layers.30.attention.v_proj.scale": "model.layers.30.self_attn.v_proj.weight_scale",
  "layers.30.attention.o_proj.weight": "model.layers.30.self_attn.o_proj.weight",
  "layers.30.attention.o_proj.scale": "model.layers.30.self_attn.o_proj.weight_scale",
  "layers.30.mlp.gate.weight": "model.layers.30.mlp.gate_proj.weight",
  "layers.30.mlp.gate.scale": "model.layers.30.mlp.gate_proj.weight_scale",
  "layers.30.mlp.up.weight": "model.layers.30.mlp.up_proj.weight",
  "layers.30.mlp.up.scale": "model.layers.30.mlp.up_proj.weight_scale",
  "layers.30.mlp.down.weight": "model.layers.30.mlp.down_proj.weight",
  "layers.30.mlp.down.scale": "model.layers.30.mlp.down_proj.weight_scale",
  "layers.31.attention_norm.weight": "model.layers.31.input_layernorm.weight",
  "layers.31.mlp_norm.weight": "model.layers.31.post_attention_layernorm.weight",
  "layers.31.attention.q_proj.weight": "model.layers.31.self_attn.q_proj.weight",
  "layers.31.attention.q_proj.scale": "model.layers.31.self_attn.q_proj.weight_scale",
  "layers.31.attention.k_proj.weight": "model.layers.31.self_attn.k_proj.weight",
  "layers.31.attention.k_proj.scale": "model.layers.31.self_attn.k_proj.weight_scale",
  "layers.31.attention.v_proj.weight": "model.layers.31.self_attn.v_proj.weight",
  "layers.31.attention.v_proj.scale": "model.layers.31.self_attn.v_proj.weight_scale",
  "layers.31.attention.o_proj.weight": "model.layers.31.self_attn.o_proj.weight",
  "layers.31.attention.o_proj.scale": "model.layers.31.self_attn.o_proj.weight_scale",
  "layers.31.mlp.gate.weight": "model.layers.31.mlp.gate_proj.weight",
  "layers.31.mlp.gate.scale": "model.layers.31.mlp.gate_proj.weight_scale",
  "layers.31.mlp.up.weight": "model.layers.31.mlp.up_proj.weight",
  "layers.31.mlp.up.scale": "model.layers.31.mlp.up_proj.weight_scale",
  "layers.31.mlp.down.weight": "model.layers.31.mlp.down_proj.weight",
  "layers.31.mlp.down.scale": "model.layers.31.mlp.down_proj.weight_scale"
}

Open on GitHub

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

Open on GitHub

nest.toml1.0 kB
toml
[model]
name = "llama-3.1-8b-instruct-fp8"
title = "Llama 3.1 8B Instruct FP8"
summary = "Llama 3.1 8B Instruct with every decoder projection in FP8 E4M3, one scale per output row: RedHatAI's FP8-dynamic checkpoint. Half the weight memory of bf16, faster single-sequence decoding."
license = "llama3.1"
family = "llama"
tags = ["text-generation", "decoder-only", "grouped-query-attention", "chat", "fp8"]

[links]
huggingface = "https://huggingface.co/RedHatAI/Meta-Llama-3.1-8B-Instruct-FP8-dynamic"
github = "https://github.com/meta-llama/llama-models"
arxiv = "https://arxiv.org/abs/2407.21783"

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

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

[check]
B = 1
S = 8

[weights]
repo = "RedHatAI/Meta-Llama-3.1-8B-Instruct-FP8-dynamic"
revision = "442e7f522277df12f53d63e8384087af31fbfc4b"
files = [
    "model-00001-of-00002.safetensors",
    "model-00002-of-00002.safetensors",
]
bindings = "bindings.json"

Open on GitHub

samples.json934 B
json
{
  "kind": "text",
  "prompt": "Explain in two sentences why the sky is blue.",
  "chat": true,
  "output": "The sky appears blue because of a phenomenon called Rayleigh scattering, where shorter (blue) wavelengths of light are scattered more than longer (red) wavelengths by the tiny molecules of gases in the Earth's atmosphere. As a result, when sunlight enters the atmosphere, the blue light is dispersed in all directions, giving the sky its blue color.",
  "reference": "The sky appears blue because of a phenomenon called Rayleigh scattering, where shorter (blue) wavelengths of light are scattered more than longer (red) wavelengths by the tiny molecules of gases in the Earth's atmosphere. As a result, when sunlight enters the atmosphere, the blue light is dispersed in all directions, giving the sky its blue color.",
  "reference_stack": "transformers (KV cache, argmax, bf16)",
  "tokens": 69,
  "agreeing_prefix": 69
}

Open on GitHub

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

Open on GitHub

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

// `std.nn.mlp::SwiGlu` over FP8 projections: the checkpoint quantizes every
// linear layer of the decoder but the output head.
pub block Fp8SwiGlu<H: Dim, Inner: Dim, T: Float = bf16> {
    sub gate: Fp8Linear<H, Inner, T>
    sub up: Fp8Linear<H, Inner, T>
    sub down: Fp8Linear<Inner, H, 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))
    }
}

Open on GitHub

rope.linnet3.5 kB
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.
//
// Llama 3.1 adds `rope_scaling` (`config.json`, type `llama3`) on top of
// Llama 3's base frequency: a piecewise rescaling of each frequency's
// wavelength that stretches the 8192-token pretraining context out to
// 131072 without retraining. Per frequency `i` (independent of position),
// with wavelength `wavelen[i] = 2 * pi / inv_freq[i]`:
//   - a wavelength longer than `original_max_position_embeddings /
//     low_freq_factor` is "low frequency": divide it by `factor`;
//   - a wavelength shorter than `original_max_position_embeddings /
//     high_freq_factor` is "high frequency": leave it alone;
//   - in between, linearly blend the unscaled and the divided-by-`factor`
//     frequency by how far `original_max_position_embeddings / wavelen[i]`
//     sits between `low_freq_factor` and `high_freq_factor`.
// The three bands are two nested `select`s on those comparisons, mirroring
// `transformers`' `torch.where` chain (`_compute_llama3_parameters`). This
// is ordinary elementwise scalar arithmetic on the `[D / 2]` frequency
// table, the same shape `inv_freq` already had, so it needed nothing beyond
// the arithmetic and `select` the rest of the standard library uses -- no
// index expression does arithmetic, since `i` only ever subscripts plainly.
module llama.rope

// Llama 3 stretched the base frequency from 10000 to 500000.
pub const THETA: f32 = 500000.0

// `rope_scaling` in `config.json`: `{"factor": 8.0, "low_freq_factor": 1.0,
// "high_freq_factor": 4.0, "original_max_position_embeddings": 8192,
// "rope_type": "llama3"}`.
pub const FACTOR: f32 = 8.0
pub const LOW_FREQ_FACTOR: f32 = 1.0
pub const HIGH_FREQ_FACTOR: f32 = 4.0
pub const ORIGINAL_MAX_POSITION_EMBEDDINGS: f32 = 8192.0
pub const PI: f32 = 3.14159265

// `(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 base_inv_freq[i] = exp(-(cast<f32>(iota<i64>(D / 2)[i]) * 2.0 / cast<f32>(D)) * log(theta))

    let wavelen[i] = (2.0 * PI) / base_inv_freq[i]
    let low_freq_wavelen = ORIGINAL_MAX_POSITION_EMBEDDINGS / LOW_FREQ_FACTOR
    let high_freq_wavelen = ORIGINAL_MAX_POSITION_EMBEDDINGS / HIGH_FREQ_FACTOR

    // The blend used in the medium band; only meaningful there, since the
    // `select` chain below only reads it when neither the low- nor the
    // high-frequency case applies. `smooth == 1` sits at the
    // `high_freq_wavelen` boundary (continuous with the unscaled band just
    // beyond it) and `smooth == 0` sits at the `low_freq_wavelen` boundary
    // (continuous with the divided-by-`FACTOR` band just beyond that), so
    // `(1 - smooth)` -- not `smooth` -- weights the scaled term.
    let smooth[i] =
        (ORIGINAL_MAX_POSITION_EMBEDDINGS / wavelen[i] - LOW_FREQ_FACTOR) /
        (HIGH_FREQ_FACTOR - LOW_FREQ_FACTOR)
    let smoothed[i] = (1.0 - smooth[i]) * (base_inv_freq[i] / FACTOR) + smooth[i] * base_inv_freq[i]

    let inv_freq[i] = select(
        wavelen[i] > low_freq_wavelen,
        base_inv_freq[i] / FACTOR,
        select(wavelen[i] < high_freq_wavelen, base_inv_freq[i], smoothed[i]),
    )

    let angles[s, i] = cast<f32>(iota<i64>(S)[s]) * inv_freq[i]
    let full = concat(angles, angles, axis = -1)
    return (cast<T>(cos(full)), cast<T>(sin(full)))
}

Open on GitHub

model-00001-of-00002.safetensorsHugging Face ↗

model-00002-of-00002.safetensorsHugging Face ↗