Skip to content

Session bank: anchor a GDN boundary at the end of the prompt head (opt-in, depends on #550) - #559

Open
jvmenen wants to merge 5 commits into
youssofal:mainfrom
jvmenen:pr/session-head-anchor
Open

jvmenen wants to merge 5 commits into
youssofal:mainfrom
jvmenen:pr/session-head-anchor

Conversation

@jvmenen

@jvmenen jvmenen commented Sep 28, 2026

Copy link
Copy Markdown

Depends on #550 (and through it on #549). This branch is #550 plus one commit; only the last commit (perf(session-bank): anchor a GDN boundary at the end of the prompt head (opt-in)) is new here. The feature itself does not need the batch-invariant lane or the in-forward boundaries: without them the anchor is an ordinary mandatory span end of the tail ladder. The dependency is in the tests: the bit-for-bit restore tests run a tiny Qwen3.5-MoE on the lane with in-forward boundaries and reuse the fixtures of tests/test_gdn_inforward_boundaries.py from #550. Once #550 is merged, the diff reduces to that commit.

Summary

Opt-in MTPLX_SESSION_HEAD_ANCHOR=1: on hybrid (GDN) models, every prefill that goes to the session bank records a GDN restore boundary exactly where the prompt's fixed head (tool schemas and system turn) ends, and boundary thinning never drops it. A new session with the same head then restores the whole head from any banked entry of an earlier session, instead of only up to the oldest kept boundary. Default off; off keeps planning and retention unchanged.

Motivation

Agent harnesses start every session with the same head: tool schemas and the system text, often several thousand tokens. Sessions diverge at the first user turn. On a hybrid model a restore needs a recurrent (GDN) state at or below the shared prefix, and the bank only holds states at recorded boundaries. Thinning keeps at most 8 boundaries per entry: the oldest record and a geometric ladder near the tail. The donor for a new session is usually a long conversation, whose kept boundaries sit near its tail, far past the head, so a new session restores only from the oldest boundary and prefills the rest of the head again.

A boundary exactly at the end of the head, kept through thinning and inheritance, makes the whole head reusable.

Change

  • New mtplx/session_head_anchor.py: turn_markers(tokenizer) finds the ChatML turn-open token (atomic <|im_start|>) and the system role token; session_head_length(prompt_ids, markers) returns the number of tokens before the first turn that is not a system turn. Read from the prompt tokens, so it is token-exact for every ChatML template, harness and endpoint (chat completions and /v1/messages). No anchor for tokenizers without ChatML turns (for example Gemma 4), for image prompts, or for a head shorter than MTPLX_SESSION_BLOCK_PREFIX_MIN_MATCH_TOKENS (no restore could use it).
  • mtplx/generation.py:
    • GdnBoundarySink: the boundary list a prefill appends to, plus its anchors. A plain list stays a sink without anchors.
    • _mandatory_prefill_edges merges the existing stable-prefix edge with the anchors; the cold loops, the warm suffix loop and _predicted_first_prefill_span plan the anchor as a span end (inside the forward where the model publishes in-forward boundary hooks).
    • _thin_gdn_boundary_records(..., keep=...), the sink's thinning and _inherited_gdn_boundaries never drop a record at an anchor.
  • mtplx/cache_bank/cold_tier.py: with the switch on, the SSD block-prefix lane restores to the exact match instead of the block-aligned one (the RAM lane already does), so an unaligned anchor is reachable after a restart. The prefix decode takes any length, and a boundary restore picks the newest persisted boundary at or below the match.
  • mtplx/runtime_options.py: session_head_anchor_enabled() (MTPLX_SESSION_HEAD_ANCHOR, default off).

How to enable: MTPLX_SESSION_HEAD_ANCHOR=1, with a session bank on a hybrid model.

Tests

New tests/test_session_head_anchor.py (43 tests):

  • Head detection: first non-system turn, consecutive system turns, no head without a system turn or a later turn, atomic turn-open only; the real chat templates of locally cached Qwen packs with and without thinking, system and tools (skipped when the tokenizer is not in ~/.mtplx/models), /v1/messages, Gemma without ChatML turns.
  • Planning and retention: stable prefix and anchors merge into mandatory edges; the anchor is recorded inside the last wide forward; thinning, the sink and inheritance keep it within the cap.
  • A bank simulation of a long chat with the real planner, retention and SessionBank: a new session restores at the anchor with the switch on, at the oldest boundary with it off; the anchor survives the SSD tier.
  • On a tiny Qwen3.5-MoE with the invariant lane (Metal): the anchor leaves the prefill bit-identical (logits, hidden, cache); a new session restored at the anchor from RAM or from the SSD gives the cold prefill's logits and hidden bit for bit; an anchor inside a warm suffix does not change the result or the forward count.
  • tests/test_qwen4_ple_first_gather_early.py (updated for the planner helper), tests/test_gdn_inforward_boundaries.py, tests/test_batch_invariant_prefill.py and tests/test_session_bank.py pass.
  • Full suite with MTPLX_CONFIG=/nonexistent: 9,535 passed, 67 skipped, 1 failed. The failure is test_laguna_model.py::test_laguna_s_2_1_ar_route_skips_qwen_performance_hooks, which fails the same way on main on a Mac with less than 85.3 GiB (addressed in Tests no longer read the user's ~/.mtplx/config.toml or depend on the Mac's memory size #535). Ruff: new files clean, no new findings in the changed files.

🤖 Generated with Claude Code

Jeroen van Menen and others added 5 commits September 27, 2026 17:32
MTPLX_BATCH_INVARIANT_PREFILL=1 (default off) makes every prefill row
independent of how many rows share its forward. On MLX 0.32 three kernel
choices follow the row count and round differently: split-K quantized
matmul for narrow outputs (router N=256, shared expert, GDN a/b, k/v),
gather_qmv against gather_qmm_rhs for the routed experts, and the unfused
against the fused SDPA below 1024 query rows. On Qwen3.6-35B-A3B the 8th
and 9th expert often tie within a bf16 step, so those differences change
the routing and move scores between block layouts.

In the prefill phase only: QuantizedLinear runs as a two-batch matmul of
at least 33 rows per batch (no split-K, one NAX tile shape), SwitchGLU
pads to 4 token rows per expert (always the sorted tiled kernel), and
mlx-lm's scaled_dot_product_attention (hooked in qwen3_next and in base,
where attention_split imports it per call) always runs the fused causal
kernel, with zero query rows in front below 9 rows. The lone final prompt
token runs on the stock kernels (stock_prefill_kernels): padded to the
lane's minimums it cost TTFT, and no caller compares that row with a wider
forward. Class swaps and one function hook; decode and verify keep the
stock kernels.

The lane is installed only where it covers every quantized projection
(affine QuantizedLinear, stock SwitchGLU with at least 16 experts) and
only on MoE models; a dense model is refused with "dense_model" unless
MTPLX_BATCH_INVARIANT_PREFILL_DENSE=1. On a refusal nothing is swapped or
hooked. /health reports the install report or the refusal reason under
degradation.batch_invariant_prefill.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Prompt scoring (/v1/completions with echo and max_tokens 0) ran its trunk
in forwards of 256 rows, the width of its logits slices. With the
batch-invariant prefill lane installed the result no longer depends on
the width, so the trunk now runs at the prefill chunk a generation would
use (the server's live setting, else the profile) and the target lm_head
is applied per 256-row slice of the post-norm rows, inside the prefill
phase: logits residency stays at 256 rows. mtp_patch exposes the head as
logits_from_post_norm, the same call __call__ uses.

Without the lane the default stays 256 rows, because on MoE models a
wider trunk changes the routing and the scores.
MTPLX_PROMPT_SCORE_TRUNK_CHUNK names a width explicitly.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
…the invariant lane

With the batch-invariant prefill lane installed, a cold prompt shorter
than the block-restore floor (MTPLX_SESSION_BLOCK_PREFIX_MIN_MATCH_TOKENS,
512) is prefilled without the tail boundary grid. On Qwen3.6-35B-A3B the
grid cuts one extra forward off the last chunk (~0.1 s, all experts read
again) whenever a session bank is attached, e.g. [0,331],[331,395] for a
396-token prompt. With the lane both layouts give bitwise the same states,
so only the cost changes; without it the cut stays because removing it
would change the routing and the scores.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
…iant lane

A restore boundary inside the last prefill chunk ends a forward on families
without in-forward hooks (the tail ladder). On Qwen3.6-35B-A3B that extra
forward reads the routed experts again: ~0.1 s per warm turn and per
first-token request with a session bank from 512 tokens.

New mtplx/gdn_inforward_boundaries.py publishes the two hooks the prefill
loops already look for (boundary_capture_scope, take_boundary_captures).
While a capture is armed, each mlx-lm GDN layer (qwen3_5, qwen3_next) runs
its stock __call__ once per segment on its own cache, which is exactly the
ladder's arithmetic for that layer, and records its state between segments.
Attention and MoE see the wide forward, so the hooks are installed only
together with the batch-invariant prefill lane (MTPLX_BATCH_INVARIANT_PREFILL=1
and the lane admitted for the model); MTPLX_GDN_BOUNDARY_INFORWARD=0
restores the ladder. The install report gains gdn_inforward_layers.

Tests (tiny quantized Qwen3.5-MoE, real cold and warm prefill loops, Metal):
boundary positions, every record leaf, logits, hidden and cache bit-identical
to the ladder with 3 forwards instead of 6-7; no hooks without the lane or
on a refused model.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
…ad (opt-in)

New agent sessions share a fixed head (tool schemas and system text) and
diverge at the first user turn, but restored only from the oldest GDN
boundary: thinning to 8 boundaries per entry keeps the oldest and records
near the tail, and the donor is usually a long conversation.

With MTPLX_SESSION_HEAD_ANCHOR=1 (default off):
- the head end is read from the prompt tokens: the first <|im_start|> turn
  that is not a system turn (mtplx/session_head_anchor.py), so it is
  token-exact for every ChatML template, harness and endpoint;
- the prefill makes it a mandatory span end (recorded in-forward where the
  model publishes boundary hooks) and the sink, inheritance and thinning
  never drop that record;
- the SSD block-prefix lane restores to the exact match instead of the
  block-aligned one, so the unaligned anchor is reachable after a restart.

Switch off: no anchors, identical planning and retention.

Tests: tests/test_session_head_anchor.py (43): head detection on fake and
cached real ChatML tokenizers, planning, thinning, sink, inheritance and
SSD; a bank simulation of a long chat; on a tiny Qwen3.5-MoE with the
invariant lane and in-forward boundaries, the anchor leaves the prefill
bit-identical and a new session restored at the anchor (RAM or SSD) gives
the cold prefill's logits bit for bit.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
@jvmenen
jvmenen requested a review from youssofal as a code owner September 28, 2026 09:58

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant