Skip to content

Llama ​

A Llama-style decoder in three files: grouped-query attention, rotary positions, SwiGLU, a KV cache, and entries from a full forward pass to generation loops and packed training. Its dtype generic T defaults to bf16. The Nest card llama-3.1-8b-instruct is this structure with the serving entries 04-serve runs, and the training examples train it.

Commands ​

bash
linnet lint --std stdlib examples/01-llama
linnet inspect --parameters --std stdlib examples/01-llama/src/lib.linnet
linnet stablehlo --std stdlib --entry decode --bind Vocab=32000 --bind H=512 --bind Heads=8 \
                 --bind KvHeads=8 --bind Inner=1376 --bind Layers=8 --bind Batch=1 \
                 --bind MaxSeq=512 --bind T=bf16 examples/01-llama/src/lib.linnet
python
from linnet.torch import load
model = load("examples/01-llama/src/lib.linnet", generics={...}, weights="weights/",
             numerics="equivalent", compile="inductor")
tokens = model.run_entry("generate", [prompt, torch.tensor(0, dtype=torch.int32)],
                         generics={"Steps": 32})

Key ideas ​

Modules and rotary tables ​

src/lib.linnet imports GroupedQueryAttention from src/attention.linnet, which imports THETA and tables from src/rope.linnet, both through crate. imports. THETA is the base frequency, a pub const. tables<S, D, T> computes the cos and sin tables at compile time from iota, so the model needs no table inputs.

Grouped-query attention ​

GroupedQueryAttention<H, Heads, KvHeads, Batch, MaxSeq, T> projects KvHeads key/value heads. std.nn.attention::grouped_attention repeats each for Heads / KvHeads query heads with broadcast_to and a reshape. The block's where clause states the divisibility the reshapes rely on.

KV cache in state ​

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

decode(x, pos) writes the new key and value at pos with std.nn.cache::write_at, a masked select, then attends over positions <= pos. The runtime keeps the caches between calls; PyTorch holds them in buffers, reset with reset_state(). Graph exports thread them in and out as extra arguments and results.

Entries ​

EntryReturns
forwardlogits for every position
next_tokenlast-position logits, sliced with S - 1
decodelogits for one token at pos, through the per-layer KV caches
generate<Steps>Steps greedy tokens from std.nn.decoding::argmax
sample<Steps>Steps tokens drawn with a key
generate_until<MaxNew>up to MaxNew greedy tokens and their count
loss_packed<P>training loss over sequences packed into one row
log_probs_packed<P>each packed position's log-probability of its target
hidden_packed<P>each packed position's final hidden state

generate loops with static for i in 0..Steps, expanded at compile time, so the whole generation is one graph. generate_until loops while running && count < MaxNew at runtime over scalar carried values, and stops once every row has produced the end token. It exports as stablehlo.while, an ONNX Loop, or a Python loop.

sample splits key into one key per step and draws with std.random::categorical. The same key gives the same tokens in PyTorch and under XLA.

loss_packed, log_probs_packed and hidden_packed take several sequences packed into one row: positions gives each token's place in its sequence and segments which sequence it belongs to, and a token attends only within its own. The output head and its loss run through std.nn.loss::linear_cross_entropy, which PyTorch computes a block of tokens at a time.

Every layer is library source ​

The layers come from the standard library: grouped_attention, rope and linear_cross_entropy are ops, and RmsNorm, SwiGlu and Embedding are blocks, all in ordinary Linnet under stdlib/. Each op has a body that defines its result. A backend may run an op as a native kernel instead, and linnet explain shows each choice and why:

text
$ linnet explain --std stdlib --numerics fast examples/01-llama/src/lib.linnet
std.nn.loss::linear_cross_entropy
  called from llama::Model.loss_packed at examples/01-llama/src/lib.linnet:102:16
  implementations:
      canonical decomposition (exact)
    * linnet.linear_cross_entropy (numerically equivalent)
        requires PyTorch: blocks of rows, the gradient computed with the loss
  selected: linnet.linear_cross_entropy
  reason: strongest implementation allowed by the numerics policy

A new layer, whether another attention, norm or loss, is a new op in a library and needs no compiler change. Every backend runs its body from the start; a native kernel for it is an optimization added later.

Source ​

linnet.toml ​

toml
[package]
name = "llama"
version = "0.1.0"
language = "0.1"

src/attention.linnet ​

linnet
// Grouped-query attention: `KvHeads` key/value heads shared by `Heads`
// query heads, through `std.nn.attention::grouped_attention`. The `where`
// clause states the divisibility the reshapes depend on; the checker proves
// every shape from it.
//
// `forward` attends over a whole sequence. `decode` takes one position and
// keeps the keys and values seen so far in two `state` members, sized by
// the block's `Batch` and `MaxSeq`: the block assigns them, the runtime
// keeps them between calls, and graph exports thread them in and out.
module llama.attention

