← back to Exo
add global reset
e17905e2951ca5b99d7c9ccbd31964fa4edd8c58 · 2024-07-15 22:35:54 -0700 · Alex Cheema
Files touched
M examples/llama3_distributed.pyM exo/inference/tinygrad/inference.pyM exo/networking/grpc/grpc_discovery.pyM 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/grpc/node_service_pb2_grpc.pyM exo/networking/peer_handle.pyM exo/orchestration/node.pyM exo/orchestration/standard_node.py
Diff
commit e17905e2951ca5b99d7c9ccbd31964fa4edd8c58
Author: Alex Cheema <alexcheema123@gmail.com>
Date: Mon Jul 15 22:35:54 2024 -0700
add global reset
---
examples/llama3_distributed.py | 12 +++---
exo/inference/tinygrad/inference.py | 2 +-
exo/networking/grpc/grpc_discovery.py | 2 +
exo/networking/grpc/grpc_peer_handle.py | 4 ++
exo/networking/grpc/grpc_server.py | 7 ++++
exo/networking/grpc/node_service.proto | 7 ++++
exo/networking/grpc/node_service_pb2.py | 32 ++++++++-------
exo/networking/grpc/node_service_pb2_grpc.py | 43 ++++++++++++++++++++
exo/networking/peer_handle.py | 5 +++
exo/orchestration/node.py | 6 +--
exo/orchestration/standard_node.py | 60 +++++++++++++++++++---------
11 files changed, 138 insertions(+), 42 deletions(-)
diff --git a/examples/llama3_distributed.py b/examples/llama3_distributed.py
index 76551e1f..cb591926 100644
--- a/examples/llama3_distributed.py
+++ b/examples/llama3_distributed.py
@@ -17,19 +17,20 @@ models = {
"mlx-community/Meta-Llama-3-70B-Instruct-4bit": Shard(model_id="mlx-community/Meta-Llama-3-70B-Instruct-4bit", start_layer=0, end_layer=0, n_layers=80)
}
-path_or_hf_repo = "mlx-community/Meta-Llama-3-70B-Instruct-4bit"
+path_or_hf_repo = "mlx-community/Meta-Llama-3-8B-Instruct-4bit"
model_path = get_model_path(path_or_hf_repo)
tokenizer_config = {}
tokenizer = load_tokenizer(model_path, tokenizer_config)
-peer2 = GRPCPeerHandle(
+peer1 = GRPCPeerHandle(
"node1",
"localhost:8080",
DeviceCapabilities(model="placeholder", chip="placeholder", memory=0)
)
-peer1 = GRPCPeerHandle(
+peer2 = GRPCPeerHandle(
"node2",
- "10.0.0.161:8080",
+ # "10.0.0.161:8080",
+ "localhost:8081",
DeviceCapabilities(model="placeholder", chip="placeholder", memory=0)
)
shard = models[path_or_hf_repo]
@@ -49,7 +50,8 @@ async def run_prompt(prompt: str):
for peer in [peer1, peer2]:
await peer.connect()
- await peer.reset_shard(shard)
+
+ await peer.global_reset(shard, set(), 2)
try:
await peer1.send_prompt(shard, prompt, request_id)
diff --git a/exo/inference/tinygrad/inference.py b/exo/inference/tinygrad/inference.py
index ae9712e5..1c009a96 100644
--- a/exo/inference/tinygrad/inference.py
+++ b/exo/inference/tinygrad/inference.py
@@ -159,7 +159,7 @@ class TinygradDynamicShardInferenceEngine(InferenceEngine):
toks = [self.tokenizer.bos_id] + encode_message("user", prompt) + encode_role("assistant")
start_pos = prefill(self.model, toks[:-1])
last_tok = toks[-1]
-
+
output_data = np.array(self.model(Tensor([[last_tok]]), start_pos, TEMPERATURE, TOP_K, TOP_P, ALPHA_F, ALPHA_P).tolist())
start_pos += 1
diff --git a/exo/networking/grpc/grpc_discovery.py b/exo/networking/grpc/grpc_discovery.py
index dc3f72d8..46783cf5 100644
--- a/exo/networking/grpc/grpc_discovery.py
+++ b/exo/networking/grpc/grpc_discovery.py
@@ -106,6 +106,8 @@ class GRPCDiscovery(Discovery):
self.peer_last_seen[peer_id] = time.time()
except Exception as e:
print(f"Error in peer discovery: {e}")
+ import traceback
+ print(traceback.format_exc())
await asyncio.sleep(self.broadcast_interval / 2)
async def _cleanup_peers(self):
diff --git a/exo/networking/grpc/grpc_peer_handle.py b/exo/networking/grpc/grpc_peer_handle.py
index 1f3e7418..05a48601 100644
--- a/exo/networking/grpc/grpc_peer_handle.py
+++ b/exo/networking/grpc/grpc_peer_handle.py
@@ -88,3 +88,7 @@ class GRPCPeerHandle(PeerHandle):
for peer_id in peers.peer_ids:
topology.add_edge(node_id, peer_id)
return topology
+
+ async def global_reset(self, base_shard: Shard, visited: set[str], max_depth: int) -> None:
+ request = node_service_pb2.GlobalResetRequest(base_shard=node_service_pb2.Shard(model_id=base_shard.model_id, start_layer=base_shard.start_layer, end_layer=base_shard.end_layer, n_layers=base_shard.n_layers), visited=visited, max_depth=max_depth)
+ await self.stub.GlobalReset(request)
diff --git a/exo/networking/grpc/grpc_server.py b/exo/networking/grpc/grpc_server.py
index f6586e6d..63267839 100644
--- a/exo/networking/grpc/grpc_server.py
+++ b/exo/networking/grpc/grpc_server.py
@@ -70,3 +70,10 @@ class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):
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)
+
+ async def GlobalReset(self, request, context):
+ base_shard = Shard(model_id=request.base_shard.model_id, start_layer=request.base_shard.start_layer, end_layer=request.base_shard.end_layer, n_layers=request.base_shard.n_layers)
+ visited = set(request.visited)
+ max_depth = request.max_depth
+ await self.node.global_reset(base_shard, visited, max_depth)
+ return node_service_pb2.Empty()
diff --git a/exo/networking/grpc/node_service.proto b/exo/networking/grpc/node_service.proto
index 38cfd30b..80295de2 100644
--- a/exo/networking/grpc/node_service.proto
+++ b/exo/networking/grpc/node_service.proto
@@ -8,6 +8,7 @@ service NodeService {
rpc ResetShard (ResetShardRequest) returns (Empty) {}
rpc GetInferenceResult (GetInferenceResultRequest) returns (InferenceResult) {}
rpc CollectTopology (CollectTopologyRequest) returns (Topology) {}
+ rpc GlobalReset (GlobalResetRequest) returns (Empty) {}
}
message Shard {
@@ -53,6 +54,12 @@ message CollectTopologyRequest {
int32 max_depth = 2;
}
+message GlobalResetRequest {
+ Shard base_shard = 1;
+ repeated string visited = 2;
+ int32 max_depth = 3;
+}
+
message Topology {
map<string, DeviceCapabilities> nodes = 1;
map<string, Peers> peer_graph = 2;
diff --git a/exo/networking/grpc/node_service_pb2.py b/exo/networking/grpc/node_service_pb2.py
index a7f51cdc..bad9a291 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\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')
+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\"a\n\x12GlobalResetRequest\x12\'\n\nbase_shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12\x0f\n\x07visited\x18\x02 \x03(\t\x12\x11\n\tmax_depth\x18\x03 \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\xd4\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\x12\x46\n\x0bGlobalReset\x12 .node_service.GlobalResetRequest\x1a\x13.node_service.Empty\"\x00\x62\x06proto3')
_globals = globals()
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
@@ -41,18 +41,20 @@ if not _descriptor._USE_C_DESCRIPTORS:
_globals['_RESETSHARDREQUEST']._serialized_end=621
_globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=623
_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
+ _globals['_GLOBALRESETREQUEST']._serialized_start=685
+ _globals['_GLOBALRESETREQUEST']._serialized_end=782
+ _globals['_TOPOLOGY']._serialized_start=785
+ _globals['_TOPOLOGY']._serialized_end=1055
+ _globals['_TOPOLOGY_NODESENTRY']._serialized_start=906
+ _globals['_TOPOLOGY_NODESENTRY']._serialized_end=984
+ _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=986
+ _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=1055
+ _globals['_PEERS']._serialized_start=1057
+ _globals['_PEERS']._serialized_end=1082
+ _globals['_DEVICECAPABILITIES']._serialized_start=1084
+ _globals['_DEVICECAPABILITIES']._serialized_end=1149
+ _globals['_EMPTY']._serialized_start=1151
+ _globals['_EMPTY']._serialized_end=1158
+ _globals['_NODESERVICE']._serialized_start=1161
+ _globals['_NODESERVICE']._serialized_end=1629
# @@protoc_insertion_point(module_scope)
diff --git a/exo/networking/grpc/node_service_pb2_grpc.py b/exo/networking/grpc/node_service_pb2_grpc.py
index 259e9938..c62dc6f0 100644
--- a/exo/networking/grpc/node_service_pb2_grpc.py
+++ b/exo/networking/grpc/node_service_pb2_grpc.py
@@ -64,6 +64,11 @@ class NodeServiceStub(object):
request_serializer=node__service__pb2.CollectTopologyRequest.SerializeToString,
response_deserializer=node__service__pb2.Topology.FromString,
_registered_method=True)
+ self.GlobalReset = channel.unary_unary(
+ '/node_service.NodeService/GlobalReset',
+ request_serializer=node__service__pb2.GlobalResetRequest.SerializeToString,
+ response_deserializer=node__service__pb2.Empty.FromString,
+ _registered_method=True)
class NodeServiceServicer(object):
@@ -99,6 +104,12 @@ class NodeServiceServicer(object):
context.set_details('Method not implemented!')
raise NotImplementedError('Method not implemented!')
+ def GlobalReset(self, request, context):
+ """Missing associated documentation comment in .proto file."""
+ context.set_code(grpc.StatusCode.UNIMPLEMENTED)
+ context.set_details('Method not implemented!')
+ raise NotImplementedError('Method not implemented!')
+
def add_NodeServiceServicer_to_server(servicer, server):
rpc_method_handlers = {
@@ -127,6 +138,11 @@ def add_NodeServiceServicer_to_server(servicer, server):
request_deserializer=node__service__pb2.CollectTopologyRequest.FromString,
response_serializer=node__service__pb2.Topology.SerializeToString,
),
+ 'GlobalReset': grpc.unary_unary_rpc_method_handler(
+ servicer.GlobalReset,
+ request_deserializer=node__service__pb2.GlobalResetRequest.FromString,
+ response_serializer=node__service__pb2.Empty.SerializeToString,
+ ),
}
generic_handler = grpc.method_handlers_generic_handler(
'node_service.NodeService', rpc_method_handlers)
@@ -272,3 +288,30 @@ class NodeService(object):
timeout,
metadata,
_registered_method=True)
+
+ @staticmethod
+ def GlobalReset(request,
+ target,
+ options=(),
+ channel_credentials=None,
+ call_credentials=None,
+ insecure=False,
+ compression=None,
+ wait_for_ready=None,
+ timeout=None,
+ metadata=None):
+ return grpc.experimental.unary_unary(
+ request,
+ target,
+ '/node_service.NodeService/GlobalReset',
+ node__service__pb2.GlobalResetRequest.SerializeToString,
+ node__service__pb2.Empty.FromString,
+ options,
+ channel_credentials,
+ insecure,
+ call_credentials,
+ compression,
+ wait_for_ready,
+ timeout,
+ metadata,
+ _registered_method=True)
diff --git a/exo/networking/peer_handle.py b/exo/networking/peer_handle.py
index e9182d56..75234578 100644
--- a/exo/networking/peer_handle.py
+++ b/exo/networking/peer_handle.py
@@ -42,5 +42,10 @@ class PeerHandle(ABC):
async def reset_shard(self, shard: Shard) -> None:
pass
+ @abstractmethod
async def collect_topology(self, visited: set[str], max_depth: int) -> Topology:
pass
+
+ @abstractmethod
+ async def global_reset(self, base_shard: Shard, visited: set[str], max_depth: int) -> None:
+ pass
diff --git a/exo/orchestration/node.py b/exo/orchestration/node.py
index 5ba5c8b2..3e02340b 100644
--- a/exo/orchestration/node.py
+++ b/exo/orchestration/node.py
@@ -26,13 +26,13 @@ class Node(ABC):
pass
@abstractmethod
- async def collect_topology(self, visited: set[str] = set(), max_depth: int = 2) -> Topology:
+ async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
pass
@abstractmethod
- async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
+ async def collect_topology(self, visited: set[str] = set(), max_depth: int = 2) -> Topology:
pass
@abstractmethod
- async def global_reset(self, visited: set[str], max_depth: int = 2) -> None:
+ async def global_reset(self, base_shard: Shard, visited: set[str] = set(), max_depth: int = 2) -> None:
pass
diff --git a/exo/orchestration/standard_node.py b/exo/orchestration/standard_node.py
index c38ed7f4..04b9a4ea 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -147,18 +147,38 @@ class StandardNode(Node):
await peer.connect()
if DEBUG >= 2: print(f"Connected to peer {peer.id()}")
+ async def periodic_topology_collection(self, interval: int):
+ while True:
+ await asyncio.sleep(interval)
+ try:
+ await self.update_peers()
+ await self.collect_topology()
+ except Exception as e:
+ print(f"Error collecting topology: {e}")
+
+ if DEBUG >= 2: print("Topology collection task executed.")
+ if DEBUG >= 2: print(f"Current topology: {self.topology}")
+
+ async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
+ 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 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=}")
+
+ prev_visited = visited.copy()
+ visited.update(p.id() for p in self.peers)
+
for peer in self.peers:
self.topology.update_node(peer.id(), peer.device_capabilities())
self.topology.add_edge(self.id, peer.id())
- if peer.id() in visited:
+ if peer.id() in prev_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...")
@@ -173,22 +193,26 @@ class StandardNode(Node):
return self.topology
- async def periodic_topology_collection(self, interval: int):
- while True:
- await asyncio.sleep(interval)
- try:
- await self.update_peers()
- await self.collect_topology()
- except Exception as e:
- print(f"Error collecting topology: {e}")
+ # TODO: unify this and collect_topology as global actions
+ async def global_reset(self, base_shard: Shard, visited: set[str] = set(), max_depth: int = 2) -> None:
+ await self.reset_shard(self.get_current_shard(base_shard))
- if DEBUG >= 2: print("Topology collection task executed.")
- if DEBUG >= 2: print(f"Current topology: {self.topology}")
+ if DEBUG >= 2: print(f"Global reset {base_shard=} {max_depth=} {visited=}")
- async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
- 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]
+ prev_visited = visited.copy()
+ visited.update(p.id() for p in self.peers)
- async def global_reset(self, max_depth: int = 2) -> None:
- pass
\ No newline at end of file
+ for peer in self.peers:
+ if peer.id() in prev_visited:
+ if DEBUG >= 2: print(f"Already visited {peer.id()}. Skipping...")
+ continue
+
+ if max_depth <= 0:
+ if DEBUG >= 2: print(f"Max depth reached. Skipping...")
+ continue
+
+ try:
+ print(f"Forwarding global reset to peer {peer.id()}")
+ await peer.global_reset(base_shard, visited, max_depth = max_depth - 1)
+ except Exception as e:
+ print(f"Error collecting topology from {peer.id()}: {e}")
← de9b89ea readme typo
·
back to Exo
·
Readme whitespace 108a904a →