Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion docs/models/edge0-8b.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
2 changes: 1 addition & 1 deletion docs/streaming.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
17 changes: 13 additions & 4 deletions src/edge0/engine/hooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,17 +9,24 @@
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.

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)
Expand All @@ -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()
Expand Down
9 changes: 6 additions & 3 deletions src/edge0/engine/ling.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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
Expand Down
5 changes: 4 additions & 1 deletion src/edge0/engine/qwen.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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:
Expand Down
52 changes: 52 additions & 0 deletions tests/test_e2e_slow.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
71 changes: 71 additions & 0 deletions tests/test_prefill_hook.py
Original file line number Diff line number Diff line change
@@ -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)) == []
Loading