use crate.rope::{THETA, tables}
use std.nn.attention::{causal_mask, grouped_attention}
use std.nn.cache::{write_at}
use std.nn.linear::{Linear}
use std.nn.rope::{rope}

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

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

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

    // `forward` over several sequences packed into one row, for training:
    // token `p` sits at `positions[p]` of sequence `segments[p]` and attends
    // to its own sequence's tokens up to that position. Many short sequences
    // go through in one call, with no padding and no cache.
    pub fn packed<P: Dim>(
        x: Tensor[1, P, H; T],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
    ) -> Tensor[1, P, H; T] {
        let (cos_table, sin_table) = tables<MaxSeq, H / Heads, T>(THETA)
        let cos_at[p, d] = cos_table[cast<i64>(positions[p]), d]
        let sin_at[p, d] = sin_table[cast<i64>(positions[p]), d]
        let q = rope(split_heads<1, P, Heads, H / Heads, T>(q_proj.forward(x)), cos_at, sin_at)
        let k = rope(split_heads<1, P, KvHeads, H / Heads, T>(k_proj.forward(x)), cos_at, sin_at)
        let v = split_heads<1, P, KvHeads, H / Heads, T>(v_proj.forward(x))
        let own[i, j] = segments[i] == segments[j] && positions[j] <= positions[i]
        let mixed = grouped_attention(q, k, v, rsqrt(cast<f32>(H / Heads)), some(own))
        return o_proj.forward(merge_heads<1, P, 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 q = rope(
            split_heads<Batch, 1, Heads, H / Heads, T>(q_proj.forward(x)),
            cos_row,
            sin_row,
        )
        let k = rope(
            split_heads<Batch, 1, KvHeads, H / Heads, T>(k_proj.forward(x)),
            cos_row,
            sin_row,
        )
        let v = split_heads<Batch, 1, KvHeads, H / Heads, T>(v_proj.forward(x))
        cache_k = write_at(cache_k, k, pos)
        cache_v = write_at(cache_v, v, pos)
        let positions = iota<i32>(MaxSeq)
        let seen[s] = positions[s] <= pos
        let mixed = grouped_attention(
            q,
            cache_k,
            cache_v,
            rsqrt(cast<f32>(H / Heads)),
            some(reshape(seen, [1, MaxSeq])),
        )
        return o_proj.forward(merge_heads<Batch, 1, 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])
}

src/lib.linnet ​

linnet
// A Llama-style decoder: RMSNorm, grouped-query attention with rotary
// positions, SwiGLU, and an untied output head. Widths, head counts, and
// depth are generic parameters; the dtype defaults to bf16 and every
// intermediate accumulation is spelled out in the standard library. The
// `decode` entry generates one token at a time from the KV caches the
// attention blocks own, sized by `Batch` and `MaxSeq`.
module llama

use crate.attention::{GroupedQueryAttention}
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.mlp::{SwiGlu}
use std.nn.norm::{RmsNorm}

pub block DecoderLayer<
    H: Dim,
    Heads: Dim,
    KvHeads: Dim,
    Inner: Dim,
    Batch: Dim,
    MaxSeq: Dim,
    T: Float,
>
where
    Heads > 0,
    KvHeads > 0,
    H % Heads == 0,
    Heads % KvHeads == 0,
    (H / Heads) % 2 == 0,
    MaxSeq > 0
{
    sub attention_norm: RmsNorm<H, T>
    sub attention: GroupedQueryAttention<H, Heads, KvHeads, Batch, MaxSeq, T>
    sub mlp_norm: RmsNorm<H, T>
    sub mlp: SwiGlu<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 packed<P: Dim>(
        x: Tensor[1, P, H; T],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
    ) -> Tensor[1, P, H; T] {
        let attended = x + attention.packed(attention_norm.forward(x), positions, segments)
        return attended + mlp.forward(mlp_norm.forward(attended))
    }
}

pub block Model<
    Vocab: Dim,
    H: Dim,
    Heads: Dim,
    KvHeads: Dim,
    Inner: Dim,
    Layers: Dim,
    Batch: Dim,
    MaxSeq: Dim,
    T: Float = bf16,
>
where
    Heads > 0,
    KvHeads > 0,
    H % Heads == 0,
    Heads % KvHeads == 0,
    (H / Heads) % 2 == 0,
    MaxSeq > 0
{
    sub embedding: Embedding<Vocab, H, T>
    sub layers: [DecoderLayer<H, Heads, KvHeads, 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))
    }

    // Training over sequences packed into one row (see the attention's
    // `packed`): `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 (a prompt, a sequence's last
    // token); `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 {
        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] {
        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] {
        return packed_states(tokens, positions, segments)
    }

    // 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)
    }

    // 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 packed_states<P: Dim>(
        tokens: Tensor[P; i32],
        positions: Tensor[P; i32],
        segments: Tensor[P; i32],
    ) -> Tensor[P, H; T] {
        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])
    }

    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
    }
}

src/rope.linnet ​

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 llama.rope

// Llama 3 stretched the base frequency from 10000 to 500000.
pub const THETA: f32 = 500000.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)))
}

Browse on GitHub

Released under the MIT License.