A 4B-parameter decoder with the Qwen3 architecture: 36 layers of width 2560, grouped-query attention (32 query heads, 8 key/value heads, head dimension 128 -- not H / Heads, which would be 80; Qwen3 sizes its projections by Heads * HeadDim, not by H), rotary positions with base 1,000,000, RMS normalization with epsilon 1e-6, and a SwiGLU MLP of width 9728. 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 tied to the token embedding.
The Linnet source is a self-contained Qwen3 package (src/lib.linnet, src/attention.linnet, src/norm.linnet, src/rope.linnet), the same one models/qwen3-8b uses -- Qwen3-4B and Qwen3-8B are the same architecture at two sizes, differing in width, MLP width, and whether the head is tied, all of which are generic parameters or a bindings.json choice, not structural. Three changes from the Qwen2 decoder models/qwen2.5-0.5b-instruct uses:
HeadDimis its own generic rather thanH / Heads: Qwen3 sizes the query, key, and value projections byHeads * HeadDimandKvHeads * HeadDim, and the two do not agree withHhere (His 2560,Heads * HeadDimis 4096).q_proj/k_proj/v_projtherefore mapHtoHeads * HeadDim/KvHeads * HeadDim, ando_projmapsHeads * HeadDimback toH, notHtoH.q_normandk_norm, insrc/norm.linnet'sRmsNormblock instantiated at widthHeadDim, called on the[B, Heads, S, HeadDim]tensorsplit_headsproduces (and the one-position slicedecodeuses), before rope.RmsNorm::forwardnormalizes only its last axis, so this normalizes each head's vector independently rather than the flattenedHeads * HeadDimprojection -- putting it before thesplit_headsreshape would normalize over all heads at once, which is a different (and wrong) computation.attention_biasis false, soq_proj,k_proj, andv_projcarry no bias --std.nn.linear::Linear's optional bias parameter already defaults tonone, so this needs no structural change, only thatbindings.jsonhas no*.biasentries for them (nor did Qwen2'so_projand MLP projections).
bindings.json maps both embedding.weight and lm_head.weight to the checkpoint's single model.embed_tokens.weight (tie_word_embeddings is true in config.json, unlike the 8B checkpoint of this family), which is how tied embeddings are expressed without a language feature for them. 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
from linnet import nest
model = nest.load("qwen3-4b", 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-00003.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
- Weights: Qwen/Qwen3-4B, Apache-2.0.
- Code: QwenLM/Qwen3.
- Paper: Qwen3 Technical Report.
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.0e-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 7.2 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
from linnet import nest
model = nest.load("qwen3-4b", backend="torch", numerics="fast", compile="inductor")
logits = model(tokens)JAX
from linnet import nest
f = nest.load("qwen3-4b", backend="jax") # StableHLO compiled by XLA
logits = f(tokens)Flax NNX
from linnet import nest
model = nest.load("qwen3-4b", backend="nnx")
logits = model(tokens)Command line
python -m linnet.nest pull qwen3-4b
linnet stablehlo --root Model --entry forward --bind Vocab=151936 --bind H=2560 --bind Heads=32 --bind KvHeads=8 --bind HeadDim=128 --bind Inner=9728 --bind Layers=36 --bind Batch=1 --bind MaxSeq=40960 --bind T=bf16 --bind B=1 --bind S=8 <model dir>/src/lib.linnet