Models / gpt-oss-20b

GPT-OSS 20B

A 21B-parameter mixture-of-experts decoder, 32 experts per layer with 4 active per token, alternating 128-position sliding-window and full causal attention with learned attention sinks, YaRN rope scaling for a 131072-token context, and expert weights read straight from the checkpoint's MXFP4 blocks and scales.

20.9B parametersgpt_ossApache-2.0text-generationdecoder-onlymixture-of-expertsgrouped-query-attentionattention-sinkssliding-window-attentionmxfp4
models/gpt-oss-20b13 files · 106.1 kB in the registry · GitHub
README.md13.8 kB
md
# GPT-OSS 20B

A 21B-parameter mixture-of-experts decoder, of which 3.6B parameters are read
per token. 24 layers of width 2880: RMS normalization with epsilon 1e-5,
grouped-query attention with 64 query heads and 8 key/value heads of dimension
64, and a feed-forward block that is 32 experts wide with 4 chosen per token.
Every projection carries a bias. The output head is not tied to the embedding.

Three things separate it from a Llama-shaped decoder, and each is a module of
its own in `src/`:

- **Attention sinks** (`src/attention.linnet`). Each head owns one learned
  logit, `self_attn.sinks`, that joins its row of attention scores as an extra
  key, takes part in the softmax, and is then discarded. The probabilities over
  real keys therefore sum to less than one, and a head can decline to attend
  anywhere. That is a change to the softmax's normalization, not a mask or a
  bias, so none of the plain attention ops can express it;
  `std.nn.attention::sink_attention` does, by putting the sink's term in the
  denominator instead of concatenating a column and slicing it off again --
  the same numbers without a `K + 1` axis. PyTorch on CUDA runs it as
  FlexAttention, the sink folded in from its log-sum-exp.
- **MXFP4 expert weights** (`std.quant`). See below.
- **YaRN rope and alternating windows** (`src/rope.linnet`, `src/lib.linnet`).
  `layer_types` in config.json alternates a 128-position sliding window with
  full causal attention, starting with sliding. A `sub` array holds one block
  type, so the repeating unit here is a `LayerPair` of one windowed layer and
  one full one; the window is a runtime bound, and the full layer is given a
  bound as wide as the sequence, which is the same mask as no window.
  Parameter paths are `pairs.<i>.sliding.*` and `pairs.<i>.full.*`, which
  `bindings.json` maps to checkpoint layers `2i` and `2i + 1`.

## MXFP4

The official repository publishes the expert weights in MXFP4, not bf16:
`quantization_config.quant_method` is `mxfp4`, and the checkpoint carries
`mlp.experts.gate_up_proj_blocks` and `..._scales` as `u8` tensors rather than
a float `gate_up_proj`. Attention, the routers, the embedding, and the head
stay bf16 (`modules_to_not_convert`).

This card binds those bytes as what they are and dequantizes them in the
source. Nothing pretends a dtype matches:

| Checkpoint tensor | Shape | Dtype | Linnet parameter |
| --- | --- | --- | --- |
| `...experts.gate_up_proj_blocks` | `[32, 5760, 90, 16]` | `u8` | `mlp.gate_up_blocks: Tensor[Experts, 2 * Inner, H / 32, 16; u8]` |
| `...experts.gate_up_proj_scales` | `[32, 5760, 90]` | `u8` | `mlp.gate_up_scales: Tensor[Experts, 2 * Inner, H / 32; u8]` |
| `...experts.down_proj_blocks` | `[32, 2880, 90, 16]` | `u8` | `mlp.down_blocks: Tensor[Experts, H, Inner / 32, 16; u8]` |
| `...experts.down_proj_scales` | `[32, 2880, 90]` | `u8` | `mlp.down_scales: Tensor[Experts, H, Inner / 32; u8]` |

A row of `N` weights is cut into `N / 32` blocks of 32. Each block is 16 bytes
holding two FP4 (E2M1) values each -- low nibble the even element, high nibble
the odd one -- plus one byte of E8M0 scale, an exponent biased by 127. A
weight is its FP4 value times `2 ** (scale - 127)`: 4.25 bits per weight, and
`90 * 32 = 2880`, the width of the projection's input.

`std.quant`'s integer formats cannot express this. Its `unpack_int4` splits bytes into nibbles
with exactly this interleaving, but reads each nibble as a two's-complement
integer in `-8..7`, where MXFP4's nibble is a sign, two exponent bits, and one
mantissa bit selecting one of `0, 0.5, 1, 1.5, 2, 3, 4, 6` and their negatives.
And `dequantize_int8`'s scale is one factor per row (`Tensor[*S; f32]`), where
MXFP4 has one per 32-element block, as a power-of-two exponent byte.
`std.quant::dequantize_mxfp4` covers both differences (it began in this card
and moved to the standard library with the ops that read MXFP4 directly). It
computes the FP4 values arithmetically rather than from a lookup table,
because the language has no tensor data in source: twice each
value is an integer in `0..12`, so the unpacking stays in one byte per weight
and the only float arithmetic is the final scaling. The block scale is built
from integer shifts rather than `exp`, so it is exact; the shifts saturate at
`2 ** 62` either way, and every scale byte in this checkpoint's 96 expert
scale tensors lies in 115..136, that is `2 ** -12` to `2 ** 9`, so the
saturation is never reached. It returns `[Experts, Rows, N]`, the
checkpoint's own layout -- `Rows` the projection's output width -- and
`src/experts.linnet` reads that layout rather than transposing it, as the
reference implementation does to suit its own parameter.

The dequantization is arithmetic in the graph, so it runs on every forward
pass and a full-depth pass materializes each layer's 530M-element gate/up
weight in turn. That is the price of reading the published bytes instead of a
dequantized mirror. A backend with a fused MXFP4 kernel can select it against
this op's body; `linnet explain` says whether one did.

## Dense routing

The reference gathers the tokens each expert was chosen for and runs the
experts one at a time. That needs indices whose *values* decide which rows are
read, which the language has no shapes for. `MixtureOfExperts.forward` in
`src/experts.linnet`, the whole-sequence reference entry, instead evaluates
every expert on every token and weighs the results by the router's
probabilities, which are exactly zero outside the top four, so the sum is the
same and costs eight times the arithmetic. The entries that serve prompts and
decoded tokens route through library ops a backend can run sparsely (below).

Both expert projections go through `std.nn.linear::linear` with the weight
flattened rather than through a contraction with a free expert axis. That is
not cosmetic. `sum[h] x[b, s, h] * w[e, o, h]` is the natural way to write the
gate/up projection, but `[Experts, 2 * Inner, H]` reshapes to
`[Experts * 2 * Inner, H]` without moving an element, so the flattened form is
the same matrix product and a backend selects its kernel for it, where the
contraction form makes an interpreter materialize the whole
`[B, S, Experts, 2 * Inner, H]` outer product before reducing it -- 85 GiB at
160 positions. For the down projection the same trick needs one more step:
weighing each expert's activation by its router probability *before* the
projection rather than its output after folds the sum over experts into the
projection's own contraction, so that too is a single matrix product.

The top-`k` selection is written without a top-k primitive: count how many
experts beat each expert (ties to the lower index, as `torch.topk` orders
them), push everything ranked `TopK` or worse to `-1e30`, and softmax over all
32. The floored entries come out of the softmax as exactly zero, so the
surviving four probabilities are the reference's softmax over its four kept
logits.

## Loading

```python
from linnet import nest

model = nest.load("gpt-oss-20b", backend="torch", numerics="fast")
logits = model(tokens)                      # Tensor[1, S; i32] -> Tensor[1, S, 201088; bf16]
```

The weights are the published `model-0000N-of-00002.safetensors` shards, which
are numbered from zero, so there are three of them. `bindings.json` maps all
459 tensors to Linnet parameter paths; every name, shape, and dtype is checked
before anything runs.

## Entries

| Entry | |
| --- | --- |
| `forward<B, S>(tokens)` | logits for a whole sequence |
| `next_token<B, S>(tokens)` | logits for the last position |
| `states<B, S>(tokens)` | the last layer's hidden states, before the norm and the head |
| `prefill<S>(tokens, pos)` | a prompt into the KV caches, logits after its last token |
| `decode(token, pos)` | one token per row through the KV caches |
| `prefill_slots<M, S>(tokens, slots, lengths)` | `M` requests' prompts into rows `slots` of the caches, in one pass |
| `decode_rows(tokens, positions)` | one token per row, each row at its own position |
| `prefill_packed<P>(tokens, rows, positions, segments, last)` | prompts packed end to end into one pass of `P` tokens, each seeing its own prompt within the window |
| `step_packed<P>(tokens, rows, positions, segments, last, step_tokens, step_positions)` | `prefill_packed`'s pass with `decode_rows`'s step for every row in it |

The caches are `state` members of the attention blocks, `Batch` rows of
`MaxSeq` positions (the card binds 1 and 4096). A sliding-window layer still
caches every position and masks all but the last 128 of them; the full layers'
window is the whole cache. `decode_rows` and the packed entries are what
`linnet.serve` batches continuously (`prefill_slots` where prompts are not
packed, as with JAX and ONNX Runtime).

A prompt's tokens take the same route. `prefill` and `prefill_slots` go
through `MixtureOfExperts.forward_routed`, where every token is a row of its
own: `std.quant::mxfp4_linear_experts_shared` for the gate and up
projections, which read the token's one input, and
`std.quant::mxfp4_combine_experts` for the down projection with the router's
weights. The ops' bodies dequantize the experts and do the dense form's
work, which XLA and ONNX Runtime run as before (the dequantization once at
load); on a Hopper GPU, PyTorch runs each as one grouped matrix product over
the tokens sorted by expert, each expert multiplying only the tokens that
chose it, straight from MXFP4 with OpenAI's `triton_kernels` installed (see
Linnet's `docs/quantization.md`) and from bf16 copies of the experts
otherwise. On an H100 with `triton_kernels`, the first token of a 512-token
prompt takes 20.5 ms with CUDA graphs (25.0 from bf16 experts), a 64-row
serving step 8.3 ms (11.5), and serving 256 requests reaches 4702 tokens per
second.

