← back to Exo
logs for file filtering, grpc_discovery -> udp_discovery
65e0488ebe7160be1d91a4e9735e0fe8bbe02ef9 · 2024-08-21 15:59:34 +0100 · Alex Cheema
Files touched
M exo/download/hf/hf_helpers.pyM exo/download/hf/hf_shard_download.pyM exo/networking/grpc/test_grpc_discovery.pyM exo/networking/peer_handle.pyR097 exo/networking/grpc/grpc_discovery.py exo/networking/udp_discovery.pyM main.py
Diff
commit 65e0488ebe7160be1d91a4e9735e0fe8bbe02ef9
Author: Alex Cheema <alexcheema123@gmail.com>
Date: Wed Aug 21 15:59:34 2024 +0100
logs for file filtering, grpc_discovery -> udp_discovery
---
exo/download/hf/hf_helpers.py | 2 ++
exo/download/hf/hf_shard_download.py | 2 +-
exo/networking/grpc/test_grpc_discovery.py | 8 ++++----
exo/networking/peer_handle.py | 7 +++----
exo/networking/{grpc/grpc_discovery.py => udp_discovery.py} | 13 +++++--------
main.py | 4 ++--
6 files changed, 17 insertions(+), 19 deletions(-)
diff --git a/exo/download/hf/hf_helpers.py b/exo/download/hf/hf_helpers.py
index 5770a634..60962774 100644
--- a/exo/download/hf/hf_helpers.py
+++ b/exo/download/hf/hf_helpers.py
@@ -235,6 +235,7 @@ async def download_repo_files(repo_id: str, revision: str = "main", progress_cal
if DEBUG >= 2: print(f"Cached file list at {cached_file_list_path}")
filtered_file_list = list(filter_repo_objects(file_list, allow_patterns=allow_patterns, ignore_patterns=ignore_patterns, key=lambda x: x["path"]))
+ if DEBUG >= 2: print(f"Filtered file list {allow_patterns=} {ignore_patterns=}\noriginal: {file_list}\nfiltered: {filtered_file_list}")
total_files = len(filtered_file_list)
total_bytes = sum(file["size"] for file in filtered_file_list)
file_progress: Dict[str, RepoFileProgressEvent] = {file["path"]: RepoFileProgressEvent(repo_id, revision, file["path"], 0, 0, file["size"], 0, timedelta(0), "not_started") for file in filtered_file_list}
@@ -353,4 +354,5 @@ def get_allow_patterns(weight_map: Dict[str, str], shard: Shard) -> List[str]:
shard_specific_patterns.append(sorted_file_names[-1])
else:
shard_specific_patterns = ["*.safetensors"]
+ if DEBUG >= 2: print(f"get_allow_patterns {weight_map=} {shard=} {shard_specific_patterns=}")
return list(set(default_patterns + shard_specific_patterns)) # Remove duplicates
diff --git a/exo/download/hf/hf_shard_download.py b/exo/download/hf/hf_shard_download.py
index fa740e04..ac0bfb38 100644
--- a/exo/download/hf/hf_shard_download.py
+++ b/exo/download/hf/hf_shard_download.py
@@ -41,7 +41,7 @@ class HFShardDownloader(ShardDownloader):
try:
await task
except asyncio.CancelledError:
- pass # This is expected when cancelling a task
+ pass
except Exception as e:
if DEBUG >= 2: print(f"Error in cancelling download {active_shard}: {e}")
traceback.print_exc()
diff --git a/exo/networking/grpc/test_grpc_discovery.py b/exo/networking/grpc/test_grpc_discovery.py
index 13372bbb..64ce33b3 100644
--- a/exo/networking/grpc/test_grpc_discovery.py
+++ b/exo/networking/grpc/test_grpc_discovery.py
@@ -1,12 +1,12 @@
import asyncio
import unittest
-from .grpc_discovery import GRPCDiscovery
+from ..udp_discovery import UDPDiscovery
-class TestGRPCDiscovery(unittest.IsolatedAsyncioTestCase):
+class TestUDPDiscovery(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
- self.node1 = GRPCDiscovery("node1", 50051, 5678, 5679)
- self.node2 = GRPCDiscovery("node2", 50052, 5679, 5678)
+ self.node1 = UDPDiscovery("node1", 50051, 5678, 5679)
+ self.node2 = UDPDiscovery("node2", 50052, 5679, 5678)
await self.node1.start()
await self.node2.start()
diff --git a/exo/networking/peer_handle.py b/exo/networking/peer_handle.py
index cf232d00..9399a94a 100644
--- a/exo/networking/peer_handle.py
+++ b/exo/networking/peer_handle.py
@@ -5,7 +5,6 @@ from exo.inference.shard import Shard
from exo.topology.device_capabilities import DeviceCapabilities
from exo.topology.topology import Topology
-
class PeerHandle(ABC):
@abstractmethod
def id(self) -> str:
@@ -36,13 +35,13 @@ class PeerHandle(ABC):
pass
@abstractmethod
- async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
+ async def send_result(self, request_id: str, result: List[int], is_finished: bool) -> None:
pass
@abstractmethod
- async def collect_topology(self, visited: set[str], max_depth: int) -> Topology:
+ async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
pass
@abstractmethod
- async def send_result(self, request_id: str, result: List[int], is_finished: bool) -> None:
+ async def collect_topology(self, visited: set[str], max_depth: int) -> Topology:
pass
diff --git a/exo/networking/grpc/grpc_discovery.py b/exo/networking/udp_discovery.py
similarity index 97%
rename from exo/networking/grpc/grpc_discovery.py
rename to exo/networking/udp_discovery.py
index 12078472..a0e84ff6 100644
--- a/exo/networking/grpc/grpc_discovery.py
+++ b/exo/networking/udp_discovery.py
@@ -2,10 +2,11 @@ import asyncio
import json
import socket
import time
+import traceback
from typing import List, Dict, Callable, Tuple, Coroutine
-from ..discovery import Discovery
-from ..peer_handle import PeerHandle
-from .grpc_peer_handle import GRPCPeerHandle
+from .discovery import Discovery
+from .peer_handle import PeerHandle
+from .grpc.grpc_peer_handle import GRPCPeerHandle
from exo.topology.device_capabilities import DeviceCapabilities, device_capabilities, UNKNOWN_DEVICE_CAPABILITIES
from exo import DEBUG_DISCOVERY
@@ -23,7 +24,7 @@ class ListenProtocol(asyncio.DatagramProtocol):
asyncio.create_task(self.on_message(data, addr))
-class GRPCDiscovery(Discovery):
+class UDPDiscovery(Discovery):
def __init__(
self,
node_id: str,
@@ -114,8 +115,6 @@ class GRPCDiscovery(Discovery):
await asyncio.sleep(self.broadcast_interval)
except Exception as e:
print(f"Error in broadcast presence: {e}")
- import traceback
-
print(traceback.format_exc())
async def on_listen_message(self, data, addr):
@@ -185,6 +184,4 @@ class GRPCDiscovery(Discovery):
await asyncio.sleep(self.broadcast_interval)
except Exception as e:
print(f"Error in cleanup peers: {e}")
- import traceback
-
print(traceback.format_exc())
diff --git a/main.py b/main.py
index 5b07a5a7..5b038103 100644
--- a/main.py
+++ b/main.py
@@ -6,7 +6,7 @@ import time
import traceback
from exo.orchestration.standard_node import StandardNode
from exo.networking.grpc.grpc_server import GRPCServer
-from exo.networking.grpc.grpc_discovery import GRPCDiscovery
+from exo.networking.udp_discovery import UDPDiscovery
from exo.topology.ring_memory_weighted_partitioning_strategy import RingMemoryWeightedPartitioningStrategy
from exo.api import ChatGPTAPI
from exo.download.shard_download import ShardDownloader, RepoProgressEvent
@@ -48,7 +48,7 @@ if args.node_port is None:
if DEBUG >= 1: print(f"Using available port: {args.node_port}")
args.node_id = args.node_id or get_or_create_node_id()
-discovery = GRPCDiscovery(args.node_id, args.node_port, args.listen_port, args.broadcast_port, discovery_timeout=args.discovery_timeout)
+discovery = UDPDiscovery(args.node_id, args.node_port, args.listen_port, args.broadcast_port, discovery_timeout=args.discovery_timeout)
chatgpt_api_endpoints=[f"http://{ip}:{args.chatgpt_api_port}/v1/chat/completions" for ip in get_all_ip_addresses()]
web_chat_urls=[f"http://{ip}:{args.chatgpt_api_port}" for ip in get_all_ip_addresses()]
if DEBUG >= 0:
← cea9b48d update mlx-lm to 0.17.0, use lru caches for kv_cache with Ro
·
back to Exo
·
add a cli that can be triggered with --run-model <model> --p e8430431 →