[object Object]

← 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

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 →