Skip to content
Open
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
29 changes: 29 additions & 0 deletions python/src/edge0/moe/spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,35 @@ def block_of(self, model, layer: int):
obj = getattr(obj, part)
return obj

def layer_exists(self, model, layer: int) -> bool:
"""True if ``layer`` resolves at block_path's layer-index segment.

Used by layer-count discovery (``install_streaming_experts`` with
``num_layers=None``) to find where the layer list ends. Only an
``AttributeError``/``IndexError`` raised while resolving the
segment templated by ``{layer}`` itself means "past the last
layer"; the same errors raised by a *different* segment further
down ``block_path`` (e.g. a fixed expert-slot index) indicate a
bug in the spec or model and are re-raised rather than read as
end-of-list.
"""
template_parts = self.block_path.split(".")
layer_pos = next(
i for i, p in enumerate(template_parts) if "{layer}" in p)
obj = model
for i, raw_part in enumerate(template_parts):
part = raw_part.format(layer=layer)
try:
if part.isdigit():
obj = obj[int(part)]
else:
obj = getattr(obj, part)
except (AttributeError, IndexError):
if i == layer_pos:
return False
raise
return True

def layer_of(self, model, layer: int):
"""Resolve the decoder layer object (the block's owner).

Expand Down
6 changes: 1 addition & 5 deletions python/src/edge0/streaming/install.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,11 +38,7 @@ def install_streaming_experts(
"""
if num_layers is None:
n = 0
while True:
try:
spec.block_of(model, n)
except AttributeError:
break
while spec.layer_exists(model, n):
n += 1
if n == 0:
raise ValueError(
Expand Down
37 changes: 37 additions & 0 deletions python/tests/test_moe_spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@

from __future__ import annotations

from types import SimpleNamespace

import pytest

from edge0.moe.spec import (MoESpec, QuantSpec, RouterKind, WeightLayout)
Expand Down Expand Up @@ -59,6 +61,41 @@ def test_block_of_digit_segments():
assert s.block_of(m, 0) is m.language_model.model.layers[0].mlp.switch_mlp


def test_layer_exists_stops_at_end_of_list():
spec = _spec()
model = _FakeModel(n=2)
assert spec.layer_exists(model, 0) is True
assert spec.layer_exists(model, 1) is True
assert spec.layer_exists(model, 2) is False
assert spec.layer_exists(_FakeModel(n=0), 0) is False


def test_layer_exists_stops_on_attribute_error():
spec = _spec(block_path="layer_{layer}")
model = SimpleNamespace(layer_0=object(), layer_1=object())
assert spec.layer_exists(model, 0) is True
assert spec.layer_exists(model, 1) is True
assert spec.layer_exists(model, 2) is False


def test_layer_exists_reraises_unrelated_index_error():
spec = _spec(block_path="layers.{layer}.experts.9")
model = SimpleNamespace(layers=[SimpleNamespace(experts=[object()])])
with pytest.raises(IndexError):
spec.layer_exists(model, 0)


@pytest.mark.parametrize("block_path", [
"missing.layers.{layer}.mlp",
"layers.{layer}.missing",
])
def test_layer_exists_reraises_unrelated_attribute_error(block_path):
spec = _spec(block_path=block_path)
model = SimpleNamespace(layers=[SimpleNamespace(mlp=object())])
with pytest.raises(AttributeError):
spec.layer_exists(model, 0)


def test_layer_of_defaults_from_block_path():
s = _spec()
m = _FakeModel()
Expand Down
46 changes: 46 additions & 0 deletions python/tests/test_streaming_install.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
"""Layer discovery and installation without tensor calculations."""

from types import SimpleNamespace

import pytest

from edge0.moe.spec import MoESpec
from edge0.streaming import install


def _spec():
return MoESpec(
num_experts=8, top_k=2, intermediate_size=64,
block_path="layers.{layer}.mlp",
)


def test_install_discovers_mixed_layers(monkeypatch):
twin = object()
resident = object()
moe = SimpleNamespace(switch_mlp=resident)
dense = SimpleNamespace()
model = SimpleNamespace(layers=[
SimpleNamespace(mlp=moe), SimpleNamespace(mlp=dense),
])
calls = []

def make_twin(shards, layer, spec, **kwargs):
calls.append((shards, layer, spec))
return twin

monkeypatch.setattr(install, "StreamingSwitchGLU", make_twin)
spec = _spec()
twins = install.install_streaming_experts(model, [], spec)

assert twins == [twin, None]
assert calls == [([], 0, spec)]
assert moe.switch_mlp is twin
assert moe._edge0_resident_switch is resident
assert not hasattr(dense, "switch_mlp")


def test_install_rejects_empty_layer_list():
model = SimpleNamespace(layers=[])
with pytest.raises(ValueError, match="resolves no layer 0"):
install.install_streaming_experts(model, [], _spec())
30 changes: 30 additions & 0 deletions python/tests/test_streaming_math.py
Original file line number Diff line number Diff line change
Expand Up @@ -344,3 +344,33 @@ def test_double_buffered_swap(layer):
ref2 = lay(x, mx.array([second], dtype=mx.int32))
assert mx.allclose(out1, ref1).item()
assert mx.allclose(out2, ref2).item()


def test_install_discovers_layer_count(tmp_path):
"""install_streaming_experts(num_layers=None) probes block_path with
increasing layer indices until it stops resolving. Layers live in a
list, so running off the end raises IndexError, not AttributeError."""
from types import SimpleNamespace

from edge0.streaming.install import install_streaming_experts

path = tmp_path / "w.safetensors"
_write_shard(path, fuse_gu=False)
spec = MoESpec(
num_experts=N_EXPERTS, top_k=4, intermediate_size=INTER,
quant=QuantSpec(bits=4, group_size=64),
layout=WeightLayout.SEPARATE,
key_template="layers.{layer}.mlp.switch_mlp",
block_path="layers.{layer}.mlp",
)
moe = SimpleNamespace(switch_mlp=object())
model = SimpleNamespace(layers=[SimpleNamespace(mlp=moe),
SimpleNamespace(mlp=SimpleNamespace()),
SimpleNamespace(mlp=SimpleNamespace())])
twins = install_streaming_experts(
model, [SafetensorsMmap(str(path))], spec, options=_options())
assert len(twins) == 3
assert isinstance(twins[0], StreamingSwitchGLU)
assert twins[1] is None and twins[2] is None # dense layers
assert moe.switch_mlp is twins[0]
twins[0].close()
Loading