From 62291e1a80e55651f0c8cd0c5be57343f81c8f1b Mon Sep 17 00:00:00 2001 From: linyubupa <403864360@qq.com> Date: Sun, 20 Sep 2026 15:49:13 +0800 Subject: [PATCH] fix(prefill): honor full_layer_prefill=False on the streaming engines `make_prefill_before_layer` used `full_n=0` both as "every layer" and as the unset/falsy value, so a hook built with `full_n=0` always fell through to `load_full_layer()`. Both streaming engines installed that hook for every multi-token prefill regardless of `full_layer_prefill`, so on the tiers that ship `prefill_hot=0` (edge0-8b) the flag was a silent no-op: disabling the whole-layer prefill still streamed all 128 experts of all 23 layers -- 4.06 GiB for a 27-token prompt, against 0.42 GiB for the routed experts alone. - hooks: take an explicit `full_layer` flag and skip the whole-layer load when it is off (the hot window, if any, still runs). - ling: pass `full_layer=opts.full_layer_prefill`, wire the hard-coded `hot_n=0` to `opts.prefill_hot`, and install the hook only when the whole-layer prefill or the hot stack is actually requested. - qwen: the same gating. Measured (M4 Pro, 24 GB, 27-token prompt, buffer cache evicted, MLX peak unchanged at 0.88-0.91 GiB): full_layer_prefill=True 5.50 s 266,354 pageins (4.06 GiB) 23 whole-layer loads full_layer_prefill=False 0.56 s 27,528 pageins (0.42 GiB) 0 whole-layer loads Warm cache the two arms are level (0.24 s vs 0.34 s), so E3b stays the right default wherever the checkpoint fits in the page cache; the on-demand path is the lever for machines that cannot hold it (16 GB class, issue #110). Tests: `tests/test_prefill_hook.py` covers the gating logic without weights, and a slow real-checkpoint regression in `tests/test_e2e_slow.py` asserts the on-demand arm never loads a whole layer while producing the same next token. --- docs/models/edge0-8b.md | 2 +- docs/streaming.md | 2 +- src/edge0/engine/hooks.py | 17 ++++++--- src/edge0/engine/ling.py | 9 +++-- src/edge0/engine/qwen.py | 5 ++- tests/test_e2e_slow.py | 52 ++++++++++++++++++++++++++++ tests/test_prefill_hook.py | 71 ++++++++++++++++++++++++++++++++++++++ 7 files changed, 148 insertions(+), 10 deletions(-) create mode 100644 tests/test_prefill_hook.py diff --git a/docs/models/edge0-8b.md b/docs/models/edge0-8b.md index f44792c..7517db4 100644 --- a/docs/models/edge0-8b.md +++ b/docs/models/edge0-8b.md @@ -34,7 +34,7 @@ This tier uses the `LayerOptions.prod_k8()` preset (aligned with the reference d - Staged decode is off (`staged=False`, `staged_sync=False`, `staged_n=8`) — deployment verification showed that staged decode degrades output on this tier, so the prerouter directly drives expert prefetch for the next token (one `stage_all` at the step boundary and one at the prefill tail). - Expert cache `cache_slots=64`, hot-expert pinning off (`hot_per_layer=0`). -- Full-layer E3b prefill (`full_layer_prefill=True`, `prefill_chunk=2048`). +- Full-layer E3b prefill (`full_layer_prefill=True`, `prefill_chunk=2048`). Setting `full_layer_prefill=False` switches this tier to on-demand prefill: for a 27-token prompt that reads ≈0.42 GiB of routed experts instead of the whole ≈4.1 GiB checkpoint (measured on an M4 Pro), which is the difference between ~0.6 s and ~5 s of cold-cache prefill — the lever for machines whose page cache cannot hold the checkpoint (issue #110). ## Usage diff --git a/docs/streaming.md b/docs/streaming.md index 7ba4647..9e9b558 100644 --- a/docs/streaming.md +++ b/docs/streaming.md @@ -50,7 +50,7 @@ Deduplicate the routing indices → `_get_bundles(unique)` builds per expert (re - `load_full_layer()`: loads the layer's 9 tensors directly (the checkpoint already stores each layer as a single stacked tensor; the mmap dtype conversion costs ≈9ms per layer, and the CPU load is hidden under the previous layer's GPU execution); - `_gather_sort` folds the batch into the token dimension → sorted gather → `_scatter_unsort` un-sorts and restores the batch dimension; -- `full_layer_prefill` turns whole-layer prefill on per tier, and `prefill_full_layers` limits it to the leading N layers (0 = every layer, the edge0-8b setting); layers outside that range go through hot/exact; +- `full_layer_prefill` turns whole-layer prefill on per tier, and `prefill_full_layers` limits it to the leading N layers (0 = every layer, the edge0-8b setting); layers outside that range go through hot/exact. With `full_layer_prefill=False` the whole-layer path is off entirely — no layer is loaded whole, whatever `prefill_full_layers` says — and prefill runs the `prefill_hot` window if one is set, else the plain per-expert on-demand path; - After use, `clear_full_layer()` frees the GPU copy, and the page cache carries the hot data. ## Sorting and Compilation diff --git a/src/edge0/engine/hooks.py b/src/edge0/engine/hooks.py index c075dc5..8e5a931 100644 --- a/src/edge0/engine/hooks.py +++ b/src/edge0/engine/hooks.py @@ -9,7 +9,8 @@ from __future__ import annotations -def make_prefill_before_layer(all_stream_layers, *, full_n: int = 0, +def make_prefill_before_layer(all_stream_layers, *, full_layer: bool = True, + full_n: int = 0, hot_n: int = 0, hot_window: int = 4): """Before-layer prefill hook: whole-layer load-drop (E3b) + the sliding hot-expert window. @@ -17,9 +18,15 @@ def make_prefill_before_layer(all_stream_layers, *, full_n: int = 0, Before layer ``li`` runs: drop layer ``li-1``'s whole-layer set (its async_eval was already submitted, so the GPU queue keeps the arrays alive until evaluated), then either load layer ``li``'s full stacked - tensors (``full_n`` unset / ``li < full_n``) or build the hot-stack - window (``hot_n``): numpy backing for layers ``li..li+ahead``, - GPU materialization for the window, dematerialize trailing layers. + tensors (``full_layer`` and ``full_n`` unset / ``li < full_n``) or + build the hot-stack window (``hot_n``): numpy backing for layers + ``li..li+ahead``, GPU materialization for the window, dematerialize + trailing layers. + + ``full_layer=False`` switches the whole-layer load-drop off (the hot + window, if requested, still runs). Callers should not install this + hook at all when both ``full_layer`` and ``hot_n`` are off, so a + disabled prefill keeps the plain per-expert on-demand path. """ w = max(1, hot_window) ahead = max(1, w // 2) @@ -44,6 +51,8 @@ def before_layer(li: int) -> None: exp_t.dematerialize_hot() if not full_n: return + if not full_layer: + return exp = all_stream_layers.get(li) if exp is not None: exp.load_full_layer() diff --git a/src/edge0/engine/ling.py b/src/edge0/engine/ling.py index 87229b9..936803c 100644 --- a/src/edge0/engine/ling.py +++ b/src/edge0/engine/ling.py @@ -151,8 +151,9 @@ def _build(self): self._prefill_before_layer = make_prefill_before_layer( self._all_stream_layers, + full_layer=bool(getattr(opts, "full_layer_prefill", False)), full_n=getattr(opts, "prefill_full_layers", 0), - hot_n=0, hot_window=1) + hot_n=getattr(opts, "prefill_hot", 0), hot_window=1) self._history_prefetch = make_history_prefetch( self._all_stream_layers, enabled=cfg.prefetch_history) @@ -208,8 +209,10 @@ def _forward(self, ids, intra_stage: bool = True) -> core.array: and opts.full_layer_prefill) h = self.model.model( inputs, cache=self.cache, - before_layer_cb=self._prefill_before_layer if prefill_multi - else None, + before_layer_cb=(self._prefill_before_layer + if (prefill_multi and (full_layer + or opts.prefill_hot)) + else None), after_layer_cb=None, async_eval_per_layer=bool(prefill_multi and full_layer), prerouter_cache=(self._pg_stager.pg_cache diff --git a/src/edge0/engine/qwen.py b/src/edge0/engine/qwen.py index cdf88c4..b22a107 100644 --- a/src/edge0/engine/qwen.py +++ b/src/edge0/engine/qwen.py @@ -141,6 +141,7 @@ def _build(self): self._prefill_before_layer = make_prefill_before_layer( self._all_stream_layers, + full_layer=bool(getattr(opts, "full_layer_prefill", False)), full_n=getattr(opts, "prefill_full_layers", 0), hot_n=opts.prefill_hot, hot_window=cfg.hot_window) self._intra_after_layer = make_intra_after_layer( @@ -159,7 +160,9 @@ def _forward(self, ids, intra_stage: bool = True) -> core.array: full_layer = bool( prefill_multi and self._all_stream_layers and opts.full_layer_prefill) - before_cb = self._prefill_before_layer if prefill_multi else None + before_cb = (self._prefill_before_layer + if (prefill_multi and (full_layer or opts.prefill_hot)) + else None) if full_layer: after_cb = None else: diff --git a/tests/test_e2e_slow.py b/tests/test_e2e_slow.py index 3e4d8b7..5b18e86 100644 --- a/tests/test_e2e_slow.py +++ b/tests/test_e2e_slow.py @@ -87,3 +87,55 @@ def test_edge0_35b_real_checkpoint(): ) assert output assert len(output) <= MAX_NEW_TOKENS + + +def test_edge0_8b_ondemand_prefill_is_honored(monkeypatch): + """``full_layer_prefill=False`` must really skip the whole-layer prefill. + + Regression guard: the flag was a no-op on this tier. The prefill hook + was installed for every multi-token prefill regardless of the flag, and + with ``full_n=0`` ("every layer") it always fell through to + ``load_full_layer()`` -- so disabling it still read all 128 experts of + all 23 layers (~4.1 GiB for a 27-token prompt) instead of the routed + experts only (~0.42 GiB, measured on an M4 Pro). + """ + from dataclasses import replace + + from edge0.backends import core + from edge0.streaming.layer import StreamingSwitchGLU + from edge0.streaming.options import LayerOptions + + model_dir = _checkpoint("EDGE0_8B_MODEL", "edge0-8b") + messages = [{"role": "user", "content": "你好,请用一句话介绍海滨城市。"}] + + # 1) tier default (whole-layer load-drop prefill) + engine = AutoEngine.from_pretrained(model_dir, name="edge0-8b") + try: + prompt_ids = _chat_ids(engine, messages) + assert len(prompt_ids) > 1, "need a multi-token prefill" + engine.reset() + engine.prefill(prompt_ids) + baseline = int(core.argmax(engine.next_logits(), axis=-1).item()) + finally: + engine.close() + + # 2) on-demand prefill: no whole-layer load may fire, same next token + loads: list[int] = [] + original = StreamingSwitchGLU.load_full_layer + + def counted(self, *args, **kwargs): + loads.append(1) + return original(self, *args, **kwargs) + + monkeypatch.setattr(StreamingSwitchGLU, "load_full_layer", counted) + engine = AutoEngine.from_pretrained( + model_dir, name="edge0-8b", + options=replace(LayerOptions.prod_k8(), full_layer_prefill=False)) + try: + engine.reset() + engine.prefill(prompt_ids) + assert loads == [], ( + f"on-demand prefill still loaded {len(loads)} whole layers") + assert int(core.argmax(engine.next_logits(), axis=-1).item()) == baseline + finally: + engine.close() diff --git a/tests/test_prefill_hook.py b/tests/test_prefill_hook.py new file mode 100644 index 0000000..9c195e1 --- /dev/null +++ b/tests/test_prefill_hook.py @@ -0,0 +1,71 @@ +"""The prefill hook must honor ``full_layer=False``. + +Regression guard for a silent no-op: ``make_prefill_before_layer`` used +``full_n=0`` to mean "every layer" *and* as the falsy "unset" value, so a +hook built with ``full_n=0`` fell straight through to +``load_full_layer()``. Because the engines installed the hook for every +multi-token prefill regardless of ``full_layer_prefill``, that flag did +nothing on the tier that ships ``prefill_hot=0``: disabling it still +streamed all 128 experts of all 23 layers (~4.1 GiB for a 27-token +prompt) instead of the routed experts only (~0.4 GiB). +""" + +from __future__ import annotations + +from edge0.engine.hooks import make_prefill_before_layer + + +class _FakeLayer: + def __init__(self) -> None: + self.calls: list = [] + + def clear_full_layer(self) -> None: + self.calls.append("clear") + + def load_full_layer(self) -> None: + self.calls.append("full") + + def load_hot_layer(self, n: int) -> None: + self.calls.append(("hot_load", n)) + + def materialize_hot(self) -> None: + self.calls.append("hot_mat") + + def dematerialize_hot(self) -> None: + self.calls.append("hot_dem") + + +def _run(n: int = 4, **kwargs): + layers = {i: _FakeLayer() for i in range(n)} + before_layer = make_prefill_before_layer(layers, **kwargs) + for li in range(n): + before_layer(li) + return layers + + +def _full_loads(layers) -> list[int]: + return sorted(li for li, exp in layers.items() if "full" in exp.calls) + + +def test_full_layer_disabled_never_loads_a_whole_layer(): + assert _full_loads(_run(full_layer=False)) == [] + + +def test_full_layer_default_still_loads_every_layer(): + # full_n=0 means "every layer" (the prod_k8 default). + assert _full_loads(_run()) == [0, 1, 2, 3] + + +def test_full_layer_honors_the_leading_layer_count(): + assert _full_loads(_run(full_layer=True, full_n=2)) == [0, 1] + + +def test_hot_window_survives_full_layer_disabled(): + layers = _run(full_layer=False, hot_n=8) + assert _full_loads(layers) == [] + assert any("hot_mat" in exp.calls for exp in layers.values()) + + +def test_hot_stack_layer_is_not_also_loaded_whole(): + # staged_k4 shape: no full layers, hot stack only. + assert _full_loads(_run(full_layer=False, hot_n=32)) == []