diff --git a/python/src/edge0/moe/spec.py b/python/src/edge0/moe/spec.py index ced4ac4..a39d403 100644 --- a/python/src/edge0/moe/spec.py +++ b/python/src/edge0/moe/spec.py @@ -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). diff --git a/python/src/edge0/streaming/install.py b/python/src/edge0/streaming/install.py index 37b0e7e..c0c84de 100644 --- a/python/src/edge0/streaming/install.py +++ b/python/src/edge0/streaming/install.py @@ -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( diff --git a/python/tests/test_moe_spec.py b/python/tests/test_moe_spec.py index b3dd18f..34a0f13 100644 --- a/python/tests/test_moe_spec.py +++ b/python/tests/test_moe_spec.py @@ -2,6 +2,8 @@ from __future__ import annotations +from types import SimpleNamespace + import pytest from edge0.moe.spec import (MoESpec, QuantSpec, RouterKind, WeightLayout) @@ -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() diff --git a/python/tests/test_streaming_install.py b/python/tests/test_streaming_install.py new file mode 100644 index 0000000..da88de3 --- /dev/null +++ b/python/tests/test_streaming_install.py @@ -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()) diff --git a/python/tests/test_streaming_math.py b/python/tests/test_streaming_math.py index eeb0f19..660d289 100644 --- a/python/tests/test_streaming_math.py +++ b/python/tests/test_streaming_math.py @@ -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()