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 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 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
68 changes: 68 additions & 0 deletions tests/test_streaming_math.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Loading