[object Object]

← back to Exo

fixed errors

da52357321a1caee4dbcb3a0e837bb73c5f97291 · 2024-10-22 21:24:16 -0700 · josh

Files touched

Diff

commit da52357321a1caee4dbcb3a0e837bb73c5f97291
Author: josh <eyasunigussie@Eyasus-MacBook-Air.local>
Date:   Tue Oct 22 21:24:16 2024 -0700

    fixed errors
---
 exo/main.py                         | 20 +-------------------
 exo/networking/udp/udp_discovery.py |  6 +++++-
 exo/orchestration/standard_node.py  | 23 ++++++++++++++++++++---
 3 files changed, 26 insertions(+), 23 deletions(-)

diff --git a/exo/main.py b/exo/main.py
index 0358ca36..e3171ad1 100644
--- a/exo/main.py
+++ b/exo/main.py
@@ -153,24 +153,6 @@ async def shutdown(signal, loop):
   await server.stop()
   loop.stop()
 
-async def select_best_inference_engine(node: StandardNode):
-  supported_engines = node.get_supported_inference_engines()
-  await node.broadcast_supported_engines(supported_engines)
-  logger.error('ABOVE and ALL')
-  logger.error("Topology inference engines pool: %s", node.get_topology_inference_engines())
-  logger.error(f'result:{node.get_topology_inference_engines()}')
-  if node.get_topology_inference_engines():
-    logger.info("Topology inference engines pool: %s", node.get_topology_inference_engines())
-    topology_inference_engines_pool = node.get_topology_inference_engines()
-    if any("tinygrad" in engines and len(engines) == 1 for engines in topology_inference_engines_pool):
-        return "tinygrad"
-    common_engine_across_peers = set.intersection(*topology_inference_engines_pool)
-    if "mlx" in common_engine_across_peers:
-        return "mlx"
-    else:
-        raise ValueError("No compatible inference engine found across all nodes")
-
-
 async def run_model_cli(node: Node, inference_engine: InferenceEngine, model_name: str, prompt: str):
   shard = model_base_shards.get(model_name, {}).get(inference_engine.__class__.__name__)
   if not shard:
@@ -210,7 +192,7 @@ async def main():
     loop.add_signal_handler(s, handle_exit)
 
   await node.start(wait_for_peers=args.wait_for_peers)
-  await select_best_inference_engine(node)
+
   if args.command == "run" or args.run_model:
     model_name = args.model_name or args.run_model
     if not model_name:
diff --git a/exo/networking/udp/udp_discovery.py b/exo/networking/udp/udp_discovery.py
index 322c9741..eb2a90c3 100644
--- a/exo/networking/udp/udp_discovery.py
+++ b/exo/networking/udp/udp_discovery.py
@@ -159,7 +159,11 @@ class UDPDiscovery(Discovery):
           if peer_id in self.known_peers: del self.known_peers[peer_id]
           return
         if peer_id in self.known_peers: self.known_peers[peer_id] = (self.known_peers[peer_id][0], self.known_peers[peer_id][1], time.time(), peer_prio)
-  
+    if message["type"] == "supported_inference_engines":
+      logger.error(f'supported_inference_engines: {message}')
+      peer_id = message["node_id"]
+      engines = message["engines"]
+      if peer_id in self.known_peers: self.known_peers[peer_id][0].topology_inference_engines_pool.append(engines)
   async def task_listen_for_peers(self):
     await asyncio.get_event_loop().create_datagram_endpoint(lambda: ListenProtocol(self.on_listen_message),
                                                             local_addr=("0.0.0.0", self.listen_port))
diff --git a/exo/orchestration/standard_node.py b/exo/orchestration/standard_node.py
index 29085807..f81cbfe1 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -362,6 +362,24 @@ class StandardNode(Node):
     self.peers = next_peers
     return len(peers_added) > 0 or len(peers_removed) > 0 or len(peers_updated) > 0
 
+  async def select_best_inference_engine(self):
+    supported_engines = self.get_supported_inference_engines()
+    await self.broadcast_supported_engines(supported_engines)
+    logger.error('ABOVE and ALL')
+    logger.error("Topology inference engines pool: %s", self.get_topology_inference_engines())
+    logger.error(f'result:{self.get_topology_inference_engines()}')
+    if self.get_topology_inference_engines():
+      logger.info("Topology inference engines pool: %s", self.get_topology_inference_engines())
+      topology_inference_engines_pool = self.get_topology_inference_engines()
+      if any("tinygrad" in engines and len(engines) == 1 for engines in topology_inference_engines_pool):
+          return "tinygrad"
+      common_engine_across_peers = set.intersection(*topology_inference_engines_pool)
+      if "mlx" in common_engine_across_peers:
+          return "mlx"
+      else:
+          raise ValueError("No compatible inference engine found across all nodes")
+
+
   async def periodic_topology_collection(self, interval: int):
     while True:
       await asyncio.sleep(interval)
@@ -369,7 +387,9 @@ class StandardNode(Node):
         did_peers_change = await self.update_peers()
         if DEBUG >= 2: print(f"{did_peers_change=}")
         if did_peers_change:
+          logger.error('peers changed, collecting topology and selecting best inference engine')
           await self.collect_topology()
+          await self.select_best_inference_engine()
       except Exception as e:
         print(f"Error collecting topology: {e}")
         traceback.print_exc()
@@ -443,15 +463,12 @@ class StandardNode(Node):
     async def send_status_to_peer(peer):
       try:
         status_dict = json.loads(status)
-        if status_dict.get("type") == "supported_inference_engines":
-          logger.error(f'broadcasting_inference_engines: {status_dict}')
         await asyncio.wait_for(peer.send_opaque_status(request_id, status), timeout=15.0)
       except asyncio.TimeoutError:
         print(f"Timeout sending opaque status to {peer.id()}")
       except Exception as e:
         print(f"Error sending opaque status to {peer.id()}: {e}")
         traceback.print_exc()
-
     await asyncio.gather(*[send_status_to_peer(peer) for peer in self.peers], return_exceptions=True)
     # in the case of opaque status, we also want to receive our own opaque statuses
     self.on_opaque_status.trigger_all(request_id, status)

← 593d810d fix to broadcast  ·  back to Exo  ·  fix to creating engines cd4d324a →