← 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
M exo/networking/grpc/grpc_peer_handle.pyM exo/networking/grpc/grpc_server.pyM exo/networking/grpc/node_service.protoM exo/networking/grpc/node_service_pb2.pyM exo/networking/peer_handle.pyM exo/orchestration/node.pyM exo/orchestration/standard_node.py
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 →