30. Capstone: Qwen3.8-Flash-Next from scratch
In this chapter
- Reading a frontier model's configuration and mapping every component to something you've built.
- The two genuinely new pieces: hyper-connections (four residual streams) and the hashed n-gram memory.
- Assembling the hybrid model, its per-request state, and an exact parity check against the official implementation.
- Loading the real checkpoint on a single machine: experts quantized while loading, a 95 GiB n-gram table left on disk, and a byte budget for each choice.
You will build
engine/flashnext.py: the complete Qwen3.8-Flash-Next text model, its loader and its memory plan, served through your Chapter 18 engine.
Time: 2-3 weeks. GPU: the tests run on a CPU; the real model needs about 70 GB of GPU or unified memory with 4-bit experts.
The target
Qwen3.8-Flash-Next is a multimodal model whose language part combines nearly every idea in Part VII. Its model card describes a 125B-parameter backbone with 6B parameters active per token, plus 51B parameters of n-gram memory, 4B parameters of multi-token-prediction heads, and a vision encoder. The text configuration (qwen4_exp_text in Transformers 5.18 and later):
| field | value | chapter |
|---|---|---|
| layers / hidden size | 48 / 2,560 | 6, 17 |
| layer pattern | 3 linear-attention layers, then 1 attention layer, ×12 | 28, 29 |
| attention: query / KV heads, head dim | 24 / 2, 256 | 17 |
| partial RoPE | 25% of each head (64 of 256 dims), θ = 10⁷ | 17, 29 |
| QSA indexer: heads × dim, block ratio, budget | 4 × 128, 4, 2,048 tokens | 29 |
| Gated DeltaNet: key / value heads, dims | 16 / 48, 128 / 128 | 28 |
| routed experts / per token, expert width | 512 / 10, 640 | 27 |
| shared expert width (sigmoid-gated) | 640 | 27 |
| residual streams / low-rank width | 4 / 320 | this chapter |
| n-gram memory: orders, hash heads per order, layer | 2-3, 8, layer 2 | this chapter |
| vocabulary / context | 248,320 / 262,144 | 4 |
Twenty-eight chapters of preparation cover all but two rows. Before writing any code, the professional habit is an architecture ledger: for every component, the equations, tensor names and shapes, and the reference you’ll check against. The official implementation (transformers/models/qwen4_exp/modeling_qwen4_exp.py) is the ground truth; this chapter’s code was written from it and is tested against it.
What’s already built
- Gated DeltaNet layers with their convolution and recurrent state, exactly Chapter 28’s
GatedDeltaNet, with the same parameter names. - Gated attention with zero-centered Q/K norms, partial RoPE, two KV heads and the QSA indexer: Chapter 29’s
GatedAttention. - MoE with 512 experts, top-10, renormalized softmax router and a sigmoid-gated shared expert: Chapter 27’s
SparseMoeBlock. - Loading sharded safetensors, meta-device construction and expert stacking: Chapters 9, 18 and 27.
Two things are new.
Hyper-connections: four residual streams
Every model so far had one residual stream: x = x + sublayer(norm(x)) (Chapter 23). Flash-Next keeps four parallel streams of width 2,560, and each sublayer learns how to read from them and how to write back (hyper-connections, Zhu et al., 2024). The embedding initializes all four streams to the same vector, and the streams drift apart as layers write to them differently.
Around each sublayer (attention or linear attention, then the MoE), a GatedResidual module:
- Normalizes each stream separately: a zero-centered RMSNorm over groups of 2,560 features. This replaces the usual pre-norm; there’s no other norm before the sublayer.
- Reads: computes per-feature weights for each stream through a low-rank bottleneck, $w = \sigma\big(W_\text{up}, \operatorname{SiLU}(W_\text{down}, \hat s / 4)\big)$, with rank 320. The sublayer’s input is the mean over streams of $w \odot \hat s$.
- Writes: computes one scalar per stream, $\gamma = 2\sigma(W_\text{inject}, \hat s / 4) \in (0, 2)$, and adds $\gamma_i \cdot \text{output}$ to stream $i$.
At the end, a final GatedResidual without the write part (combine=False) reads one vector from the four streams for the LM head. There’s no final norm.
class GatedResidual(nn.Module):
"""Hyper-connection around one sublayer. (Your engine: Chapter 30)
streams [B, T, hc*D] -> normalize each stream (grouped zero-centered RMSNorm);
read weights w = sigmoid(up(silu(down(normed) / hc))) [B, T, hc, D]
sublayer input = mean over streams of w * normed [B, T, D]
write weights = 2 * sigmoid(inject(normed) / hc) [B, T, hc]
After the sublayer: streams + write_weights[..., None] * output[..., None, :].
"""
def __init__(self, hidden, hc, rank, eps=1e-6, combine=True):
super().__init__()
self.hc, self.hidden = hc, hidden
self.hc_norm = ZeroCenteredRMSNorm(hc * hidden, eps, group_size=hidden)
self.input_mix_weight_down = nn.Linear(hc * hidden, rank, bias=False)
self.input_mix_weight_up = nn.Linear(rank, hc * hidden, bias=False)
self.block_inject_weight = nn.Linear(hc * hidden, hc, bias=False) if combine else None
def read(self, streams):
"""(Your engine: Chapter 30)"""
normed = self.hc_norm(streams)
weights = torch.sigmoid(self.input_mix_weight_up(F.silu(self.input_mix_weight_down(normed) / self.hc)))
mixed = (weights.unflatten(-1, (self.hc, self.hidden)) * normed.unflatten(-1, (self.hc, self.hidden))).mean(-2)
if self.block_inject_weight is None:
return mixed, None
return mixed, 2 * torch.sigmoid(self.block_inject_weight(normed) / self.hc)
@staticmethod
def write(streams, output, inject):
return streams + (output.unsqueeze(-2) * inject.unsqueeze(-1)).flatten(-2)
Why bother? With one stream, each layer’s output is added with weight 1, and depth works through a single shared channel. Multiple streams with learned read and write weights let the network keep some information untouched by later layers and vary how strongly each layer contributes, which in the hyper-connection papers improves training stability and quality at a cost of a few small matrices per layer. For inference, it means the hidden state is four times wider between layers: 20 KiB per token instead of 5 KiB in BF16, which matters for activation memory during prefill and not at all for the KV cache.
The n-gram memory
The second new piece is a huge lookup table indexed by the last few tokens. The idea: many next-token facts depend only on the immediately preceding tokens (“New York” → “City”), and a model shouldn’t spend attention and MLP compute rediscovering them. A hashed table can store an embedding for every frequent bigram and trigram, and the model looks it up in $O(1)$.
Hashing n-grams
For position $t$, the bigram is $(x_{t-1}, x_t)$ and the trigram $(x_{t-2}, x_{t-1}, x_t)$. There are $248{,}320^3$ possible trigrams, far too many to store, so each n-gram is hashed into a table of about 20 million rows:
$$ h = \Big(\bigoplus_{i=0}^{n-1} x_{t-i} \cdot m_i\Big) \bmod p , $$
with odd multipliers $m_i$ derived from a seed with splitmix64, XOR ($\oplus$) to combine, and a prime table size $p$. Hash collisions are inevitable, so each n-gram order uses 8 independent hash heads, each with its own prime size (the 16 smallest primes above 20 million), each returning a 160-dimensional row. Concatenated, the 16 rows form one 2,560-dimensional embedding. Two n-grams colliding in one head are very unlikely to collide in all eight.
MASK64 = (1 << 64) - 1
def splitmix64(value):
value = (value + 0x9E3779B97F4A7C15) & MASK64
value = ((value ^ (value >> 30)) * 0xBF58476D1CE4E5B9) & MASK64
value = ((value ^ (value >> 27)) * 0x94D049BB133111EB) & MASK64
return (value ^ (value >> 31)) & MASK64
def layer_multipliers(vocab, ngram_size, ple_index, seed):
"""Odd multipliers, small enough that token_id * multiplier never overflows int64."""
half = max(1, (((1 << 63) - 1) // max(vocab, 1)) // 2)
base = seed + 10007 * ple_index
return [2 * (splitmix64((base + 0x9E3779B97F4A7C15 * (i + 1)) & MASK64) % half) + 1 for i in range(ngram_size)]
def next_primes(start, count):
"""The `count` smallest primes greater than `start` (one distinct table size per hash head)."""
def is_prime(n):
if n < 2 or (n % 2 == 0 and n != 2):
return n == 2
return all(n % d for d in range(3, math.isqrt(n) + 1, 2))
primes, n = [], start
while len(primes) < count:
n += 1
if is_prime(n):
primes.append(n)
return primes
Two details that a careless implementation gets wrong:
- Document boundaries. An n-gram must not span an end-of-sequence token: after an EOS, the “previous token” positions read as EOS.
shiftedcomputes, for every position, how far back the current document starts, and substitutes EOS beyond it. - Carried history. In decode, the n-gram of the new token needs the previous two tokens, which belong to earlier calls. The layer’s state carries them.
class NGramEmbedding(nn.Module):
"""Hashed bigram and trigram embeddings, heads_per_ngram independent hashes per order. (Your engine: Chapter 30)
The n-gram ending at token t is (x[t-n+1], ..., x[t]); its id for one head is
(XOR over i of x[t-i] * m_i) mod p_head, offset into one big table. N-grams never cross an
EOS: positions before the most recent EOS read as EOS.
"""
def __init__(self, cfg, ple_index):
super().__init__()
self.n, self.heads_per = cfg.ngram_size, cfg.heads_per_ngram
self.heads = (cfg.ngram_size - 1) * cfg.heads_per_ngram
self.eos = cfg.eos_token_id
sizes = next_primes(cfg.ngram_vocab_size_base - 1, (ple_index + 1) * self.heads)[ple_index * self.heads:]
self._sizes = sizes
offsets = [sum(sizes[:i]) for i in range(len(sizes))]
padded = -(-sum(sizes) // cfg.make_ngram_vocab_size_divisible_by) * cfg.make_ngram_vocab_size_divisible_by
self.register_buffer("layer_multipliers", torch.tensor(layer_multipliers(cfg.vocab_size, self.n, ple_index, cfg.seed)))
self.register_buffer("ngram_heads_vocab_sizes", torch.tensor(sizes))
self.register_buffer("ngram_heads_offsets", torch.tensor(offsets))
self.ngram_embedding = nn.Embedding(padded, cfg.ple_embed_dim // self.heads)
def shifted(self, history, shift):
"""history[t - shift], or EOS when that position is before the current document. (Your engine: Chapter 30)"""
if shift == 0:
return history
b, length = history.shape
pos = torch.arange(length, device=history.device)
eos_at = torch.where(history == self.eos, pos, -1)
last_eos_before = torch.cat((eos_at.new_full((b, 1), -1), eos_at.cummax(1).values[:, :-1]), dim=1)
in_document = pos - (last_eos_before + 1)
source = (pos - shift).clamp_min(0).expand(b, -1)
valid = (in_document >= shift) & ((pos - shift) >= 0)
return torch.where(valid, history.gather(1, source), torch.full_like(history, self.eos))
def forward(self, ids, context=None):
"""ids [B, T]; context = the previous n-1 token ids (EOS at the start). Returns (emb [B, T, E], new context). (Your engine: Chapter 30)"""
if context is None:
context = ids.new_full((ids.shape[0], self.n - 1), self.eos)
history = torch.cat((context, ids.long()), dim=1)
shifted = [self.shifted(history, s) for s in range(self.n)]
blocks = []
for order in range(2, self.n + 1):
mixed = shifted[0] * self.layer_multipliers[0]
for position in range(1, order):
mixed = mixed ^ (shifted[position] * self.layer_multipliers[position])
first = (order - 2) * self.heads_per
sizes = self.ngram_heads_vocab_sizes[first:first + self.heads_per]
blocks.append(mixed[..., None] % sizes + self.ngram_heads_offsets[first:first + self.heads_per])
index = torch.cat(blocks, dim=-1)[:, -ids.shape[1]:]
table = self.ngram_embedding.weight
rows = table[index.to(table.device)].to(ids.device) # the table may live on the host
return rows.flatten(-2), history[:, -(self.n - 1):]
Injecting it: per-layer embeddings
The lookup result is injected into the residual streams at layer 2 by a PLELayer (“per-layer embedding”):
- A key projection of the n-gram embedding (one per stream) and a query from the normalized streams give a gate per stream: their dot product over $\sqrt{d}$, passed through a signed square root (to tame large values while keeping the sign) and a sigmoid.
- A value projection of the n-gram embedding, scaled by each stream’s gate, is the addition to that stream.
- A dilated causal convolution (kernel 4, dilation 3, depthwise) over recent gated values adds local context, with its own carried state of the last 9 positions.
class PLELayer(nn.Module):
"""Per-layer embedding: gate the n-gram value into each residual stream, then add a dilated
causal depthwise convolution over recent gated values. (Your engine: Chapter 30)"""
def __init__(self, cfg, ple_index):
super().__init__()
d, hc, e = cfg.hidden_size, cfg.hc_count, cfg.ple_embed_dim
self.hidden, self.hc, self.dilation = d, hc, cfg.ngram_size
self.state_len = (cfg.ple_conv_kernel_size - 1) * cfg.ngram_size
self.ple_embedding = NGramEmbedding(cfg, ple_index)
self.key_proj = nn.Linear(e, hc * d, bias=False)
self.value_proj = nn.Linear(e, d, bias=False)
self.norm_key = ZeroCenteredRMSNorm(hc * d, cfg.rms_norm_eps, group_size=d)
self.norm_query = ZeroCenteredRMSNorm(hc * d, cfg.rms_norm_eps, group_size=d)
self.norm_conv = ZeroCenteredRMSNorm(hc * d, cfg.rms_norm_eps, group_size=d)
self.conv1d = nn.Conv1d(hc * d, hc * d, cfg.ple_conv_kernel_size, groups=hc * d,
dilation=cfg.ngram_size, bias=False)
def forward(self, streams, ids, state=None):
"""state = (ngram context ids, conv history). Returns (addition to streams, new state). (Your engine: Chapter 30)"""
context, conv_state = state if state is not None else (None, None)
embedding, context = self.ple_embedding(ids, context)
key = self.norm_key(self.key_proj(embedding)).unflatten(-1, (self.hc, self.hidden))
value = self.value_proj(embedding)
query = self.norm_query(streams).unflatten(-1, (self.hc, self.hidden))
gate = (key * query).sum(-1, keepdim=True) / math.sqrt(self.hidden)
gate = gate.abs().clamp_min(1e-6).sqrt() * gate.sign() # signed square root
gated = (torch.sigmoid(gate) * value.unsqueeze(-2)).flatten(-2) # [B, T, hc*D]
x = self.norm_conv(gated).transpose(1, 2)
if conv_state is None:
conv_state = x.new_zeros(x.shape[0], x.shape[1], self.state_len)
joined = torch.cat((conv_state, x), dim=-1)
conv = F.silu(self.conv1d(joined)).transpose(1, 2)
return gated + conv, (context, joined[..., -self.state_len:])
Why it’s cheap at inference: each token reads 16 rows of 160 BF16 values, about 5 KB, from a table of 51 billion parameters. The table’s size is a storage problem, not a bandwidth problem, so it can live in host memory or even on an SSD. Closely related published work: DeepSeek’s Engram conditional memory (2026).
The layer and the model
class FlashNextLayer(nn.Module):
def __init__(self, cfg, index):
super().__init__()
self.kind = cfg.layer_types[index]
if self.kind == "linear_attention":
self.linear_attn = GatedDeltaNet(cfg.hidden_size, cfg.linear_num_key_heads, cfg.linear_num_value_heads,
cfg.linear_key_head_dim, cfg.linear_value_head_dim,
cfg.linear_conv_kernel_dim, cfg.rms_norm_eps, cfg.output_gate_type)
else:
indexer = None
if cfg.indexer_n_heads:
indexer = QSAIndexer(cfg.hidden_size, cfg.indexer_n_heads, cfg.indexer_head_dim,
cfg.indexer_budget, cfg.indexer_compress_ratio, cfg.rms_norm_eps)
self.self_attn = GatedAttention(cfg.hidden_size, cfg.num_attention_heads, cfg.num_key_value_heads,
cfg.head_dim, cfg.rotary_dim, cfg.rope_theta, cfg.rms_norm_eps, indexer)
self.mlp = SparseMoeBlock(cfg.hidden_size, cfg.num_experts, cfg.num_experts_per_tok, cfg.moe_intermediate_size,
cfg.shared_expert_intermediate_size, cfg.norm_topk_prob)
ple_index = cfg.ple_layer_ids.index(index + 1) if (index + 1) in cfg.ple_layer_ids else None
self.ple = PLELayer(cfg, ple_index) if ple_index is not None else None
self.attn_hyper_connection = GatedResidual(cfg.hidden_size, cfg.hc_count, cfg.hc_lowrank, cfg.rms_norm_eps)
self.mlp_hyper_connection = GatedResidual(cfg.hidden_size, cfg.hc_count, cfg.hc_lowrank, cfg.rms_norm_eps)
def forward(self, streams, ids, positions, state):
"""state = {"mixer": layer state, "ple": PLE state}. Returns (streams, new state). (Your engine: Chapter 30)"""
new_state = {}
if self.ple is not None:
addition, new_state["ple"] = self.ple(streams, ids, state.get("ple"))
streams = streams + addition
x, inject = self.attn_hyper_connection.read(streams)
if self.kind == "linear_attention":
y, new_state["mixer"] = self.linear_attn(x, state.get("mixer"))
else:
y, new_state["mixer"] = self.self_attn(x, positions, state.get("mixer"))
streams = GatedResidual.write(streams, y, inject)
x, inject = self.mlp_hyper_connection.read(streams)
streams = GatedResidual.write(streams, self.mlp(x), inject)
return streams, new_state
Each layer: inject the n-gram memory (layer 2 only), read from the streams, run the mixer (Gated DeltaNet or gated sparse attention), write, read again, run the MoE, write.
class HybridState:
"""Everything one sequence carries between calls. Updates return a new object, so keeping the
old one is a free snapshot: speculative decoding rolls back by simply not adopting the new state."""
def __init__(self, layers=None, length=0):
self.layers = layers or []
self.length = length
class HybridCache:
"""Adapts the functional HybridState to the engines' mutable-cache convention (Chapters 18-19):
model(ids, cache) updates cache.state in place and returns logits only."""
def __init__(self):
self.state = None
@property
def length(self):
return 0 if self.state is None else self.state.length
def snapshot(self):
return self.state # states are never modified in place: this is a free copy
def restore(self, state):
self.state = state
def truncate(self, length):
raise NotImplementedError("A recurrent state cannot be truncated: restore a snapshot instead")
class FlashNextBackbone(nn.Module):
def __init__(self, cfg):
super().__init__()
self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.hidden_size)
self.layers = nn.ModuleList(FlashNextLayer(cfg, i) for i in range(cfg.num_hidden_layers))
self.hyper_connection_mixer = GatedResidual(cfg.hidden_size, cfg.hc_count, cfg.hc_lowrank,
cfg.rms_norm_eps, combine=False)
class FlashNext(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
self.model = FlashNextBackbone(cfg)
self.lm_head = nn.Linear(cfg.hidden_size, cfg.vocab_size, bias=False)
@property
def context_limit(self):
return self.cfg.max_position_embeddings
def new_cache(self, batch, capacity):
if capacity > self.context_limit:
raise ValueError("Requested cache exceeds the context limit")
return HybridCache()
def forward(self, ids, state=None):
"""ids [B, T] continuing `state` -> (logits, new HybridState); with a HybridCache, logits only."""
if isinstance(state, HybridCache):
logits, state.state = self.step(ids, state.state)
return logits
return self.step(ids, state)
def step(self, ids, state=None):
"""ids [B, T] continuing `state` -> (logits [B, T, V], new HybridState). (Your engine: Chapter 30)"""
state = state or HybridState([{} for _ in self.model.layers])
positions = torch.arange(state.length, state.length + ids.shape[1], device=ids.device)
streams = self.model.embed_tokens(ids).repeat(1, 1, self.cfg.hc_count) # every stream starts equal
new_layers = []
for layer, layer_state in zip(self.model.layers, state.layers):
streams, layer_state = layer(streams, ids, positions, layer_state)
new_layers.append(layer_state)
hidden, _ = self.model.hyper_connection_mixer.read(streams) # collapse streams
return self.lm_head(hidden), HybridState(new_layers, state.length + ids.shape[1])
The state of a request
One Flash-Next request carries four kinds of state, all owned by HybridState:
| state | where | size per sequence |
|---|---|---|
| recurrent matrices + conv history | 36 linear-attention layers | 108 MiB, fixed |
| K, V and indexer keys | 12 attention layers | 27 KiB per token (KV 24 KiB, indexer 3 KiB): 0.84 GiB at 32k |
| last 2 token IDs + PLE conv history | layer 2 | a few KB, fixed |
| position | model | one integer |
step never modifies a state in place: it returns a new HybridState. Keeping the old one is therefore a free snapshot, which is exactly what speculative decoding needs, since the recurrent state can’t be truncated (Chapter 28). The milestone test runs a branch from a state, discards it, runs another from the same state, and checks they’re identical. For the engines of Chapters 18-19, new_cache returns a HybridCache, a small mutable wrapper with snapshot and restore, so LLM.stream serves Flash-Next unchanged.
Exactly equal to the official implementation
The tests build a tiny random Flash-Next with Transformers’ Qwen4ExpForCausalLM: 8 layers, 2 of them attention with a QSA budget small enough to matter, 8 experts, two PLE layers with sharded n-gram tables. Every norm and convolution weight is randomized, since zero-initialized ones would hide bugs. The tests save it as a checkpoint and load it with your loader:
max abs logit difference: 1.49e-08 (logits of magnitude 0.12; FP32)
and check that:
- a full forward over 2 × 41 tokens with an EOS in the middle matches the reference;
- prefill of 13 tokens, a chunk of 16, then 11 single-token steps equals one full forward;
- the state is a free snapshot;
- the real configuration’s memory plan has the published sizes;
- the n-gram tables memory-mapped from disk give bit-identical logits, and 4-bit experts stay close;
- your Chapter 18
LLMengine generates the same tokens as the model’s own greedy loop.
Getting there surfaced real bugs, each now a lesson in an earlier chapter: a transposed triangle in the chunked delta rule (28) and topk tie-breaking under ReLU zeros (29). A third, a config field renamed between library versions (27), would have broken any loader that trusted defaults.
Loading the real checkpoint
@torch.no_grad()
def load_flashnext(directory, device="cpu", dtype=torch.bfloat16, ngram_device="cpu", expert_bits=None,
ngram_mmap=False, group_size=128):
"""Load the text model from a local snapshot of the official checkpoint. (Your engine: Chapter 30)
Renames multimodal prefixes, skips the vision tower and MTP head, and stacks per-expert
tensors. expert_bits=4 or 8 quantizes each layer's experts as soon as they are read;
ngram_mmap=True leaves the n-gram tables on disk (MappedRows); otherwise they are loaded to
`ngram_device` (host memory by default: lookups touch a few rows per token).
"""
raw = json.loads((Path(directory) / "config.json").read_text())
cfg = FlashNextConfig.from_hf(raw)
with torch.device("meta"):
model = FlashNext(cfg)
for layer in model.model.layers: # never allocate what will be stored differently
if expert_bits:
layer.mlp.experts = nn.Module()
if ngram_mmap and layer.ple is not None:
del layer.ple.ple_embedding.ngram_embedding
model = model.to_empty(device=device).to(dtype)
for index, layer in enumerate(model.model.layers):
if layer.ple is not None: # buffers are derived, not learned: recompute them
fresh = _ngram_buffers(cfg, cfg.ple_layer_ids.index(index + 1))
for name, value in fresh.items():
setattr(layer.ple.ple_embedding, name, value.to(device))
if not ngram_mmap:
layer.ple.ple_embedding.ngram_embedding.to(ngram_device)
params = dict(model.named_parameters())
rename = lambda name: name.replace("model.language_model.", "model.")
is_ngram = lambda name: ".ngram_embedding." in name
experts, ngram_shards = {}, {}
skip = (lambda name: is_ngram(name)) if ngram_mmap else None
for name, value in snapshot_tensors(directory, skip=skip):
name = rename(name)
if name.startswith(("model.visual.", "mtp.", "model.mtp")) or name.endswith(
("layer_multipliers", "ngram_heads_vocab_sizes", "ngram_heads_offsets")):
continue
if ".mlp.experts." in name:
parts = name.split(".")
layer = int(parts[2])
experts.setdefault(layer, {})[(int(parts[5]), parts[6])] = value
if len(experts[layer]) == 3 * cfg.num_experts: # a whole layer has arrived
_install_experts(model, cfg, layer, experts.pop(layer), expert_bits, group_size, device, dtype)
continue
if is_ngram(name):
prefix, _, shard = name.partition(".shard_")
ngram_shards.setdefault(prefix.removesuffix(".weight"), {})[int(shard.split(".")[0]) if shard else 0] = value
continue
if name not in params:
raise ValueError(f"Unexpected tensor {name}")
assign(params[name], value, name)
if experts:
raise ValueError(f"Incomplete experts for layers {sorted(experts)}")
if ngram_mmap:
for original, path in tensor_files(directory).items():
name = rename(original)
if is_ngram(name):
prefix, _, shard = name.partition(".shard_")
ngram_shards.setdefault(prefix.removesuffix(".weight"), {})[
int(shard.split(".")[0]) if shard else 0] = map_tensor(path, original)
for prefix, shards in ngram_shards.items():
ordered = [shards[i] for i in sorted(shards)]
owner = model.get_submodule(prefix.rsplit(".", 1)[0])
if ngram_mmap:
owner.ngram_embedding = MappedRows(ordered, dtype)
else:
assign(params[prefix + ".weight"], torch.cat(ordered, dim=0), prefix)
return model.eval()
def _ngram_buffers(cfg, ple_index):
with torch.device("meta"):
probe = NGramEmbedding(cfg, ple_index) # meta: the big table is not allocated
return {"layer_multipliers": torch.tensor(layer_multipliers(cfg.vocab_size, cfg.ngram_size, ple_index, cfg.seed)),
"ngram_heads_vocab_sizes": torch.tensor(probe._sizes),
"ngram_heads_offsets": torch.tensor([sum(probe._sizes[:i]) for i in range(len(probe._sizes))])}
def _install_experts(model, cfg, layer, tensors, bits, group_size, device, dtype):
gate_up, down = stack_expert_tensors(tensors, cfg.num_experts, cfg.moe_intermediate_size)
block = model.model.layers[layer].mlp
if bits:
block.experts = QuantizedExperts(gate_up.to(device, dtype), down.to(device, dtype), bits, group_size)
else:
assign(block.experts.gate_up_proj, gate_up, "gate_up_proj")
assign(block.experts.down_proj, down, "down_proj")
The loader streams the safetensors shards one at a time and handles the checkpoint’s quirks: the multimodal prefix model.language_model. is renamed, the vision tower and MTP weights are skipped, per-expert tensors are stacked, and n-gram tables stored as shards are concatenated. Two options make it fit on one machine:
class QuantizedExperts(nn.Module):
"""Routed experts stored as groupwise INT4/INT8 codes (Chapter 20). An expert is dequantized
only when a token is routed to it, so a BF16 copy of all experts never exists."""
def __init__(self, gate_up, down, bits=4, group_size=128):
super().__init__()
self.bits, self.dtype, self.shapes, self.groups = bits, gate_up.dtype, {}, {}
for name, w in (("gate_up", gate_up), ("down", down)):
experts, rows, cols = w.shape
codes, scales = quantize_groupwise(w.reshape(experts * rows, cols), bits, group_size)
codes = pack_int4(codes).view(experts, -1) if bits == 4 else codes.view(experts, rows, cols)
self.register_buffer(f"{name}_codes", codes)
self.register_buffer(f"{name}_scales", scales.view(experts, rows, -1).to(torch.float16))
self.shapes[name], self.groups[name] = (rows, cols), group_size
self.num_experts = gate_up.shape[0]
def weight(self, name, e):
codes = getattr(self, f"{name}_codes")[e]
if self.bits == 4:
codes = unpack_int4(codes, self.shapes[name])
scales = getattr(self, f"{name}_scales")[e].float()
return dequantize_groupwise(codes, scales, self.groups[name]).to(self.dtype)
def expert(self, e, x):
gate, up = F.linear(x, self.weight("gate_up", e)).chunk(2, dim=-1)
return F.linear(F.silu(gate) * up, self.weight("down", e))
forward_loop = Experts.forward_loop
forward_grouped = Experts.forward_grouped
class MappedRows:
"""An embedding table left on disk: rows are read through memory maps when looked up, and the
operating system's page cache keeps the hot ones in RAM. Lookups touch a few rows per token,
so the 95 GiB n-gram table never has to fit in memory."""
def __init__(self, shards, dtype):
self.shards, self.dtype, self.device = shards, dtype, torch.device("cpu")
sizes = torch.tensor([0] + [t.shape[0] for t in shards])
self.starts = sizes.cumsum(0)
@property
def weight(self):
return self
def __getitem__(self, index):
flat = index.reshape(-1).cpu()
shard = torch.searchsorted(self.starts, flat, right=True) - 1
out = torch.empty(flat.numel(), self.shards[0].shape[1], dtype=self.dtype)
for s in shard.unique().tolist():
hit = shard == s
out[hit] = self.shards[s][flat[hit] - self.starts[s]].to(self.dtype)
return out.view(*index.shape, -1)
expert_bits=4replaces each layer’s experts withQuantizedExpertsas soon as that layer’s 1,536 expert tensors have arrived: the BF16 experts of more than one layer never exist at once, and the model is built on the meta device so their full-size placeholders are never allocated. An expert is dequantized only when a token is routed to it (a W4A16 grouped kernel, Chapters 20 and 27, is the fast version).ngram_mmap=Truenever reads the n-gram shards into memory.MappedRowswraps memory maps of the safetensors files; a lookup reads a few pages and the operating system caches the hot ones.
The memory plan
def memory_plan(cfg, weight_bits=16, expert_bits=4, ngram_bits=16, context=32768, kv_bits=16):
"""Rough byte budget for the text model: what must be resident, and where."""
d, l = cfg.hidden_size, cfg.num_hidden_layers
n_attn = sum(t != "linear_attention" for t in cfg.layer_types)
n_lin = l - n_attn
expert = cfg.num_experts * 3 * d * cfg.moe_intermediate_size * l
shared = 3 * d * cfg.shared_expert_intermediate_size * l + l * (cfg.num_experts + 1) * d
kd, vd = cfg.linear_num_key_heads * cfg.linear_key_head_dim, cfg.linear_num_value_heads * cfg.linear_value_head_dim
linear = n_lin * (d * (2 * kd + vd) + d * vd + vd * d + 2 * d * cfg.linear_num_value_heads)
attn = n_attn * (d * cfg.num_attention_heads * cfg.head_dim * 3 + 2 * d * cfg.num_key_value_heads * cfg.head_dim)
hc = (2 * l + 1) * 2 * cfg.hc_count * d * cfg.hc_lowrank
embed = 2 * cfg.vocab_size * d
ngram_rows = sum(next_primes(cfg.ngram_vocab_size_base - 1, (cfg.ngram_size - 1) * cfg.heads_per_ngram * len(cfg.ple_layer_ids)))
ngram = ngram_rows * cfg.ple_embed_dim // ((cfg.ngram_size - 1) * cfg.heads_per_ngram)
kv = 2 * n_attn * context * cfg.num_key_value_heads * cfg.head_dim * kv_bits // 8
recurrent = n_lin * cfg.linear_num_value_heads * cfg.linear_key_head_dim * cfg.linear_value_head_dim * 4
gib = 1024 ** 3
return {
"routed_experts_GiB": expert * expert_bits / 8 / gib,
"dense_weights_GiB": (shared + linear + attn + hc + embed) * weight_bits / 8 / gib,
"ngram_tables_GiB": ngram * ngram_bits / 8 / gib,
"kv_cache_GiB_per_sequence": kv / gib,
"recurrent_state_MiB_per_sequence": recurrent / 1024 ** 2,
"active_params_per_token_B": (cfg.num_experts_per_tok * 3 * d * cfg.moe_intermediate_size * l
+ shared + linear + attn + hc) / 1e9,
}
python run.py flashnext
{"expert_bits": 16, "context": 32768, "routed_experts_GiB": 225.0, "dense_weights_GiB": 9.11, "ngram_tables_GiB": 95.37, "kv_cache_GiB_per_sequence": 0.75, "recurrent_state_MiB_per_sequence": 108.0, "active_params_per_token_B": 5.98}
{"expert_bits": 8, "context": 32768, "routed_experts_GiB": 112.5, ...}
{"expert_bits": 4, "context": 32768, "routed_experts_GiB": 56.25, ...}
{"expert_bits": 4, "context": 262144, "routed_experts_GiB": 56.25, ..., "kv_cache_GiB_per_sequence": 6.0, ...}
Putting it together for three machines (estimates from the plan, not measurements):
| machine | experts | dense weights | n-gram table | resident total | decode ceiling |
|---|---|---|---|---|---|
| DGX Spark (128 GB unified, 273 GB/s) | 4-bit, in memory | BF16 | memory-mapped from NVMe | ~68 GiB | ~28 tokens/s |
| same, dense weights in INT8 | 4-bit | INT8 | memory-mapped | ~63 GiB | ~50 tokens/s |
| H100 80 GB (3,350 GB/s) | 4-bit | BF16 | host RAM | ~68 GiB | ~350 tokens/s |
| RTX 4090 24 GB + 128 GB host RAM | 4-bit, offloaded to host (stretch) | BF16 | host RAM | 11 GiB on GPU | PCIe-bound |
The ceilings use Chapter 10’s rule with the bytes read per token: 2.36B active expert parameters at 4 bits (1.2 GB) plus 4.25B other active parameters including the LM head at 2 bytes (8.5 GB), about 9.7 GB per token in BF16, or 5.4 GB with INT8 dense weights. The dense part, not the experts, dominates decode traffic once the experts are 4-bit: quantize it next.
Run it
hf download Qwen/Qwen3.8-Flash-Next --local-dir models/Qwen3.8-Flash-Next # ~360 GB on disk
python run.py chat --model-dir models/Qwen3.8-Flash-Next --expert-bits 4 --ngram-mmap \
--prompt "Explain hyper-connections in two sentences." --new-tokens 128
Important
This book’s code was validated against the official implementation on tiny random checkpoints with the official architecture (Appendix F); the full checkpoint was not run in the validation environment. On a real machine, climb Chapter 18’s ladder of evidence again: compare your logits with Transformers’ on a few hundred tokens in BF16, layer by layer if they differ, before trusting generations.
Multi-token prediction
The checkpoint also ships about 4B parameters of multi-token prediction (MTP) heads, which both your loader and Transformers skip. In the DeepSeek-V3 style, an MTP module takes the final hidden state at position $t$ and the embedding of token $t+1$ and predicts token $t+2$, reusing the model’s embedding and head. At inference, it’s a built-in draft model for speculative decoding (Chapter 26): one cheap extra module proposes the next token, and the main model verifies. With HybridState snapshots, your speculative loop needs only one change: instead of truncating caches after a rejection, restore the snapshot and re-run the accepted tokens. That’s the first stretch exercise.
Build it
Engine milestone 30: the capstone. Implement in engine/flashnext.py: GatedResidual.read, NGramEmbedding.shifted and forward, PLELayer.forward, FlashNextLayer.forward, FlashNext.step and load_flashnext (the configuration, hashing constants, quantized experts, mapped tables, state classes and memory plan are provided).
pytest tests/test_ch30_flashnext.py
python run.py flashnext --impl engine
The tests check the hashing constants, logit parity with Transformers’ Qwen4ExpForCausalLM (including an EOS mid-sequence), that chunked prefill plus decode equals the full forward, free snapshots, the real model’s memory plan, memory-mapped and quantized loading, and generation through your LLM engine. When they pass, your engine runs Qwen3.8-Flash-Next.
Stretch exercises
- ★★ Speculative decoding for Flash-Next: change
speculative_generateto snapshot and restoreHybridCacheinstead of truncating, and use a smaller model with the same tokenizer as the draft. Verify greedy outputs are unchanged. Where:speculative_generateinengine/speculative.py, usingHybridCache.snapshot/restoreinengine/flashnext.py. - ★★ Quantize the dense weights to INT8 (attention, linear attention, shared experts, LM head) and measure the logit error and decode speed against the plan’s prediction. Where: weight installation in
load_flashnextinengine/flashnext.py, usingengine.quant. - ★★★ Load and use the MTP head as a draft: read its weights (prefix
mtp.) and implement its forward from the configuration and the DeepSeek-V3 report’s description, then measure acceptance rates. Where: add an MTP module and load its weights inengine/flashnext.py; call it fromengine/speculative.py. - ★★★ Expert offloading for a 24 GB GPU: keep
QuantizedExpertscodes in pinned host memory and copy only the routed experts per layer, overlapping the copy for layer $\ell+1$ with compute for layer $\ell$ using CUDA streams. Report tokens/s against the PCIe bound. Where:QuantizedExpertsand expert installation inengine/flashnext.py. - ★★★ Continuous batching for a hybrid model: give each request a
HybridState, batch the linear-attention layers (every state has the same shape) and the attention layers (ragged KV), and verify solo-equivalence as in Chapter 24. Where: newexperiments/hybrid_batching.py, adaptingengine.scheduler.ContinuousBatchingEngineforengine.flashnext.HybridState.
Check your understanding
- Which Flash-Next components come from which earlier chapters, and which two are new?
- How do the read and write weights of a hyper-connection differ in shape and range?
- Why does the n-gram memory use 8 hash heads with different prime sizes?
- Why can a 51B-parameter table live on disk without slowing decode much?
- Why is a functional (never modified in place) state convenient for speculative decoding?
- After quantizing the experts to 4 bits, what dominates the bytes read per decode token?
Going deeper
- The Qwen3.8-Flash-Next model card and configuration, and
modeling_qwen4_exp.pyin Transformers 5.18+. - Zhu et al., Hyper-Connections (2024), and the follow-ups on manifold-constrained hyper-connections; DeepSeek-AI’s Engram (2026) for hashed n-gram memory; the DeepSeek-V3 technical report for MTP.
- The Qwen3-Next and Qwen3.5 model cards, the direct ancestors of this architecture (3:1 Gated DeltaNet / gated attention, sigmoid-gated shared expert).
- vLLM and SGLang’s model files for Qwen3-Next, to see how production engines batch a hybrid model’s states.