[object Object]

← back to Exo

fix to broadcast

593d810db82e13d75b38c5fad53487f319409174 · 2024-10-22 19:46:46 -0700 · josh

Files touched

Diff

commit 593d810db82e13d75b38c5fad53487f319409174
Author: josh <eyasunigussie@Eyasus-MacBook-Air.local>
Date:   Tue Oct 22 19:46:46 2024 -0700

    fix to broadcast
---
 exo/main.py                         | 3 +--
 exo/networking/udp/udp_discovery.py | 7 +------
 exo/orchestration/standard_node.py  | 7 +++++++
 3 files changed, 9 insertions(+), 8 deletions(-)

diff --git a/exo/main.py b/exo/main.py
index 9331b6cb..0358ca36 100644
--- a/exo/main.py
+++ b/exo/main.py
@@ -210,13 +210,12 @@ 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:
       print("Error: Model name is required when using 'run' command or --run-model")
       return
-    await select_best_inference_engine(node)
     await run_model_cli(node, inference_engine, model_name, args.prompt)
   else:
     asyncio.create_task(api.run(port=args.chatgpt_api_port))  # Start the API server as a non-blocking task
diff --git a/exo/networking/udp/udp_discovery.py b/exo/networking/udp/udp_discovery.py
index 98fb6b3a..322c9741 100644
--- a/exo/networking/udp/udp_discovery.py
+++ b/exo/networking/udp/udp_discovery.py
@@ -159,12 +159,7 @@ 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)
-    elif 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 aef5bc7e..29085807 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -15,7 +15,9 @@ from exo import DEBUG
 from exo.helpers import AsyncCallbackSystem
 from exo.viz.topology_viz import TopologyViz
 from exo.download.hf.hf_helpers import RepoProgressEvent
+import logging
 
+logger = logging.getLogger(__name__)
 
 class StandardNode(Node):
   def __init__(
@@ -91,7 +93,9 @@ class StandardNode(Node):
       "node_id": self.id,
       "engines": supported_engines
     })
+    logger.error(f'broadcast_supported_engines: {status_message}')
     await self.broadcast_opaque_status("", status_message)
+    logger.error(f'broadcast_supported_engines: done')
 
   def get_topology_inference_engines(self) -> List[str]:
     return self.topology_inference_engines_pool
@@ -438,6 +442,9 @@ 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()}")

← f4a5562c added logger  ·  back to Exo  ·  fixed errors da523573 →