A decoded token reads the experts differently again. `decode` routes
through `MixtureOfExperts.forward_topk`: only each token's four chosen
experts multiply it, an eighth of the dense form's work, and they are read
in MXFP4 as the checkpoint has them, through `std.quant::mxfp4_experts_shared`
(the gate and up projections, whose four experts share the token's input)
and `std.quant::mxfp4_experts` (the down projection, whose inputs differ per
expert). Their bodies dequantize, gather and multiply: XLA and ONNX Runtime
dequantize once at load and fuse the rest as before. PyTorch on CUDA runs a
Triton kernel that reads each chosen expert's four-bit bytes in place, a
quarter of what their 16-bit weights would be: on an H100 it decodes at 270
tokens per second with CUDA graphs, where reading the dequantized experts it
decoded at 202.

A server's batch of rows, `decode_rows`, routes through `forward_routed` as a
prompt does. In 64 rows nearly every expert is chosen by someone, so the
weights read are the same as the dense form's, but each expert multiplies
only the rows that chose it, an eighth of the arithmetic.

`sink_attention` groups the 64 query heads by the key/value head they share
(`[B, 8, 8, Q, D]`) and contracts them with the cached keys and values as
they are, rather than copying the keys and values out to every query head
first, which for a batch of 64 rows was eight copies of every cache on every
step. Together, serving 256 requests with 64 in flight went from 918 to 2129
tokens per second on an H100 (with CUDA graphs), the first tokens unchanged;
a single request decodes at 274 rather than 268, and XLA is unchanged.

With the op in the standard library, PyTorch's fast numerics run it as
FlexAttention (or, where FlexAttention does not fit, as bf16 products with
the softmax in f32) rather than as the f32 index notation: on one H100 a
512-token prompt takes 26.5 ms rather than 33.9, a decoding step 3.37 ms
rather than 3.85 (297 tokens per second), and serving goes from 1999 to
2327 tokens per second. JAX and ONNX run the same body as before.

`prefill_packed` packs the prompts end to end, each token seeing its own
prompt within the layer's window: one mask for the whole pass, which
FlexAttention takes, where `prefill_slots` pads each prompt and gives each
its own mask. A pass of 2048 prompt tokens then takes 62.7 ms rather than
148 as four padded prompts of 512, and serving reaches 3144 tokens per
second (with the engine capturing each pass size during warmup).

## Validation

Two layers of the real checkpoint -- layer 0 sliding, layer 1 full -- with all
32 experts, run in f32 against `transformers` 5.17 at 160 positions, so the
128-position window is narrower than the sequence and the two layers really do
mask differently:

| | |
| --- | --- |
| max abs logit difference | 1.88e-05 |
| mean abs logit difference | 1.47e-06 |
| argmax agreement | 160 / 160 positions |

`std.quant::dequantize_mxfp4` is separately bitwise equal to
`transformers.integrations.mxfp4.convert_moe_packed_tensors` over random
blocks and scales, for every one of the sixteen FP4 codes.

One trap for anyone repeating this. `transformers` dequantizes MXFP4 to
**bf16 regardless of the dtype asked for** -- `convert_moe_packed_tensors`
defaults to `torch.bfloat16`, so `dtype=torch.float32` still leaves
`experts.gate_up_proj` in bf16 while the rest of the model is f32. Comparing
against that gives 9.5e-2, which looks like an architecture bug and is not
one: the dequantized *values* are exact in bf16 (an FP4 value needs three
mantissa bits), so only the accumulation differs. Casting the reference's
expert parameters to f32 after loading is what produces the numbers above.

Not measured here, and left for a GPU run: the whole 24 layers in bf16 against
the same reference, and any timing. Both are memory- rather than
correctness-bound. The two-layer f32 run above peaked at 15.8 GiB of RSS on
the Linnet side, almost all of it the dequantized expert weights and the
interpreter's copies of them, which is why full depth is not a CPU job.

## Provenance

- Weights: [openai/gpt-oss-20b](https://huggingface.co/openai/gpt-oss-20b), Apache-2.0.
- Code: [openai/gpt-oss](https://github.com/openai/gpt-oss).
- Paper: [gpt-oss-120b & gpt-oss-20b Model Card](https://arxiv.org/abs/2508.10925).

Open on GitHub

bench.json11.7 kB
json
{
  "date": "2026-09-28T20:56:55+00:00",
  "environment": {
    "python": "3.12.3",
    "machine": "x86_64",
    "torch": "2.14.0",
    "jax": "0.11.2",
    "transformers": "5.17.0",
    "diffusers": "0.40.0",
    "onnxruntime": "1.30.0",
    "linnet": "0.1.0",
    "gpu": "NVIDIA H100 80GB HBM3",
    "gpus": "1",
    "driver": "580.126.09"
  },
  "workload": {
    "device": "cuda",
    "prompt_tokens": 512,
    "new_tokens": 128,
    "batch": 1,
    "warmup": 3,
    "iters": 10,
    "seed": 0,
    "offload_gib": 8.0,
    "serve_requests": 256,
    "serve_concurrency": 64,
    "serve_prompt_min": 128,
    "serve_prompt_max": 512,
    "serve_new": 128
  },
  "methods": [
    {
      "method": "transformers (eager)",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 41.88473429530859,
        "decode_tok_s": 44.89978494395782,
        "load_s": 114.69324987567961,
        "peak_vram_mib": 41386.0
      },
      "max_abs_diff": 0.0,
      "notes": "eager attention: no SDPA for this architecture",
      "error": null,
      "key": "transformers-eager"
    },
    {
      "method": "transformers (torch.compile, static cache)",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 25.567395612597466,
        "decode_tok_s": 100.31750366467506,
        "load_s": 125.57700556889176,
        "peak_vram_mib": 41286.0
      },
      "max_abs_diff": 0.1875,
      "notes": "eager attention: no SDPA for this architecture",
      "error": null,
      "key": "transformers-compile"
    },
    {
      "method": "vLLM",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 17.842570319771767,
        "decode_tok_s": 298.7910163802649,
        "load_s": 50.225351843982935,
        "first_token": 220.0,
        "peak_vram_mib": 70830.9
      },
      "max_abs_diff": null,
      "notes": "reserves a KV-cache pool up front (gpu_memory_utilization 0.85), so its memory is a setting, not a need",
      "error": null,
      "key": "vllm"
    },
    {
      "method": "Linnet torch (generated source)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 89.98703537508845,
        "decode_tok_s": 25.48134065162181,
        "load_s": 2.537807548418641,
        "peak_vram_mib": 54578.0
      },
      "max_abs_diff": 0.1953125,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-torch"
    },
    {
      "method": "Linnet torch (CUDA graphs)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 14.10851813852787,
        "decode_tok_s": 367.66697068942375,
        "load_s": 3.20823784917593,
        "peak_vram_mib": 29442.9
      },
      "max_abs_diff": 0.1875,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-cudagraphs"
    },
    {
      "method": "Linnet torch (torch.compile, inductor)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 34.81742553412914,
        "decode_tok_s": 199.93736969305772,
        "load_s": 3.114900317043066,
        "peak_vram_mib": 29392.9
      },
      "max_abs_diff": 0.1875,
      "notes": "KV cache compiled for 768 positions",
      "error": null,
      "key": "linnet-inductor"
    },
    {
      "method": "Linnet JAX (XLA, StableHLO)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 121.5835502371192,
        "decode_tok_s": 140.91474858291332,
        "load_s": 8.614564770832658,
        "peak_vram_mib": 15212.4,
        "driver_vram_mib": 33406.9
      },
      "max_abs_diff": 0.203125,
      "notes": "KV cache compiled for 768 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "linnet-jax"
    },
    {
      "method": "Linnet JAX (XLA, generated source)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 25.16063302755356,
        "decode_tok_s": 298.20299727781423,
        "load_s": 9.19578356295824,
        "peak_vram_mib": 51298.3,
        "driver_vram_mib": 66713.8
      },
      "max_abs_diff": 0.109375,
      "notes": "KV cache compiled for 768 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "linnet-jax-source"
    },
    {
      "method": "KerasHub (JAX)",
      "kind": "reference",
      "metrics": {
        "ttft_ms": 72.83677533268929,
        "decode_tok_s": 67.37037337686414,
        "load_s": 2104.777321504429,
        "peak_vram_mib": 45163.0,
        "driver_vram_mib": 61452.0
      },
      "max_abs_diff": 8.9140625,
      "notes": "peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "keras-hub"
    },
    {
      "method": "linnet-onnx-f32",
      "kind": "linnet",
      "metrics": {
        "peak_vram_mib": 79578.0
      },
      "max_abs_diff": null,
      "notes": "",
      "error": "RuntimeError: Error in execution: Non-zero status code returned while running Mul node. Name:'' Status Message: /onnxruntime_src/onnxruntime/core/framework/bfc_arena.cc:360 void* onnxruntime::BFCArena::AllocateRawInternal(size_t, bool, onnxruntime::Stream*) Failed to allocate memory for requested buffer of size 2123366400\n",
      "key": "linnet-onnx-f32"
    },
    {
      "method": "linnet-onnx-f32-trt",
      "kind": "linnet",
      "metrics": {},
      "max_abs_diff": null,
      "notes": "",
      "error": "exit -9: [5] Failed to import initializer: In node -1 with name:  and operator:  (parseGraph): UNSUPPORTED_NODE: Assertion failed: ctx->getWeightsContext().convertOnnxWeights(initializer, &weights): Failed to import initializer: v582\u001b[m",
      "key": "linnet-onnx-f32-trt"
    },
    {
      "method": "Linnet ONNX f16 -> ONNX Runtime (CUDA)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 107.29587823152542,
        "decode_tok_s": 21.419122280183412,
        "load_s": 0.23715465515851974,
        "peak_vram_mib": 70018.9
      },
      "max_abs_diff": 0.1923828125,
      "notes": "f16; ONNX Runtime's CUDA execution provider; first calls build the sessions; the argmax taken in the graph, each step replayed as a CUDA graph on the CUDA provider",
      "error": null,
      "key": "linnet-onnx-f16"
    },
    {
      "method": "linnet-onnx-f16-trt",
      "kind": "linnet",
      "metrics": {
        "peak_vram_mib": 70006.0
      },
      "max_abs_diff": null,
      "notes": "",
      "error": "Fail: [ONNXRuntimeError] : 1 : FAIL : TensorRT EP failed to create engine from network for fused node: TensorrtExecutionProvider_TRTKernel_graph_main_8376132635522534149_0_0",
      "key": "linnet-onnx-f16-trt"
    },
    {
      "method": "Linnet ONNX bf16 -> ONNX Runtime (CUDA)",
      "kind": "linnet",
      "metrics": {
        "ttft_ms": 112.95818444341421,
        "decode_tok_s": 21.219097580307174,
        "load_s": 0.2055520098656416,
        "peak_vram_mib": 69494.9
      },
      "max_abs_diff": 0.1875,
      "notes": "bf16; ONNX Runtime's CUDA execution provider; first calls build the sessions; the argmax taken in the graph, each step replayed as a CUDA graph on the CUDA provider",
      "error": null,
      "key": "linnet-onnx-bf16"
    },
    {
      "method": "linnet-onnx-bf16-trt",
      "kind": "linnet",
      "metrics": {
        "peak_vram_mib": 69482.0
      },
      "max_abs_diff": null,
      "notes": "",
      "error": "Fail: [ONNXRuntimeError] : 1 : FAIL : TensorRT EP failed to create engine from network for fused node: TensorrtExecutionProvider_TRTKernel_graph_main_11065519248658182765_0_0",
      "key": "linnet-onnx-bf16-trt"
    },
    {
      "method": "vLLM (offline, continuous batching)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 4313.237144994805,
        "serve_s": 7.59707822650671,
        "load_s": 79.9682690333575,
        "peak_vram_mib": 70152.8
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at gpu_memory_utilization 0.85",
      "error": null,
      "key": "serve-vllm"
    },
    {
      "method": "serve-transformers",
      "kind": "reference",
      "metrics": {
        "peak_vram_mib": 79444.0
      },
      "max_abs_diff": null,
      "notes": "",
      "error": "RuntimeError: generate_batch returned no tokens for any request",
      "key": "serve-transformers"
    },
    {
      "method": "serve-keras-hub",
      "kind": "reference",
      "metrics": {
        "peak_vram_mib": 61426.0
      },
      "max_abs_diff": null,
      "notes": "",
      "error": "JaxRuntimeError: NOT_FOUND: Failed to get configs for: 2 out of 100 instructions. See logs for all failures. Example failure: \nAll configs failed during profiling or were excluded from selection.\nFailures (19):\nEXECUTION FAILED: RESOURCE_EXHAUSTED: Out of memory while trying to allocate 14.08GiB with allocator GPU_0_bfc on device 0. [tf-allocator-allocation-error='']\nEXECUTION FAILED: RESOURCE_EXH",
      "key": "serve-keras-hub"
    },
    {
      "method": "Triton Inference Server (vLLM backend)",
      "kind": "reference",
      "metrics": {
        "serve_tok_s": 1701.2363790417264,
        "serve_s": 19.261285735294223,
        "load_s": 106.10278656892478,
        "serve_ttft_ms": 8388.68260756135,
        "peak_vram_mib": 69816.4
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed, first tokens timed at the client",
      "error": null,
      "key": "serve-triton-vllm"
    },
    {
      "method": "Linnet torch (linnet.serve, CUDA graphs)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 6030.689685359162,
        "serve_s": 5.433541055768728,
        "load_s": 7.278862033039331,
        "serve_ttft_ms": 2302.384303882718,
        "peak_vram_mib": 31330.2
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions",
      "error": null,
      "key": "serve-linnet-torch"
    },
    {
      "method": "Linnet jax (linnet.serve, XLA)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 3117.03775450011,
        "serve_s": 10.512545108795166,
        "load_s": 11.517576947808266,
        "serve_ttft_ms": 4653.521720319986,
        "peak_vram_mib": 53388.5,
        "driver_vram_mib": 67890.9
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
      "error": null,
      "key": "serve-linnet-jax"
    },
    {
      "method": "serve-linnet-onnx",
      "kind": "reference",
      "metrics": {
        "peak_vram_mib": 81028.0
      },
      "max_abs_diff": null,
      "notes": "",
      "error": "RuntimeError: Error in execution: Non-zero status code returned while running Einsum node. Name:'' Status Message: /onnxruntime_src/onnxruntime/core/framework/bfc_arena.cc:360 void* onnxruntime::BFCArena::AllocateRawInternal(size_t, bool, onnxruntime::Stream*) Failed to allocate memory for requested buffer of size 301989888\n",
      "key": "serve-linnet-onnx"
    },
    {
      "method": "Triton Inference Server (Python backend, linnet.serve)",
      "kind": "linnet",
      "metrics": {
        "serve_tok_s": 5839.18617468094,
        "serve_s": 5.611740920692682,
        "load_s": 705.3848326392472,
        "serve_ttft_ms": 2377.982411533594,
        "peak_vram_mib": 31914.4
      },
      "max_abs_diff": null,
      "notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed, first tokens timed at the client",
      "error": null,
      "key": "serve-triton-linnet"
    }
  ],
  "reference": "transformers-eager"
}

Open on GitHub

