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
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.linnetfrom 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
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
| Entry | Returns |
|---|---|
forward | logits for every position |
next_token | last-position logits, sliced with S - 1 |
decode | logits 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:
$ 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 policyA 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
[package]
name = "llama"
version = "0.1.0"
language = "0.1"src/attention.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
// 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
// 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)))
}