← back to Exo
Register original layer in CustomMlxLayer (#1229)
9e2179c848e937fe20a6a52daeb07cbf32cf5c9c · 2026-01-20 18:20:01 +0000 · rltakashige
## Motivation
Kimi K2 Thinking Pipeline RDMA was broken before.
## Why It Works
No clue tbh
## Test Plan
### Manual Testing
Kimi K2 Thinking and GPT OSS work at the same time on Pipeline RDMA.
Needs exo bench to check more thoroughly
### Automated Testing
Layer composition tests still pass.
Files touched
M src/exo/worker/engines/mlx/auto_parallel.py
Diff
commit 9e2179c848e937fe20a6a52daeb07cbf32cf5c9c
Author: rltakashige <rl.takashige@gmail.com>
Date: Tue Jan 20 18:20:01 2026 +0000
Register original layer in CustomMlxLayer (#1229)
## Motivation
Kimi K2 Thinking Pipeline RDMA was broken before.
## Why It Works
No clue tbh
## Test Plan
### Manual Testing
Kimi K2 Thinking and GPT OSS work at the same time on Pipeline RDMA.
Needs exo bench to check more thoroughly
### Automated Testing
Layer composition tests still pass.
---
src/exo/worker/engines/mlx/auto_parallel.py | 6 +++---
1 file changed, 3 insertions(+), 3 deletions(-)
diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py
index 7e5562a1..cb0ea2a4 100644
--- a/src/exo/worker/engines/mlx/auto_parallel.py
+++ b/src/exo/worker/engines/mlx/auto_parallel.py
@@ -83,11 +83,11 @@ class CustomMlxLayer(nn.Module):
def __init__(self, original_layer: _LayerCallable):
super().__init__()
- object.__setattr__(self, "_original_layer", original_layer)
+ dict.__setitem__(self, "_original_layer", original_layer) # pyright: ignore[reportUnknownMemberType]
@property
def original_layer(self) -> _LayerCallable:
- return cast(_LayerCallable, object.__getattribute__(self, "_original_layer"))
+ return cast(_LayerCallable, self["_original_layer"])
# Calls __getattr__ for any attributes not found on nn.Module (e.g. use_sliding)
if not TYPE_CHECKING:
@@ -96,7 +96,7 @@ class CustomMlxLayer(nn.Module):
try:
return super().__getattr__(name)
except AttributeError:
- original_layer = object.__getattribute__(self, "_original_layer")
+ original_layer = cast(_LayerCallable, self["_original_layer"])
return getattr(original_layer, name)
← 22b5d836 swap all instances of model_id: str for model_id: ModelId (#
·
back to Exo
·
Fix GPT OSS tensor sharding with upstream MLX LM (#1223) 75846470 →