Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

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):

fieldvaluechapter
layers / hidden size48 / 2,5606, 17
layer pattern3 linear-attention layers, then 1 attention layer, ×1228, 29
attention: query / KV heads, head dim24 / 2, 25617
partial RoPE25% of each head (64 of 256 dims), θ = 10⁷17, 29
QSA indexer: heads × dim, block ratio, budget4 × 128, 4, 2,048 tokens29
Gated DeltaNet: key / value heads, dims16 / 48, 128 / 12828
routed experts / per token, expert width512 / 10, 64027
shared expert width (sigmoid-gated)64027
residual streams / low-rank width4 / 320this chapter
n-gram memory: orders, hash heads per order, layer2-3, 8, layer 2this chapter
vocabulary / context248,320 / 262,1444

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:

  1. 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.
  2. 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$.
  3. 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. shifted computes, 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”):

  1. 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.
  2. A value projection of the n-gram embedding, scaled by each stream’s gate, is the addition to that stream.
  3. 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).

Explore: hashing n-grams

Type tokens and an EOS. See each position's bigram and trigram, the hash heads' table rows, and where the document boundary replaces history with EOS.

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:

statewheresize per sequence
recurrent matrices + conv history36 linear-attention layers108 MiB, fixed
K, V and indexer keys12 attention layers27 KiB per token (KV 24 KiB, indexer 3 KiB): 0.84 GiB at 32k
last 2 token IDs + PLE conv historylayer 2a few KB, fixed
positionmodelone 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 LLM engine 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=4 replaces each layer’s experts with QuantizedExperts as 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=True never reads the n-gram shards into memory. MappedRows wraps 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):

machineexpertsdense weightsn-gram tableresident totaldecode ceiling
DGX Spark (128 GB unified, 273 GB/s)4-bit, in memoryBF16memory-mapped from NVMe~68 GiB~28 tokens/s
same, dense weights in INT84-bitINT8memory-mapped~63 GiB~50 tokens/s
H100 80 GB (3,350 GB/s)4-bitBF16host RAM~68 GiB~350 tokens/s
RTX 4090 24 GB + 128 GB host RAM4-bit, offloaded to host (stretch)BF16host RAM11 GiB on GPUPCIe-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

  1. ★★ Speculative decoding for Flash-Next: change speculative_generate to snapshot and restore HybridCache instead of truncating, and use a smaller model with the same tokenizer as the draft. Verify greedy outputs are unchanged. Where: speculative_generate in engine/speculative.py, using HybridCache.snapshot / restore in engine/flashnext.py.
  2. ★★ 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_flashnext in engine/flashnext.py, using engine.quant.
  3. ★★★ 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 in engine/flashnext.py; call it from engine/speculative.py.
  4. ★★★ Expert offloading for a 24 GB GPU: keep QuantizedExperts codes 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: QuantizedExperts and expert installation in engine/flashnext.py.
  5. ★★★ 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: new experiments/hybrid_batching.py, adapting engine.scheduler.ContinuousBatchingEngine for engine.flashnext.HybridState.

Check your understanding

  1. Which Flash-Next components come from which earlier chapters, and which two are new?
  2. How do the read and write weights of a hyper-connection differ in shape and range?
  3. Why does the n-gram memory use 8 hash heads with different prime sizes?
  4. Why can a 51B-parameter table live on disk without slowing decode much?
  5. Why is a functional (never modified in place) state convenient for speculative decoding?
  6. 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.py in 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.