Models / qwen3-8b

Qwen3 8B

An 8B-parameter Qwen3 decoder with grouped-query attention, per-head RMS normalization of the query and key projections, unbiased projections, and an untied output head.

8.2B parametersqwen3Apache-2.0text-generationdecoder-onlygrouped-query-attentionqk-norm

An 8.2B-parameter decoder with the Qwen3 architecture: 36 layers of width 4096, grouped-query attention (32 query heads, 8 key/value heads, head dimension 128 -- not H / Heads, which is 128 here only by coincidence; Qwen3-4B's head dimension is the same 128 while its width is 2560), rotary positions with base 1,000,000, RMS normalization with epsilon 1e-6, and a SwiGLU MLP of width 12288. The query and key projections carry no bias (attention_bias is false in config.json, unlike Qwen2.5); the change from Qwen2 that does matter here is QK-norm, an RMS normalization of every head's head_dim-wide slice of the query and key projections, applied after the projection is split into heads but before rope turns it. The output head is untied from the token embedding for this checkpoint.

The Linnet source is a self-contained Qwen3 package (src/lib.linnet, src/attention.linnet, src/norm.linnet, src/rope.linnet), adapted from the Qwen2 decoder models/qwen2.5-0.5b-instruct uses. Three changes from that source:

  • HeadDim is its own generic rather than H / Heads: Qwen3 sizes the query, key, and value projections by Heads * HeadDim and KvHeads * HeadDim, and the two need not agree with H (they don't, in the 4B checkpoint of this family). q_proj/k_proj/v_proj therefore map H to Heads * HeadDim / KvHeads * HeadDim, and o_proj maps Heads * HeadDim back to H, not H to H.
  • q_norm and k_norm, in src/norm.linnet's RmsNorm block instantiated at width HeadDim, called on the [B, Heads, S, HeadDim] tensor split_heads produces (and the one-position slice decode uses), before rope. RmsNorm::forward normalizes only its last axis, so this normalizes each head's vector independently rather than the flattened Heads * HeadDim projection -- putting it before the split_heads reshape would normalize over all heads at once, which is a different (and wrong) computation.
  • attention_bias is false, so q_proj, k_proj, and v_proj carry no bias -- std.nn.linear::Linear's optional bias parameter already defaults to none, so this needs no structural change, only that bindings.json has no *.bias entries for them (nor did Qwen2's o_proj and MLP projections).

bindings.json maps lm_head.weight to the checkpoint's own lm_head.weight tensor (tie_word_embeddings is false in config.json). The epsilon Qwen3 uses for RMSNorm (1e-6) is tighter than the stdlib RmsNorm block's default (1e-5), so src/norm.linnet defines its own RmsNorm block that calls std.nn.norm::rms_norm directly with the checkpoint's epsilon; the same block, instantiated at H and at HeadDim, serves both the per-layer norms and QK-norm.

Loading ​

python
from linnet import nest

model = nest.load("qwen3-8b", backend="torch", numerics="fast", compile="inductor")
logits = model(tokens)                      # Tensor[1, S; i32] -> Tensor[1, S, 151936; bf16]

The weights are the published, sharded model-0000N-of-00005.safetensors (bf16); bindings.json maps Linnet parameter paths to their 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 = 40960)
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
prefill_paged<P, Rows, Pages>(tokens, positions, rows, slots, last, table)for serving from pages: prompts packed end to end, each token written at its place in the pool and attending over its row's pages up to its position (chunks, shared prefixes)
decode_paged<Rows, Pages>(tokens, positions, table)decode_rows for serving from pages: loaded with Batch = 1, the caches' one row is a pool of MaxSeq positions in pages of PageSize (64), and row b's positions lie in the pages table[b] lists
step_paged<P, Rows, Pages>(...)prefill_paged and decode_paged in one pass
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

Shards ​

Shards (1 unless bound) splits the model across devices: one shard holds Heads / Shards query heads, KvHeads / Shards key and value heads and Inner / Shards hidden units, with the KV caches to match, and the attention's and the MLP's output projections end in std.nn.parallel::all_reduce, the sum over the shards. lm_head holds Vocab / Shards rows of the vocabulary, and std.nn.parallel::all_gather sets the shards' slices of the logits side by side. On one device the sum and the gather are the value itself. linnet.torch.load(..., tensor_parallel=mesh) binds Shards to the mesh size and gives each process its part of the checkpoint. The training entries train split as well: each split computation reads its input through std.nn.parallel::shared, and loss_packed and log_probs_packed run over the vocabulary's parts (std.nn.loss::split_cross_entropy, split_token_log_probs).

Provenance ​

The model is large enough that comparing the full checkpoint against transformers in one pass on shared CPU hardware is impractical, so this was validated on CPU with both sides truncated to the first 2 layers (num_hidden_layers = 2 on the transformers side, generics = {"Layers": 2} on this one, both loaded in f32, real weights read from the published checkpoint): on a chat-templated prompt, the Linnet forward pass matches transformers logits to a maximum absolute difference of 5.2e-6 on the final position (well inside the usual 2e-3), with the top-1 next token agreeing exactly; peak RSS for the comparison was 10.8 GiB. This proves the architecture -- every weight this checkpoint has, read correctly, through the right shapes -- but not the whole depth. A second pass, comparing the full model in bf16 (top-1 agreement and the logits' overall magnitude of difference), is deferred to a batched GPU run across every Nest model and has not been done here; treat that pass as outstanding, not as having passed.

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("qwen3-8b", backend="torch", numerics="fast", compile="inductor")
logits = model(tokens)

JAX ​

python
from linnet import nest

f = nest.load("qwen3-8b", backend="jax")          # StableHLO compiled by XLA
logits = f(tokens)

Flax NNX ​

python
from linnet import nest

model = nest.load("qwen3-8b", backend="nnx")
logits = model(tokens)

Command line ​

bash
python -m linnet.nest pull qwen3-8b
linnet stablehlo --root Model --entry forward --bind Vocab=151936 --bind H=4096 --bind Heads=32 --bind KvHeads=8 --bind HeadDim=128 --bind Inner=12288 --bind Layers=36 --bind Batch=1 --bind MaxSeq=40960 --bind T=bf16 --bind B=1 --bind S=8 <model dir>/src/lib.linnet