← back to Exo
one token at a time
f55a53ae7e18956ce11b36fe1042024cd42d6865 · 2024-12-14 21:06:41 +0000 · Alex Cheema
Files touched
M exo/api/chatgpt_api.pyM exo/main.pyM exo/networking/grpc/grpc_peer_handle.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.py
Diff
commit f55a53ae7e18956ce11b36fe1042024cd42d6865
Author: Alex Cheema <alexcheema123@gmail.com>
Date: Sat Dec 14 21:06:41 2024 +0000
one token at a time
---
exo/api/chatgpt_api.py | 27 +++++-------
exo/main.py | 2 +-
exo/networking/grpc/grpc_peer_handle.py | 16 ++-----
exo/networking/grpc/node_service.proto | 16 ++-----
exo/networking/grpc/node_service_pb2.py | 66 +++++++++++++---------------
exo/networking/grpc/node_service_pb2_grpc.py | 63 +++++---------------------
exo/networking/peer_handle.py | 6 +--
exo/orchestration/node.py | 32 ++++++--------
8 files changed, 73 insertions(+), 155 deletions(-)
diff --git a/exo/api/chatgpt_api.py b/exo/api/chatgpt_api.py
index 7aa8e595..1775c4fe 100644
--- a/exo/api/chatgpt_api.py
+++ b/exo/api/chatgpt_api.py
@@ -367,15 +367,11 @@ class ChatGPTAPI:
)
await response.prepare(request)
- async def stream_result(_request_id: str, tokens: List[int], is_finished: bool):
- prev_last_tokens_len = self.prev_token_lens.get(_request_id, 0)
- self.prev_token_lens[_request_id] = max(prev_last_tokens_len, len(tokens))
- new_tokens = tokens[prev_last_tokens_len:]
+ async def stream_result(_request_id: str, token: int, is_finished: bool):
finish_reason = None
eos_token_id = tokenizer.special_tokens_map.get("eos_token_id") if hasattr(tokenizer, "_tokenizer") and isinstance(tokenizer._tokenizer,
AutoTokenizer) else getattr(tokenizer, "eos_token_id", None)
- if len(new_tokens) > 0 and new_tokens[-1] == eos_token_id:
- new_tokens = new_tokens[:-1]
+ if token == eos_token_id:
if is_finished:
finish_reason = "stop"
if is_finished and not finish_reason:
@@ -386,7 +382,7 @@ class ChatGPTAPI:
tokenizer,
prompt,
request_id,
- new_tokens,
+ [token],
stream,
finish_reason,
"chat.completion",
@@ -398,12 +394,12 @@ class ChatGPTAPI:
if DEBUG >= 2: print(f"Error streaming completion: {e}")
if DEBUG >= 2: traceback.print_exc()
- def on_result(_request_id: str, tokens: List[int], is_finished: bool):
- if _request_id == request_id: self.stream_tasks[_request_id] = asyncio.create_task(stream_result(_request_id, tokens, is_finished))
+ def on_result(_request_id: str, token: int, is_finished: bool):
+ if _request_id == request_id: self.stream_tasks[_request_id] = asyncio.create_task(stream_result(_request_id, token, is_finished))
return _request_id == request_id and is_finished
- _, tokens, _ = await callback.wait(on_result, timeout=self.response_timeout)
+ _, token, _ = await callback.wait(on_result, timeout=self.response_timeout)
if request_id in self.stream_tasks: # in case there is still a stream task running, wait for it to complete
if DEBUG >= 2: print("Pending stream task. Waiting for stream task to complete.")
try:
@@ -413,19 +409,18 @@ class ChatGPTAPI:
await response.write_eof()
return response
else:
- _, tokens, _ = await callback.wait(
- lambda _request_id, tokens, is_finished: _request_id == request_id and is_finished,
+ _, token, _ = await callback.wait(
+ lambda _request_id, token, is_finished: _request_id == request_id and is_finished,
timeout=self.response_timeout,
)
finish_reason = "length"
eos_token_id = tokenizer.special_tokens_map.get("eos_token_id") if isinstance(getattr(tokenizer, "_tokenizer", None), AutoTokenizer) else tokenizer.eos_token_id
- if DEBUG >= 2: print(f"Checking if end of tokens result {tokens[-1]=} is {eos_token_id=}")
- if tokens[-1] == eos_token_id:
- tokens = tokens[:-1]
+ if DEBUG >= 2: print(f"Checking if end of tokens result {token=} is {eos_token_id=}")
+ if token == eos_token_id:
finish_reason = "stop"
- return web.json_response(generate_completion(chat_request, tokenizer, prompt, request_id, tokens, stream, finish_reason, "chat.completion"))
+ return web.json_response(generate_completion(chat_request, tokenizer, prompt, request_id, [token], stream, finish_reason, "chat.completion"))
except asyncio.TimeoutError:
return web.json_response({"detail": "Response generation timed out"}, status=408)
except Exception as e:
diff --git a/exo/main.py b/exo/main.py
index 3c4fe02e..3c02872b 100644
--- a/exo/main.py
+++ b/exo/main.py
@@ -152,7 +152,7 @@ api = ChatGPTAPI(
default_model=args.default_model
)
node.on_token.register("update_topology_viz").on_next(
- lambda req_id, tokens, __: topology_viz.update_prompt_output(req_id, inference_engine.tokenizer.decode(tokens)) if topology_viz and hasattr(inference_engine, "tokenizer") else None
+ lambda req_id, token, __: topology_viz.update_prompt_output(req_id, inference_engine.tokenizer.decode([token])) if topology_viz and hasattr(inference_engine, "tokenizer") else None
)
def preemptively_start_download(request_id: str, opaque_status: str):
diff --git a/exo/networking/grpc/grpc_peer_handle.py b/exo/networking/grpc/grpc_peer_handle.py
index 35b6968a..350036ce 100644
--- a/exo/networking/grpc/grpc_peer_handle.py
+++ b/exo/networking/grpc/grpc_peer_handle.py
@@ -147,16 +147,6 @@ class GRPCPeerHandle(PeerHandle):
return np.frombuffer(response.tensor_data, dtype=np.dtype(response.dtype)).reshape(response.shape)
- async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
- request = node_service_pb2.GetInferenceResultRequest(request_id=request_id)
- response = await self.stub.GetInferenceResult(request)
- if response.tensor is None:
- return None, response.is_finished
- return (
- np.frombuffer(response.tensor.tensor_data, dtype=np.dtype(response.tensor.dtype)).reshape(response.tensor.shape),
- response.is_finished,
- )
-
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)
@@ -174,9 +164,9 @@ class GRPCPeerHandle(PeerHandle):
topology.add_edge(node_id, conn.to_id, conn.description)
return topology
- async def send_result(self, request_id: str, result: List[int], is_finished: bool) -> None:
- request = node_service_pb2.SendResultRequest(request_id=request_id, result=result, is_finished=is_finished)
- await self.stub.SendResult(request)
+ async def send_new_token(self, request_id: str, token: int, is_finished: bool) -> None:
+ request = node_service_pb2.SendNewTokenRequest(request_id=request_id, token=token, is_finished=is_finished)
+ await self.stub.SendNewToken(request)
async def send_opaque_status(self, request_id: str, status: str) -> None:
request = node_service_pb2.SendOpaqueStatusRequest(request_id=request_id, status=status)
diff --git a/exo/networking/grpc/node_service.proto b/exo/networking/grpc/node_service.proto
index 83f1f03c..3f18f51b 100644
--- a/exo/networking/grpc/node_service.proto
+++ b/exo/networking/grpc/node_service.proto
@@ -6,9 +6,8 @@ service NodeService {
rpc SendPrompt (PromptRequest) returns (Tensor) {}
rpc SendTensor (TensorRequest) returns (Tensor) {}
rpc SendExample (ExampleRequest) returns (Loss) {}
- rpc GetInferenceResult (GetInferenceResultRequest) returns (InferenceResult) {}
rpc CollectTopology (CollectTopologyRequest) returns (Topology) {}
- rpc SendResult (SendResultRequest) returns (Empty) {}
+ rpc SendNewToken (SendNewTokenRequest) returns (Empty) {}
rpc SendOpaqueStatus (SendOpaqueStatusRequest) returns (Empty) {}
rpc HealthCheck (HealthCheckRequest) returns (HealthCheckResponse) {}
}
@@ -45,15 +44,6 @@ message Loss {
float loss = 1;
optional Tensor grads = 2;
}
-
-message GetInferenceResultRequest {
- string request_id = 1;
-}
-
-message InferenceResult {
- optional Tensor tensor = 1;
- bool is_finished = 2;
-}
message Tensor {
bytes tensor_data = 1;
@@ -93,9 +83,9 @@ message DeviceCapabilities {
DeviceFlops flops = 4;
}
-message SendResultRequest {
+message SendNewTokenRequest {
string request_id = 1;
- repeated int32 result = 2;
+ int32 token = 2;
bool is_finished = 3;
}
diff --git a/exo/networking/grpc/node_service_pb2.py b/exo/networking/grpc/node_service_pb2.py
index 3e34e365..64c14517 100644
--- a/exo/networking/grpc/node_service_pb2.py
+++ b/exo/networking/grpc/node_service_pb2.py
@@ -24,7 +24,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\"\xde\x01\n\x0e\x45xampleRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12%\n\x07\x65xample\x18\x02 \x01(\x0b\x32\x14.node_service.Tensor\x12$\n\x06target\x18\x03 \x01(\x0b\x32\x14.node_service.Tensor\x12$\n\x06length\x18\x04 \x01(\x0b\x32\x14.node_service.Tensor\x12\r\n\x05train\x18\x05 \x01(\x08\x12\x17\n\nrequest_id\x18\x06 \x01(\tH\x00\x88\x01\x01\x42\r\n\x0b_request_id\"H\n\x04Loss\x12\x0c\n\x04loss\x18\x01 \x01(\x02\x12(\n\x05grads\x18\x02 \x01(\x0b\x32\x14.node_service.TensorH\x00\x88\x01\x01\x42\x08\n\x06_grads\"/\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\"<\n\x16\x43ollectTopologyRequest\x12\x0f\n\x07visited\x18\x01 \x03(\t\x12\x11\n\tmax_depth\x18\x02 \x01(\x05\"\x98\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\x1aO\n\x0ePeerGraphEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12,\n\x05value\x18\x02 \x01(\x0b\x32\x1d.node_service.PeerConnections:\x02\x38\x01\"I\n\x0ePeerConnection\x12\r\n\x05to_id\x18\x01 \x01(\t\x12\x18\n\x0b\x64\x65scription\x18\x02 \x01(\tH\x00\x88\x01\x01\x42\x0e\n\x0c_description\"D\n\x0fPeerConnections\x12\x31\n\x0b\x63onnections\x18\x01 \x03(\x0b\x32\x1c.node_service.PeerConnection\"7\n\x0b\x44\x65viceFlops\x12\x0c\n\x04\x66p32\x18\x01 \x01(\x01\x12\x0c\n\x04\x66p16\x18\x02 \x01(\x01\x12\x0c\n\x04int8\x18\x03 \x01(\x01\"k\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\x12(\n\x05\x66lops\x18\x04 \x01(\x0b\x32\x19.node_service.DeviceFlops\"L\n\x11SendResultRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06result\x18\x02 \x03(\x05\x12\x13\n\x0bis_finished\x18\x03 \x01(\x08\"=\n\x17SendOpaqueStatusRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t\"\x14\n\x12HealthCheckRequest\")\n\x13HealthCheckResponse\x12\x12\n\nis_healthy\x18\x01 \x01(\x08\"\x07\n\x05\x45mpty2\xf7\x04\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\x41\n\x0bSendExample\x12\x1c.node_service.ExampleRequest\x1a\x12.node_service.Loss\"\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\x44\n\nSendResult\x12\x1f.node_service.SendResultRequest\x1a\x13.node_service.Empty\"\x00\x12P\n\x10SendOpaqueStatus\x12%.node_service.SendOpaqueStatusRequest\x1a\x13.node_service.Empty\"\x00\x12T\n\x0bHealthCheck\x12 .node_service.HealthCheckRequest\x1a!.node_service.HealthCheckResponse\"\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\"\xde\x01\n\x0e\x45xampleRequest\x12\"\n\x05shard\x18\x01 \x01(\x0b\x32\x13.node_service.Shard\x12%\n\x07\x65xample\x18\x02 \x01(\x0b\x32\x14.node_service.Tensor\x12$\n\x06target\x18\x03 \x01(\x0b\x32\x14.node_service.Tensor\x12$\n\x06length\x18\x04 \x01(\x0b\x32\x14.node_service.Tensor\x12\r\n\x05train\x18\x05 \x01(\x08\x12\x17\n\nrequest_id\x18\x06 \x01(\tH\x00\x88\x01\x01\x42\r\n\x0b_request_id\"H\n\x04Loss\x12\x0c\n\x04loss\x18\x01 \x01(\x02\x12(\n\x05grads\x18\x02 \x01(\x0b\x32\x14.node_service.TensorH\x00\x88\x01\x01\x42\x08\n\x06_grads\";\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\"<\n\x16\x43ollectTopologyRequest\x12\x0f\n\x07visited\x18\x01 \x03(\t\x12\x11\n\tmax_depth\x18\x02 \x01(\x05\"\x98\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\x1aO\n\x0ePeerGraphEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12,\n\x05value\x18\x02 \x01(\x0b\x32\x1d.node_service.PeerConnections:\x02\x38\x01\"I\n\x0ePeerConnection\x12\r\n\x05to_id\x18\x01 \x01(\t\x12\x18\n\x0b\x64\x65scription\x18\x02 \x01(\tH\x00\x88\x01\x01\x42\x0e\n\x0c_description\"D\n\x0fPeerConnections\x12\x31\n\x0b\x63onnections\x18\x01 \x03(\x0b\x32\x1c.node_service.PeerConnection\"7\n\x0b\x44\x65viceFlops\x12\x0c\n\x04\x66p32\x18\x01 \x01(\x01\x12\x0c\n\x04\x66p16\x18\x02 \x01(\x01\x12\x0c\n\x04int8\x18\x03 \x01(\x01\"k\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\x12(\n\x05\x66lops\x18\x04 \x01(\x0b\x32\x19.node_service.DeviceFlops\"M\n\x13SendNewTokenRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\r\n\x05token\x18\x02 \x01(\x05\x12\x13\n\x0bis_finished\x18\x03 \x01(\x08\"=\n\x17SendOpaqueStatusRequest\x12\x12\n\nrequest_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t\"\x14\n\x12HealthCheckRequest\")\n\x13HealthCheckResponse\x12\x12\n\nis_healthy\x18\x01 \x01(\x08\"\x07\n\x05\x45mpty2\x9b\x04\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\x41\n\x0bSendExample\x12\x1c.node_service.ExampleRequest\x1a\x12.node_service.Loss\"\x00\x12Q\n\x0f\x43ollectTopology\x12$.node_service.CollectTopologyRequest\x1a\x16.node_service.Topology\"\x00\x12H\n\x0cSendNewToken\x12!.node_service.SendNewTokenRequest\x1a\x13.node_service.Empty\"\x00\x12P\n\x10SendOpaqueStatus\x12%.node_service.SendOpaqueStatusRequest\x1a\x13.node_service.Empty\"\x00\x12T\n\x0bHealthCheck\x12 .node_service.HealthCheckRequest\x1a!.node_service.HealthCheckResponse\"\x00\x62\x06proto3')
_globals = globals()
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
@@ -45,38 +45,34 @@ if not _descriptor._USE_C_DESCRIPTORS:
_globals['_EXAMPLEREQUEST']._serialized_end=585
_globals['_LOSS']._serialized_start=587
_globals['_LOSS']._serialized_end=659
- _globals['_GETINFERENCERESULTREQUEST']._serialized_start=661
- _globals['_GETINFERENCERESULTREQUEST']._serialized_end=708
- _globals['_INFERENCERESULT']._serialized_start=710
- _globals['_INFERENCERESULT']._serialized_end=802
- _globals['_TENSOR']._serialized_start=804
- _globals['_TENSOR']._serialized_end=863
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=865
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=925
- _globals['_TOPOLOGY']._serialized_start=928
- _globals['_TOPOLOGY']._serialized_end=1208
- _globals['_TOPOLOGY_NODESENTRY']._serialized_start=1049
- _globals['_TOPOLOGY_NODESENTRY']._serialized_end=1127
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=1129
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=1208
- _globals['_PEERCONNECTION']._serialized_start=1210
- _globals['_PEERCONNECTION']._serialized_end=1283
- _globals['_PEERCONNECTIONS']._serialized_start=1285
- _globals['_PEERCONNECTIONS']._serialized_end=1353
- _globals['_DEVICEFLOPS']._serialized_start=1355
- _globals['_DEVICEFLOPS']._serialized_end=1410
- _globals['_DEVICECAPABILITIES']._serialized_start=1412
- _globals['_DEVICECAPABILITIES']._serialized_end=1519
- _globals['_SENDRESULTREQUEST']._serialized_start=1521
- _globals['_SENDRESULTREQUEST']._serialized_end=1597
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_start=1599
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_end=1660
- _globals['_HEALTHCHECKREQUEST']._serialized_start=1662
- _globals['_HEALTHCHECKREQUEST']._serialized_end=1682
- _globals['_HEALTHCHECKRESPONSE']._serialized_start=1684
- _globals['_HEALTHCHECKRESPONSE']._serialized_end=1725
- _globals['_EMPTY']._serialized_start=1727
- _globals['_EMPTY']._serialized_end=1734
- _globals['_NODESERVICE']._serialized_start=1737
- _globals['_NODESERVICE']._serialized_end=2368
+ _globals['_TENSOR']._serialized_start=661
+ _globals['_TENSOR']._serialized_end=720
+ _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=722
+ _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=782
+ _globals['_TOPOLOGY']._serialized_start=785
+ _globals['_TOPOLOGY']._serialized_end=1065
+ _globals['_TOPOLOGY_NODESENTRY']._serialized_start=906
+ _globals['_TOPOLOGY_NODESENTRY']._serialized_end=984
+ _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=986
+ _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=1065
+ _globals['_PEERCONNECTION']._serialized_start=1067
+ _globals['_PEERCONNECTION']._serialized_end=1140
+ _globals['_PEERCONNECTIONS']._serialized_start=1142
+ _globals['_PEERCONNECTIONS']._serialized_end=1210
+ _globals['_DEVICEFLOPS']._serialized_start=1212
+ _globals['_DEVICEFLOPS']._serialized_end=1267
+ _globals['_DEVICECAPABILITIES']._serialized_start=1269
+ _globals['_DEVICECAPABILITIES']._serialized_end=1376
+ _globals['_SENDNEWTOKENREQUEST']._serialized_start=1378
+ _globals['_SENDNEWTOKENREQUEST']._serialized_end=1455
+ _globals['_SENDOPAQUESTATUSREQUEST']._serialized_start=1457
+ _globals['_SENDOPAQUESTATUSREQUEST']._serialized_end=1518
+ _globals['_HEALTHCHECKREQUEST']._serialized_start=1520
+ _globals['_HEALTHCHECKREQUEST']._serialized_end=1540
+ _globals['_HEALTHCHECKRESPONSE']._serialized_start=1542
+ _globals['_HEALTHCHECKRESPONSE']._serialized_end=1583
+ _globals['_EMPTY']._serialized_start=1585
+ _globals['_EMPTY']._serialized_end=1592
+ _globals['_NODESERVICE']._serialized_start=1595
+ _globals['_NODESERVICE']._serialized_end=2134
# @@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 9cf336ce..a4d42387 100644
--- a/exo/networking/grpc/node_service_pb2_grpc.py
+++ b/exo/networking/grpc/node_service_pb2_grpc.py
@@ -49,19 +49,14 @@ class NodeServiceStub(object):
request_serializer=node__service__pb2.ExampleRequest.SerializeToString,
response_deserializer=node__service__pb2.Loss.FromString,
_registered_method=True)
- self.GetInferenceResult = channel.unary_unary(
- '/node_service.NodeService/GetInferenceResult',
- request_serializer=node__service__pb2.GetInferenceResultRequest.SerializeToString,
- response_deserializer=node__service__pb2.InferenceResult.FromString,
- _registered_method=True)
self.CollectTopology = channel.unary_unary(
'/node_service.NodeService/CollectTopology',
request_serializer=node__service__pb2.CollectTopologyRequest.SerializeToString,
response_deserializer=node__service__pb2.Topology.FromString,
_registered_method=True)
- self.SendResult = channel.unary_unary(
- '/node_service.NodeService/SendResult',
- request_serializer=node__service__pb2.SendResultRequest.SerializeToString,
+ self.SendNewToken = channel.unary_unary(
+ '/node_service.NodeService/SendNewToken',
+ request_serializer=node__service__pb2.SendNewTokenRequest.SerializeToString,
response_deserializer=node__service__pb2.Empty.FromString,
_registered_method=True)
self.SendOpaqueStatus = channel.unary_unary(
@@ -97,19 +92,13 @@ class NodeServiceServicer(object):
context.set_details('Method not implemented!')
raise NotImplementedError('Method not implemented!')
- def GetInferenceResult(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 CollectTopology(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 SendResult(self, request, context):
+ def SendNewToken(self, request, context):
"""Missing associated documentation comment in .proto file."""
context.set_code(grpc.StatusCode.UNIMPLEMENTED)
context.set_details('Method not implemented!')
@@ -145,19 +134,14 @@ def add_NodeServiceServicer_to_server(servicer, server):
request_deserializer=node__service__pb2.ExampleRequest.FromString,
response_serializer=node__service__pb2.Loss.SerializeToString,
),
- 'GetInferenceResult': grpc.unary_unary_rpc_method_handler(
- servicer.GetInferenceResult,
- request_deserializer=node__service__pb2.GetInferenceResultRequest.FromString,
- response_serializer=node__service__pb2.InferenceResult.SerializeToString,
- ),
'CollectTopology': grpc.unary_unary_rpc_method_handler(
servicer.CollectTopology,
request_deserializer=node__service__pb2.CollectTopologyRequest.FromString,
response_serializer=node__service__pb2.Topology.SerializeToString,
),
- 'SendResult': grpc.unary_unary_rpc_method_handler(
- servicer.SendResult,
- request_deserializer=node__service__pb2.SendResultRequest.FromString,
+ 'SendNewToken': grpc.unary_unary_rpc_method_handler(
+ servicer.SendNewToken,
+ request_deserializer=node__service__pb2.SendNewTokenRequest.FromString,
response_serializer=node__service__pb2.Empty.SerializeToString,
),
'SendOpaqueStatus': grpc.unary_unary_rpc_method_handler(
@@ -262,33 +246,6 @@ class NodeService(object):
metadata,
_registered_method=True)
- @staticmethod
- def GetInferenceResult(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/GetInferenceResult',
- node__service__pb2.GetInferenceResultRequest.SerializeToString,
- node__service__pb2.InferenceResult.FromString,
- options,
- channel_credentials,
- insecure,
- call_credentials,
- compression,
- wait_for_ready,
- timeout,
- metadata,
- _registered_method=True)
-
@staticmethod
def CollectTopology(request,
target,
@@ -317,7 +274,7 @@ class NodeService(object):
_registered_method=True)
@staticmethod
- def SendResult(request,
+ def SendNewToken(request,
target,
options=(),
channel_credentials=None,
@@ -330,8 +287,8 @@ class NodeService(object):
return grpc.experimental.unary_unary(
request,
target,
- '/node_service.NodeService/SendResult',
- node__service__pb2.SendResultRequest.SerializeToString,
+ '/node_service.NodeService/SendNewToken',
+ node__service__pb2.SendNewTokenRequest.SerializeToString,
node__service__pb2.Empty.FromString,
options,
channel_credentials,
diff --git a/exo/networking/peer_handle.py b/exo/networking/peer_handle.py
index 45d37c4b..14d0c24f 100644
--- a/exo/networking/peer_handle.py
+++ b/exo/networking/peer_handle.py
@@ -48,11 +48,7 @@ class PeerHandle(ABC):
pass
@abstractmethod
- async def send_result(self, request_id: str, result: List[int], is_finished: bool) -> None:
- pass
-
- @abstractmethod
- async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
+ async def send_new_token(self, request_id: str, token: int, is_finished: bool) -> None:
pass
@abstractmethod
diff --git a/exo/orchestration/node.py b/exo/orchestration/node.py
index 131b75e5..f54744b9 100644
--- a/exo/orchestration/node.py
+++ b/exo/orchestration/node.py
@@ -47,7 +47,7 @@ class Node:
self.max_generate_tokens = max_generate_tokens
self.topology_viz = topology_viz
self.default_sample_temperature = default_sample_temperature
- self._on_token = AsyncCallbackSystem[str, Tuple[str, List[int], bool]]()
+ self._on_token = AsyncCallbackSystem[str, Tuple[str, int, bool]]()
self._on_opaque_status = AsyncCallbackSystem[str, Tuple[str, str]]()
self._on_opaque_status.register("node_status").on_next(self.on_node_status)
self.node_download_progress: Dict[str, RepoProgressEvent] = {}
@@ -122,10 +122,9 @@ class Node:
self.buffered_token_output[request_id][0].append(token.item())
is_finished = token.item() == self.inference_engine.tokenizer.eos_token_id or is_finished or len(self.buffered_token_output[request_id][0]) >= self.max_generate_tokens
if DEBUG >= 2: print(f"[{request_id}] result size: {result.size}, is finished: {is_finished}, buffered tokens: {len(self.buffered_token_output[request_id][0])}")
- asyncio.create_task(self.broadcast_result(request_id, *self.buffered_token_output[request_id]))
forward = token.reshape(1, -1)
- self.trigger_on_token_callbacks(request_id, self.buffered_token_output[request_id][0], is_finished)
- asyncio.create_task(self.broadcast_result(request_id, self.buffered_token_output[request_id][0], is_finished))
+ self.trigger_on_token_callbacks(request_id, token.item(), is_finished)
+ asyncio.create_task(self.broadcast_new_token(request_id, token.item(), is_finished))
else:
forward = result
@@ -549,11 +548,6 @@ class Node:
print(f"Error collecting topology: {e}")
traceback.print_exc()
- 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], max_depth: int = 4) -> Topology:
next_topology = Topology()
next_topology.update_node(self.id, self.device_capabilities)
@@ -590,28 +584,28 @@ class Node:
return self.topology
@property
- def on_token(self) -> AsyncCallbackSystem[str, Tuple[str, List[int], bool]]:
+ def on_token(self) -> AsyncCallbackSystem[str, Tuple[str, int, bool]]:
return self._on_token
@property
def on_opaque_status(self) -> AsyncCallbackSystem[str, Tuple[str, str]]:
return self._on_opaque_status
- def trigger_on_token_callbacks(self, request_id: str, tokens: List[int], is_finished: bool) -> None:
- if DEBUG >= 2: print(f"Triggering all on_token callbacks with {request_id=} num_tokens={len(tokens)} {is_finished=}")
- self.on_token.trigger_all(request_id, tokens, is_finished)
+ def trigger_on_token_callbacks(self, request_id: str, token: int, is_finished: bool) -> None:
+ if DEBUG >= 2: print(f"Triggering all on_token callbacks with {request_id=} {token=} {is_finished=}")
+ self.on_token.trigger_all(request_id, token, is_finished)
- async def broadcast_result(self, request_id: str, result: List[int], is_finished: bool) -> None:
- async def send_result_to_peer(peer):
+ async def broadcast_new_token(self, request_id: str, token: int, is_finished: bool) -> None:
+ async def send_new_token_to_peer(peer):
try:
- await asyncio.wait_for(peer.send_result(request_id, result, is_finished), timeout=15.0)
+ await asyncio.wait_for(peer.send_new_token(request_id, token, is_finished), timeout=15.0)
except asyncio.TimeoutError:
- print(f"Timeout broadcasting result to {peer.id()}")
+ print(f"Timeout broadcasting new token to {peer.id()}")
except Exception as e:
- print(f"Error broadcasting result to {peer.id()}: {e}")
+ print(f"Error broadcasting new token to {peer.id()}: {e}")
traceback.print_exc()
- await asyncio.gather(*[send_result_to_peer(peer) for peer in self.peers], return_exceptions=True)
+ await asyncio.gather(*[send_new_token_to_peer(peer) for peer in self.peers], return_exceptions=True)
async def broadcast_opaque_status(self, request_id: str, status: str) -> None:
if DEBUG >= 8: print(f"Broadcasting opaque status: {request_id=} {status=}")
← 470f961f Only collect topology if peers changed
·
back to Exo
·
fix SendNewToken cb4615c9 →