bindings.json37.3 kB
json
{
  "embed_tokens.weight": "model.embed_tokens.weight",
  "pairs.0.sliding.input_layernorm.weight": "model.layers.0.input_layernorm.weight",
  "pairs.0.sliding.self_attn.q_proj.weight": "model.layers.0.self_attn.q_proj.weight",
  "pairs.0.sliding.self_attn.q_proj.bias": "model.layers.0.self_attn.q_proj.bias",
  "pairs.0.sliding.self_attn.k_proj.weight": "model.layers.0.self_attn.k_proj.weight",
  "pairs.0.sliding.self_attn.k_proj.bias": "model.layers.0.self_attn.k_proj.bias",
  "pairs.0.sliding.self_attn.v_proj.weight": "model.layers.0.self_attn.v_proj.weight",
  "pairs.0.sliding.self_attn.v_proj.bias": "model.layers.0.self_attn.v_proj.bias",
  "pairs.0.sliding.self_attn.o_proj.weight": "model.layers.0.self_attn.o_proj.weight",
  "pairs.0.sliding.self_attn.o_proj.bias": "model.layers.0.self_attn.o_proj.bias",
  "pairs.0.sliding.self_attn.sinks": "model.layers.0.self_attn.sinks",
  "pairs.0.sliding.post_attention_layernorm.weight": "model.layers.0.post_attention_layernorm.weight",
  "pairs.0.sliding.mlp.router.weight": "model.layers.0.mlp.router.weight",
  "pairs.0.sliding.mlp.router.bias": "model.layers.0.mlp.router.bias",
  "pairs.0.sliding.mlp.gate_up_blocks": "model.layers.0.mlp.experts.gate_up_proj_blocks",
  "pairs.0.sliding.mlp.gate_up_scales": "model.layers.0.mlp.experts.gate_up_proj_scales",
  "pairs.0.sliding.mlp.gate_up_bias": "model.layers.0.mlp.experts.gate_up_proj_bias",
  "pairs.0.sliding.mlp.down_blocks": "model.layers.0.mlp.experts.down_proj_blocks",
  "pairs.0.sliding.mlp.down_scales": "model.layers.0.mlp.experts.down_proj_scales",
  "pairs.0.sliding.mlp.down_bias": "model.layers.0.mlp.experts.down_proj_bias",
  "pairs.0.full.input_layernorm.weight": "model.layers.1.input_layernorm.weight",
  "pairs.0.full.self_attn.q_proj.weight": "model.layers.1.self_attn.q_proj.weight",
  "pairs.0.full.self_attn.q_proj.bias": "model.layers.1.self_attn.q_proj.bias",
  "pairs.0.full.self_attn.k_proj.weight": "model.layers.1.self_attn.k_proj.weight",
  "pairs.0.full.self_attn.k_proj.bias": "model.layers.1.self_attn.k_proj.bias",
  "pairs.0.full.self_attn.v_proj.weight": "model.layers.1.self_attn.v_proj.weight",
  "pairs.0.full.self_attn.v_proj.bias": "model.layers.1.self_attn.v_proj.bias",
  "pairs.0.full.self_attn.o_proj.weight": "model.layers.1.self_attn.o_proj.weight",
  "pairs.0.full.self_attn.o_proj.bias": "model.layers.1.self_attn.o_proj.bias",
  "pairs.0.full.self_attn.sinks": "model.layers.1.self_attn.sinks",
  "pairs.0.full.post_attention_layernorm.weight": "model.layers.1.post_attention_layernorm.weight",
  "pairs.0.full.mlp.router.weight": "model.layers.1.mlp.router.weight",
  "pairs.0.full.mlp.router.bias": "model.layers.1.mlp.router.bias",
  "pairs.0.full.mlp.gate_up_blocks": "model.layers.1.mlp.experts.gate_up_proj_blocks",
  "pairs.0.full.mlp.gate_up_scales": "model.layers.1.mlp.experts.gate_up_proj_scales",
  "pairs.0.full.mlp.gate_up_bias": "model.layers.1.mlp.experts.gate_up_proj_bias",
  "pairs.0.full.mlp.down_blocks": "model.layers.1.mlp.experts.down_proj_blocks",
  "pairs.0.full.mlp.down_scales": "model.layers.1.mlp.experts.down_proj_scales",
  "pairs.0.full.mlp.down_bias": "model.layers.1.mlp.experts.down_proj_bias",
  "pairs.1.sliding.input_layernorm.weight": "model.layers.2.input_layernorm.weight",
  "pairs.1.sliding.self_attn.q_proj.weight": "model.layers.2.self_attn.q_proj.weight",
  "pairs.1.sliding.self_attn.q_proj.bias": "model.layers.2.self_attn.q_proj.bias",
  "pairs.1.sliding.self_attn.k_proj.weight": "model.layers.2.self_attn.k_proj.weight",
  "pairs.1.sliding.self_attn.k_proj.bias": "model.layers.2.self_attn.k_proj.bias",
  "pairs.1.sliding.self_attn.v_proj.weight": "model.layers.2.self_attn.v_proj.weight",
  "pairs.1.sliding.self_attn.v_proj.bias": "model.layers.2.self_attn.v_proj.bias",
  "pairs.1.sliding.self_attn.o_proj.weight": "model.layers.2.self_attn.o_proj.weight",
  "pairs.1.sliding.self_attn.o_proj.bias": "model.layers.2.self_attn.o_proj.bias",
  "pairs.1.sliding.self_attn.sinks": "model.layers.2.self_attn.sinks",
  "pairs.1.sliding.post_attention_layernorm.weight": "model.layers.2.post_attention_layernorm.weight",
  "pairs.1.sliding.mlp.router.weight": "model.layers.2.mlp.router.weight",
  "pairs.1.sliding.mlp.router.bias": "model.layers.2.mlp.router.bias",
  "pairs.1.sliding.mlp.gate_up_blocks": "model.layers.2.mlp.experts.gate_up_proj_blocks",
  "pairs.1.sliding.mlp.gate_up_scales": "model.layers.2.mlp.experts.gate_up_proj_scales",
  "pairs.1.sliding.mlp.gate_up_bias": "model.layers.2.mlp.experts.gate_up_proj_bias",
  "pairs.1.sliding.mlp.down_blocks": "model.layers.2.mlp.experts.down_proj_blocks",
  "pairs.1.sliding.mlp.down_scales": "model.layers.2.mlp.experts.down_proj_scales",
  "pairs.1.sliding.mlp.down_bias": "model.layers.2.mlp.experts.down_proj_bias",
  "pairs.1.full.input_layernorm.weight": "model.layers.3.input_layernorm.weight",
  "pairs.1.full.self_attn.q_proj.weight": "model.layers.3.self_attn.q_proj.weight",
  "pairs.1.full.self_attn.q_proj.bias": "model.layers.3.self_attn.q_proj.bias",
  "pairs.1.full.self_attn.k_proj.weight": "model.layers.3.self_attn.k_proj.weight",
  "pairs.1.full.self_attn.k_proj.bias": "model.layers.3.self_attn.k_proj.bias",
  "pairs.1.full.self_attn.v_proj.weight": "model.layers.3.self_attn.v_proj.weight",
  "pairs.1.full.self_attn.v_proj.bias": "model.layers.3.self_attn.v_proj.bias",
  "pairs.1.full.self_attn.o_proj.weight": "model.layers.3.self_attn.o_proj.weight",
  "pairs.1.full.self_attn.o_proj.bias": "model.layers.3.self_attn.o_proj.bias",
  "pairs.1.full.self_attn.sinks": "model.layers.3.self_attn.sinks",
  "pairs.1.full.post_attention_layernorm.weight": "model.layers.3.post_attention_layernorm.weight",
  "pairs.1.full.mlp.router.weight": "model.layers.3.mlp.router.weight",
  "pairs.1.full.mlp.router.bias": "model.layers.3.mlp.router.bias",
  "pairs.1.full.mlp.gate_up_blocks": "model.layers.3.mlp.experts.gate_up_proj_blocks",
  "pairs.1.full.mlp.gate_up_scales": "model.layers.3.mlp.experts.gate_up_proj_scales",
  "pairs.1.full.mlp.gate_up_bias": "model.layers.3.mlp.experts.gate_up_proj_bias",
  "pairs.1.full.mlp.down_blocks": "model.layers.3.mlp.experts.down_proj_blocks",
  "pairs.1.full.mlp.down_scales": "model.layers.3.mlp.experts.down_proj_scales",
  "pairs.1.full.mlp.down_bias": "model.layers.3.mlp.experts.down_proj_bias",
  "pairs.2.sliding.input_layernorm.weight": "model.layers.4.input_layernorm.weight",
  "pairs.2.sliding.self_attn.q_proj.weight": "model.layers.4.self_attn.q_proj.weight",
  "pairs.2.sliding.self_attn.q_proj.bias": "model.layers.4.self_attn.q_proj.bias",
  "pairs.2.sliding.self_attn.k_proj.weight": "model.layers.4.self_attn.k_proj.weight",
  "pairs.2.sliding.self_attn.k_proj.bias": "model.layers.4.self_attn.k_proj.bias",
  "pairs.2.sliding.self_attn.v_proj.weight": "model.layers.4.self_attn.v_proj.weight",
  "pairs.2.sliding.self_attn.v_proj.bias": "model.layers.4.self_attn.v_proj.bias",
  "pairs.2.sliding.self_attn.o_proj.weight": "model.layers.4.self_attn.o_proj.weight",
  "pairs.2.sliding.self_attn.o_proj.bias": "model.layers.4.self_attn.o_proj.bias",
  "pairs.2.sliding.self_attn.sinks": "model.layers.4.self_attn.sinks",
  "pairs.2.sliding.post_attention_layernorm.weight": "model.layers.4.post_attention_layernorm.weight",
  "pairs.2.sliding.mlp.router.weight": "model.layers.4.mlp.router.weight",
  "pairs.2.sliding.mlp.router.bias": "model.layers.4.mlp.router.bias",
  "pairs.2.sliding.mlp.gate_up_blocks": "model.layers.4.mlp.experts.gate_up_proj_blocks",
  "pairs.2.sliding.mlp.gate_up_scales": "model.layers.4.mlp.experts.gate_up_proj_scales",
  "pairs.2.sliding.mlp.gate_up_bias": "model.layers.4.mlp.experts.gate_up_proj_bias",
  "pairs.2.sliding.mlp.down_blocks": "model.layers.4.mlp.experts.down_proj_blocks",
  "pairs.2.sliding.mlp.down_scales": "model.layers.4.mlp.experts.down_proj_scales",
  "pairs.2.sliding.mlp.down_bias": "model.layers.4.mlp.experts.down_proj_bias",
  "pairs.2.full.input_layernorm.weight": "model.layers.5.input_layernorm.weight",
  "pairs.2.full.self_attn.q_proj.weight": "model.layers.5.self_attn.q_proj.weight",
  "pairs.2.full.self_attn.q_proj.bias": "model.layers.5.self_attn.q_proj.bias",
  "pairs.2.full.self_attn.k_proj.weight": "model.layers.5.self_attn.k_proj.weight",
  "pairs.2.full.self_attn.k_proj.bias": "model.layers.5.self_attn.k_proj.bias",
  "pairs.2.full.self_attn.v_proj.weight": "model.layers.5.self_attn.v_proj.weight",
  "pairs.2.full.self_attn.v_proj.bias": "model.layers.5.self_attn.v_proj.bias",
  "pairs.2.full.self_attn.o_proj.weight": "model.layers.5.self_attn.o_proj.weight",
  "pairs.2.full.self_attn.o_proj.bias": "model.layers.5.self_attn.o_proj.bias",
  "pairs.2.full.self_attn.sinks": "model.layers.5.self_attn.sinks",
  "pairs.2.full.post_attention_layernorm.weight": "model.layers.5.post_attention_layernorm.weight",
  "pairs.2.full.mlp.router.weight": "model.layers.5.mlp.router.weight",
  "pairs.2.full.mlp.router.bias": "model.layers.5.mlp.router.bias",
  "pairs.2.full.mlp.gate_up_blocks": "model.layers.5.mlp.experts.gate_up_proj_blocks",
  "pairs.2.full.mlp.gate_up_scales": "model.layers.5.mlp.experts.gate_up_proj_scales",
  "pairs.2.full.mlp.gate_up_bias": "model.layers.5.mlp.experts.gate_up_proj_bias",
  "pairs.2.full.mlp.down_blocks": "model.layers.5.mlp.experts.down_proj_blocks",
  "pairs.2.full.mlp.down_scales": "model.layers.5.mlp.experts.down_proj_scales",
  "pairs.2.full.mlp.down_bias": "model.layers.5.mlp.experts.down_proj_bias",
  "pairs.3.sliding.input_layernorm.weight": "model.layers.6.input_layernorm.weight",
  "pairs.3.sliding.self_attn.q_proj.weight": "model.layers.6.self_attn.q_proj.weight",
  "pairs.3.sliding.self_attn.q_proj.bias": "model.layers.6.self_attn.q_proj.bias",
  "pairs.3.sliding.self_attn.k_proj.weight": "model.layers.6.self_attn.k_proj.weight",
  "pairs.3.sliding.self_attn.k_proj.bias": "model.layers.6.self_attn.k_proj.bias",
  "pairs.3.sliding.self_attn.v_proj.weight": "model.layers.6.self_attn.v_proj.weight",
  "pairs.3.sliding.self_attn.v_proj.bias": "model.layers.6.self_attn.v_proj.bias",
  "pairs.3.sliding.self_attn.o_proj.weight": "model.layers.6.self_attn.o_proj.weight",
  "pairs.3.sliding.self_attn.o_proj.bias": "model.layers.6.self_attn.o_proj.bias",
  "pairs.3.sliding.self_attn.sinks": "model.layers.6.self_attn.sinks",
  "pairs.3.sliding.post_attention_layernorm.weight": "model.layers.6.post_attention_layernorm.weight",
  "pairs.3.sliding.mlp.router.weight": "model.layers.6.mlp.router.weight",
  "pairs.3.sliding.mlp.router.bias": "model.layers.6.mlp.router.bias",
  "pairs.3.sliding.mlp.gate_up_blocks": "model.layers.6.mlp.experts.gate_up_proj_blocks",
  "pairs.3.sliding.mlp.gate_up_scales": "model.layers.6.mlp.experts.gate_up_proj_scales",
  "pairs.3.sliding.mlp.gate_up_bias": "model.layers.6.mlp.experts.gate_up_proj_bias",
  "pairs.3.sliding.mlp.down_blocks": "model.layers.6.mlp.experts.down_proj_blocks",
  "pairs.3.sliding.mlp.down_scales": "model.layers.6.mlp.experts.down_proj_scales",
  "pairs.3.sliding.mlp.down_bias": "model.layers.6.mlp.experts.down_proj_bias",
  "pairs.3.full.input_layernorm.weight": "model.layers.7.input_layernorm.weight",
  "pairs.3.full.self_attn.q_proj.weight": "model.layers.7.self_attn.q_proj.weight",
  "pairs.3.full.self_attn.q_proj.bias": "model.layers.7.self_attn.q_proj.bias",
  "pairs.3.full.self_attn.k_proj.weight": "model.layers.7.self_attn.k_proj.weight",
  "pairs.3.full.self_attn.k_proj.bias": "model.layers.7.self_attn.k_proj.bias",
  "pairs.3.full.self_attn.v_proj.weight": "model.layers.7.self_attn.v_proj.weight",
  "pairs.3.full.self_attn.v_proj.bias": "model.layers.7.self_attn.v_proj.bias",
  "pairs.3.full.self_attn.o_proj.weight": "model.layers.7.self_attn.o_proj.weight",
  "pairs.3.full.self_attn.o_proj.bias": "model.layers.7.self_attn.o_proj.bias",
  "pairs.3.full.self_attn.sinks": "model.layers.7.self_attn.sinks",
  "pairs.3.full.post_attention_layernorm.weight": "model.layers.7.post_attention_layernorm.weight",
  "pairs.3.full.mlp.router.weight": "model.layers.7.mlp.router.weight",
  "pairs.3.full.mlp.router.bias": "model.layers.7.mlp.router.bias",
  "pairs.3.full.mlp.gate_up_blocks": "model.layers.7.mlp.experts.gate_up_proj_blocks",
  "pairs.3.full.mlp.gate_up_scales": "model.layers.7.mlp.experts.gate_up_proj_scales",
  "pairs.3.full.mlp.gate_up_bias": "model.layers.7.mlp.experts.gate_up_proj_bias",
  "pairs.3.full.mlp.down_blocks": "model.layers.7.mlp.experts.down_proj_blocks",
  "pairs.3.full.mlp.down_scales": "model.layers.7.mlp.experts.down_proj_scales",
  "pairs.3.full.mlp.down_bias": "model.layers.7.mlp.experts.down_proj_bias",
  "pairs.4.sliding.input_layernorm.weight": "model.layers.8.input_layernorm.weight",
  "pairs.4.sliding.self_attn.q_proj.weight": "model.layers.8.self_attn.q_proj.weight",
  "pairs.4.sliding.self_attn.q_proj.bias": "model.layers.8.self_attn.q_proj.bias",
  "pairs.4.sliding.self_attn.k_proj.weight": "model.layers.8.self_attn.k_proj.weight",
  "pairs.4.sliding.self_attn.k_proj.bias": "model.layers.8.self_attn.k_proj.bias",
  "pairs.4.sliding.self_attn.v_proj.weight": "model.layers.8.self_attn.v_proj.weight",
  "pairs.4.sliding.self_attn.v_proj.bias": "model.layers.8.self_attn.v_proj.bias",
  "pairs.4.sliding.self_attn.o_proj.weight": "model.layers.8.self_attn.o_proj.weight",
  "pairs.4.sliding.self_attn.o_proj.bias": "model.layers.8.self_attn.o_proj.bias",
  "pairs.4.sliding.self_attn.sinks": "model.layers.8.self_attn.sinks",
  "pairs.4.sliding.post_attention_layernorm.weight": "model.layers.8.post_attention_layernorm.weight",
  "pairs.4.sliding.mlp.router.weight": "model.layers.8.mlp.router.weight",
  "pairs.4.sliding.mlp.router.bias": "model.layers.8.mlp.router.bias",
  "pairs.4.sliding.mlp.gate_up_blocks": "model.layers.8.mlp.experts.gate_up_proj_blocks",
  "pairs.4.sliding.mlp.gate_up_scales": "model.layers.8.mlp.experts.gate_up_proj_scales",
  "pairs.4.sliding.mlp.gate_up_bias": "model.layers.8.mlp.experts.gate_up_proj_bias",
  "pairs.4.sliding.mlp.down_blocks": "model.layers.8.mlp.experts.down_proj_blocks",
  "pairs.4.sliding.mlp.down_scales": "model.layers.8.mlp.experts.down_proj_scales",
  "pairs.4.sliding.mlp.down_bias": "model.layers.8.mlp.experts.down_proj_bias",
  "pairs.4.full.input_layernorm.weight": "model.layers.9.input_layernorm.weight",
  "pairs.4.full.self_attn.q_proj.weight": "model.layers.9.self_attn.q_proj.weight",
  "pairs.4.full.self_attn.q_proj.bias": "model.layers.9.self_attn.q_proj.bias",
  "pairs.4.full.self_attn.k_proj.weight": "model.layers.9.self_attn.k_proj.weight",
  "pairs.4.full.self_attn.k_proj.bias": "model.layers.9.self_attn.k_proj.bias",
  "pairs.4.full.self_attn.v_proj.weight": "model.layers.9.self_attn.v_proj.weight",
  "pairs.4.full.self_attn.v_proj.bias": "model.layers.9.self_attn.v_proj.bias",
  "pairs.4.full.self_attn.o_proj.weight": "model.layers.9.self_attn.o_proj.weight",
  "pairs.4.full.self_attn.o_proj.bias": "model.layers.9.self_attn.o_proj.bias",
  "pairs.4.full.self_attn.sinks": "model.layers.9.self_attn.sinks",
  "pairs.4.full.post_attention_layernorm.weight": "model.layers.9.post_attention_layernorm.weight",
  "pairs.4.full.mlp.router.weight": "model.layers.9.mlp.router.weight",
  "pairs.4.full.mlp.router.bias": "model.layers.9.mlp.router.bias",
  "pairs.4.full.mlp.gate_up_blocks": "model.layers.9.mlp.experts.gate_up_proj_blocks",
  "pairs.4.full.mlp.gate_up_scales": "model.layers.9.mlp.experts.gate_up_proj_scales",
  "pairs.4.full.mlp.gate_up_bias": "model.layers.9.mlp.experts.gate_up_proj_bias",
  "pairs.4.full.mlp.down_blocks": "model.layers.9.mlp.experts.down_proj_blocks",
  "pairs.4.full.mlp.down_scales": "model.layers.9.mlp.experts.down_proj_scales",
  "pairs.4.full.mlp.down_bias": "model.layers.9.mlp.experts.down_proj_bias",
  "pairs.5.sliding.input_layernorm.weight": "model.layers.10.input_layernorm.weight",
  "pairs.5.sliding.self_attn.q_proj.weight": "model.layers.10.self_attn.q_proj.weight",
  "pairs.5.sliding.self_attn.q_proj.bias": "model.layers.10.self_attn.q_proj.bias",
  "pairs.5.sliding.self_attn.k_proj.weight": "model.layers.10.self_attn.k_proj.weight",
  "pairs.5.sliding.self_attn.k_proj.bias": "model.layers.10.self_attn.k_proj.bias",
  "pairs.5.sliding.self_attn.v_proj.weight": "model.layers.10.self_attn.v_proj.weight",
  "pairs.5.sliding.self_attn.v_proj.bias": "model.layers.10.self_attn.v_proj.bias",
  "pairs.5.sliding.self_attn.o_proj.weight": "model.layers.10.self_attn.o_proj.weight",
  "pairs.5.sliding.self_attn.o_proj.bias": "model.layers.10.self_attn.o_proj.bias",
  "pairs.5.sliding.self_attn.sinks": "model.layers.10.self_attn.sinks",
  "pairs.5.sliding.post_attention_layernorm.weight": "model.layers.10.post_attention_layernorm.weight",
  "pairs.5.sliding.mlp.router.weight": "model.layers.10.mlp.router.weight",
  "pairs.5.sliding.mlp.router.bias": "model.layers.10.mlp.router.bias",
  "pairs.5.sliding.mlp.gate_up_blocks": "model.layers.10.mlp.experts.gate_up_proj_blocks",
  "pairs.5.sliding.mlp.gate_up_scales": "model.layers.10.mlp.experts.gate_up_proj_scales",
  "pairs.5.sliding.mlp.gate_up_bias": "model.layers.10.mlp.experts.gate_up_proj_bias",
  "pairs.5.sliding.mlp.down_blocks": "model.layers.10.mlp.experts.down_proj_blocks",
  "pairs.5.sliding.mlp.down_scales": "model.layers.10.mlp.experts.down_proj_scales",
  "pairs.5.sliding.mlp.down_bias": "model.layers.10.mlp.experts.down_proj_bias",
  "pairs.5.full.input_layernorm.weight": "model.layers.11.input_layernorm.weight",
  "pairs.5.full.self_attn.q_proj.weight": "model.layers.11.self_attn.q_proj.weight",
  "pairs.5.full.self_attn.q_proj.bias": "model.layers.11.self_attn.q_proj.bias",
  "pairs.5.full.self_attn.k_proj.weight": "model.layers.11.self_attn.k_proj.weight",
  "pairs.5.full.self_attn.k_proj.bias": "model.layers.11.self_attn.k_proj.bias",
  "pairs.5.full.self_attn.v_proj.weight": "model.layers.11.self_attn.v_proj.weight",
  "pairs.5.full.self_attn.v_proj.bias": "model.layers.11.self_attn.v_proj.bias",
  "pairs.5.full.self_attn.o_proj.weight": "model.layers.11.self_attn.o_proj.weight",
  "pairs.5.full.self_attn.o_proj.bias": "model.layers.11.self_attn.o_proj.bias",
  "pairs.5.full.self_attn.sinks": "model.layers.11.self_attn.sinks",
  "pairs.5.full.post_attention_layernorm.weight": "model.layers.11.post_attention_layernorm.weight",
  "pairs.5.full.mlp.router.weight": "model.layers.11.mlp.router.weight",
  "pairs.5.full.mlp.router.bias": "model.layers.11.mlp.router.bias",
  "pairs.5.full.mlp.gate_up_blocks": "model.layers.11.mlp.experts.gate_up_proj_blocks",
  "pairs.5.full.mlp.gate_up_scales": "model.layers.11.mlp.experts.gate_up_proj_scales",
  "pairs.5.full.mlp.gate_up_bias": "model.layers.11.mlp.experts.gate_up_proj_bias",
  "pairs.5.full.mlp.down_blocks": "model.layers.11.mlp.experts.down_proj_blocks",
  "pairs.5.full.mlp.down_scales": "model.layers.11.mlp.experts.down_proj_scales",
  "pairs.5.full.mlp.down_bias": "model.layers.11.mlp.experts.down_proj_bias",
  "pairs.6.sliding.input_layernorm.weight": "model.layers.12.input_layernorm.weight",
  "pairs.6.sliding.self_attn.q_proj.weight": "model.layers.12.self_attn.q_proj.weight",
  "pairs.6.sliding.self_attn.q_proj.bias": "model.layers.12.self_attn.q_proj.bias",
  "pairs.6.sliding.self_attn.k_proj.weight": "model.layers.12.self_attn.k_proj.weight",
  "pairs.6.sliding.self_attn.k_proj.bias": "model.layers.12.self_attn.k_proj.bias",
  "pairs.6.sliding.self_attn.v_proj.weight": "model.layers.12.self_attn.v_proj.weight",
  "pairs.6.sliding.self_attn.v_proj.bias": "model.layers.12.self_attn.v_proj.bias",
  "pairs.6.sliding.self_attn.o_proj.weight": "model.layers.12.self_attn.o_proj.weight",
  "pairs.6.sliding.self_attn.o_proj.bias": "model.layers.12.self_attn.o_proj.bias",
  "pairs.6.sliding.self_attn.sinks": "model.layers.12.self_attn.sinks",
  "pairs.6.sliding.post_attention_layernorm.weight": "model.layers.12.post_attention_layernorm.weight",
  "pairs.6.sliding.mlp.router.weight": "model.layers.12.mlp.router.weight",
  "pairs.6.sliding.mlp.router.bias": "model.layers.12.mlp.router.bias",
  "pairs.6.sliding.mlp.gate_up_blocks": "model.layers.12.mlp.experts.gate_up_proj_blocks",
  "pairs.6.sliding.mlp.gate_up_scales": "model.layers.12.mlp.experts.gate_up_proj_scales",
  "pairs.6.sliding.mlp.gate_up_bias": "model.layers.12.mlp.experts.gate_up_proj_bias",
  "pairs.6.sliding.mlp.down_blocks": "model.layers.12.mlp.experts.down_proj_blocks",
  "pairs.6.sliding.mlp.down_scales": "model.layers.12.mlp.experts.down_proj_scales",
  "pairs.6.sliding.mlp.down_bias": "model.layers.12.mlp.experts.down_proj_bias",
  "pairs.6.full.input_layernorm.weight": "model.layers.13.input_layernorm.weight",
  "pairs.6.full.self_attn.q_proj.weight": "model.layers.13.self_attn.q_proj.weight",
  "pairs.6.full.self_attn.q_proj.bias": "model.layers.13.self_attn.q_proj.bias",
  "pairs.6.full.self_attn.k_proj.weight": "model.layers.13.self_attn.k_proj.weight",
  "pairs.6.full.self_attn.k_proj.bias": "model.layers.13.self_attn.k_proj.bias",
  "pairs.6.full.self_attn.v_proj.weight": "model.layers.13.self_attn.v_proj.weight",
  "pairs.6.full.self_attn.v_proj.bias": "model.layers.13.self_attn.v_proj.bias",
  "pairs.6.full.self_attn.o_proj.weight": "model.layers.13.self_attn.o_proj.weight",
  "pairs.6.full.self_attn.o_proj.bias": "model.layers.13.self_attn.o_proj.bias",
  "pairs.6.full.self_attn.sinks": "model.layers.13.self_attn.sinks",
  "pairs.6.full.post_attention_layernorm.weight": "model.layers.13.post_attention_layernorm.weight",
  "pairs.6.full.mlp.router.weight": "model.layers.13.mlp.router.weight",
  "pairs.6.full.mlp.router.bias": "model.layers.13.mlp.router.bias",
  "pairs.6.full.mlp.gate_up_blocks": "model.layers.13.mlp.experts.gate_up_proj_blocks",
  "pairs.6.full.mlp.gate_up_scales": "model.layers.13.mlp.experts.gate_up_proj_scales",
  "pairs.6.full.mlp.gate_up_bias": "model.layers.13.mlp.experts.gate_up_proj_bias",
  "pairs.6.full.mlp.down_blocks": "model.layers.13.mlp.experts.down_proj_blocks",
  "pairs.6.full.mlp.down_scales": "model.layers.13.mlp.experts.down_proj_scales",
  "pairs.6.full.mlp.down_bias": "model.layers.13.mlp.experts.down_proj_bias",
  "pairs.7.sliding.input_layernorm.weight": "model.layers.14.input_layernorm.weight",
  "pairs.7.sliding.self_attn.q_proj.weight": "model.layers.14.self_attn.q_proj.weight",
  "pairs.7.sliding.self_attn.q_proj.bias": "model.layers.14.self_attn.q_proj.bias",
  "pairs.7.sliding.self_attn.k_proj.weight": "model.layers.14.self_attn.k_proj.weight",
  "pairs.7.sliding.self_attn.k_proj.bias": "model.layers.14.self_attn.k_proj.bias",
  "pairs.7.sliding.self_attn.v_proj.weight": "model.layers.14.self_attn.v_proj.weight",
  "pairs.7.sliding.self_attn.v_proj.bias": "model.layers.14.self_attn.v_proj.bias",
  "pairs.7.sliding.self_attn.o_proj.weight": "model.layers.14.self_attn.o_proj.weight",
  "pairs.7.sliding.self_attn.o_proj.bias": "model.layers.14.self_attn.o_proj.bias",
  "pairs.7.sliding.self_attn.sinks": "model.layers.14.self_attn.sinks",
  "pairs.7.sliding.post_attention_layernorm.weight": "model.layers.14.post_attention_layernorm.weight",
  "pairs.7.sliding.mlp.router.weight": "model.layers.14.mlp.router.weight",
  "pairs.7.sliding.mlp.router.bias": "model.layers.14.mlp.router.bias",
  "pairs.7.sliding.mlp.gate_up_blocks": "model.layers.14.mlp.experts.gate_up_proj_blocks",
  "pairs.7.sliding.mlp.gate_up_scales": "model.layers.14.mlp.experts.gate_up_proj_scales",
  "pairs.7.sliding.mlp.gate_up_bias": "model.layers.14.mlp.experts.gate_up_proj_bias",
  "pairs.7.sliding.mlp.down_blocks": "model.layers.14.mlp.experts.down_proj_blocks",
  "pairs.7.sliding.mlp.down_scales": "model.layers.14.mlp.experts.down_proj_scales",
  "pairs.7.sliding.mlp.down_bias": "model.layers.14.mlp.experts.down_proj_bias",
  "pairs.7.full.input_layernorm.weight": "model.layers.15.input_layernorm.weight",
  "pairs.7.full.self_attn.q_proj.weight": "model.layers.15.self_attn.q_proj.weight",
  "pairs.7.full.self_attn.q_proj.bias": "model.layers.15.self_attn.q_proj.bias",
  "pairs.7.full.self_attn.k_proj.weight": "model.layers.15.self_attn.k_proj.weight",
  "pairs.7.full.self_attn.k_proj.bias": "model.layers.15.self_attn.k_proj.bias",
  "pairs.7.full.self_attn.v_proj.weight": "model.layers.15.self_attn.v_proj.weight",
  "pairs.7.full.self_attn.v_proj.bias": "model.layers.15.self_attn.v_proj.bias",
  "pairs.7.full.self_attn.o_proj.weight": "model.layers.15.self_attn.o_proj.weight",
  "pairs.7.full.self_attn.o_proj.bias": "model.layers.15.self_attn.o_proj.bias",
  "pairs.7.full.self_attn.sinks": "model.layers.15.self_attn.sinks",
  "pairs.7.full.post_attention_layernorm.weight": "model.layers.15.post_attention_layernorm.weight",
  "pairs.7.full.mlp.router.weight": "model.layers.15.mlp.router.weight",
  "pairs.7.full.mlp.router.bias": "model.layers.15.mlp.router.bias",
  "pairs.7.full.mlp.gate_up_blocks": "model.layers.15.mlp.experts.gate_up_proj_blocks",
  "pairs.7.full.mlp.gate_up_scales": "model.layers.15.mlp.experts.gate_up_proj_scales",
  "pairs.7.full.mlp.gate_up_bias": "model.layers.15.mlp.experts.gate_up_proj_bias",
  "pairs.7.full.mlp.down_blocks": "model.layers.15.mlp.experts.down_proj_blocks",
  "pairs.7.full.mlp.down_scales": "model.layers.15.mlp.experts.down_proj_scales",
  "pairs.7.full.mlp.down_bias": "model.layers.15.mlp.experts.down_proj_bias",
  "pairs.8.sliding.input_layernorm.weight": "model.layers.16.input_layernorm.weight",
  "pairs.8.sliding.self_attn.q_proj.weight": "model.layers.16.self_attn.q_proj.weight",
  "pairs.8.sliding.self_attn.q_proj.bias": "model.layers.16.self_attn.q_proj.bias",
  "pairs.8.sliding.self_attn.k_proj.weight": "model.layers.16.self_attn.k_proj.weight",
  "pairs.8.sliding.self_attn.k_proj.bias": "model.layers.16.self_attn.k_proj.bias",
  "pairs.8.sliding.self_attn.v_proj.weight": "model.layers.16.self_attn.v_proj.weight",
  "pairs.8.sliding.self_attn.v_proj.bias": "model.layers.16.self_attn.v_proj.bias",
  "pairs.8.sliding.self_attn.o_proj.weight": "model.layers.16.self_attn.o_proj.weight",
  "pairs.8.sliding.self_attn.o_proj.bias": "model.layers.16.self_attn.o_proj.bias",
  "pairs.8.sliding.self_attn.sinks": "model.layers.16.self_attn.sinks",
  "pairs.8.sliding.post_attention_layernorm.weight": "model.layers.16.post_attention_layernorm.weight",
  "pairs.8.sliding.mlp.router.weight": "model.layers.16.mlp.router.weight",
  "pairs.8.sliding.mlp.router.bias": "model.layers.16.mlp.router.bias",
  "pairs.8.sliding.mlp.gate_up_blocks": "model.layers.16.mlp.experts.gate_up_proj_blocks",
  "pairs.8.sliding.mlp.gate_up_scales": "model.layers.16.mlp.experts.gate_up_proj_scales",
  "pairs.8.sliding.mlp.gate_up_bias": "model.layers.16.mlp.experts.gate_up_proj_bias",
  "pairs.8.sliding.mlp.down_blocks": "model.layers.16.mlp.experts.down_proj_blocks",
  "pairs.8.sliding.mlp.down_scales": "model.layers.16.mlp.experts.down_proj_scales",
  "pairs.8.sliding.mlp.down_bias": "model.layers.16.mlp.experts.down_proj_bias",
  "pairs.8.full.input_layernorm.weight": "model.layers.17.input_layernorm.weight",
  "pairs.8.full.self_attn.q_proj.weight": "model.layers.17.self_attn.q_proj.weight",
  "pairs.8.full.self_attn.q_proj.bias": "model.layers.17.self_attn.q_proj.bias",
  "pairs.8.full.self_attn.k_proj.weight": "model.layers.17.self_attn.k_proj.weight",
  "pairs.8.full.self_attn.k_proj.bias": "model.layers.17.self_attn.k_proj.bias",
  "pairs.8.full.self_attn.v_proj.weight": "model.layers.17.self_attn.v_proj.weight",
  "pairs.8.full.self_attn.v_proj.bias": "model.layers.17.self_attn.v_proj.bias",
  "pairs.8.full.self_attn.o_proj.weight": "model.layers.17.self_attn.o_proj.weight",
  "pairs.8.full.self_attn.o_proj.bias": "model.layers.17.self_attn.o_proj.bias",
  "pairs.8.full.self_attn.sinks": "model.layers.17.self_attn.sinks",
  "pairs.8.full.post_attention_layernorm.weight": "model.layers.17.post_attention_layernorm.weight",
  "pairs.8.full.mlp.router.weight": "model.layers.17.mlp.router.weight",
  "pairs.8.full.mlp.router.bias": "model.layers.17.mlp.router.bias",
  "pairs.8.full.mlp.gate_up_blocks": "model.layers.17.mlp.experts.gate_up_proj_blocks",
  "pairs.8.full.mlp.gate_up_scales": "model.layers.17.mlp.experts.gate_up_proj_scales",
  "pairs.8.full.mlp.gate_up_bias": "model.layers.17.mlp.experts.gate_up_proj_bias",
  "pairs.8.full.mlp.down_blocks": "model.layers.17.mlp.experts.down_proj_blocks",
  "pairs.8.full.mlp.down_scales": "model.layers.17.mlp.experts.down_proj_scales",
  "pairs.8.full.mlp.down_bias": "model.layers.17.mlp.experts.down_proj_bias",
  "pairs.9.sliding.input_layernorm.weight": "model.layers.18.input_layernorm.weight",
  "pairs.9.sliding.self_attn.q_proj.weight": "model.layers.18.self_attn.q_proj.weight",
  "pairs.9.sliding.self_attn.q_proj.bias": "model.layers.18.self_attn.q_proj.bias",
  "pairs.9.sliding.self_attn.k_proj.weight": "model.layers.18.self_attn.k_proj.weight",
  "pairs.9.sliding.self_attn.k_proj.bias": "model.layers.18.self_attn.k_proj.bias",
  "pairs.9.sliding.self_attn.v_proj.weight": "model.layers.18.self_attn.v_proj.weight",
  "pairs.9.sliding.self_attn.v_proj.bias": "model.layers.18.self_attn.v_proj.bias",
  "pairs.9.sliding.self_attn.o_proj.weight": "model.layers.18.self_attn.o_proj.weight",
  "pairs.9.sliding.self_attn.o_proj.bias": "model.layers.18.self_attn.o_proj.bias",
  "pairs.9.sliding.self_attn.sinks": "model.layers.18.self_attn.sinks",
  "pairs.9.sliding.post_attention_layernorm.weight": "model.layers.18.post_attention_layernorm.weight",
  "pairs.9.sliding.mlp.router.weight": "model.layers.18.mlp.router.weight",
  "pairs.9.sliding.mlp.router.bias": "model.layers.18.mlp.router.bias",
  "pairs.9.sliding.mlp.gate_up_blocks": "model.layers.18.mlp.experts.gate_up_proj_blocks",
  "pairs.9.sliding.mlp.gate_up_scales": "model.layers.18.mlp.experts.gate_up_proj_scales",
  "pairs.9.sliding.mlp.gate_up_bias": "model.layers.18.mlp.experts.gate_up_proj_bias",
  "pairs.9.sliding.mlp.down_blocks": "model.layers.18.mlp.experts.down_proj_blocks",
  "pairs.9.sliding.mlp.down_scales": "model.layers.18.mlp.experts.down_proj_scales",
  "pairs.9.sliding.mlp.down_bias": "model.layers.18.mlp.experts.down_proj_bias",
  "pairs.9.full.input_layernorm.weight": "model.layers.19.input_layernorm.weight",
  "pairs.9.full.self_attn.q_proj.weight": "model.layers.19.self_attn.q_proj.weight",
  "pairs.9.full.self_attn.q_proj.bias": "model.layers.19.self_attn.q_proj.bias",
  "pairs.9.full.self_attn.k_proj.weight": "model.layers.19.self_attn.k_proj.weight",
  "pairs.9.full.self_attn.k_proj.bias": "model.layers.19.self_attn.k_proj.bias",
  "pairs.9.full.self_attn.v_proj.weight": "model.layers.19.self_attn.v_proj.weight",
  "pairs.9.full.self_attn.v_proj.bias": "model.layers.19.self_attn.v_proj.bias",
  "pairs.9.full.self_attn.o_proj.weight": "model.layers.19.self_attn.o_proj.weight",
  "pairs.9.full.self_attn.o_proj.bias": "model.layers.19.self_attn.o_proj.bias",
  "pairs.9.full.self_attn.sinks": "model.layers.19.self_attn.sinks",
  "pairs.9.full.post_attention_layernorm.weight": "model.layers.19.post_attention_layernorm.weight",
  "pairs.9.full.mlp.router.weight": "model.layers.19.mlp.router.weight",
  "pairs.9.full.mlp.router.bias": "model.layers.19.mlp.router.bias",
  "pairs.9.full.mlp.gate_up_blocks": "model.layers.19.mlp.experts.gate_up_proj_blocks",
  "pairs.9.full.mlp.gate_up_scales": "model.layers.19.mlp.experts.gate_up_proj_scales",
  "pairs.9.full.mlp.gate_up_bias": "model.layers.19.mlp.experts.gate_up_proj_bias",
  "pairs.9.full.mlp.down_blocks": "model.layers.19.mlp.experts.down_proj_blocks",
  "pairs.9.full.mlp.down_scales": "model.layers.19.mlp.experts.down_proj_scales",
  "pairs.9.full.mlp.down_bias": "model.layers.19.mlp.experts.down_proj_bias",
  "pairs.10.sliding.input_layernorm.weight": "model.layers.20.input_layernorm.weight",
  "pairs.10.sliding.self_attn.q_proj.weight": "model.layers.20.self_attn.q_proj.weight",
  "pairs.10.sliding.self_attn.q_proj.bias": "model.layers.20.self_attn.q_proj.bias",
  "pairs.10.sliding.self_attn.k_proj.weight": "model.layers.20.self_attn.k_proj.weight",
  "pairs.10.sliding.self_attn.k_proj.bias": "model.layers.20.self_attn.k_proj.bias",
  "pairs.10.sliding.self_attn.v_proj.weight": "model.layers.20.self_attn.v_proj.weight",
  "pairs.10.sliding.self_attn.v_proj.bias": "model.layers.20.self_attn.v_proj.bias",
  "pairs.10.sliding.self_attn.o_proj.weight": "model.layers.20.self_attn.o_proj.weight",
  "pairs.10.sliding.self_attn.o_proj.bias": "model.layers.20.self_attn.o_proj.bias",
  "pairs.10.sliding.self_attn.sinks": "model.layers.20.self_attn.sinks",
  "pairs.10.sliding.post_attention_layernorm.weight": "model.layers.20.post_attention_layernorm.weight",
  "pairs.10.sliding.mlp.router.weight": "model.layers.20.mlp.router.weight",
  "pairs.10.sliding.mlp.router.bias": "model.layers.20.mlp.router.bias",
  "pairs.10.sliding.mlp.gate_up_blocks": "model.layers.20.mlp.experts.gate_up_proj_blocks",
  "pairs.10.sliding.mlp.gate_up_scales": "model.layers.20.mlp.experts.gate_up_proj_scales",
  "pairs.10.sliding.mlp.gate_up_bias": "model.layers.20.mlp.experts.gate_up_proj_bias",
  "pairs.10.sliding.mlp.down_blocks": "model.layers.20.mlp.experts.down_proj_blocks",
  "pairs.10.sliding.mlp.down_scales": "model.layers.20.mlp.experts.down_proj_scales",
  "pairs.10.sliding.mlp.down_bias": "model.layers.20.mlp.experts.down_proj_bias",
  "pairs.10.full.input_layernorm.weight": "model.layers.21.input_layernorm.weight",
  "pairs.10.full.self_attn.q_proj.weight": "model.layers.21.self_attn.q_proj.weight",
  "pairs.10.full.self_attn.q_proj.bias": "model.layers.21.self_attn.q_proj.bias",
  "pairs.10.full.self_attn.k_proj.weight": "model.layers.21.self_attn.k_proj.weight",
  "pairs.10.full.self_attn.k_proj.bias": "model.layers.21.self_attn.k_proj.bias",
  "pairs.10.full.self_attn.v_proj.weight": "model.layers.21.self_attn.v_proj.weight",
  "pairs.10.full.self_attn.v_proj.bias": "model.layers.21.self_attn.v_proj.bias",
  "pairs.10.full.self_attn.o_proj.weight": "model.layers.21.self_attn.o_proj.weight",
  "pairs.10.full.self_attn.o_proj.bias": "model.layers.21.self_attn.o_proj.bias",
  "pairs.10.full.self_attn.sinks": "model.layers.21.self_attn.sinks",
  "pairs.10.full.post_attention_layernorm.weight": "model.layers.21.post_attention_layernorm.weight",
  "pairs.10.full.mlp.router.weight": "model.layers.21.mlp.router.weight",
  "pairs.10.full.mlp.router.bias": "model.layers.21.mlp.router.bias",
  "pairs.10.full.mlp.gate_up_blocks": "model.layers.21.mlp.experts.gate_up_proj_blocks",
  "pairs.10.full.mlp.gate_up_scales": "model.layers.21.mlp.experts.gate_up_proj_scales",
  "pairs.10.full.mlp.gate_up_bias": "model.layers.21.mlp.experts.gate_up_proj_bias",
  "pairs.10.full.mlp.down_blocks": "model.layers.21.mlp.experts.down_proj_blocks",
  "pairs.10.full.mlp.down_scales": "model.layers.21.mlp.experts.down_proj_scales",
  "pairs.10.full.mlp.down_bias": "model.layers.21.mlp.experts.down_proj_bias",
  "pairs.11.sliding.input_layernorm.weight": "model.layers.22.input_layernorm.weight",
  "pairs.11.sliding.self_attn.q_proj.weight": "model.layers.22.self_attn.q_proj.weight",
  "pairs.11.sliding.self_attn.q_proj.bias": "model.layers.22.self_attn.q_proj.bias",
  "pairs.11.sliding.self_attn.k_proj.weight": "model.layers.22.self_attn.k_proj.weight",
  "pairs.11.sliding.self_attn.k_proj.bias": "model.layers.22.self_attn.k_proj.bias",
  "pairs.11.sliding.self_attn.v_proj.weight": "model.layers.22.self_attn.v_proj.weight",
  "pairs.11.sliding.self_attn.v_proj.bias": "model.layers.22.self_attn.v_proj.bias",
  "pairs.11.sliding.self_attn.o_proj.weight": "model.layers.22.self_attn.o_proj.weight",
  "pairs.11.sliding.self_attn.o_proj.bias": "model.layers.22.self_attn.o_proj.bias",
  "pairs.11.sliding.self_attn.sinks": "model.layers.22.self_attn.sinks",
  "pairs.11.sliding.post_attention_layernorm.weight": "model.layers.22.post_attention_layernorm.weight",
  "pairs.11.sliding.mlp.router.weight": "model.layers.22.mlp.router.weight",
  "pairs.11.sliding.mlp.router.bias": "model.layers.22.mlp.router.bias",
  "pairs.11.sliding.mlp.gate_up_blocks": "model.layers.22.mlp.experts.gate_up_proj_blocks",
  "pairs.11.sliding.mlp.gate_up_scales": "model.layers.22.mlp.experts.gate_up_proj_scales",
  "pairs.11.sliding.mlp.gate_up_bias": "model.layers.22.mlp.experts.gate_up_proj_bias",
  "pairs.11.sliding.mlp.down_blocks": "model.layers.22.mlp.experts.down_proj_blocks",
  "pairs.11.sliding.mlp.down_scales": "model.layers.22.mlp.experts.down_proj_scales",
  "pairs.11.sliding.mlp.down_bias": "model.layers.22.mlp.experts.down_proj_bias",
  "pairs.11.full.input_layernorm.weight": "model.layers.23.input_layernorm.weight",
  "pairs.11.full.self_attn.q_proj.weight": "model.layers.23.self_attn.q_proj.weight",
  "pairs.11.full.self_attn.q_proj.bias": "model.layers.23.self_attn.q_proj.bias",
  "pairs.11.full.self_attn.k_proj.weight": "model.layers.23.self_attn.k_proj.weight",
  "pairs.11.full.self_attn.k_proj.bias": "model.layers.23.self_attn.k_proj.bias",
  "pairs.11.full.self_attn.v_proj.weight": "model.layers.23.self_attn.v_proj.weight",
  "pairs.11.full.self_attn.v_proj.bias": "model.layers.23.self_attn.v_proj.bias",
  "pairs.11.full.self_attn.o_proj.weight": "model.layers.23.self_attn.o_proj.weight",
  "pairs.11.full.self_attn.o_proj.bias": "model.layers.23.self_attn.o_proj.bias",
  "pairs.11.full.self_attn.sinks": "model.layers.23.self_attn.sinks",
  "pairs.11.full.post_attention_layernorm.weight": "model.layers.23.post_attention_layernorm.weight",
  "pairs.11.full.mlp.router.weight": "model.layers.23.mlp.router.weight",
  "pairs.11.full.mlp.router.bias": "model.layers.23.mlp.router.bias",
  "pairs.11.full.mlp.gate_up_blocks": "model.layers.23.mlp.experts.gate_up_proj_blocks",
  "pairs.11.full.mlp.gate_up_scales": "model.layers.23.mlp.experts.gate_up_proj_scales",
  "pairs.11.full.mlp.gate_up_bias": "model.layers.23.mlp.experts.gate_up_proj_bias",
  "pairs.11.full.mlp.down_blocks": "model.layers.23.mlp.experts.down_proj_blocks",
  "pairs.11.full.mlp.down_scales": "model.layers.23.mlp.experts.down_proj_scales",
  "pairs.11.full.mlp.down_bias": "model.layers.23.mlp.experts.down_proj_bias",
  "norm.weight": "model.norm.weight",
  "lm_head.weight": "lm_head.weight"
}

