← back to Exo
fix to broadcast
593d810db82e13d75b38c5fad53487f319409174 · 2024-10-22 19:46:46 -0700 · josh
Files touched
M exo/main.pyM exo/networking/udp/udp_discovery.pyM exo/orchestration/standard_node.py
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 →