[object Object]

← back to Exo

keep track of already visited peers in global operations: collect_topology

d2184f583aa7757ee94fce1613742494c4b7d7c6 · 2024-07-15 15:43:00 -0700 · Alex Cheema

Files touched

Diff

commit d2184f583aa7757ee94fce1613742494c4b7d7c6
Author: Alex Cheema <alexcheema123@gmail.com>
Date:   Mon Jul 15 15:43:00 2024 -0700

    keep track of already visited peers in global operations: collect_topology
---
 exo/networking/grpc/grpc_peer_handle.py |  4 ++--
 exo/networking/grpc/grpc_server.py      |  3 ++-
 exo/networking/grpc/node_service.proto  |  3 ++-
 exo/networking/grpc/node_service_pb2.py | 32 ++++++++++++++++----------------
 exo/networking/peer_handle.py           |  2 +-
 exo/orchestration/node.py               |  6 +++++-
 exo/orchestration/standard_node.py      | 28 ++++++++++++++++++++--------
 7 files changed, 48 insertions(+), 30 deletions(-)

diff --git a/exo/networking/grpc/grpc_peer_handle.py b/exo/networking/grpc/grpc_peer_handle.py
index 3ad44a95..1f3e7418 100644
--- a/exo/networking/grpc/grpc_peer_handle.py
+++ b/exo/networking/grpc/grpc_peer_handle.py
@@ -77,8 +77,8 @@ class GRPCPeerHandle(PeerHandle):
         request = node_service_pb2.ResetShardRequest(shard=node_service_pb2.Shard(model_id=shard.model_id, start_layer=shard.start_layer, end_layer=shard.end_layer, n_layers=shard.n_layers))
         await self.stub.ResetShard(request)
 
-    async def collect_topology(self, max_depth: int) -> Topology:
-        request = node_service_pb2.CollectTopologyRequest(max_depth=max_depth)
+    async def collect_topology(self, visited: set[str], max_depth: int) -> Topology:
+        request = node_service_pb2.CollectTopologyRequest(visited=visited, max_depth=max_depth)
         response = await self.stub.CollectTopology(request)
         topology = Topology()
         for node_id, capabilities in response.nodes.items():
diff --git a/exo/networking/grpc/grpc_server.py b/exo/networking/grpc/grpc_server.py
index a81b3ce4..f6586e6d 100644
--- a/exo/networking/grpc/grpc_server.py
+++ b/exo/networking/grpc/grpc_server.py
@@ -65,7 +65,8 @@ class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):
 
     async def CollectTopology(self, request, context):
         max_depth = request.max_depth
-        topology = await self.node.collect_topology(max_depth)
+        visited = set(request.visited)
+        topology = await self.node.collect_topology(visited, max_depth)
         nodes = {node_id: node_service_pb2.DeviceCapabilities(model=cap.model, chip=cap.chip, memory=cap.memory) for node_id, cap in topology.nodes.items()}
         peer_graph = {node_id: node_service_pb2.Peers(peer_ids=peers) for node_id, peers in topology.peer_graph.items()}
         return node_service_pb2.Topology(nodes=nodes, peer_graph=peer_graph)