Open on GitHub

linnet.toml62 B
toml
[package]
name = "gpt_oss"
version = "0.1.0"
language = "0.1"

Open on GitHub

nest.toml1.5 kB
toml
[model]
# The experts are MXFP4, two 4-bit weights to a stored byte, so counting the
# checkpoint's elements would give 12.0B. Counted properly from the headers:
# every packed block byte is two parameters, the per-block scales are none.
parameters = 20914757184
name = "gpt-oss-20b"
title = "GPT-OSS 20B"
summary = "A 21B-parameter mixture-of-experts decoder, 32 experts per layer with 4 active per token, alternating 128-position sliding-window and full causal attention with learned attention sinks, YaRN rope scaling for a 131072-token context, and expert weights read straight from the checkpoint's MXFP4 blocks and scales."
license = "Apache-2.0"
family = "gpt_oss"
tags = [
    "text-generation",
    "decoder-only",
    "mixture-of-experts",
    "grouped-query-attention",
    "attention-sinks",
    "sliding-window-attention",
    "mxfp4",
]

[links]
huggingface = "https://huggingface.co/openai/gpt-oss-20b"
github = "https://github.com/openai/gpt-oss"
arxiv = "https://arxiv.org/abs/2508.10925"
homepage = "https://openai.com/index/introducing-gpt-oss/"

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

