[object Object]

← back to Exo

Fixed MLX import blocking native Windows execution of exo. (Not Final)

6737e36e238ebb9b72f497b642f1e813f0416a80 · 2025-01-14 20:35:21 -0500 · Sandesh Bharadwaj

Files touched

Diff

commit 6737e36e238ebb9b72f497b642f1e813f0416a80
Author: Sandesh Bharadwaj <sndshvnktsh@gmail.com>
Date:   Tue Jan 14 20:35:21 2025 -0500

    Fixed MLX import blocking native Windows execution of exo. (Not Final)
---
 exo/api/chatgpt_api.py                  | 8 +++++++-
 exo/networking/grpc/grpc_peer_handle.py | 7 ++++++-
 exo/networking/grpc/grpc_server.py      | 8 +++++++-
 3 files changed, 20 insertions(+), 3 deletions(-)

diff --git a/exo/api/chatgpt_api.py b/exo/api/chatgpt_api.py
index 1928541d..18ef292e 100644
--- a/exo/api/chatgpt_api.py
+++ b/exo/api/chatgpt_api.py
@@ -21,7 +21,13 @@ from PIL import Image
 import numpy as np
 import base64
 from io import BytesIO
-import mlx.core as mx
+import platform
+
+if platform.system().lower() == "darwin" and platform.machine().lower() == "arm64":
+  import mlx.core as mx
+else:
+  import numpy as mx
+
 import tempfile
 from exo.download.hf.hf_shard_download import HFShardDownloader
 import shutil
diff --git a/exo/networking/grpc/grpc_peer_handle.py b/exo/networking/grpc/grpc_peer_handle.py
index e37a5f86..e3c10549 100644
--- a/exo/networking/grpc/grpc_peer_handle.py
+++ b/exo/networking/grpc/grpc_peer_handle.py
@@ -12,7 +12,12 @@ from exo.topology.topology import Topology
 from exo.topology.device_capabilities import DeviceCapabilities, DeviceFlops
 from exo.helpers import DEBUG
 import json
-import mlx.core as mx
+import platform
+
+if platform.system().lower() == "darwin" and platform.machine().lower() == "arm64":
+  import mlx.core as mx
+else:
+  import numpy as mx
 
 class GRPCPeerHandle(PeerHandle):
   def __init__(self, _id: str, address: str, desc: str, device_capabilities: DeviceCapabilities):
diff --git a/exo/networking/grpc/grpc_server.py b/exo/networking/grpc/grpc_server.py
index 6a12530a..6b8e388d 100644
--- a/exo/networking/grpc/grpc_server.py
+++ b/exo/networking/grpc/grpc_server.py
@@ -3,13 +3,19 @@ from concurrent import futures
 import numpy as np
 from asyncio import CancelledError
 
+import platform
+
 from . import node_service_pb2
 from . import node_service_pb2_grpc
 from exo import DEBUG
 from exo.inference.shard import Shard
 from exo.orchestration import Node
 import json
-import mlx.core as mx
+
+if platform.system().lower() == "darwin" and platform.machine().lower() == "arm64":
+  import mlx.core as mx
+else:
+  import numpy as mx
 
 
 class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):

← fcc699a5 fix  ·  back to Exo  ·  Add AMD GPU querying + Windows device capabilities df3624d2 →