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
3 changes: 2 additions & 1 deletion docs/models/edge0-8b.md
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ Optional arguments:

- `--no-prerouter`: disable the prerouter (`prerouter=None`).
- `--no-lora`: disable LoRA (`lora=""`).
- `--prefill-ondemand`: prefill through per-expert on-demand loads instead of the E3b whole-layer path. A 27-token prompt then reads ≈0.4–0.8 GiB of routed experts instead of the whole ≈4.1 GiB checkpoint (measured cold-cache prefill: 0.6–1.1 s vs 5.3 s on an M4 Pro). This is the setting for a machine whose page cache cannot hold the checkpoint — 16 GB class, see issue #110.
- `--flask`: switch to the Flask transport (requires flask to be installed; supports SSE streaming).

For single-turn chat in the terminal, use `chat` instead:
Expand Down Expand Up @@ -159,4 +160,4 @@ engine = AutoEngine.from_pretrained(
)
```

The corresponding CLI overrides are `--no-prerouter` / `--no-lora` (see `_engine_kwargs` in `src/edge0/cli.py`). To override engine parameters, call `Ling8BConfig.from_pretrained(model_dir, **overrides)` directly.
The corresponding CLI override is `--prefill-ondemand` (config switch `prefill_ondemand`, applied to whatever preset the tier ships), alongside `--no-prerouter` / `--no-lora` (see `_engine_kwargs` in `src/edge0/cli.py`). To override engine parameters directly, call `Ling8BConfig.from_pretrained(model_dir, **overrides)`.
3 changes: 3 additions & 0 deletions docs/streaming.md
Original file line number Diff line number Diff line change
Expand Up @@ -73,10 +73,13 @@ With `use_compile`, the staged/exact paths are wrapped in `mx.compile`:
| `load_threads` / `prefetch_threads` | Build / prefetch thread counts | 8 / 4 |
| `full_layer_prefill` / `prefill_full_layers` | Whole-layer prefill loading / number of leading layers | False / 0 |
| `prefill_hot` | hot stack size during prefill | 0 |
| `warm_willneed` | Kernel bulk readahead (`madvise WILLNEED`) over the expert ranges a prefetch/stage is about to touch | False |
| `use_compile` / `top_k` | compile wrapping / routing top-k override | True / None |

Presets: `staged_k4()` (edge0-35b: staged decode with 4 slots, prefill hot stack 32, on-demand prefill), `prod_k8()` (edge0-8b: the reference deployment profile — staged decode off, E3b whole-layer prefill), and `staged_k8()` (the plain K=8 staged variant). Both tiers share `cache_slots=64`.

The whole-layer prefill is the fastest path **when the checkpoint stays in the page cache** (warm 27-token prefill: 0.24 s vs 0.37 s on-demand on an M4 Pro), and the slowest one when it does not (cold: 5.3 s / 4.06 GiB read vs 0.6-1.1 s / 0.4-0.8 GiB; the on-demand figure varies with how many distinct experts the prompt routes to). `edge0 demo|chat|serve --prefill-ondemand` selects the on-demand path for machines in the second group.

## Why Whole-Layer Loading Is Also Fast

The MoE weights in the checkpoint are already single per-layer stacked tensors (such as `[256, 512, 256]`), so a whole-layer load is 9 direct reads + a dtype view, not 256×9 per-expert builds. CPU loading and GPU execution are overlapped through the `before_layer_cb` + `async_eval_per_layer` pipeline.
Expand Down
48 changes: 28 additions & 20 deletions src/edge0/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,8 @@ def _engine_kwargs(args) -> dict:
kw["lora"] = ""
if getattr(args, "history_slots", False):
kw["history_slots"] = True
if getattr(args, "prefill_ondemand", False):
kw["prefill_ondemand"] = True
return kw


Expand Down Expand Up @@ -218,7 +220,24 @@ def cmd_convert(args) -> int:
return 0


def main(argv: list[str] | None = None) -> int:
def _add_engine_flags(p) -> None:
"""Engine-construction flags shared by demo / chat / serve."""
p.add_argument("--no-prerouter", action="store_true")
p.add_argument("--no-lora", action="store_true")
p.add_argument("--history-slots", action="store_true",
help="legacy staging: also fill staged slots from history "
"for layers whose route is not a prerouter prediction "
"(zeroes routed experts outside the slot set)")
p.add_argument("--prefill-ondemand", action="store_true",
help="prefill through per-expert on-demand loads instead of "
"the tier's whole-layer (E3b) path: reads only the "
"routed experts (~0.4-0.8 GiB instead of the whole "
"~4.1 GiB checkpoint for edge0-8b), which is what a "
"machine whose page cache cannot hold the checkpoint "
"needs")


def _build_parser() -> argparse.ArgumentParser:
ap = argparse.ArgumentParser(prog="edge0", description=__doc__)
sub = ap.add_subparsers(dest="cmd", required=True)

Expand All @@ -237,12 +256,7 @@ def main(argv: list[str] | None = None) -> int:
help="max tokens to generate (default: tier config)")
p.add_argument("--show-thinking", action="store_true",
help="print the model's reasoning block too")
p.add_argument("--no-prerouter", action="store_true")
p.add_argument("--no-lora", action="store_true")
p.add_argument("--history-slots", action="store_true",
help="legacy staging: also fill staged slots from history "
"for layers whose route is not a prerouter prediction "
"(zeroes routed experts outside the slot set)")
_add_engine_flags(p)
p.set_defaults(fn=cmd_demo)

p = sub.add_parser(
Expand All @@ -257,12 +271,7 @@ def main(argv: list[str] | None = None) -> int:
help="max tokens to generate (default: tier config)")
p.add_argument("--show-thinking", action="store_true",
help="print the model's reasoning block too")
p.add_argument("--no-prerouter", action="store_true")
p.add_argument("--no-lora", action="store_true")
p.add_argument("--history-slots", action="store_true",
help="legacy staging: also fill staged slots from history "
"for layers whose route is not a prerouter prediction "
"(zeroes routed experts outside the slot set)")
_add_engine_flags(p)
p.set_defaults(fn=cmd_chat)

p = sub.add_parser(
Expand All @@ -276,19 +285,18 @@ def main(argv: list[str] | None = None) -> int:
p.add_argument("--port", type=int, default=8000)
p.add_argument("--flask", action="store_true",
help="use the Flask transport (needs flask installed)")
p.add_argument("--no-prerouter", action="store_true")
p.add_argument("--no-lora", action="store_true")
p.add_argument("--history-slots", action="store_true",
help="legacy staging: also fill staged slots from history "
"for layers whose route is not a prerouter prediction "
"(zeroes routed experts outside the slot set)")
_add_engine_flags(p)
p.set_defaults(fn=cmd_serve)

p = sub.add_parser("convert-adapters",
help="one-shot legacy npz -> safetensors migration")
p.set_defaults(fn=cmd_convert)

args = ap.parse_args(argv)
return ap


def main(argv: list[str] | None = None) -> int:
args = _build_parser().parse_args(argv)
return args.fn(args)


Expand Down
27 changes: 26 additions & 1 deletion src/edge0/models/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,16 @@ class ModelConfig:
hot_window: int = 4
intra_staging: bool = False
prefetch_history: bool = True
prefill_ondemand: bool = False # force the per-expert on-demand prefill
# (full_layer_prefill=False) whatever the
# tier preset says. The lever for
# machines whose page cache cannot hold
# the checkpoint: a whole-layer prefill
# streams every expert of every layer
# (~4.1 GiB for edge0-8b, whatever the
# prompt length), the on-demand path only
# the routed experts (~0.4-0.8 GiB for a
# 27-token prompt).
port: int = 8000
# acceptance profile (measured on the production Mac)
target_tok_s: float = 0.0
Expand All @@ -86,13 +96,28 @@ def from_pretrained(cls, model_dir: str | None = None, **overrides):
f"unknown {cls.__name__} override {key!r} "
f"(known: {sorted(f.name for f in fields(base))})")
base = replace(base, **{key: value})
return base
return apply_prefill_switches(base)

@classmethod
def _defaults(cls, model_dir: str) -> "ModelConfig":
raise NotImplementedError


def apply_prefill_switches(cfg: "ModelConfig") -> "ModelConfig":
"""Map the prefill convenience switch onto ``cfg.options``.

``prefill_ondemand`` forces the per-expert on-demand prefill path
(``full_layer_prefill=False``) whatever the tier preset says; it is
exposed as ``edge0 demo|chat|serve --prefill-ondemand``. While it stays
False the tier preset is returned untouched, so the presets keep their
measured defaults.
"""
opts = getattr(cfg, "options", None)
if opts is None or not getattr(cfg, "prefill_ondemand", False):
return cfg
return replace(cfg, options=replace(opts, full_layer_prefill=False))


def resolve_prerouter(pspec: PrerouterSpec, weights_file: str) -> PrerouterSpec:
"""Fill the prerouter weights path (specs are frozen; only used when
the caller didn't already set one)."""
Expand Down
33 changes: 33 additions & 0 deletions tests/test_cli_flags.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
"""CLI wiring for the prefill switches (no GPU needed)."""

from __future__ import annotations

import pytest

from edge0.cli import _build_parser, _engine_kwargs


@pytest.mark.parametrize("cmd", ["demo", "chat", "serve"])
def test_prefill_ondemand_reaches_the_engine(cmd):
args = _build_parser().parse_args([cmd, "edge0-8b", "--prefill-ondemand"])
assert _engine_kwargs(args) == {"prefill_ondemand": True}


@pytest.mark.parametrize("cmd", ["demo", "chat", "serve"])
def test_prefill_ondemand_defaults_off(cmd):
args = _build_parser().parse_args([cmd, "edge0-8b"])
assert "prefill_ondemand" not in _engine_kwargs(args)


def test_existing_engine_flags_still_map():
args = _build_parser().parse_args(
["serve", "edge0-8b", "--no-prerouter", "--no-lora",
"--history-slots"])
assert _engine_kwargs(args) == {"prerouter": None, "lora": "",
"history_slots": True}


def test_serve_keeps_its_own_flags():
args = _build_parser().parse_args(
["serve", "edge0-8b", "--host", "0.0.0.0", "--port", "9"])
assert (args.host, args.port) == ("0.0.0.0", 9)
39 changes: 39 additions & 0 deletions tests/test_e2e_slow.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,45 @@ def test_edge0_35b_real_checkpoint():
assert len(output) <= MAX_NEW_TOKENS


def test_edge0_8b_prefill_ondemand_switch_skips_whole_layers(monkeypatch):
"""The `--prefill-ondemand` path (config switch -> engine) really takes
the on-demand prefill, and leaves the next token identical."""
from edge0.backends import core
from edge0.streaming.layer import StreamingSwitchGLU

model_dir = _checkpoint("EDGE0_8B_MODEL", "edge0-8b")
messages = [{"role": "user", "content": "你好,请用一句话介绍海滨城市。"}]

engine = AutoEngine.from_pretrained(model_dir, name="edge0-8b")
try:
ids = _chat_ids(engine, messages)
assert len(ids) > 1, "need a multi-token prefill"
engine.reset()
engine.prefill(ids)
baseline = int(core.argmax(engine.next_logits(), axis=-1).item())
finally:
engine.close()

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",
prefill_ondemand=True)
try:
engine.reset()
engine.prefill(ids)
assert loads == [], (
f"--prefill-ondemand still loaded {len(loads)} whole layers")
assert int(core.argmax(engine.next_logits(), axis=-1).item()) == baseline
finally:
engine.close()


def test_edge0_8b_ondemand_prefill_is_honored(monkeypatch):
"""``full_layer_prefill=False`` must really skip the whole-layer prefill.

Expand Down
22 changes: 22 additions & 0 deletions tests/test_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,28 @@ def test_override_and_reject():
AutoConfig.from_pretrained(name="edge0-8b", bogus_field=1)


def test_prefill_ondemand_switch_maps_onto_the_tier_preset():
"""`--prefill-ondemand` adjusts the tier's options without replacing the
preset, and defaults to leaving it alone."""
base = AutoConfig.from_pretrained(name="edge0-8b")
assert base.prefill_ondemand is False
assert base.options.full_layer_prefill is True # E3b, the tier default

ondemand = AutoConfig.from_pretrained(name="edge0-8b",
prefill_ondemand=True)
assert ondemand.options.full_layer_prefill is False
assert ondemand.options.prefill_full_layers == 0 # everything else kept
assert ondemand.options.prefill_hot == base.options.prefill_hot
assert ondemand.options.warm_willneed == base.options.warm_willneed
assert ondemand.moe_spec == base.moe_spec

# edge0-35b already prefills on demand (prefill_hot=32): no-op, not a
# preset rewrite.
q35 = AutoConfig.from_pretrained(name="edge0-35b", prefill_ondemand=True)
assert q35.options.full_layer_prefill is False
assert q35.options.prefill_hot == 32


def test_demo_defaults():
"""The demo entry points run each tier's showcase configuration."""
from edge0.registry import demo_kwargs
Expand Down
Loading