[generics]
Vocab = 201088
H = 2880
Heads = 64
KvHeads = 8
HeadDim = 64
Inner = 2880
Experts = 32
TopK = 4
Layers = 24
Window = 128
Batch = 1
MaxSeq = 4096
T = "bf16"

[check]
B = 1
S = 8

[weights]
repo = "openai/gpt-oss-20b"
revision = "6cee5e81ee83917806bbde320786a8fb61efebee"
files = [
    "model-00000-of-00002.safetensors",
    "model-00001-of-00002.safetensors",
    "model-00002-of-00002.safetensors",
]
bindings = "bindings.json"

Open on GitHub

samples.json1.2 kB
json
{
  "kind": "text",
  "prompt": "Explain in two sentences why the sky is blue.",
  "chat": true,
  "output": "analysisThe user asks: \"Explain in two sentences why the sky is blue.\" They want a concise explanation in exactly two sentences. So we need to produce two sentences. The explanation: Rayleigh scattering of sunlight by air molecules, shorter wavelengths scatter more, so blue dominates. Also mention that the sky appears blue because of scattering of shorter wavelengths. So produce two sentences. Ensure exactly two sentences. Let's craft: \"The sky appears blue because sunlight is scattered by the Earth's",
  "reference": "analysisThe user asks: \"Explain in two sentences why the sky is blue.\" They want a concise explanation in two sentences. So we need to produce exactly two sentences. The explanation: Rayleigh scattering of sunlight by air molecules, shorter wavelengths scatter more, making sky appear blue. Also mention that the sun's light is white, but blue is scattered. So produce two sentences. Ensure exactly two sentences. Let's craft: \"The sky appears blue because sunlight, which contains",
  "reference_stack": "transformers (KV cache, argmax, bf16)",
  "tokens": 96,
  "agreeing_prefix": 24
}