diff --git a/exo/networking/grpc/node_service.proto b/exo/networking/grpc/node_service.proto
index 4c31d581..38cfd30b 100644
--- a/exo/networking/grpc/node_service.proto
+++ b/exo/networking/grpc/node_service.proto
@@ -49,7 +49,8 @@ message ResetShardRequest {
 }
 
 message CollectTopologyRequest {
-  int32 max_depth = 1;
+  repeated string visited = 1;
+  int32 max_depth = 2;
 }
 
 message Topology {
diff --git a/exo/networking/grpc/node_service_pb2.py b/exo/networking/grpc/node_service_pb2.py
index 2694191f..a7f51cdc 100644
--- a/exo/networking/grpc/node_service_pb2.py
+++ b/exo/networking/grpc/node_service_pb2.py
@@ -14,7 +14,7 @@ _sym_db = _symbol_database.Default()
 
 
 
-DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x12node_service.proto\x12\x0cnode_service\"S\n\x05Shard\x12\x10\n\x08model_id\x18\x01 \x01(\t\x12\x13\n\x0bstart_layer\x18\x02 \x01(\x05\x12\x11\n\tend_layer\x18\x03 \x01(\x05\x12\x10\n\x08n_layers\x18\x04 \x01(\x05\"k\n\rPromptRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12\x0e\n\x06prompt\x18\x02 \x01(\t\x12\x17\n\nrequest_id\x18\x03 \x01(\tH\x00\x88\x01\x01\x42\r\n\x0b_request_id\"\x81\x01\n\rTensorRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12$\n\x06tensor\x18\x02 \x01(\x0b\x32\x14.node_service.Tensor\x12\x17\n\nrequest_id\x18\x03 \x01(\tH\x00\x88\x01\x01\x42\r\n\x0b_request_id\"/\n\x19GetInferenceResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\"\\\n\x0fInferenceResult\x12)\n\x06tensor\x18\x01 \x01(\x0b\x32\x14.node_service.TensorH\x00\x88\x01\x01\x12\x13\n\x0bis_finished\x18\x02 \x01(\x08\x42\t\n\x07_tensor\";\n\x06Tensor\x12\x13\n\x0btensor_data\x18\x01 \x01(\x0c\x12\r\n\x05shape\x18\x02 \x03(\x05\x12\r\n\x05\x64type\x18\x03 \x01(\t\"7\n\x11ResetShardRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\"+\n\x16\x43ollectTopologyRequest\x12\x11\n\tmax_depth\x18\x01 \x01(\x05\"\x8e\x02\n\x08Topology\x12\x30\n\x05nodes\x18\x01 \x03(\x0b\x32!.node_service.Topology.NodesEntry\x12\x39\n\npeer_graph\x18\x02 \x03(\x0b\x32%.node_service.Topology.PeerGraphEntry\x1aN\n\nNodesEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12/\n\x05value\x18\x02 \x01(\x0b\x32 .node_service.DeviceCapabilities:\x02\x38\x01\x1a\x45\n\x0ePeerGraphEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\"\n\x05value\x18\x02 \x01(\x0b\x32\x13.node_service.Peers:\x02\x38\x01\"\x19\n\x05Peers\x12\x10\n\x08peer_ids\x18\x01 \x03(\t\"A\n\x12\x44\x65viceCapabilities\x12\r\n\x05model\x18\x01 \x01(\t\x12\x0c\n\x04\x63hip\x18\x02 \x01(\t\x12\x0e\n\x06memory\x18\x03 \x01(\x05\"\x07\n\x05\x45mpty2\x8c\x03\n\x0bNodeService\x12\x41\n\nSendPrompt\x12\x1b.node_service.PromptRequest\x1a\x14.node_service.Tensor\"\x00\x12\x41\n\nSendTensor\x12\x1b.node_service.TensorRequest\x1a\x14.node_service.Tensor\"\x00\x12\x44\n\nResetShard\x12\x1f.node_service.ResetShardRequest\x1a\x13.node_service.Empty\"\x00\x12^\n\x12GetInferenceResult\x12\'.node_service.GetInferenceResultRequest\x1a\x1d.node_service.InferenceResult\"\x00\x12Q\n\x0f\x43ollectTopology\x12$.node_service.CollectTopologyRequest\x1a\x16.node_service.Topology\"\x00\x62\x06proto3')
+DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x12node_service.proto\x12\x0cnode_service\"S\n\x05Shard\x12\x10\n\x08model_id\x18\x01 \x01(\t\x12\x13\n\x0bstart_layer\x18\x02 \x01(\x05\x12\x11\n\tend_layer\x18\x03 \x01(\x05\x12\x10\n\x08n_layers\x18\x04 \x01(\x05\"k\n\rPromptRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12\x0e\n\x06prompt\x18\x02 \x01(\t\x12\x17\n\nrequest_id\x18\x03 \x01(\tH\x00\x88\x01\x01\x42\r\n\x0b_request_id\"\x81\x01\n\rTensorRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12$\n\x06tensor\x18\x02 \x01(\x0b\x32\x14.node_service.Tensor\x12\x17\n\nrequest_id\x18\x03 \x01(\tH\x00\x88\x01\x01\x42\r\n\x0b_request_id\"/\n\x19GetInferenceResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\"\\\n\x0fInferenceResult\x12)\n\x06tensor\x18\x01 \x01(\x0b\x32\x14.node_service.TensorH\x00\x88\x01\x01\x12\x13\n\x0bis_finished\x18\x02 \x01(\x08\x42\t\n\x07_tensor\";\n\x06Tensor\x12\x13\n\x0btensor_data\x18\x01 \x01(\x0c\x12\r\n\x05shape\x18\x02 \x03(\x05\x12\r\n\x05\x64type\x18\x03 \x01(\t\"7\n\x11ResetShardRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\"<\n\x16\x43ollectTopologyRequest\x12\x0f\n\x07visited\x18\x01 \x03(\t\x12\x11\n\tmax_depth\x18\x02 \x01(\x05\"\x8e\x02\n\x08Topology\x12\x30\n\x05nodes\x18\x01 \x03(\x0b\x32!.node_service.Topology.NodesEntry\x12\x39\n\npeer_graph\x18\x02 \x03(\x0b\x32%.node_service.Topology.PeerGraphEntry\x1aN\n\nNodesEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12/\n\x05value\x18\x02 \x01(\x0b\x32 .node_service.DeviceCapabilities:\x02\x38\x01\x1a\x45\n\x0ePeerGraphEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\"\n\x05value\x18\x02 \x01(\x0b\x32\x13.node_service.Peers:\x02\x38\x01\"\x19\n\x05Peers\x12\x10\n\x08peer_ids\x18\x01 \x03(\t\"A\n\x12\x44\x65viceCapabilities\x12\r\n\x05model\x18\x01 \x01(\t\x12\x0c\n\x04\x63hip\x18\x02 \x01(\t\x12\x0e\n\x06memory\x18\x03 \x01(\x05\"\x07\n\x05\x45mpty2\x8c\x03\n\x0bNodeService\x12\x41\n\nSendPrompt\x12\x1b.node_service.PromptRequest\x1a\x14.node_service.Tensor\"\x00\x12\x41\n\nSendTensor\x12\x1b.node_service.TensorRequest\x1a\x14.node_service.Tensor\"\x00\x12\x44\n\nResetShard\x12\x1f.node_service.ResetShardRequest\x1a\x13.node_service.Empty\"\x00\x12^\n\x12GetInferenceResult\x12\'.node_service.GetInferenceResultRequest\x1a\x1d.node_service.InferenceResult\"\x00\x12Q\n\x0f\x43ollectTopology\x12$.node_service.CollectTopologyRequest\x1a\x16.node_service.Topology\"\x00\x62\x06proto3')
 
 _globals = globals()
 _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
@@ -40,19 +40,19 @@ if not _descriptor._USE_C_DESCRIPTORS:
   _globals['_RESETSHARDREQUEST']._serialized_start=566
   _globals['_RESETSHARDREQUEST']._serialized_end=621
   _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=623
-  _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=666
-  _globals['_TOPOLOGY']._serialized_start=669
-  _globals['_TOPOLOGY']._serialized_end=939
-  _globals['_TOPOLOGY_NODESENTRY']._serialized_start=790
-  _globals['_TOPOLOGY_NODESENTRY']._serialized_end=868
-  _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=870
-  _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=939
-  _globals['_PEERS']._serialized_start=941
-  _globals['_PEERS']._serialized_end=966
-  _globals['_DEVICECAPABILITIES']._serialized_start=968
-  _globals['_DEVICECAPABILITIES']._serialized_end=1033
-  _globals['_EMPTY']._serialized_start=1035
-  _globals['_EMPTY']._serialized_end=1042
-  _globals['_NODESERVICE']._serialized_start=1045
-  _globals['_NODESERVICE']._serialized_end=1441
+  _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=683
+  _globals['_TOPOLOGY']._serialized_start=686
+  _globals['_TOPOLOGY']._serialized_end=956
+  _globals['_TOPOLOGY_NODESENTRY']._serialized_start=807
+  _globals['_TOPOLOGY_NODESENTRY']._serialized_end=885
+  _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=887
+  _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=956
+  _globals['_PEERS']._serialized_start=958
+  _globals['_PEERS']._serialized_end=983
+  _globals['_DEVICECAPABILITIES']._serialized_start=985
+  _globals['_DEVICECAPABILITIES']._serialized_end=1050
+  _globals['_EMPTY']._serialized_start=1052
+  _globals['_EMPTY']._serialized_end=1059
+  _globals['_NODESERVICE']._serialized_start=1062
+  _globals['_NODESERVICE']._serialized_end=1458
 # @@protoc_insertion_point(module_scope)
diff --git a/exo/networking/peer_handle.py b/exo/networking/peer_handle.py
index c3c2c6c5..e9182d56 100644
--- a/exo/networking/peer_handle.py
+++ b/exo/networking/peer_handle.py
@@ -42,5 +42,5 @@ class PeerHandle(ABC):
     async def reset_shard(self, shard: Shard) -> None:
         pass
 
-    async def collect_topology(self, max_depth: int) -> Topology:
+    async def collect_topology(self, visited: set[str], max_depth: int) -> Topology:
         pass
diff --git a/exo/orchestration/node.py b/exo/orchestration/node.py
index 98415923..5ba5c8b2 100644
--- a/exo/orchestration/node.py
+++ b/exo/orchestration/node.py
@@ -26,9 +26,13 @@ class Node(ABC):
         pass
 
     @abstractmethod
-    async def collect_topology(self, max_depth: int = 2) -> Topology:
+    async def collect_topology(self, visited: set[str] = set(), max_depth: int = 2) -> Topology:
         pass
 
     @abstractmethod
     async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
         pass
+
+    @abstractmethod
+    async def global_reset(self, visited: set[str], max_depth: int = 2) -> None:
+        pass
diff --git a/exo/orchestration/standard_node.py b/exo/orchestration/standard_node.py
index 029f00c4..c38ed7f4 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -147,20 +147,29 @@ class StandardNode(Node):
                 await peer.connect()
                 if DEBUG >= 2: print(f"Connected to peer {peer.id()}")
 
-    async def collect_topology(self, max_depth: int = 4) -> Topology:
+    async def collect_topology(self, visited: set[str] = set(), max_depth: int = 4) -> Topology:
         self.topology.update_node(self.id, self.device_capabilities)
 
+        if DEBUG >= 2: print(f"Collecting topoloy {max_depth=} {visited=}")
         for peer in self.peers:
             self.topology.update_node(peer.id(), peer.device_capabilities())
             self.topology.add_edge(self.id, peer.id())
 
-            if max_depth > 0:
-                try:
-                    other_topology = await peer.collect_topology(max_depth = max_depth - 1)
-                    if DEBUG >= 2: print(f"Collected topology from: {peer.id()}: {other_topology}")
-                    self.topology.merge(other_topology)
-                except Exception as e:
-                    print(f"Error collecting topology from {peer.id()}: {e}")
+            if peer.id() in visited:
+                if DEBUG >= 2: print(f"Already visited {peer.id()}. Skipping...")
+                continue
+            visited.add(peer.id())
+
+            if max_depth <= 0:
+                if DEBUG >= 2: print(f"Max depth reached. Skipping...")
+                continue
+
+            try:
+                other_topology = await peer.collect_topology(visited, max_depth = max_depth - 1)
+                if DEBUG >= 2: print(f"Collected topology from: {peer.id()}: {other_topology}")
+                self.topology.merge(other_topology)
+            except Exception as e:
+                print(f"Error collecting topology from {peer.id()}: {e}")
 
         return self.topology
 
@@ -180,3 +189,6 @@ class StandardNode(Node):
         if request_id not in self.buffered_token_output:
             return None, False
         return np.array(self.buffered_token_output[request_id][0]), self.buffered_token_output[request_id][1]
+
+    async def global_reset(self, max_depth: int = 2) -> None:
+        pass
\ No newline at end of file

← 4502da5b readme bug notice  ·  back to Exo  ·  known issues section in readme 199eeb03 →