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
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
- Weights: openai-community/gpt2, MIT.
- Code: openai/gpt-2.
- Announcement: Better language models and their implications.
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
from linnet import nest
model = nest.load("gpt2", backend="torch", numerics="fast", compile="inductor")
logits = model(tokens)JAX
from linnet import nest
f = nest.load("gpt2", backend="jax") # StableHLO compiled by XLA
logits = f(tokens)Flax NNX
from linnet import nest
model = nest.load("gpt2", backend="nnx")
logits = model(tokens)Command line
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