Open on GitHub

src
attention.linnet13.2 kB
linnet
// Attention with sinks, the part of gpt-oss no existing attention op can
// stand in for.
//
// A sink is one learned logit per query head. It joins the row of attention
// scores as an extra key, takes part in the softmax, and is then dropped: the
// probabilities over real keys sum to less than one, and the head can decline
// to attend anywhere by putting its mass on the sink. That changes the
// normalization, so it is not a mask, a bias, or anything
// `std.nn.attention::attention` can be handed -- hence
// `std.nn.attention::sink_attention`, which is otherwise the same scaled
// dot-product attention with grouped key/value heads.
//
// The layers alternate between a 128-position sliding window and full causal
// attention (`layer_types` in config.json). `window_mask` covers both: the
// window is a runtime bound, and a bound at least as large as the key length
// leaves plain causal attention.
//
// `decode`, `prefill`, `decode_rows`, and `prefill_slots` keep the keys and
// values in two `state` caches of `MaxSeq` positions per row, as the other
// decoders in the zoo do; a sliding-window layer still caches every position
// and masks all but the last `window` of them.
module gpt_oss.attention

use crate.rope::{tables}
use std.nn.attention::{sink_attention}
use std.nn.cache::{write_at, write_rows, write_slots, write_span, write_tokens}
use std.nn.linear::{Linear}
use std.nn.rope::{rope, rope_rows}

// Keys at or before each query position and no further back than `window`
// positions, counting the query itself. `K - Q` aligns the query block with
// the end of the keys, as `std.nn.attention::causal_mask` does.
pub fn window_mask<Q: Dim, K: Dim>(window: i32) -> Tensor[Q, K; bool] {
    let queries = iota<i32>(Q)
    let keys = iota<i32>(K)
    let offset = cast<i32>(K - Q)
    let mask[q, k] = keys[k] <= queries[q] + offset && keys[k] > queries[q] + offset - window
    return mask
}

// gpt-oss puts a bias on all four projections (`attention_bias` is true) and
// its head dimension is 64, which `H / Heads` does not give: 2880 / 64 is 45.
// `HeadDim` is therefore its own generic rather than derived.
pub block SinkAttention<
    H: Dim,
    Heads: Dim,
    KvHeads: Dim,
    HeadDim: Dim,
    Batch: Dim,
    MaxSeq: Dim,
    T: Float,
>
where
    Heads > 0,
    KvHeads > 0,
    Heads % KvHeads == 0,
    HeadDim % 2 == 0,
    MaxSeq > 0
{
    sub q_proj: Linear<H, Heads * HeadDim, T>
    sub k_proj: Linear<H, KvHeads * HeadDim, T>
    sub v_proj: Linear<H, KvHeads * HeadDim, T>
    sub o_proj: Linear<Heads * HeadDim, H, T>

    param sinks: Tensor[Heads; T]

    state cache_k: Tensor[Batch, KvHeads, MaxSeq, HeadDim; T]
    state cache_v: Tensor[Batch, KvHeads, MaxSeq, HeadDim; T]

    pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T], window: i32) -> Tensor[B, S, H; T] {
        let (cos_table, sin_table) = tables<S, HeadDim, T>()
        let q = rope(
            split_heads<B, S, Heads, HeadDim, T>(q_proj.forward(x)),
            cos_table,
            sin_table,
        )
        let k = rope(
            split_heads<B, S, KvHeads, HeadDim, T>(k_proj.forward(x)),
            cos_table,
            sin_table,
        )
        let v = split_heads<B, S, KvHeads, HeadDim, T>(v_proj.forward(x))
        let mixed = sink_attention(
            q,
            k,
            v,
            sinks,
            rsqrt(cast<f32>(HeadDim)),
            broadcast_to(reshape(window_mask<S, S>(window), [1, S, S]), [B, S, S]),
        )
        return o_proj.forward(merge_heads<B, S, Heads, HeadDim, T>(mixed))
    }

    // One new position `pos` for every row: its key and value go into the
    // caches there, and the query attends over the cached positions in its
    // window.
    pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32, window: i32) -> Tensor[Batch, 1, H; T] {
        let (cos_table, sin_table) = tables<MaxSeq, HeadDim, T>()
        let cos_at[d] = cos_table[cast<i64>(pos), d]
        let sin_at[d] = sin_table[cast<i64>(pos), d]
        let cos_row = reshape(cos_at, [1, HeadDim])
        let sin_row = reshape(sin_at, [1, HeadDim])
        let q = rope(split_heads<Batch, 1, Heads, HeadDim, T>(q_proj.forward(x)), cos_row, sin_row)
        let k = rope(
            split_heads<Batch, 1, KvHeads, HeadDim, T>(k_proj.forward(x)),
            cos_row,
            sin_row,
        )
        let v = split_heads<Batch, 1, KvHeads, HeadDim, T>(v_proj.forward(x))
        cache_k = write_at(cache_k, k, pos)
        cache_v = write_at(cache_v, v, pos)
        let slots = iota<i32>(MaxSeq)
        let seen[s] = slots[s] <= pos && slots[s] > pos - window
        let mixed = sink_attention(
            q,
            cache_k,
            cache_v,
            sinks,
            rsqrt(cast<f32>(HeadDim)),
            broadcast_to(reshape(seen, [1, 1, MaxSeq]), [Batch, 1, MaxSeq]),
        )
        return o_proj.forward(merge_heads<Batch, 1, Heads, HeadDim, T>(mixed))
    }

    // `S` new positions from `pos`, as a prompt arrives: written into the
    // caches as one span, each query attending over the cached positions in
    // its window.
    pub fn prefill<S: Dim>(
        x: Tensor[Batch, S, H; T],
        pos: i32,
        window: i32,
    ) -> Tensor[Batch, S, H; T]
    where S > 0 {
        let (cos_table, sin_table) = tables<MaxSeq, HeadDim, T>()
        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, HeadDim, T>(q_proj.forward(x)),
            cos_rows,
            sin_rows,
        )
        let k = rope(
            split_heads<Batch, S, KvHeads, HeadDim, T>(k_proj.forward(x)),
            cos_rows,
            sin_rows,
        )
        let v = split_heads<Batch, S, KvHeads, HeadDim, T>(v_proj.forward(x))
        cache_k = write_span(cache_k, k, pos)
        cache_v = write_span(cache_v, v, pos)
        let slots = iota<i32>(MaxSeq)
        let queries = iota<i32>(S)
        let seen[i, s] = slots[s] <= pos + queries[i] && slots[s] > pos + queries[i] - window
        let mixed = sink_attention(
            q,
            cache_k,
            cache_v,
            sinks,
            rsqrt(cast<f32>(HeadDim)),
            broadcast_to(reshape(seen, [1, S, MaxSeq]), [Batch, S, MaxSeq]),
        )
        return o_proj.forward(merge_heads<Batch, S, Heads, HeadDim, T>(mixed))
    }

    // One new position per row, each at its own `positions[b]`: the step a
    // server takes for every request it is decoding together.
    pub fn decode_rows(
        x: Tensor[Batch, 1, H; T],
        positions: Tensor[Batch; i32],
        window: i32,
    ) -> Tensor[Batch, 1, H; T] {
        let (cos_table, sin_table) = tables<MaxSeq, HeadDim, T>()
        let cos_at[b, d] = cos_table[cast<i64>(positions[b]), d]
        let sin_at[b, d] = sin_table[cast<i64>(positions[b]), d]
        let cos_rows = reshape(cos_at, [Batch, 1, HeadDim])
        let sin_rows = reshape(sin_at, [Batch, 1, HeadDim])
        let q = rope_rows(
            split_heads<Batch, 1, Heads, HeadDim, T>(q_proj.forward(x)),
            cos_rows,
            sin_rows,
        )
        let k = rope_rows(
            split_heads<Batch, 1, KvHeads, HeadDim, T>(k_proj.forward(x)),
            cos_rows,
            sin_rows,
        )
        let v = split_heads<Batch, 1, KvHeads, HeadDim, T>(v_proj.forward(x))
        cache_k = write_rows(cache_k, k, positions)
        cache_v = write_rows(cache_v, v, positions)
        let slots = iota<i32>(MaxSeq)
        let seen[b, s] = slots[s] <= positions[b] && slots[s] > positions[b] - window
        let mixed = sink_attention(
            q,
            cache_k,
            cache_v,
            sinks,
            rsqrt(cast<f32>(HeadDim)),
            reshape(seen, [Batch, 1, MaxSeq]),
        )
        return o_proj.forward(merge_heads<Batch, 1, Heads, HeadDim, T>(mixed))
    }

    // `M` requests' prompts of `S` tokens, from their first positions, into
    // rows `slots` of the caches; each prompt's queries attend among
    // themselves, within the window.
    pub fn prefill_slots<M: Dim, S: Dim>(
        x: Tensor[M, S, H; T],
        slots: Tensor[M; i32],
        window: i32,
    ) -> Tensor[M, S, H; T]
    where
        M > 0,
        S > 0
    {
        let (cos_table, sin_table) = tables<S, HeadDim, T>()
        let q = rope(split_heads<M, S, Heads, HeadDim, T>(q_proj.forward(x)), cos_table, sin_table)
        let k = rope(
            split_heads<M, S, KvHeads, HeadDim, T>(k_proj.forward(x)),
            cos_table,
            sin_table,
        )
        let v = split_heads<M, S, KvHeads, HeadDim, T>(v_proj.forward(x))
        cache_k = write_slots(cache_k, k, slots, 0)
        cache_v = write_slots(cache_v, v, slots, 0)
        let mixed = sink_attention(
            q,
            k,
            v,
            sinks,
            rsqrt(cast<f32>(HeadDim)),
            broadcast_to(reshape(window_mask<S, S>(window), [1, S, S]), [M, S, S]),
        )
        return o_proj.forward(merge_heads<M, S, Heads, HeadDim, T>(mixed))
    }

    // Prompts packed end to end, `P` tokens in all: token `p` is position
    // `positions[p]` of prompt `segments[p]`, whose row of the caches is
    // `rows[p]`. Each token's key and value go to its row at its position,
    // and its query attends to the tokens of its own prompt up to itself,
    // within the window.
    pub fn prefill_packed<P: Dim>(
        x: Tensor[1, P, H; T],
        rows: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
        window: i32,
    ) -> Tensor[1, P, H; T]
    where P > 0 {
        let (cos_table, sin_table) = tables<MaxSeq, HeadDim, T>()
        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, HeadDim, T>(q_proj.forward(x)), cos_at, sin_at)
        let k = rope(split_heads<1, P, KvHeads, HeadDim, T>(k_proj.forward(x)), cos_at, sin_at)
        let v = split_heads<1, P, KvHeads, HeadDim, T>(v_proj.forward(x))
        cache_k = write_tokens(cache_k, k, rows, positions)
        cache_v = write_tokens(cache_v, v, rows, positions)
        let own[i, j] =
            segments[i] == segments[j] &&
            positions[j] <= positions[i] &&
            positions[j] > positions[i] - window
        let mixed = sink_attention(
            q,
            k,
            v,
            sinks,
            rsqrt(cast<f32>(HeadDim)),
            reshape(own, [1, P, P]),
        )
        return o_proj.forward(merge_heads<1, P, Heads, HeadDim, T>(mixed))
    }

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

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

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

