Models / gpt2

GPT-2 (124M)

OpenAI's 124M-parameter GPT-2: learned positions, pre-norm blocks, a fused QKV projection, and a head tied to the token embedding.

124.4M parametersgpt2MITtext-generationdecoder-only

OpenAI's smallest GPT-2: 12 pre-norm blocks of width 768 with 12 heads, learned absolute positions up to 1024, a fused QKV projection split by slicing, the tanh GELU, and an output head tied to the token embedding. The Linnet source keeps GPT-2's [in, out] projection layout (Conv1D in the reference code), so the published checkpoint loads as it is.

Loading ​

python
from linnet import nest

model = nest.load("gpt2", backend="torch", numerics="equivalent")
logits = model(tokens)                      # Tensor[B, S; i32] -> Tensor[B, S, 50257; f32]

bindings.json maps Linnet parameter paths (blocks.0.attn.qkv.weight) to the checkpoint's names (h.0.attn.c_attn.weight). Sequences longer than 1024 positions are rejected at compile time (where S <= MaxPositions).

Entries ​

Entry
forward<B, S>(tokens)logits for a whole sequence
decode(token, pos)one token through the KV caches (Batch = 1, MaxSeq = 1024)
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
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

Provenance ​

The Linnet forward pass matches transformers logits to 2e-2 absolute (see python/linnet/tests/torch/test_hf_checkpoints.py in the Linnet repository).

Backends ​

The same source and checkpoint in each framework; numerics and compile are the loaders' options.

PyTorch ​

python
from linnet import nest

model = nest.load("gpt2", backend="torch", numerics="fast", compile="inductor")
logits = model(tokens)

JAX ​

python
from linnet import nest

f = nest.load("gpt2", backend="jax")          # StableHLO compiled by XLA
logits = f(tokens)

Flax NNX ​

python
from linnet import nest

model = nest.load("gpt2", backend="nnx")
logits = model(tokens)

Command line ​

bash
python -m linnet.nest pull gpt2
linnet stablehlo --root Model --entry forward --bind Vocab=50257 --bind MaxPositions=1024 --bind H=768 --bind Heads=12 --bind Layers=12 --bind Batch=1 --bind MaxSeq=1024 --bind T=f32 --bind B=1 --bind S=8 <model dir>/gpt2.linnet