diff --git a/src/edge0/moe/spec.py b/src/edge0/moe/spec.py index 176c43c..6d0d92c 100644 --- a/src/edge0/moe/spec.py +++ b/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/src/edge0/streaming/install.py b/src/edge0/streaming/install.py index 37b0e7e..c0c84de 100644 --- a/src/edge0/streaming/install.py +++ b/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/tests/test_streaming_math.py b/tests/test_streaming_math.py index e77154e..0d2bcbd 100644 --- a/tests/test_streaming_math.py +++ b/tests/test_streaming_math.py @@ -299,3 +299,71 @@ 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() + + +def test_layer_exists_stops_on_attribute_error(): + """The AttributeError half of ``layer_exists``'s except clause is + reachable independently of IndexError: a family whose layer container + is attribute-based (no list, no ``__getitem__``) runs off the end via + a plain missing attribute, not an out-of-range index.""" + from types import SimpleNamespace + + spec = MoESpec( + num_experts=N_EXPERTS, top_k=4, intermediate_size=INTER, + key_template="layer_{layer}.mlp.switch_mlp", + 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 # no `layer_2` attribute + + +def test_layer_exists_reraises_unrelated_index_error(): + """An IndexError from a segment *other* than the ``{layer}`` slot is a + real bug (a malformed block_path or a broken model), not end-of-list, + and must propagate instead of being read as "past the last layer" -- + otherwise install_streaming_experts(num_layers=None) would silently + under-count layers instead of surfacing the break.""" + from types import SimpleNamespace + + spec = MoESpec( + num_experts=N_EXPERTS, top_k=4, intermediate_size=INTER, + key_template="layers.{layer}.experts.9.switch_mlp", + block_path="layers.{layer}.experts.9", + ) + # `layer=0` is in range for `layers`, but the fixed trailing index `9` + # is out of range for `experts` -- unrelated to layer-count discovery. + model = SimpleNamespace(layers=[SimpleNamespace(experts=[object()])]) + with pytest.raises(IndexError): + spec.layer_exists(model, 0)