Open on GitHub

experts.linnet10.3 kB
linnet
// The mixture-of-experts feed-forward block: a router that picks four of
// thirty-two experts per token, and the experts themselves, whose weights the
// checkpoint stores in MXFP4.
//
// The reference gathers the tokens each expert was chosen for and runs the
// experts one at a time. That needs indices whose values decide which rows
// are read, and the language has no data-dependent shapes. The equivalent
// written here evaluates every expert on every token and weighs the results by
// the router's probabilities, which are exactly zero outside the top four, so
// the sum is the same. It costs eight times the arithmetic of the sparse form
// and is the honest way to say it with fixed shapes.
module gpt_oss.experts

use std.nn.activations::{sigmoid}
use std.nn.linear::{linear, Linear}
use std.quant::{
    dequantize_mxfp4,
    mxfp4_combine_experts,
    mxfp4_experts,
    mxfp4_experts_shared,
    mxfp4_linear_experts_shared,
}
use std.nn.softmax::{softmax}

// `swiglu_limit` in config.json, and the gate's sigmoid slope.
pub const LIMIT: f32 = 7.0
pub const ALPHA: f32 = 1.702

pub block MixtureOfExperts<
    H: Dim,
    Inner: Dim,
    Experts: Dim,
    TopK: Dim,
    Blocks: Dim,
    InnerBlocks: Dim,
    T: Float,
>
where
    Experts > 0,
    TopK > 0,
    H == Blocks * 32,
    Inner == InnerBlocks * 32
{
    sub router: Linear<H, Experts, T>

    // The gate and the up projection share one weight, interleaved along the
    // output axis: even columns gate, odd columns multiply. Both halves are
    // `Inner` wide, so the stored width is `2 * Inner`.
    param gate_up_blocks: Tensor[Experts, 2 * Inner, Blocks, 16; u8]
    param gate_up_scales: Tensor[Experts, 2 * Inner, Blocks; u8]
    param gate_up_bias: Tensor[Experts, 2 * Inner; T]

    param down_blocks: Tensor[Experts, H, InnerBlocks, 16; u8]
    param down_scales: Tensor[Experts, H, InnerBlocks; u8]
    param down_bias: Tensor[Experts, H; T]

    pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T]) -> Tensor[B, S, H; T] {
        let scores = route<B, S>(router.forward(x))
        let gate_up_weight = dequantize_mxfp4<Experts, 2 * Inner, Blocks, T>(
            gate_up_blocks,
            gate_up_scales,
        )
        // One linear layer over every expert at once. `[Experts, 2 * Inner, H]`
        // flattens to `[Experts * 2 * Inner, H]` without moving anything, so
        // this is an ordinary matrix product and a backend selects its kernel
        // for it. Written instead as a contraction with a free `Experts` axis,
        // `sum[h] x[b, s, h] * w[e, o, h]`, an interpreter has to materialize
        // the whole `[B, S, Experts, 2 * Inner, H]` outer product before
        // reducing it -- 85 GiB at 160 positions.
        let wide = linear(
            x,
            reshape(gate_up_weight, [Experts * 2 * Inner, H]),
            some(reshape(gate_up_bias, [Experts * 2 * Inner])),
        )
        let gate_up = reshape(wide, [B, S, Experts, 2 * Inner])
        // gpt-oss clamps the gate from above only and the multiplier from both
        // sides, then adds one to the multiplier, so the branch contributes
        // `glu` even where the up projection is zero. `swiglu_limit` in
        // config.json is the bound; 1.702 is the gate's sigmoid slope.
        let gate = min(gate_up[..., 0::2], cast<T>(LIMIT))
        let up = min(max(gate_up[..., 1::2], cast<T>(-LIMIT)), cast<T>(LIMIT))
        let activated = (up + 1.0) * (gate * sigmoid(gate * cast<T>(ALPHA)))
        // Weighing each expert's activation by its router probability here,
        // before the down projection, rather than weighing its output after,
        // folds the sum over experts into the projection's own contraction:
        // one matrix product again rather than a per-expert one.
        let weighted[b, s, e, i] = activated[b, s, e, i] * scores[b, s, e]
        let down_weight = dequantize_mxfp4<Experts, H, InnerBlocks, T>(down_blocks, down_scales)
        let mixed = linear(
            reshape(weighted, [B, S, Experts * Inner]),
            reshape(permute(down_weight, [1, 0, 2]), [H, Experts * Inner]),
        )
        // The down bias carries the same router weights, which makes summing
        // it over the experts a matrix product too.
        return mixed + linear(scores, permute(down_bias, [1, 0]))
    }

    // One token per row through its `TopK` experts only, `TopK / Experts` of
    // the dense form's work (an eighth here), the experts read in MXFP4 as
    // the checkpoint has them: `std.quant::mxfp4_experts_shared` for the gate
    // and up projections, whose four experts share the token's input, and
    // `std.quant::mxfp4_experts` for the down projection, whose inputs
    // differ per expert. Their bodies dequantize and gather, which XLA and
    // ONNX Runtime run as before (the dequantization, weight-only, once at
    // load); on CUDA, PyTorch reads the four experts' four-bit bytes in place,
    // a quarter of what their 16-bit weights would be. A server's batch of
    // rows takes `forward_routed` instead, as a prompt does.
    pub fn forward_topk<B: Dim>(x: Tensor[B, 1, H; T]) -> Tensor[B, 1, H; T] {
        let expert = experts_of<B>(x)
        let weight = weights_of<B>(x)
        let gate_up = mxfp4_experts_shared<B, TopK, Experts, 2 * Inner, Blocks, T>(
            x,
            gate_up_blocks,
            gate_up_scales,
            expert,
        )
        let activated = activate<B>(gate_up, expert)
        let weighted[b, k, i] = activated[b, k, i] * weight[b, k]
        let down = mxfp4_experts<B, TopK, Experts, H, InnerBlocks, T>(
            weighted,
            down_blocks,
            down_scales,
            expert,
        )
        let mixed[b, h] = sum<f32>[k] cast<f32>(down[b, k, h])
        return reshape(cast<T>(mixed) + down_bias_of<B>(weight, expert), [B, 1, H])
    }

    // A prompt's tokens through their `TopK` experts only, every token a row
    // of its own: the gate and up projections
    // `std.quant::mxfp4_linear_experts_shared`, the down projection with the
    // router's weights `mxfp4_combine_experts`. Their bodies dequantize the
    // experts and do the dense form's work, which any backend runs (the
    // dequantization once at load); PyTorch on a Hopper GPU multiplies each
    // expert by the tokens that chose it instead, an eighth of the work, with
    // `triton_kernels`' MXFP4 product where it is installed.
    pub fn forward_routed<B: Dim, S: Dim>(x: Tensor[B, S, H; T]) -> Tensor[B, S, H; T] {
        let rows = reshape(x, [B * S, 1, H])
        let expert = experts_of<B * S>(rows)
        let weight = weights_of<B * S>(rows)
        let gate_up = mxfp4_linear_experts_shared<B * S, TopK, Experts, 2 * Inner, Blocks, T>(
            reshape(x, [B * S, H]),
            gate_up_blocks,
            gate_up_scales,
            expert,
        )
        let activated = activate<B * S>(gate_up, expert)
        let mixed = mxfp4_combine_experts<B * S, TopK, Experts, H, InnerBlocks, T>(
            activated,
            down_blocks,
            down_scales,
            expert,
            weight,
        )
        return reshape(mixed + down_bias_of<B * S>(weight, expert), [B, S, H])
    }

    // Each row's `TopK` experts, the highest-ranked first.
    fn experts_of<R: Dim>(x: Tensor[R, 1, H; T]) -> Tensor[R, TopK; i64] {
        let ranked = ranks<R, 1>(router.forward(x))
        let slots = iota<i32>(Experts)
        let places = iota<i32>(TopK)
        let chosen[r, k] = sum<i32>[e] select(ranked[r, 0, e] == places[k], slots[e], 0)
        let expert[r, k] = cast<i64>(chosen[r, k])
        return expert
    }

    // The router's probability for each of those experts.
    fn weights_of<R: Dim>(x: Tensor[R, 1, H; T]) -> Tensor[R, TopK; T] {
        let logits = router.forward(x)
        let scores = route<R, 1>(logits)
        let ranked = ranks<R, 1>(logits)
        let places = iota<i32>(TopK)
        let weight[r, k] =
            sum[e] select(ranked[r, 0, e] == places[k], scores[r, 0, e], cast<T>(0.0))
        return weight
    }

    // The chosen experts' gate and up products with their biases, through
    // the gated activation.
    fn activate<R: Dim>(
        gate_up: Tensor[R, TopK, 2 * Inner; T],
        expert: Tensor[R, TopK; i64],
    ) -> Tensor[R, TopK, Inner; T] {
        let up_bias[r, k, o] = gate_up_bias[expert[r, k], o]
        let biased = gate_up + up_bias
        let gate = min(biased[..., 0::2], cast<T>(LIMIT))
        let up = min(max(biased[..., 1::2], cast<T>(-LIMIT)), cast<T>(LIMIT))
        return (up + 1.0) * (gate * sigmoid(gate * cast<T>(ALPHA)))
    }

    // The chosen experts' down biases, weighed by the router.
    fn down_bias_of<R: Dim>(
        weight: Tensor[R, TopK; T],
        expert: Tensor[R, TopK; i64],
    ) -> Tensor[R, H; T] {
        // Each chosen expert's bias row, then weighed and summed: indexing
        // `down_bias` inside the sum over `k` builds an `[R, H, TopK]` grid of
        // indices first, many times the size of the hidden states.
        let taken[r, k, h] = down_bias[expert[r, k], h]
        let bias[r, h] = sum[k] weight[r, k] * taken[r, k, h]
        return bias
    }

    // The router's probabilities, with zeros outside the top `TopK`.
    //
    // The reference takes the `TopK` largest logits, softmaxes those alone,
    // and scatters them back. There is no top-k primitive here and no need
    // for one: count how many experts beat each expert, ties going to the
    // lower index as `torch.topk` orders them, push everything ranked `TopK`
    // or worse below the softmax's floor, and softmax over all `Experts`. The
    // floored entries come out as exactly zero, so the surviving `TopK`
    // probabilities are the reference's.
    fn route<B: Dim, S: Dim>(logits: Tensor[B, S, Experts; T]) -> Tensor[B, S, Experts; T] {
        let beaten = ranks<B, S>(logits)
        let kept[b, s, e] = select(beaten[b, s, e] < cast<i32>(TopK), logits[b, s, e], -1e30)
        return softmax(kept)
    }

    // Each expert's rank among a token's router logits: 0 for the largest.
    fn ranks<B: Dim, S: Dim>(logits: Tensor[B, S, Experts; T]) -> Tensor[B, S, Experts; i32] {
        let slots = iota<i32>(Experts)
        let outranks[b, s, e, j] =
            logits[b, s, j] > logits[b, s, e] ||
            (logits[b, s, j] == logits[b, s, e] && slots[j] < slots[e])
        let beaten[b, s, e] = sum<i32>[j] cast<i32>(outranks[b, s, e, j])
        return beaten
    }
}

Open on GitHub

