models/phi-3-mini-4k-instructREADME.md5.8 kB
md
# Phi-3 Mini 4K Instruct
A 3.8B-parameter decoder: RMS normalization, full (non-grouped) multi-head
attention with rotary positions, a SwiGLU MLP, and an untied output head, no
bias anywhere. `config.json` sets `num_attention_heads` and
`num_key_value_heads` both to 32, so there is no grouped-query attention to
model. Two checkpoint tensors are fused and the Linnet source keeps them
fused, slicing the projected result instead of splitting the weight:
`self_attn.qkv_proj.weight` is one `[3 * H, H]` tensor (query, then key,
then value, along the output axis) and `mlp.gate_up_proj.weight` is one
`[2 * Inner, H]` tensor (gate, then up). Both keep the ordinary `nn.Linear`
`[Out, In]` layout, so `bindings.json` binds them with no transpose.
## Determining the gate/up order
`gate_up_proj`'s halves are not distinguishable from the shape alone. The
reference `Phi3MLP.forward` in `transformers/models/phi3/modeling_phi3.py`
reads:
```python
up_states = self.gate_up_proj(hidden_states)
gate, up_states = up_states.chunk(2, dim=-1)
up_states = up_states * self.activation_fn(gate)
```
`chunk(2, dim=-1)` splits the output axis in storage order, so the first
`Inner` rows of the weight produce `gate` (the half passed through SiLU) and
the second `Inner` rows produce the value `gate` multiplies. The same
ordering (query, key, value) holds for `qkv_proj`, read from
`Phi3Attention.forward`'s `qkv[..., :query_pos]` /
`qkv[..., query_pos : query_pos + num_key_value_heads * head_dim]` /
`qkv[..., query_pos + num_key_value_heads * head_dim :]`. A wrong gate/up
split still produces fluent-looking text, since SwiGLU is not symmetric but
both halves are learned linear projections of the same input -- see
"Validation" below for the check that actually catches it.
## Sliding window
`config.json` sets `sliding_window` to 2047, and `transformers` masks every
layer with it: a query attends to itself and the 2046 positions before it.
The Linnet source does the same (`WINDOW`, `window_mask`), in every entry;
the caches still hold every position, and a sequence at or under the window
sees plain causal attention.
## Loading
```python
from linnet import nest
model = nest.load("phi-3-mini-4k-instruct", backend="torch", numerics="fast")
logits = model(tokens) # Tensor[1, S; i32] -> Tensor[1, S, 32064; bf16]
```
The weights are the published checkpoint (bf16, two shards);
`bindings.json` maps Linnet parameter paths to its tensor names. Every name,
shape, and dtype is checked before anything runs.
## Entries
| Entry | |
| --- | --- |
| `forward<B, S>(tokens)` | logits for a whole sequence |
| `next_token<B, S>(tokens)` | logits for the last position |
| `decode(token, pos)` | one token through the KV caches (`Batch = 1`, `MaxSeq = 4096`) |
| `prefill<S>(tokens, pos)` | a whole prompt through the KV caches in one pass; logits after its last token, `decode` continues at `pos + S` |
| `prefill_slots<M, S>(tokens, slots, lengths)` | `M` requests' prompts into rows `slots` of the caches in one pass (each padded to `S`, its first `lengths[m]` tokens real), as they join a batch being served |
| `prefill_packed<P>(tokens, rows, positions, segments, last)` | several requests' prompts packed end to end into one pass of `P` tokens, each token given its cache row, its position, and its prompt: no padding between prompts, and each prompt sees only itself |
| `decode_rows(tokens, positions)` | one token for every row of the caches, each at its own position: the step `linnet.serve` takes for continuous batching |
| `generate<Steps>(token, pos)` | greedy decoding in the graph |
| `sample<Steps>(token, pos, key, temperature)` | sampling with `std.random` |
| `generate_until<MaxNew>(token, pos, eos)` | decoding until an end token |
## Provenance
- Weights: [microsoft/Phi-3-mini-4k-instruct](https://huggingface.co/microsoft/Phi-3-mini-4k-instruct), MIT.
- Code: [transformers `models/phi3`](https://github.com/huggingface/transformers/tree/main/src/transformers/models/phi3).
- Paper: [Phi-3 Technical Report](https://arxiv.org/abs/2404.14219).
## Validation
A 3.8B model in f32 is too large to compare whole against `transformers` on
this machine, so validation ran truncated, against `transformers` under
`torch.no_grad()` with the tokenizer's chat template:
- **Truncated, f32, tight tolerance.** Both sides read only the first two
layers of the published checkpoint: the Linnet side through
`nest.load(..., generics={"Layers": 2, "T": "f32"})` bound from a small
local file holding only the 15 tensors the first two layers and the
embedding/head need, cast to f32 from the published bf16 checkpoint and
never touching the other 30 layers; the reference through
`AutoModelForCausalLM.from_pretrained(..., config=<num_hidden_layers=2>,
torch_dtype=torch.float32)`. Both were fed the same 10-token prompt built
from `tokenizer.apply_chat_template`. This proves the fused-QKV slicing,
the gate/up split, RoPE, and normalization are wired correctly, since a
wrong gate/up split still produces fluent-looking logits but would not
survive a 2e-3 tolerance here.
**Max abs difference in logits: 7.15e-6.** Argmax (next-token prediction)
matches exactly. Peak RSS for the whole run (reference built and released
before the Linnet model, per the shared-machine memory discipline): 2.77
GiB.
- **Whole model, bf16, top-1 agreement at full depth.** Deferred to a batch
GPU run (RunPod) covering every model in this pass, alongside the other
additions to this registry; not yet run for this model. Until that run
lands, correctness at full depth (32 layers, e.g. rope behavior far past
the truncated pass) is not directly checked, only implied by the
truncated pass and by `check`'s shape/dtype verification against the full
checkpoint's headers.
See `python/linnet/tests/torch/test_hf_checkpoints.py` in the Linnet
repository for the harness this follows.bench.json12.6 kB
json
{
"date": "2026-09-28T17:45:00+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": 19.263076595962048,
"decode_tok_s": 70.9409414289243,
"load_s": 4.614230884239078,
"peak_vram_mib": 8336.0
},
"max_abs_diff": 0.0,
"notes": "",
"error": null,
"key": "transformers-eager"
},
{
"method": "transformers (torch.compile, static cache)",
"kind": "reference",
"metrics": {
"ttft_ms": 12.431768700480461,
"decode_tok_s": 177.48252341974285,
"load_s": 2.6776997353881598,
"peak_vram_mib": 8592.0
},
"max_abs_diff": 0.15625,
"notes": "",
"error": null,
"key": "transformers-compile"
},
{
"method": "vLLM",
"kind": "reference",
"metrics": {
"ttft_ms": 11.908145621418953,
"decode_tok_s": 242.46941408271408,
"load_s": 67.52199644595385,
"first_token": 4146.0,
"peak_vram_mib": 68418.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": 12.253515422344208,
"decode_tok_s": 88.88517237513541,
"load_s": 1.9124931693077087,
"peak_vram_mib": 8396.0
},
"max_abs_diff": 0.265625,
"notes": "KV cache compiled for 768 positions",
"error": null,
"key": "linnet-torch"
},
{
"method": "Linnet torch (CUDA graphs)",
"kind": "linnet",
"metrics": {
"ttft_ms": 7.048172876238823,
"decode_tok_s": 264.45888355622213,
"load_s": 1.507862500846386,
"peak_vram_mib": 8508.9
},
"max_abs_diff": 0.140625,
"notes": "KV cache compiled for 768 positions",
"error": null,
"key": "linnet-cudagraphs"
},
{
"method": "Linnet torch (torch.compile, inductor)",
"kind": "linnet",
"metrics": {
"ttft_ms": 7.904881611466408,
"decode_tok_s": 211.15610923373725,
"load_s": 1.5232672169804573,
"peak_vram_mib": 8388.9
},
"max_abs_diff": 0.140625,
"notes": "KV cache compiled for 768 positions",
"error": null,
"key": "linnet-inductor"
},
{
"method": "Linnet JAX (XLA, StableHLO)",
"kind": "linnet",
"metrics": {
"ttft_ms": 9.758176282048225,
"decode_tok_s": 292.0402633518271,
"load_s": 4.986512377858162,
"peak_vram_mib": 8468.5,
"driver_vram_mib": 17018.9
},
"max_abs_diff": 0.125,
"notes": "KV cache compiled for 768 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
"error": null,
"key": "linnet-jax"
},
{
"method": "Linnet JAX (XLA, generated source)",
"kind": "linnet",
"metrics": {
"ttft_ms": 9.269213303923607,
"decode_tok_s": 291.748279996101,
"load_s": 4.9374614879488945,
"peak_vram_mib": 8468.5,
"driver_vram_mib": 17018.9
},
"max_abs_diff": 0.3125,
"notes": "KV cache compiled for 768 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
"error": null,
"key": "linnet-jax-source"
},
{
"method": "Linnet ONNX f32 -> ONNX Runtime (CUDA)",
"kind": "linnet",
"metrics": {
"ttft_ms": 31.667771749198437,
"decode_tok_s": 75.15848793048104,
"load_s": 0.3668599259108305,
"peak_vram_mib": 20596.0
},
"max_abs_diff": 0.13383817672729492,
"notes": "f32; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
"error": null,
"key": "linnet-onnx-f32"
},
{
"method": "Linnet ONNX f32 -> ONNX Runtime (TensorRT)",
"kind": "linnet",
"metrics": {
"ttft_ms": 30.343933030962944,
"decode_tok_s": 50.115414353319785,
"load_s": 0.3913907874375582,
"peak_vram_mib": 43010.0
},
"max_abs_diff": 0.13777589797973633,
"notes": "f32; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
"error": null,
"key": "linnet-onnx-f32-trt"
},
{
"method": "Linnet ONNX f16 -> ONNX Runtime (CUDA)",
"kind": "linnet",
"metrics": {
"ttft_ms": 22.195509634912014,
"decode_tok_s": 95.13229779240991,
"load_s": 0.3609587009996176,
"peak_vram_mib": 10918.0
},
"max_abs_diff": 0.15625,
"notes": "f16; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
"error": null,
"key": "linnet-onnx-f16"
},
{
"method": "Linnet ONNX f16 -> ONNX Runtime (TensorRT)",
"kind": "linnet",
"metrics": {
"ttft_ms": 15.986711718142033,
"decode_tok_s": 84.37322085280859,
"load_s": 0.3601016737520695,
"peak_vram_mib": 22674.0
},
"max_abs_diff": 14.5,
"notes": "f16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
"error": null,
"key": "linnet-onnx-f16-trt"
},
{
"method": "Linnet ONNX bf16 -> ONNX Runtime (CUDA)",
"kind": "linnet",
"metrics": {
"ttft_ms": 22.813155315816402,
"decode_tok_s": 93.819816730942,
"load_s": 0.4100595396012068,
"peak_vram_mib": 10020.0
},
"max_abs_diff": 0.125,
"notes": "bf16; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
"error": null,
"key": "linnet-onnx-bf16"
},
{
"method": "Linnet ONNX bf16 -> ONNX Runtime (TensorRT)",
"kind": "linnet",
"metrics": {
"ttft_ms": 16.092244535684586,
"decode_tok_s": 83.32864937080885,
"load_s": 0.3686938416212797,
"peak_vram_mib": 22154.0
},
"max_abs_diff": 0.1875,
"notes": "bf16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax",
"error": null,
"key": "linnet-onnx-bf16-trt"
},
{
"method": "vLLM (offline, continuous batching)",
"kind": "reference",
"metrics": {
"serve_tok_s": 5583.255700745566,
"serve_s": 5.868977126665413,
"load_s": 34.46890168543905,
"peak_vram_mib": 70000.8
},
"max_abs_diff": null,
"notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at gpu_memory_utilization 0.85",
"error": null,
"key": "serve-vllm"
},
{
"method": "transformers (generate_batch, continuous batching)",
"kind": "reference",
"metrics": {
"serve_tok_s": 520.9591019781975,
"serve_s": 62.89937132410705,
"load_s": 2.3650451321154833,
"peak_vram_mib": 80980.0
},
"max_abs_diff": null,
"notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; paged|sdpa attention; reserves a paged KV-cache pool up front, so its memory is a setting, not a need; no per-request timestamps",
"error": null,
"key": "serve-transformers"
},
{
"method": "Triton Inference Server (vLLM backend)",
"kind": "reference",
"metrics": {
"serve_tok_s": 1968.2018462248288,
"serve_s": 16.648698944598436,
"load_s": 51.05910853296518,
"serve_ttft_ms": 7874.6669590473175,
"peak_vram_mib": 70432.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": 6667.727671793196,
"serve_s": 4.914417866617441,
"load_s": 4.80376348271966,
"serve_ttft_ms": 2131.3224397599697,
"peak_vram_mib": 24598.9
},
"max_abs_diff": null,
"notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions",
"error": null,
"key": "serve-linnet-torch"
},
{
"method": "Linnet jax (linnet.serve, XLA)",
"kind": "linnet",
"metrics": {
"serve_tok_s": 5177.762430023286,
"serve_s": 6.328602449968457,
"load_s": 7.269423228688538,
"serve_ttft_ms": 2687.9857759922743,
"peak_vram_mib": 26643.1,
"driver_vram_mib": 33452.0
},
"max_abs_diff": null,
"notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions",
"error": null,
"key": "serve-linnet-jax"
},
{
"method": "Linnet onnx (linnet.serve, ONNX Runtime, f16)",
"kind": "linnet",
"metrics": {
"serve_tok_s": 3526.291795707058,
"serve_s": 9.292481138370931,
"load_s": 0.2581939958035946,
"serve_ttft_ms": 4343.968306668103,
"peak_vram_mib": 49972.0
},
"max_abs_diff": null,
"notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions",
"error": null,
"key": "serve-linnet-onnx"
},
{
"method": "Triton Inference Server (Python backend, linnet.serve)",
"kind": "linnet",
"metrics": {
"serve_tok_s": 6328.143729444483,
"serve_s": 5.178137760609388,
"load_s": 290.22266763448715,
"serve_ttft_ms": 2123.587576672435,
"peak_vram_mib": 25540.4
},
"max_abs_diff": null,
"notes": "256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed, first tokens timed at the client",
"error": null,
"key": "serve-triton-linnet"
},
{
"method": "Linnet -> vLLM",
"kind": "linnet",
"metrics": {
"ttft_ms": 11.993551161140203,
"decode_tok_s": 250.91478809572004,
"load_s": 50.27317692618817,
"first_token": 4146.0,
"peak_vram_mib": 68492.4,
"driver_vram_mib": 68492.4
},
"max_abs_diff": null,
"notes": "linnet.hf.export (8 s), then vLLM on the exported checkpoint; reserves a KV-cache pool up front (gpu_memory_utilization 0.85), so its memory is a setting, not a need",
"error": null,
"key": "linnet-vllm"
},
{
"method": "Linnet -> vLLM (offline, continuous batching)",
"kind": "linnet",
"metrics": {
"serve_tok_s": 5587.451650467343,
"serve_s": 5.864569762721658,
"load_s": 32.60180378332734,
"peak_vram_mib": 70000.8,
"driver_vram_mib": 70000.8
},
"max_abs_diff": null,
"notes": "linnet.hf.export (0 s), then vLLM (offline, continuous batching) on the exported checkpoint; 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at gpu_memory_utilization 0.85",
"error": null,
"key": "serve-linnet-vllm"
},
{
"method": "Linnet -> llama.cpp (GGUF, f16)",
"kind": "linnet",
"metrics": {
"ttft_ms": 16.75357837680428,
"decode_tok_s": 231.08384,
"load_s": 27.896604729816318,
"peak_vram_mib": 8090.8
},
"max_abs_diff": null,
"notes": "linnet.gguf.export, then llama-bench with every layer on the GPU; the first token is the prompt at llama-bench's prompt rate",
"error": null,
"key": "linnet-llamacpp"
}
],
"reference": "transformers-eager"
}bindings.json14.8 kB
json
{
"embedding.weight": "model.embed_tokens.weight",
"layers.0.attention_norm.weight": "model.layers.0.input_layernorm.weight",
"layers.0.attention.qkv.weight": "model.layers.0.self_attn.qkv_proj.weight",
"layers.0.attention.o_proj.weight": "model.layers.0.self_attn.o_proj.weight",
"layers.0.mlp_norm.weight": "model.layers.0.post_attention_layernorm.weight",
"layers.0.mlp.gate_up.weight": "model.layers.0.mlp.gate_up_proj.weight",
"layers.0.mlp.down.weight": "model.layers.0.mlp.down_proj.weight",
"layers.1.attention_norm.weight": "model.layers.1.input_layernorm.weight",
"layers.1.attention.qkv.weight": "model.layers.1.self_attn.qkv_proj.weight",
"layers.1.attention.o_proj.weight": "model.layers.1.self_attn.o_proj.weight",
"layers.1.mlp_norm.weight": "model.layers.1.post_attention_layernorm.weight",
"layers.1.mlp.gate_up.weight": "model.layers.1.mlp.gate_up_proj.weight",
"layers.1.mlp.down.weight": "model.layers.1.mlp.down_proj.weight",
"layers.2.attention_norm.weight": "model.layers.2.input_layernorm.weight",
"layers.2.attention.qkv.weight": "model.layers.2.self_attn.qkv_proj.weight",
"layers.2.attention.o_proj.weight": "model.layers.2.self_attn.o_proj.weight",
"layers.2.mlp_norm.weight": "model.layers.2.post_attention_layernorm.weight",
"layers.2.mlp.gate_up.weight": "model.layers.2.mlp.gate_up_proj.weight",
"layers.2.mlp.down.weight": "model.layers.2.mlp.down_proj.weight",
"layers.3.attention_norm.weight": "model.layers.3.input_layernorm.weight",
"layers.3.attention.qkv.weight": "model.layers.3.self_attn.qkv_proj.weight",
"layers.3.attention.o_proj.weight": "model.layers.3.self_attn.o_proj.weight",
"layers.3.mlp_norm.weight": "model.layers.3.post_attention_layernorm.weight",
"layers.3.mlp.gate_up.weight": "model.layers.3.mlp.gate_up_proj.weight",
"layers.3.mlp.down.weight": "model.layers.3.mlp.down_proj.weight",
"layers.4.attention_norm.weight": "model.layers.4.input_layernorm.weight",
"layers.4.attention.qkv.weight": "model.layers.4.self_attn.qkv_proj.weight",
"layers.4.attention.o_proj.weight": "model.layers.4.self_attn.o_proj.weight",
"layers.4.mlp_norm.weight": "model.layers.4.post_attention_layernorm.weight",
"layers.4.mlp.gate_up.weight": "model.layers.4.mlp.gate_up_proj.weight",
"layers.4.mlp.down.weight": "model.layers.4.mlp.down_proj.weight",
"layers.5.attention_norm.weight": "model.layers.5.input_layernorm.weight",
"layers.5.attention.qkv.weight": "model.layers.5.self_attn.qkv_proj.weight",
"layers.5.attention.o_proj.weight": "model.layers.5.self_attn.o_proj.weight",
"layers.5.mlp_norm.weight": "model.layers.5.post_attention_layernorm.weight",
"layers.5.mlp.gate_up.weight": "model.layers.5.mlp.gate_up_proj.weight",
"layers.5.mlp.down.weight": "model.layers.5.mlp.down_proj.weight",
"layers.6.attention_norm.weight": "model.layers.6.input_layernorm.weight",
"layers.6.attention.qkv.weight": "model.layers.6.self_attn.qkv_proj.weight",
"layers.6.attention.o_proj.weight": "model.layers.6.self_attn.o_proj.weight",
"layers.6.mlp_norm.weight": "model.layers.6.post_attention_layernorm.weight",
"layers.6.mlp.gate_up.weight": "model.layers.6.mlp.gate_up_proj.weight",
"layers.6.mlp.down.weight": "model.layers.6.mlp.down_proj.weight",
"layers.7.attention_norm.weight": "model.layers.7.input_layernorm.weight",
"layers.7.attention.qkv.weight": "model.layers.7.self_attn.qkv_proj.weight",
"layers.7.attention.o_proj.weight": "model.layers.7.self_attn.o_proj.weight",
"layers.7.mlp_norm.weight": "model.layers.7.post_attention_layernorm.weight",
"layers.7.mlp.gate_up.weight": "model.layers.7.mlp.gate_up_proj.weight",
"layers.7.mlp.down.weight": "model.layers.7.mlp.down_proj.weight",
"layers.8.attention_norm.weight": "model.layers.8.input_layernorm.weight",
"layers.8.attention.qkv.weight": "model.layers.8.self_attn.qkv_proj.weight",
"layers.8.attention.o_proj.weight": "model.layers.8.self_attn.o_proj.weight",
"layers.8.mlp_norm.weight": "model.layers.8.post_attention_layernorm.weight",
"layers.8.mlp.gate_up.weight": "model.layers.8.mlp.gate_up_proj.weight",
"layers.8.mlp.down.weight": "model.layers.8.mlp.down_proj.weight",
"layers.9.attention_norm.weight": "model.layers.9.input_layernorm.weight",
"layers.9.attention.qkv.weight": "model.layers.9.self_attn.qkv_proj.weight",
"layers.9.attention.o_proj.weight": "model.layers.9.self_attn.o_proj.weight",
"layers.9.mlp_norm.weight": "model.layers.9.post_attention_layernorm.weight",
"layers.9.mlp.gate_up.weight": "model.layers.9.mlp.gate_up_proj.weight",
"layers.9.mlp.down.weight": "model.layers.9.mlp.down_proj.weight",
"layers.10.attention_norm.weight": "model.layers.10.input_layernorm.weight",
"layers.10.attention.qkv.weight": "model.layers.10.self_attn.qkv_proj.weight",
"layers.10.attention.o_proj.weight": "model.layers.10.self_attn.o_proj.weight",
"layers.10.mlp_norm.weight": "model.layers.10.post_attention_layernorm.weight",
"layers.10.mlp.gate_up.weight": "model.layers.10.mlp.gate_up_proj.weight",
"layers.10.mlp.down.weight": "model.layers.10.mlp.down_proj.weight",
"layers.11.attention_norm.weight": "model.layers.11.input_layernorm.weight",
"layers.11.attention.qkv.weight": "model.layers.11.self_attn.qkv_proj.weight",
"layers.11.attention.o_proj.weight": "model.layers.11.self_attn.o_proj.weight",
"layers.11.mlp_norm.weight": "model.layers.11.post_attention_layernorm.weight",
"layers.11.mlp.gate_up.weight": "model.layers.11.mlp.gate_up_proj.weight",
"layers.11.mlp.down.weight": "model.layers.11.mlp.down_proj.weight",
"layers.12.attention_norm.weight": "model.layers.12.input_layernorm.weight",
"layers.12.attention.qkv.weight": "model.layers.12.self_attn.qkv_proj.weight",
"layers.12.attention.o_proj.weight": "model.layers.12.self_attn.o_proj.weight",
"layers.12.mlp_norm.weight": "model.layers.12.post_attention_layernorm.weight",
"layers.12.mlp.gate_up.weight": "model.layers.12.mlp.gate_up_proj.weight",
"layers.12.mlp.down.weight": "model.layers.12.mlp.down_proj.weight",
"layers.13.attention_norm.weight": "model.layers.13.input_layernorm.weight",
"layers.13.attention.qkv.weight": "model.layers.13.self_attn.qkv_proj.weight",
"layers.13.attention.o_proj.weight": "model.layers.13.self_attn.o_proj.weight",
"layers.13.mlp_norm.weight": "model.layers.13.post_attention_layernorm.weight",
"layers.13.mlp.gate_up.weight": "model.layers.13.mlp.gate_up_proj.weight",
"layers.13.mlp.down.weight": "model.layers.13.mlp.down_proj.weight",
"layers.14.attention_norm.weight": "model.layers.14.input_layernorm.weight",
"layers.14.attention.qkv.weight": "model.layers.14.self_attn.qkv_proj.weight",
"layers.14.attention.o_proj.weight": "model.layers.14.self_attn.o_proj.weight",
"layers.14.mlp_norm.weight": "model.layers.14.post_attention_layernorm.weight",
"layers.14.mlp.gate_up.weight": "model.layers.14.mlp.gate_up_proj.weight",
"layers.14.mlp.down.weight": "model.layers.14.mlp.down_proj.weight",
"layers.15.attention_norm.weight": "model.layers.15.input_layernorm.weight",
"layers.15.attention.qkv.weight": "model.layers.15.self_attn.qkv_proj.weight",
"layers.15.attention.o_proj.weight": "model.layers.15.self_attn.o_proj.weight",
"layers.15.mlp_norm.weight": "model.layers.15.post_attention_layernorm.weight",
"layers.15.mlp.gate_up.weight": "model.layers.15.mlp.gate_up_proj.weight",
"layers.15.mlp.down.weight": "model.layers.15.mlp.down_proj.weight",
"layers.16.attention_norm.weight": "model.layers.16.input_layernorm.weight",
"layers.16.attention.qkv.weight": "model.layers.16.self_attn.qkv_proj.weight",
"layers.16.attention.o_proj.weight": "model.layers.16.self_attn.o_proj.weight",
"layers.16.mlp_norm.weight": "model.layers.16.post_attention_layernorm.weight",
"layers.16.mlp.gate_up.weight": "model.layers.16.mlp.gate_up_proj.weight",
"layers.16.mlp.down.weight": "model.layers.16.mlp.down_proj.weight",
"layers.17.attention_norm.weight": "model.layers.17.input_layernorm.weight",
"layers.17.attention.qkv.weight": "model.layers.17.self_attn.qkv_proj.weight",
"layers.17.attention.o_proj.weight": "model.layers.17.self_attn.o_proj.weight",
"layers.17.mlp_norm.weight": "model.layers.17.post_attention_layernorm.weight",
"layers.17.mlp.gate_up.weight": "model.layers.17.mlp.gate_up_proj.weight",
"layers.17.mlp.down.weight": "model.layers.17.mlp.down_proj.weight",
"layers.18.attention_norm.weight": "model.layers.18.input_layernorm.weight",
"layers.18.attention.qkv.weight": "model.layers.18.self_attn.qkv_proj.weight",
"layers.18.attention.o_proj.weight": "model.layers.18.self_attn.o_proj.weight",
"layers.18.mlp_norm.weight": "model.layers.18.post_attention_layernorm.weight",
"layers.18.mlp.gate_up.weight": "model.layers.18.mlp.gate_up_proj.weight",
"layers.18.mlp.down.weight": "model.layers.18.mlp.down_proj.weight",
"layers.19.attention_norm.weight": "model.layers.19.input_layernorm.weight",
"layers.19.attention.qkv.weight": "model.layers.19.self_attn.qkv_proj.weight",
"layers.19.attention.o_proj.weight": "model.layers.19.self_attn.o_proj.weight",
"layers.19.mlp_norm.weight": "model.layers.19.post_attention_layernorm.weight",
"layers.19.mlp.gate_up.weight": "model.layers.19.mlp.gate_up_proj.weight",
"layers.19.mlp.down.weight": "model.layers.19.mlp.down_proj.weight",
"layers.20.attention_norm.weight": "model.layers.20.input_layernorm.weight",
"layers.20.attention.qkv.weight": "model.layers.20.self_attn.qkv_proj.weight",
"layers.20.attention.o_proj.weight": "model.layers.20.self_attn.o_proj.weight",
"layers.20.mlp_norm.weight": "model.layers.20.post_attention_layernorm.weight",
"layers.20.mlp.gate_up.weight": "model.layers.20.mlp.gate_up_proj.weight",
"layers.20.mlp.down.weight": "model.layers.20.mlp.down_proj.weight",
"layers.21.attention_norm.weight": "model.layers.21.input_layernorm.weight",
"layers.21.attention.qkv.weight": "model.layers.21.self_attn.qkv_proj.weight",
"layers.21.attention.o_proj.weight": "model.layers.21.self_attn.o_proj.weight",
"layers.21.mlp_norm.weight": "model.layers.21.post_attention_layernorm.weight",
"layers.21.mlp.gate_up.weight": "model.layers.21.mlp.gate_up_proj.weight",
"layers.21.mlp.down.weight": "model.layers.21.mlp.down_proj.weight",
"layers.22.attention_norm.weight": "model.layers.22.input_layernorm.weight",
"layers.22.attention.qkv.weight": "model.layers.22.self_attn.qkv_proj.weight",
"layers.22.attention.o_proj.weight": "model.layers.22.self_attn.o_proj.weight",
"layers.22.mlp_norm.weight": "model.layers.22.post_attention_layernorm.weight",
"layers.22.mlp.gate_up.weight": "model.layers.22.mlp.gate_up_proj.weight",
"layers.22.mlp.down.weight": "model.layers.22.mlp.down_proj.weight",
"layers.23.attention_norm.weight": "model.layers.23.input_layernorm.weight",
"layers.23.attention.qkv.weight": "model.layers.23.self_attn.qkv_proj.weight",
"layers.23.attention.o_proj.weight": "model.layers.23.self_attn.o_proj.weight",
"layers.23.mlp_norm.weight": "model.layers.23.post_attention_layernorm.weight",
"layers.23.mlp.gate_up.weight": "model.layers.23.mlp.gate_up_proj.weight",
"layers.23.mlp.down.weight": "model.layers.23.mlp.down_proj.weight",
"layers.24.attention_norm.weight": "model.layers.24.input_layernorm.weight",
"layers.24.attention.qkv.weight": "model.layers.24.self_attn.qkv_proj.weight",
"layers.24.attention.o_proj.weight": "model.layers.24.self_attn.o_proj.weight",
"layers.24.mlp_norm.weight": "model.layers.24.post_attention_layernorm.weight",
"layers.24.mlp.gate_up.weight": "model.layers.24.mlp.gate_up_proj.weight",
"layers.24.mlp.down.weight": "model.layers.24.mlp.down_proj.weight",
"layers.25.attention_norm.weight": "model.layers.25.input_layernorm.weight",
"layers.25.attention.qkv.weight": "model.layers.25.self_attn.qkv_proj.weight",
"layers.25.attention.o_proj.weight": "model.layers.25.self_attn.o_proj.weight",
"layers.25.mlp_norm.weight": "model.layers.25.post_attention_layernorm.weight",
"layers.25.mlp.gate_up.weight": "model.layers.25.mlp.gate_up_proj.weight",
"layers.25.mlp.down.weight": "model.layers.25.mlp.down_proj.weight",
"layers.26.attention_norm.weight": "model.layers.26.input_layernorm.weight",
"layers.26.attention.qkv.weight": "model.layers.26.self_attn.qkv_proj.weight",
"layers.26.attention.o_proj.weight": "model.layers.26.self_attn.o_proj.weight",
"layers.26.mlp_norm.weight": "model.layers.26.post_attention_layernorm.weight",
"layers.26.mlp.gate_up.weight": "model.layers.26.mlp.gate_up_proj.weight",
"layers.26.mlp.down.weight": "model.layers.26.mlp.down_proj.weight",
"layers.27.attention_norm.weight": "model.layers.27.input_layernorm.weight",
"layers.27.attention.qkv.weight": "model.layers.27.self_attn.qkv_proj.weight",
"layers.27.attention.o_proj.weight": "model.layers.27.self_attn.o_proj.weight",
"layers.27.mlp_norm.weight": "model.layers.27.post_attention_layernorm.weight",
"layers.27.mlp.gate_up.weight": "model.layers.27.mlp.gate_up_proj.weight",
"layers.27.mlp.down.weight": "model.layers.27.mlp.down_proj.weight",
"layers.28.attention_norm.weight": "model.layers.28.input_layernorm.weight",
"layers.28.attention.qkv.weight": "model.layers.28.self_attn.qkv_proj.weight",
"layers.28.attention.o_proj.weight": "model.layers.28.self_attn.o_proj.weight",
"layers.28.mlp_norm.weight": "model.layers.28.post_attention_layernorm.weight",
"layers.28.mlp.gate_up.weight": "model.layers.28.mlp.gate_up_proj.weight",
"layers.28.mlp.down.weight": "model.layers.28.mlp.down_proj.weight",
"layers.29.attention_norm.weight": "model.layers.29.input_layernorm.weight",
"layers.29.attention.qkv.weight": "model.layers.29.self_attn.qkv_proj.weight",
"layers.29.attention.o_proj.weight": "model.layers.29.self_attn.o_proj.weight",
"layers.29.mlp_norm.weight": "model.layers.29.post_attention_layernorm.weight",
"layers.29.mlp.gate_up.weight": "model.layers.29.mlp.gate_up_proj.weight",
"layers.29.mlp.down.weight": "model.layers.29.mlp.down_proj.weight",
"layers.30.attention_norm.weight": "model.layers.30.input_layernorm.weight",
"layers.30.attention.qkv.weight": "model.layers.30.self_attn.qkv_proj.weight",
"layers.30.attention.o_proj.weight": "model.layers.30.self_attn.o_proj.weight",
"layers.30.mlp_norm.weight": "model.layers.30.post_attention_layernorm.weight",
"layers.30.mlp.gate_up.weight": "model.layers.30.mlp.gate_up_proj.weight",
"layers.30.mlp.down.weight": "model.layers.30.mlp.down_proj.weight",
"layers.31.attention_norm.weight": "model.layers.31.input_layernorm.weight",
"layers.31.attention.qkv.weight": "model.layers.31.self_attn.qkv_proj.weight",
"layers.31.attention.o_proj.weight": "model.layers.31.self_attn.o_proj.weight",
"layers.31.mlp_norm.weight": "model.layers.31.post_attention_layernorm.weight",
"layers.31.mlp.gate_up.weight": "model.layers.31.mlp.gate_up_proj.weight",
"layers.31.mlp.down.weight": "model.layers.31.mlp.down_proj.weight",
"norm.weight": "model.norm.weight",
"lm_head.weight": "lm_head.weight"
}linnet.toml59 B
toml
[package]
name = "phi3"
version = "0.1.0"
language = "0.1"nest.toml908 B
toml
[model]
name = "phi-3-mini-4k-instruct"
title = "Phi-3 Mini 4K Instruct"
summary = "A 3.8B-parameter decoder with fused QKV and gate/up projections, full (non-grouped) multi-head attention, RMS normalization, SwiGLU, and an untied output head, instruction-tuned for a 4K context."
license = "MIT"
family = "phi3"
tags = ["text-generation", "decoder-only", "chat"]
[links]
huggingface = "https://huggingface.co/microsoft/Phi-3-mini-4k-instruct"
arxiv = "https://arxiv.org/abs/2404.14219"
[source]
path = "src/lib.linnet"
root = "Model"
entry = "forward"
[generics]
Vocab = 32064
H = 3072
Heads = 32
Inner = 8192
Layers = 32
Batch = 1
MaxSeq = 4096
T = "bf16"
[check]
B = 1
S = 8
[weights]
repo = "microsoft/Phi-3-mini-4k-instruct"
revision = "f39ac1d28e925b323eae81227eaba4464caced4e"
files = [
"model-00001-of-00002.safetensors",
"model-00002-of-00002.safetensors",
]
bindings = "bindings.json"samples.json930 B
json
{
"kind": "text",
"prompt": "Explain in two sentences why the sky is blue.",
"chat": true,
"output": "The sky appears blue due to a phenomenon called Rayleigh scattering, where molecules and small particles in the Earth's atmosphere scatter sunlight in all directions, and blue light is scattered more than other colors because it travels as shorter, smaller waves. This scattered blue light is what we see when we look up at the sky on a clear day.",
"reference": "The sky appears blue due to a phenomenon called Rayleigh scattering, where molecules and small particles in the Earth's atmosphere scatter sunlight in all directions, and blue light is scattered more than other colors because it travels as shorter, smaller waves. This scattered blue light is what we see when we look up at the sky on a clear day.",
"reference_stack": "transformers (KV cache, argmax, bf16)",
"tokens": 76,
"agreeing_prefix": 76
}src
attention.linnet15.4 kB
linnet
// Full (non-grouped) multi-head attention with rotary positions. Phi-3
// fuses the query, key, and value projections into one checkpoint tensor.
//
// `forward` attends over a whole sequence. `decode` takes one position and
// keeps the keys and values seen so far in two `state` members, sized by
// the block's `Batch` and `MaxSeq`: the block assigns them, the runtime
// keeps them between calls, and graph exports thread them in and out.
// `prefill` is `decode` for a whole prompt: every position in one pass,
// slicing the fused `qkv_proj` result the same way `decode` does.
module phi3.attention
use crate.rope::{THETA, tables}
use std.nn.attention::{attention, attention_rows}
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}
// `qkv_proj.weight` is one `[3 * H, H]` tensor: `num_key_value_heads` equals
// `num_attention_heads` in `config.json` (32 and 32), so this is not
// grouped-query attention and the fused projection is exactly three
// `H`-wide slices, queries first. The reference `Phi3Attention.forward` in
// `transformers/models/phi3/modeling_phi3.py` slices the projection with
//
// query_states = qkv[..., :query_pos]
// key_states = qkv[..., query_pos : query_pos + num_key_value_heads * head_dim]
// value_states = qkv[..., query_pos + num_key_value_heads * head_dim :]
//
// which is query, then key, then value, in that order. `qkv_proj.weight`
// keeps the ordinary `nn.Linear` `[Out, In]` layout, so `Linear<H, 3 * H, T>`
// binds the checkpoint tensor with no transpose (`std.nn.linear::Linear`
// already expects `[Out, In]`), and `forward` slices the projected result
// the way `models/gpt2` slices its fused QKV, rather than splitting the
// weight in `bindings.json`.
// `sliding_window` in `config.json`: a query attends to itself and the
// `WINDOW - 1` positions before it, as `transformers` masks every layer.
// Every position is still cached; the mask leaves out all but the last
// `WINDOW` of them, which a sequence under that length never notices.
pub const WINDOW: i32 = 2047
// 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>() -> 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
}
pub block Attention<H: Dim, Heads: Dim, Batch: Dim, MaxSeq: Dim, T: Float>
where
Heads > 0,
H % Heads == 0,
(H / Heads) % 2 == 0,
MaxSeq > 0
{
sub qkv: Linear<H, 3 * H, T>
sub o_proj: Linear<H, H, T>
state cache_k: Tensor[Batch, Heads, MaxSeq, H / Heads; T]
state cache_v: Tensor[Batch, Heads, MaxSeq, H / Heads; T]
pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T]) -> Tensor[B, S, H; T] {
let (cos_table, sin_table) = tables<S, H / Heads, T>(THETA)
let projected = qkv.forward(x)
let q = rope(
split_heads<B, S, Heads, H / Heads, T>(projected[:, :, 0:H]),
cos_table,
sin_table,
)
let k = rope(
split_heads<B, S, Heads, H / Heads, T>(projected[:, :, H:2 * H]),
cos_table,
sin_table,
)
let v = split_heads<B, S, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
let mixed = attention(
q,
k,
v,
rsqrt(cast<f32>(H / Heads)),
some(window_mask<S, S>()),
)
return o_proj.forward(merge_heads<B, S, Heads, H / Heads, T>(mixed))
}
// One new position `pos`: its key and value are written into the caches
// at that position (a masked write, no scatter), and the query attends
// over every cached position up to and including it.
pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32) -> Tensor[Batch, 1, H; T] {
let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
let cos_at[d] = cos_table[cast<i64>(pos), d]
let sin_at[d] = sin_table[cast<i64>(pos), d]
let cos_row = reshape(cos_at, [1, H / Heads])
let sin_row = reshape(sin_at, [1, H / Heads])
let projected = qkv.forward(x)
let q = rope(
split_heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, 0:H]),
cos_row,
sin_row,
)
let k = rope(
split_heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, H:2 * H]),
cos_row,
sin_row,
)
let v = split_heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
cache_k = write_at(cache_k, k, pos)
cache_v = write_at(cache_v, v, pos)
let positions = iota<i32>(MaxSeq)
let seen[s] = positions[s] <= pos && positions[s] > pos - WINDOW
let mixed = attention(
q,
cache_k,
cache_v,
rsqrt(cast<f32>(H / Heads)),
some(reshape(seen, [1, MaxSeq])),
)
return o_proj.forward(merge_heads<Batch, 1, Heads, H / Heads, T>(mixed))
}
// `S` new positions starting at `pos`, as a prompt arrives: their keys
// and values go into the caches as one span, and each query attends over
// every cached position up to and including its own. The fused
// projection is sliced over `S` positions the same way `decode` slices
// it over one. A prompt costs one pass instead of one per token, which
// is most of the time to the first generated token.
pub fn prefill<S: Dim>(x: Tensor[Batch, S, H; T], pos: i32) -> Tensor[Batch, S, H; T]
where S > 0 {
let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
let rows[i] = cast<i64>(pos) + iota<i64>(S)[i]
let cos_rows[i, d] = cos_table[rows[i], d]
let sin_rows[i, d] = sin_table[rows[i], d]
let projected = qkv.forward(x)
let q = rope(
split_heads<Batch, S, Heads, H / Heads, T>(projected[:, :, 0:H]),
cos_rows,
sin_rows,
)
let k = rope(
split_heads<Batch, S, Heads, H / Heads, T>(projected[:, :, H:2 * H]),
cos_rows,
sin_rows,
)
let v = split_heads<Batch, S, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
cache_k = write_span(cache_k, k, pos)
cache_v = write_span(cache_v, v, pos)
let positions = iota<i32>(MaxSeq)
let queries = iota<i32>(S)
let seen[i, s] =
positions[s] <= pos + queries[i] && positions[s] > pos + queries[i] - WINDOW
let mixed = attention(q, cache_k, cache_v, rsqrt(cast<f32>(H / Heads)), some(seen))
return o_proj.forward(merge_heads<Batch, S, Heads, H / Heads, T>(mixed))
}
// One new position per sequence, each at its own `positions[b]`: the
// step a server takes for every request it is decoding together. Row `b`
// of the caches is written at `positions[b]`, and its query attends over
// that row up to there.
pub fn decode_rows(
x: Tensor[Batch, 1, H; T],
positions: Tensor[Batch; i32],
) -> Tensor[Batch, 1, H; T] {
let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
let cos_at[b, d] = cos_table[cast<i64>(positions[b]), d]
let sin_at[b, d] = sin_table[cast<i64>(positions[b]), d]
let cos_rows = reshape(cos_at, [Batch, 1, H / Heads])
let sin_rows = reshape(sin_at, [Batch, 1, H / Heads])
let projected = qkv.forward(x)
let q = rope_rows(
split_heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, 0:H]),
cos_rows,
sin_rows,
)
let k = rope_rows(
split_heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, H:2 * H]),
cos_rows,
sin_rows,
)
let v = split_heads<Batch, 1, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
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 = attention_rows(
q,
cache_k,
cache_v,
rsqrt(cast<f32>(H / Heads)),
reshape(seen, [Batch, 1, MaxSeq]),
)
return o_proj.forward(merge_heads<Batch, 1, Heads, H / Heads, T>(mixed))
}
// `M` requests' prompts of `S` tokens, from their first positions, into
// rows `slots` of the caches: each prompt's keys and values fill its row,
// and its queries attend causally among themselves. The rest of a row
// keeps what an earlier request left there; decoding never reads past its
// own position, so nothing needs clearing.
pub fn prefill_slots<M: Dim, S: Dim>(
x: Tensor[M, S, H; T],
slots: Tensor[M; i32],
) -> Tensor[M, S, H; T]
where
M > 0,
S > 0
{
let (cos_table, sin_table) = tables<S, H / Heads, T>(THETA)
let projected = qkv.forward(x)
let q = rope(
split_heads<M, S, Heads, H / Heads, T>(projected[:, :, 0:H]),
cos_table,
sin_table,
)
let k = rope(
split_heads<M, S, Heads, H / Heads, T>(projected[:, :, H:2 * H]),
cos_table,
sin_table,
)
let v = split_heads<M, S, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
cache_k = write_slots(cache_k, k, slots, 0)
cache_v = write_slots(cache_v, v, slots, 0)
let mixed = attention(
q,
k,
v,
rsqrt(cast<f32>(H / Heads)),
some(window_mask<S, S>()),
)
return o_proj.forward(merge_heads<M, S, Heads, H / Heads, T>(mixed))
}
// Prompts packed end to end, `P` tokens in all: token `p` is position
// `positions[p]` of prompt `segments[p]`, whose row of the caches is
// `rows[p]`. Each token's key and value go to its row at its position, and
// its query attends to the tokens of its own prompt up to itself -- no
// padding between prompts, and none of the pass's other prompts seen.
pub fn prefill_packed<P: Dim>(
x: Tensor[1, P, H; T],
rows: Tensor[P; i32],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
) -> Tensor[1, P, H; T]
where P > 0 {
let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
let cos_at[p, d] = cos_table[cast<i64>(positions[p]), d]
let sin_at[p, d] = sin_table[cast<i64>(positions[p]), d]
let projected = qkv.forward(x)
let q = rope(
split_heads<1, P, Heads, H / Heads, T>(projected[:, :, 0:H]),
cos_at,
sin_at,
)
let k = rope(
split_heads<1, P, Heads, H / Heads, T>(projected[:, :, H:2 * H]),
cos_at,
sin_at,
)
let v = split_heads<1, P, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
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 = attention(
q,
k,
v,
rsqrt(cast<f32>(H / Heads)),
some(own),
)
return o_proj.forward(merge_heads<1, P, Heads, H / Heads, T>(mixed))
}
// `prefill_packed` for training: every packed token's output, and no
// cache written.
pub fn packed<P: Dim>(
x: Tensor[1, P, H; T],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
) -> Tensor[1, P, H; T]
where P > 0 {
let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
let cos_at[p, d] = cos_table[cast<i64>(positions[p]), d]
let sin_at[p, d] = sin_table[cast<i64>(positions[p]), d]
let projected = qkv.forward(x)
let q = rope(
split_heads<1, P, Heads, H / Heads, T>(projected[:, :, 0:H]),
cos_at,
sin_at,
)
let k = rope(
split_heads<1, P, Heads, H / Heads, T>(projected[:, :, H:2 * H]),
cos_at,
sin_at,
)
let v = split_heads<1, P, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
let own[i, j] =
segments[i] == segments[j] &&
positions[j] <= positions[i] &&
positions[j] > positions[i] - WINDOW
let mixed = attention(
q,
k,
v,
rsqrt(cast<f32>(H / Heads)),
some(own),
)
return o_proj.forward(merge_heads<1, P, Heads, H / Heads, T>(mixed))
}
// `prefill_packed`'s prompts, `P` tokens, then one step for each of the
// `Batch` rows, at `step_positions[b]`: every token's key and value go to
// its row at its position, the prompts attend among themselves as in
// `prefill_packed`, and each step over its row as in `decode_rows`.
pub fn step_packed<P: Dim>(
x: Tensor[1, P + Batch, H; T],
rows: Tensor[P; i32],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
step_positions: Tensor[Batch; i32],
) -> Tensor[1, P + Batch, H; T]
where P > 0 {
let every_at = concat(positions, step_positions, axis = 0)
let every_row = concat(rows, iota<i32>(Batch), axis = 0)
let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
let cos_at[p, d] = cos_table[cast<i64>(every_at[p]), d]
let sin_at[p, d] = sin_table[cast<i64>(every_at[p]), d]
let projected = qkv.forward(x)
let q = rope(
split_heads<1, P + Batch, Heads, H / Heads, T>(projected[:, :, 0:H]),
cos_at,
sin_at,
)
let k = rope(
split_heads<1, P + Batch, Heads, H / Heads, T>(projected[:, :, H:2 * H]),
cos_at,
sin_at,
)
let v = split_heads<1, P + Batch, Heads, H / Heads, T>(projected[:, :, 2 * H:3 * H])
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 = attention(
q[:, :, 0:P, :],
k[:, :, 0:P, :],
v[:, :, 0:P, :],
rsqrt(cast<f32>(H / Heads)),
some(own),
)
let slots = iota<i32>(MaxSeq)
let seen[b, s] = slots[s] <= step_positions[b] && slots[s] > step_positions[b] - WINDOW
let steps = attention_rows(
permute(q[:, :, P:P + Batch, :], [2, 1, 0, 3]),
cache_k,
cache_v,
rsqrt(cast<f32>(H / Heads)),
reshape(seen, [Batch, 1, MaxSeq]),
)
let mixed = concat(prompts, permute(steps, [2, 1, 0, 3]), axis = 2)
return o_proj.forward(merge_heads<1, P + Batch, Heads, H / Heads, T>(mixed))
}
}
// `[B, S, N * D]` -> `[B, N, S, D]`.
fn split_heads<B: Dim, S: Dim, N: Dim, D: Dim, T: Float>(
x: Tensor[B, S, N * D; T],
) -> Tensor[B, N, S, D; T] {
return permute(reshape(x, [B, S, N, D]), [0, 2, 1, 3])
}
fn merge_heads<B: Dim, S: Dim, N: Dim, D: Dim, T: Float>(
x: Tensor[B, N, S, D; T],
) -> Tensor[B, S, N * D; T] {
return reshape(permute(x, [0, 2, 1, 3]), [B, S, N * D])
}lib.linnet14.5 kB
linnet
// Phi-3 Mini 4K Instruct: RMSNorm, full (non-grouped) multi-head attention
// with rotary positions and a fused QKV projection, a SwiGLU MLP with a
// fused gate/up projection, and an untied output head. No projection in the
// checkpoint carries a bias. Widths, head count, and depth are generic
// parameters; the dtype defaults to bf16 and every intermediate
// accumulation is spelled out in the standard library. The `decode` entry
// generates one token at a time from the KV caches the attention blocks
// own, sized by `Batch` and `MaxSeq`.
module phi3
use crate.attention::{Attention}
use crate.mlp::{FusedSwiGlu}
use std.nn.decoding::{argmax}
use std.random::{categorical, split}
use std.nn.embedding::{Embedding}
use std.nn.linear::{Linear}
use std.nn.loss::{linear_cross_entropy, linear_token_log_probs}
use std.nn.norm::{RmsNorm}
pub block DecoderLayer<H: Dim, Heads: Dim, Inner: Dim, Batch: Dim, MaxSeq: Dim, T: Float>
where
Heads > 0,
H % Heads == 0,
(H / Heads) % 2 == 0,
MaxSeq > 0
{
sub attention_norm: RmsNorm<H, T>
sub attention: Attention<H, Heads, Batch, MaxSeq, T>
sub mlp_norm: RmsNorm<H, T>
sub mlp: FusedSwiGlu<H, Inner, T>
pub fn forward<B: Dim, S: Dim>(x: Tensor[B, S, H; T]) -> Tensor[B, S, H; T] {
let attended = x + attention.forward(attention_norm.forward(x))
return attended + mlp.forward(mlp_norm.forward(attended))
}
pub fn decode(x: Tensor[Batch, 1, H; T], pos: i32) -> Tensor[Batch, 1, H; T] {
let attended = x + attention.decode(attention_norm.forward(x), pos)
return attended + mlp.forward(mlp_norm.forward(attended))
}
pub fn prefill<S: Dim>(x: Tensor[Batch, S, H; T], pos: i32) -> Tensor[Batch, S, H; T]
where S > 0 {
let attended = x + attention.prefill(attention_norm.forward(x), pos)
return attended + mlp.forward(mlp_norm.forward(attended))
}
pub fn decode_rows(
x: Tensor[Batch, 1, H; T],
positions: Tensor[Batch; i32],
) -> Tensor[Batch, 1, H; T] {
let attended = x + attention.decode_rows(attention_norm.forward(x), positions)
return attended + mlp.forward(mlp_norm.forward(attended))
}
pub fn prefill_slots<M: Dim, S: Dim>(
x: Tensor[M, S, H; T],
slots: Tensor[M; i32],
) -> Tensor[M, S, H; T]
where
M > 0,
S > 0
{
let attended = x + attention.prefill_slots(attention_norm.forward(x), slots)
return attended + mlp.forward(mlp_norm.forward(attended))
}
pub fn prefill_packed<P: Dim>(
x: Tensor[1, P, H; T],
rows: Tensor[P; i32],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
) -> Tensor[1, P, H; T]
where P > 0 {
let attended =
x + attention.prefill_packed(attention_norm.forward(x), rows, positions, segments)
return attended + mlp.forward(mlp_norm.forward(attended))
}
pub fn packed<P: Dim>(
x: Tensor[1, P, H; T],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
) -> Tensor[1, P, H; T]
where P > 0 {
let attended = x + attention.packed(attention_norm.forward(x), positions, segments)
return attended + mlp.forward(mlp_norm.forward(attended))
}
pub fn step_packed<P: Dim>(
x: Tensor[1, P + Batch, H; T],
rows: Tensor[P; i32],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
step_positions: Tensor[Batch; i32],
) -> Tensor[1, P + Batch, H; T]
where P > 0 {
let attended =
x +
attention.step_packed(
attention_norm.forward(x),
rows,
positions,
segments,
step_positions,
)
return attended + mlp.forward(mlp_norm.forward(attended))
}
}
pub block Model<
Vocab: Dim,
H: Dim,
Heads: Dim,
Inner: Dim,
Layers: Dim,
Batch: Dim,
MaxSeq: Dim,
T: Float = bf16,
>
where
Heads > 0,
H % Heads == 0,
(H / Heads) % 2 == 0,
MaxSeq > 0
{
sub embedding: Embedding<Vocab, H, T>
sub layers: [DecoderLayer<H, Heads, Inner, Batch, MaxSeq, T>; Layers]
sub norm: RmsNorm<H, T>
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(hidden(tokens))
}
// Logits for the last position only, as decoding needs: the final norm
// and projection run on one row per sequence.
pub entry next_token<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, Vocab; T]
where S > 0 {
let last = hidden_states(tokens)[:, S - 1, :]
return lm_head.forward(norm.forward(last))
}
// Logits for one new token per sequence at position `pos`, attending
// over the positions decoded before it through the layers' KV caches.
// Feed a prompt token by token, then the tokens the model produces.
pub entry decode(token: Tensor[Batch, 1; i32], pos: i32) -> Tensor[Batch, Vocab; T] {
return step(token, pos)
}
// A prompt of `S` tokens at `pos`, written into the KV caches in one
// pass. Returns the logits after its last token, so `decode` continues
// at `pos + S`: the time to the first token is this one call.
pub entry prefill<S: Dim>(tokens: Tensor[Batch, S; i32], pos: i32) -> Tensor[Batch, Vocab; T]
where S > 0 {
var x = embedding.forward(tokens)
static for layer in layers {
x = layer.prefill(x, pos)
}
return 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 = embedding.forward(tokens)
static for layer in layers {
x = layer.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 = embedding.forward(tokens)
static for layer in layers {
x = layer.prefill_slots(x, slots)
}
let last[b, h] = x[b, cast<i64>(lengths[b]) - 1, h]
return lm_head.forward(norm.forward(last))
}
// Training over sequences packed into one row, as `prefill_packed`
// packs prompts: `sum_p weights[p] * -log p(targets[p])` for the token
// after each position (`std.nn.loss::linear_cross_entropy` with the
// output head). A weight of 0 leaves a position out; `mask / count`
// gives the mean.
pub entry loss_packed<P: Dim>(
tokens: Tensor[P; i32],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
targets: Tensor[P; i64],
weights: Tensor[P; f32],
) -> f32
where P > 0 {
let states = packed_states(tokens, positions, segments)
return linear_cross_entropy(states, lm_head.weight, targets, weights)
}
// Each packed position's log-probability of `targets[p]`, as a
// policy-gradient update weighs them.
pub entry log_probs_packed<P: Dim>(
tokens: Tensor[P; i32],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
targets: Tensor[P; i64],
) -> Tensor[P; f32]
where P > 0 {
let states = packed_states(tokens, positions, segments)
return linear_token_log_probs(states, lm_head.weight, targets)
}
// Each packed position's final hidden state, after the last norm: what
// the output head (or a value or reward head) reads.
pub entry hidden_packed<P: Dim>(
tokens: Tensor[P; i32],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
) -> Tensor[P, H; T]
where P > 0 {
return packed_states(tokens, positions, segments)
}
fn packed_states<P: Dim>(
tokens: Tensor[P; i32],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
) -> Tensor[P, H; T]
where P > 0 {
var x = embedding.forward(reshape(tokens, [1, P]))
static for layer in layers {
x = layer.packed(x, positions, segments)
}
return reshape(norm.forward(x), [P, H])
}
// Prompts of several requests packed end to end into one pass of `P`
// tokens, with no padding between them: token `p` is position
// `positions[p]` of prompt `segments[p]`, and its key and value go to row
// `rows[p]` of the caches. A pack may end in padding that writes where no
// request reads (the engine uses position `MaxSeq - 1`, which decoding
// never reaches) with a segment of its own. Returns the logits after
// prompt `m`'s last token, `last[m]`, for up to `Batch` prompts; rows
// past the pass's prompts repeat whichever token `last` names.
pub entry prefill_packed<P: Dim>(
tokens: Tensor[P; i32],
rows: Tensor[P; i32],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
last: Tensor[Batch; i32],
) -> Tensor[Batch, Vocab; T]
where P > 0 {
var x = embedding.forward(reshape(tokens, [1, P]))
static for layer in layers {
x = layer.prefill_packed(x, rows, positions, segments)
}
let ends[m, h] = x[0, cast<i64>(last[m]), h]
return 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. A row the prompts fill steps along,
// written where no request reads (the engine uses `MaxSeq - 1`). Returns
// the logits after each prompt, as `prefill_packed`, then each row's step.
pub entry step_packed<P: Dim>(
tokens: Tensor[P; i32],
rows: Tensor[P; i32],
positions: Tensor[P; i32],
segments: Tensor[P; i32],
last: Tensor[Batch; i32],
step_tokens: Tensor[Batch, 1; i32],
step_positions: Tensor[Batch; i32],
) -> Tensor[2 * Batch, Vocab; T]
where P > 0 {
let every = concat(tokens, reshape(step_tokens, [Batch]), axis = 0)
var x = embedding.forward(reshape(every, [1, P + Batch]))
static for layer in layers {
x = layer.step_packed(x, rows, positions, segments, step_positions)
}
let ends[m, h] = x[0, cast<i64>(last[m]), h]
return lm_head.forward(norm.forward(concat(ends, x[0, P:P + Batch, :], axis = 0)))
}
// Greedy generation inside the graph: `Steps` tokens after `token` at
// `pos`, each decode step feeding the caches and its `argmax` the next
// step. The range loop is expanded at compile time, so the whole
// generation exports as one program.
pub entry generate<Steps: Dim>(
token: Tensor[Batch, 1; i32],
pos: i32,
) -> Tensor[Batch, Steps; i32]
where Steps > 0 {
var current = token
var produced = fill<i32>([Batch, Steps], 0)
let slots = iota<i32>(Steps)
static for i in 0..Steps {
let next = argmax(step(current, pos + cast<i32>(i)))
let written[b, s] = select(slots[s] == cast<i32>(i), next[b], produced[b, s])
produced = written
current = reshape(next, [Batch, 1])
}
return produced
}
// Sampling with an explicit key: `Steps` tokens drawn from the softmax
// of the logits divided by `temperature`, one derived key per step, so
// the same key gives the same text on every backend.
pub entry sample<Steps: Dim>(
token: Tensor[Batch, 1; i32],
pos: i32,
key: Tensor[2; i64],
temperature: f32,
) -> Tensor[Batch, Steps; i32]
where Steps > 0 {
var current = token
var produced = fill<i32>([Batch, Steps], 0)
let slots = iota<i32>(Steps)
let keys = split<Steps>(key)
static for i in 0..Steps {
let logits = cast<f32>(step(current, pos + cast<i32>(i))) / temperature
let key_i[j] = keys[i, j]
let next = categorical(key_i, logits)
let written[b, s] = select(slots[s] == cast<i32>(i), next[b], produced[b, s])
produced = written
current = reshape(next, [Batch, 1])
}
return produced
}
// Greedy generation that stops early: at most `MaxNew` tokens, or fewer
// once every sequence has produced `eos`. A runtime `while` loop over
// scalar state; the count of tokens produced is returned with them, and
// positions past it hold zeros.
pub entry generate_until<MaxNew: Dim>(
token: Tensor[Batch, 1; i32],
pos: i32,
eos: i32,
) -> (Tensor[Batch, MaxNew; i32], i32)
where MaxNew > 0 {
var current = token
var produced = fill<i32>([Batch, MaxNew], 0)
var count: i32 = 0
var running = true
let slots = iota<i32>(MaxNew)
while running && count < MaxNew {
let next = argmax(step(current, pos + count))
let written[b, s] = select(slots[s] == count, next[b], produced[b, s])
produced = written
current = reshape(next, [Batch, 1])
count = count + 1
let finished = all[b] (next[b] == eos)
running = !finished
}
return (produced, count)
}
fn step(token: Tensor[Batch, 1; i32], pos: i32) -> Tensor[Batch, Vocab; T] {
var x = embedding.forward(token)
static for layer in layers {
x = layer.decode(x, pos)
}
return lm_head.forward(norm.forward(x[:, 0, :]))
}
fn hidden<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, H; T] {
return norm.forward(hidden_states(tokens))
}
fn hidden_states<B: Dim, S: Dim>(tokens: Tensor[B, S; i32]) -> Tensor[B, S, H; T] {
var x = embedding.forward(tokens)
static for layer in layers {
x = layer.forward(x)
}
return x
}
}mlp.linnet1.5 kB
linnet
// Phi-3's gated MLP, with the gate and up projections fused into one
// checkpoint tensor.
module phi3.mlp
use std.nn.activations::{silu}
use std.nn.linear::{Linear}
// `gate_up_proj.weight` is one `[2 * Inner, H]` tensor computing the gate
// and up projections side by side; `down_proj.weight` is separate. Which
// half is which is not obvious from the shape, so it is worth stating how
// it was determined: the reference `Phi3MLP.forward` in
// `transformers/models/phi3/modeling_phi3.py` reads
//
// up_states = self.gate_up_proj(hidden_states)
// gate, up_states = up_states.chunk(2, dim=-1)
// up_states = up_states * self.activation_fn(gate)
//
// so `chunk(2, dim=-1)` splits the output axis in the order the tensor is
// stored: the first `Inner` rows of `gate_up_proj.weight` produce the gate
// (the half passed through the activation), and the second `Inner` rows
// produce the value that gate multiplies. Both `gate_up_proj.weight` and
// `down_proj.weight` keep the ordinary `nn.Linear` `[Out, In]` layout, so
// `Linear<H, 2 * Inner, T>` binds the checkpoint tensor with no transpose
// and `forward` slices the projected result rather than the weight.
pub block FusedSwiGlu<H: Dim, Inner: Dim, T: Float = bf16> {
sub gate_up: Linear<H, 2 * Inner, T>
sub down: Linear<Inner, H, T>
pub fn forward<*S: Shape>(x: Tensor[*S, H; T]) -> Tensor[*S, H; T] {
let projected = gate_up.forward(x)
let gate = projected[..., 0:Inner]
let up = projected[..., Inner:2 * Inner]
return down.forward(silu(gate) * up)
}
}rope.linnet1.0 kB
linnet
// Rotary position tables computed from the compile-time sequence length:
// the positions and frequencies come from `iota`, so a model needs no
// precomputed inputs and the tables have exactly the length of the sequence.
module phi3.rope
// Phi-3 Mini 4K uses the standard rotary embedding: base frequency 10000,
// and `rope_scaling` is `null` in `config.json`, so no long-context scaling
// applies. `partial_rotary_factor` is not set either, so the whole head
// dimension rotates, as `rope` below assumes.
pub const THETA: f32 = 10000.0
// `(cos, sin)` tables of shape `[S, D]`, each frequency repeated for both
// halves of the head dimension as `rope` expects.
pub fn tables<S: Dim, D: Dim, T: Float>(theta: f32) -> (Tensor[S, D; T], Tensor[S, D; T])
where D % 2 == 0 {
let inv_freq[i] = exp(-(cast<f32>(iota<i64>(D / 2)[i]) * 2.0 / cast<f32>(D)) * log(theta))
let angles[s, i] = cast<f32>(iota<i64>(S)[s]) * inv_freq[i]
let full = concat(angles, angles, axis = -1)
return (cast<T>(cos(full)), cast<T>(sin(full)))
}