lib.linnet14.3 kB
linnet
// gpt-oss-20b: a mixture-of-experts decoder.
//
// Each of the 24 layers is RMSNorm, attention, RMSNorm, and a
// mixture-of-experts feed-forward block, with residual connections around the
// two halves. What separates it from a Llama-shaped decoder is in the three
// modules beside this one:
//
//   - `gpt_oss.experts`  32 experts per layer, 4 active per token, their
//                        weights stored in MXFP4 and read through
//                        `std.quant` rather than bound as floats.
//   - `gpt_oss.attention` attention sinks -- a learned logit per head that
//                        joins the softmax and is then dropped -- and layers
//                        alternating between a 128-position sliding window
//                        and full causal attention.
//   - `gpt_oss.rope`     YaRN-scaled rotary positions, which stretch a 4096
//                        training context to 131072.
//
// Only 3.6B of the 20.9B parameters are read per token, but every expert's
// weight is part of the graph: see `gpt_oss.experts` for why the routing is
// written densely in `forward`, and how the other entries route each token
// to its own experts.
module gpt_oss

use crate.attention::{SinkAttention}
use crate.experts::{MixtureOfExperts}
use std.nn.embedding::{Embedding}
use std.nn.linear::{Linear}
use std.nn.norm::{RmsNorm}

pub block DecoderLayer<
    H: Dim,
    Heads: Dim,
    KvHeads: Dim,
    HeadDim: Dim,
    Inner: Dim,
    Experts: Dim,
    TopK: Dim,
    Blocks: Dim,
    InnerBlocks: Dim,
    Batch: Dim,
    MaxSeq: Dim,
    T: Float,
>
where
    Heads > 0,
    KvHeads > 0,
    Heads % KvHeads == 0,
    HeadDim % 2 == 0,
    Experts > 0,
    TopK > 0,
    H == Blocks * 32,
    Inner == InnerBlocks * 32,
    MaxSeq > 0
{
    sub input_layernorm: RmsNorm<H, T>
    sub self_attn: SinkAttention<H, Heads, KvHeads, HeadDim, Batch, MaxSeq, T>
    sub post_attention_layernorm: RmsNorm<H, T>
    sub mlp: MixtureOfExperts<H, Inner, Experts, TopK, Blocks, InnerBlocks, T>

    pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T], window: i32) -> Tensor[B, S, H; T] {
        let attended = x + self_attn.forward(input_layernorm.forward(x), window)
        return attended + mlp.forward(post_attention_layernorm.forward(attended))
    }

    // One token per row: its experts gathered rather than all of them run.
    pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32, window: i32) -> Tensor[Batch, 1, H; T] {
        let attended = x + self_attn.decode(input_layernorm.forward(x), pos, window)
        return attended + mlp.forward_topk(post_attention_layernorm.forward(attended))
    }

    pub fn prefill<S: Dim>(
        x: Tensor[Batch, S, H; T],
        pos: i32,
        window: i32,
    ) -> Tensor[Batch, S, H; T]
    where S > 0 {
        let attended = x + self_attn.prefill(input_layernorm.forward(x), pos, window)
        return attended + mlp.forward_routed<Batch, S>(post_attention_layernorm.forward(attended))
    }

    pub fn decode_rows(
        x: Tensor[Batch, 1, H; T],
        positions: Tensor[Batch; i32],
        window: i32,
    ) -> Tensor[Batch, 1, H; T] {
        let attended = x + self_attn.decode_rows(input_layernorm.forward(x), positions, window)
        return attended + mlp.forward_routed<Batch, 1>(post_attention_layernorm.forward(attended))
    }

    pub fn prefill_slots<M: Dim, S: Dim>(
        x: Tensor[M, S, H; T],
        slots: Tensor[M; i32],
        window: i32,
    ) -> Tensor[M, S, H; T]
    where
        M > 0,
        S > 0
    {
        let attended = x + self_attn.prefill_slots(input_layernorm.forward(x), slots, window)
        return attended + mlp.forward_routed<M, S>(post_attention_layernorm.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],
        window: i32,
    ) -> Tensor[1, P, H; T]
    where P > 0 {
        let normed = input_layernorm.forward(x)
        let attended = x + self_attn.prefill_packed(normed, rows, positions, segments, window)
        return attended + mlp.forward_routed<1, P>(post_attention_layernorm.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],
        window: i32,
    ) -> Tensor[1, P + Batch, H; T]
    where P > 0 {
        let normed = input_layernorm.forward(x)
        let attended =
            x + self_attn.step_packed(normed, rows, positions, segments, step_positions, window)
        return attended +
            mlp.forward_routed<1, P + Batch>(post_attention_layernorm.forward(attended))
    }
}

// `layer_types` in config.json alternates `sliding_attention` and
// `full_attention`, starting with sliding. A `sub` array holds one block type,
// and the window is what differs, so the pair is the repeating unit: the first
// layer of each pair sees `Window` positions, the second sees the whole
// prefix. A window at least as wide as the sequence is the same mask as no
// window at all, so the second layer is given `S`.
pub block LayerPair<
    H: Dim,
    Heads: Dim,
    KvHeads: Dim,
    HeadDim: Dim,
    Inner: Dim,
    Experts: Dim,
    TopK: Dim,
    Window: Dim,
    Blocks: Dim,
    InnerBlocks: Dim,
    Batch: Dim,
    MaxSeq: Dim,
    T: Float,
>
where
    Heads > 0,
    KvHeads > 0,
    Heads % KvHeads == 0,
    HeadDim % 2 == 0,
    Experts > 0,
    TopK > 0,
    H == Blocks * 32,
    Inner == InnerBlocks * 32,
    MaxSeq > 0
{
    sub sliding: DecoderLayer<H, Heads, KvHeads, HeadDim, Inner, Experts, TopK, Blocks, InnerBlocks, Batch, MaxSeq, T>
    sub full: DecoderLayer<H, Heads, KvHeads, HeadDim, Inner, Experts, TopK, Blocks, InnerBlocks, Batch, MaxSeq, T>

    pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T]) -> Tensor[B, S, H; T] {
        let windowed = sliding.forward(x, cast<i32>(Window))
        return full.forward(windowed, cast<i32>(S))
    }

    // The cached entries: the full layer's window is the whole cache.
    pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32) -> Tensor[Batch, 1, H; T] {
        let windowed = sliding.decode(x, pos, cast<i32>(Window))
        return full.decode(windowed, pos, cast<i32>(MaxSeq))
    }

    pub fn prefill<S: Dim>(x: Tensor[Batch, S, H; T], pos: i32) -> Tensor[Batch, S, H; T]
    where S > 0 {
        let windowed = sliding.prefill(x, pos, cast<i32>(Window))
        return full.prefill(windowed, pos, cast<i32>(MaxSeq))
    }

    pub fn decode_rows(
        x: Tensor[Batch, 1, H; T],
        positions: Tensor[Batch; i32],
    ) -> Tensor[Batch, 1, H; T] {
        let windowed = sliding.decode_rows(x, positions, cast<i32>(Window))
        return full.decode_rows(windowed, positions, cast<i32>(MaxSeq))
    }

    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 windowed = sliding.prefill_slots(x, slots, cast<i32>(Window))
        return full.prefill_slots(windowed, slots, cast<i32>(S))
    }

    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 windowed = sliding.prefill_packed(x, rows, positions, segments, cast<i32>(Window))
        return full.prefill_packed(windowed, rows, positions, segments, cast<i32>(MaxSeq))
    }

    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 windowed = sliding.step_packed(
            x,
            rows,
            positions,
            segments,
            step_positions,
            cast<i32>(Window),
        )
        return full.step_packed(
            windowed,
            rows,
            positions,
            segments,
            step_positions,
            cast<i32>(MaxSeq),
        )
    }
}

pub block Model<
    Vocab: Dim,
    H: Dim,
    Heads: Dim,
    KvHeads: Dim,
    HeadDim: Dim,
    Inner: Dim,
    Experts: Dim,
    TopK: Dim,
    Layers: Dim,
    Window: Dim,
    Batch: Dim,
    MaxSeq: Dim,
    T: Float = bf16,
>
where
    Heads > 0,
    KvHeads > 0,
    Heads % KvHeads == 0,
    HeadDim % 2 == 0,
    Experts > 0,
    TopK > 0,
    Layers % 2 == 0,
    H % 32 == 0,
    Inner % 32 == 0,
    MaxSeq > 0
{
    sub embed_tokens: Embedding<Vocab, H, T>
    sub pairs: [LayerPair<H, Heads, KvHeads, HeadDim, Inner, Experts, TopK, Window, H / 32, Inner /
        32, Batch, MaxSeq, T>; Layers / 2]
    sub norm: RmsNorm<H, T>
    // gpt-oss does not tie the output head to the embedding
    // (`tie_word_embeddings` is false); the checkpoint carries its own
    // `lm_head.weight`.
    sub lm_head: Linear<H, Vocab, T>

    // Logits for every position.
    pub entry forward<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, Vocab; T] {
        return lm_head.forward(norm.forward(hidden(tokens)))
    }

    // Logits for the last position only, as sampling one token needs: the
    // final norm and the projection over the vocabulary run on one row per
    // sequence instead of `S` of them.
    pub entry next_token<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, Vocab; T]
    where S > 0 {
        return lm_head.forward(norm.forward(hidden(tokens)[:, S - 1, :]))
    }

    // Logits for one new token per sequence at position `pos`, attending
    // over the positions decoded before it through the layers' KV caches.
    pub entry decode(token: Tensor[Batch, 1; i32], pos: i32) -> Tensor[Batch, Vocab; T] {
        var x = embed_tokens.forward(token)
        static for pair in pairs {
            x = pair.decode(x, pos)
        }
        return lm_head.forward(norm.forward(x[:, 0, :]))
    }

    // 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`.
    pub entry prefill<S: Dim>(tokens: Tensor[Batch, S; i32], pos: i32) -> Tensor[Batch, Vocab; T]
    where S > 0 {
        var x = embed_tokens.forward(tokens)
        static for pair in pairs {
            x = pair.prefill(x, pos)
        }
        return lm_head.forward(norm.forward(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 = embed_tokens.forward(tokens)
        static for pair in pairs {
            x = pair.decode_rows(x, positions)
        }
        return lm_head.forward(norm.forward(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 = embed_tokens.forward(tokens)
        static for pair in pairs {
            x = pair.prefill_slots(x, slots)
        }
        let last[b, h] = x[b, cast<i64>(lengths[b]) - 1, h]
        return lm_head.forward(norm.forward(last))
    }

    // 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.
    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 = embed_tokens.forward(reshape(tokens, [1, P]))
        static for pair in pairs {
            x = pair.prefill_packed(x, rows, positions, segments)
        }
        let ends[m, h] = x[0, cast<i64>(last[m]), h]
        return lm_head.forward(norm.forward(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. Returns the logits after each prompt,
    // 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 = embed_tokens.forward(reshape(every, [1, P + Batch]))
        static for pair in pairs {
            x = pair.step_packed(x, rows, positions, segments, step_positions)
        }
        let ends[m, h] = x[0, cast<i64>(last[m]), h]
        return lm_head.forward(norm.forward(concat(ends, x[0, P:P + Batch, :], axis = 0)))
    }

    // The final layer's hidden states, before the norm and the head.
    pub entry states<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, H; T] {
        return hidden(tokens)
    }

    fn hidden<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, H; T] {
        var x = embed_tokens.forward(tokens)
        static for pair in pairs {
            x = pair.forward(x)
        }
        return x
    }
}

Open on GitHub

rope.linnet2.7 kB
linnet
// Rotary positions as gpt-oss scales them: YaRN, from `rope_scaling` in
// config.json.
//
// Plain rotary embedding gives dimension `i` of a head the frequency
// `theta ** (-2i / D)`. A model trained on 4096 positions and served at
// 131072 cannot use those frequencies unchanged: the slow ones have never
// seen a full rotation at the longer range. YaRN divides the slow
// frequencies by the context ratio (interpolation) and leaves the fast ones
// alone (extrapolation), ramping linearly between the two dimensions where a
// wavelength first covers `beta_fast` and `beta_slow` rotations of the
// original context. It then scales cos and sin by a factor above one, which
// compensates for the entropy the interpolation takes out of the attention
// logits.
module gpt_oss.rope

// config.json: `rope_theta`, and `rope_scaling`'s `factor`,
// `original_max_position_embeddings`, `beta_fast`, and `beta_slow`.
pub const THETA: f32 = 150000.0
pub const FACTOR: f32 = 32.0
pub const ORIGINAL_CONTEXT: f32 = 4096.0
pub const BETA_FAST: f32 = 32.0
pub const BETA_SLOW: f32 = 1.0

const TWO_PI: f32 = 6.2831853071795865

// The (fractional) head dimension whose wavelength spans `rotations` turns of
// the original context. `truncate` is false in this config, so the bounds are
// not rounded to whole dimensions.
fn correction_dim(rotations: f32, dim: f32) -> f32 {
    return dim * log(ORIGINAL_CONTEXT / (rotations * TWO_PI)) / (2.0 * log(THETA))
}

// `(cos, sin)` tables of shape `[S, D]`, each frequency repeated over both
// halves of the head dimension so `std.nn.rope::rope` can pair them. The
// reference keeps the tables at `D / 2` and splits the head in half itself;
// the two forms compute the same rotation.
pub fn tables<S: Dim, D: Dim, T: Float>() -> (Tensor[S, D; T], Tensor[S, D; T])
where D % 2 == 0 {
    let index = cast<f32>(iota<i64>(D / 2))
    // `theta ** (2i / D)`, the wavelength of dimension `i`.
    let wavelength = exp(index * (2.0 / cast<f32>(D)) * log(THETA))
    let extrapolated = 1.0 / wavelength
    let interpolated = 1.0 / (FACTOR * wavelength)
    let low = max(correction_dim(BETA_FAST, cast<f32>(D)), 0.0)
    let high = min(correction_dim(BETA_SLOW, cast<f32>(D)), cast<f32>(D) - 1.0)
    // 0 below `low` (pure extrapolation), 1 above `high` (pure
    // interpolation), linear in between.
    let ramp = min(max((index - low) / (high - low), 0.0), 1.0)
    let inv_freq = interpolated * ramp + extrapolated * (1.0 - ramp)
    let angles[s, i] = cast<f32>(iota<i64>(S)[s]) * inv_freq[i]
    let full = concat(angles, angles, axis = -1)
    // The YaRN attention factor, `0.1 * log(factor) + 1`.
    let attention_factor = 0.1 * log(FACTOR) + 1.0
    return (cast<T>(cos(full) * attention_factor), cast<T>(sin(full) * attention_factor))
}

Open on GitHub

model-00000-of-00002.safetensorsHugging Face ↗

model-00001-of-00002.safetensorsHugging Face ↗

model-00002-of-00002.safetensorsHugging Face ↗