← back to Exo
per-request kv cache, remove all explicit reset functionality as it wasnt used. fixes #67
20847844709b24bdf16fe38de63f7319cd81edaa · 2024-07-25 17:09:34 -0700 · Alex Cheema
Files touched
M examples/llama3_distributed.pyM exo/api/chatgpt_api.pyM exo/inference/debug_inference_engine.pyM exo/inference/inference_engine.pyM exo/inference/mlx/sharded_inference_engine.pyM exo/inference/mlx/sharded_model.pyM exo/inference/test_inference_engine.pyM exo/inference/tinygrad/inference.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 20847844709b24bdf16fe38de63f7319cd81edaa
Author: Alex Cheema <alexcheema123@gmail.com>
Date: Thu Jul 25 17:09:34 2024 -0700
per-request kv cache, remove all explicit reset functionality as it wasnt used. fixes #67
---
examples/llama3_distributed.py | 1 -
exo/api/chatgpt_api.py | 4 +-
exo/inference/debug_inference_engine.py | 17 ++----
exo/inference/inference_engine.py | 8 +--
exo/inference/mlx/sharded_inference_engine.py | 12 ++--
exo/inference/mlx/sharded_model.py | 11 ++--
exo/inference/test_inference_engine.py | 17 ++----
exo/inference/tinygrad/inference.py | 10 +---
exo/networking/grpc/grpc_peer_handle.py | 8 ---
exo/networking/grpc/grpc_server.py | 14 -----
exo/networking/grpc/node_service.proto | 12 ----
exo/networking/grpc/node_service_pb2.py | 50 +++++++---------
exo/networking/grpc/node_service_pb2_grpc.py | 86 ---------------------------
exo/networking/peer_handle.py | 8 ---
exo/orchestration/node.py | 8 ---
exo/orchestration/standard_node.py | 35 +----------
16 files changed, 55 insertions(+), 246 deletions(-)
diff --git a/examples/llama3_distributed.py b/examples/llama3_distributed.py
index 83661a04..ed6474de 100644
--- a/examples/llama3_distributed.py
+++ b/examples/llama3_distributed.py
@@ -50,7 +50,6 @@ async def run_prompt(prompt: str):
)
await peer2.connect()
- await peer2.global_reset(shard, set(), 2)
try:
await peer2.send_prompt(shard, prompt, request_id)
diff --git a/exo/api/chatgpt_api.py b/exo/api/chatgpt_api.py
index e7d8401f..ffcc90b5 100644
--- a/exo/api/chatgpt_api.py
+++ b/exo/api/chatgpt_api.py
@@ -13,7 +13,7 @@ from exo.inference.shard import Shard
from exo.orchestration import Node
shard_mappings = {
- # llama
+ ### llama
"llama-3.1-8b": {
"MLXDynamicShardInferenceEngine": Shard(model_id="mlx-community/Meta-Llama-3.1-8B-Instruct-4bit", start_layer=0, end_layer=0, n_layers=32),
},
@@ -31,7 +31,7 @@ shard_mappings = {
"MLXDynamicShardInferenceEngine": Shard(model_id="mlx-community/Meta-Llama-3-70B-Instruct-4bit", start_layer=0, end_layer=0, n_layers=80),
"TinygradDynamicShardInferenceEngine": Shard(model_id="llama3-70b-sfr", start_layer=0, end_layer=0, n_layers=80),
},
- # mistral
+ ### mistral
"mistral-nemo": {
"MLXDynamicShardInferenceEngine": Shard(model_id="mlx-community/Mistral-Nemo-Instruct-2407-4bit", start_layer=0, end_layer=0, n_layers=40),
},
diff --git a/exo/inference/debug_inference_engine.py b/exo/inference/debug_inference_engine.py
index b90044aa..8853b986 100644
--- a/exo/inference/debug_inference_engine.py
+++ b/exo/inference/debug_inference_engine.py
@@ -12,18 +12,13 @@ async def test_inference_engine(inference_engine_1: InferenceEngine, inference_e
_tokenizer = Tokenizer(str(Path(model_id) / "tokenizer.model"))
prompt = "In a single word only, what is the last name of the president of the United States? "
- resp_full, inference_state_full, _ = await inference_engine_1.infer_prompt(shard=Shard(model_id=model_id, start_layer=0, end_layer=31, n_layers=32), prompt=prompt)
- next_resp_full, next_inference_state_full, _ = await inference_engine_1.infer_tensor(shard=Shard(model_id=model_id, start_layer=0, end_layer=31, n_layers=32), input_data=resp_full, inference_state=inference_state_full)
+ resp_full, inference_state_full, _ = await inference_engine_1.infer_prompt("A", shard=Shard(model_id=model_id, start_layer=0, end_layer=31, n_layers=32), prompt=prompt)
+ next_resp_full, next_inference_state_full, _ = await inference_engine_1.infer_tensor("A", shard=Shard(model_id=model_id, start_layer=0, end_layer=31, n_layers=32), input_data=resp_full, inference_state=inference_state_full)
- await inference_engine_1.reset_shard(shard=Shard(model_id=model_id, start_layer=0, end_layer=30, n_layers=32))
- resp1, inference_state_1, _ = await inference_engine_1.infer_prompt(shard=Shard(model_id=model_id, start_layer=0, end_layer=30, n_layers=32), prompt=prompt)
-
- await inference_engine_2.reset_shard(shard=Shard(model_id=model_id, start_layer=31, end_layer=31, n_layers=32))
- resp2, inference_state_2, _ = await inference_engine_2.infer_tensor(shard=Shard(model_id=model_id, start_layer=31, end_layer=31, n_layers=32), input_data=resp1, inference_state=inference_state_1)
-
- # don't reset the second time
- resp3, inference_state_3, _ = await inference_engine_1.infer_tensor(shard=Shard(model_id=model_id, start_layer=0, end_layer=30, n_layers=32), input_data=resp2, inference_state=inference_state_2)
- resp4, inference_state_4, _ = await inference_engine_2.infer_tensor(shard=Shard(model_id=model_id, start_layer=31, end_layer=31, n_layers=32), input_data=resp3, inference_state=inference_state_3)
+ resp1, inference_state_1, _ = await inference_engine_1.infer_prompt("B", shard=Shard(model_id=model_id, start_layer=0, end_layer=30, n_layers=32), prompt=prompt)
+ resp2, inference_state_2, _ = await inference_engine_2.infer_tensor("B", shard=Shard(model_id=model_id, start_layer=31, end_layer=31, n_layers=32), input_data=resp1, inference_state=inference_state_1)
+ resp3, inference_state_3, _ = await inference_engine_1.infer_tensor("B", shard=Shard(model_id=model_id, start_layer=0, end_layer=30, n_layers=32), input_data=resp2, inference_state=inference_state_2)
+ resp4, inference_state_4, _ = await inference_engine_2.infer_tensor("B", shard=Shard(model_id=model_id, start_layer=31, end_layer=31, n_layers=32), input_data=resp3, inference_state=inference_state_3)
print(f"{resp2=}")
print(f"full: {_tokenizer.decode(resp_full)}")
diff --git a/exo/inference/inference_engine.py b/exo/inference/inference_engine.py
index 2fb4d243..d8763e51 100644
--- a/exo/inference/inference_engine.py
+++ b/exo/inference/inference_engine.py
@@ -6,13 +6,9 @@ from .shard import Shard
class InferenceEngine(ABC):
@abstractmethod
- async def infer_tensor(self, shard: Shard, input_data: np.ndarray, inference_state: Optional[str] = None) -> Tuple[np.ndarray, str, bool]:
+ async def infer_tensor(self, request_id: str, shard: Shard, input_data: np.ndarray, inference_state: Optional[str] = None) -> Tuple[np.ndarray, str, bool]:
pass
@abstractmethod
- async def infer_prompt(self, shard: Shard, prompt: str, inference_state: Optional[str] = None) -> Tuple[np.ndarray, str, bool]:
- pass
-
- @abstractmethod
- async def reset_shard(self, shard: Shard):
+ async def infer_prompt(self, request_id: str, shard: Shard, prompt: str, inference_state: Optional[str] = None) -> Tuple[np.ndarray, str, bool]:
pass
diff --git a/exo/inference/mlx/sharded_inference_engine.py b/exo/inference/mlx/sharded_inference_engine.py
index 2ab83eee..38051483 100644
--- a/exo/inference/mlx/sharded_inference_engine.py
+++ b/exo/inference/mlx/sharded_inference_engine.py
@@ -10,20 +10,16 @@ class MLXDynamicShardInferenceEngine(InferenceEngine):
def __init__(self):
self.shard = None
- async def infer_prompt(self, shard: Shard, prompt: str, inference_state: Optional[str] = None) -> (np.ndarray, str, bool):
+ async def infer_prompt(self, request_id: str, shard: Shard, prompt: str, inference_state: Optional[str] = None) -> (np.ndarray, str, bool):
await self.ensure_shard(shard)
- output_data: np.ndarray = np.array(self.stateful_sharded_model.step(mx.array(self.tokenizer.encode(prompt))))
+ output_data: np.ndarray = np.array(self.stateful_sharded_model.step(request_id, mx.array(self.tokenizer.encode(prompt))))
return output_data, "", output_data.size == 1 and output_data.item() == self.tokenizer.eos_token_id
- async def infer_tensor(self, shard: Shard, input_data: np.ndarray, inference_state: Optional[str] = None) -> (np.ndarray, str, bool):
+ async def infer_tensor(self, request_id: str, shard: Shard, input_data: np.ndarray, inference_state: Optional[str] = None) -> (np.ndarray, str, bool):
await self.ensure_shard(shard)
- output_data: np.ndarray = np.array(self.stateful_sharded_model.step(mx.array(input_data)))
+ output_data: np.ndarray = np.array(self.stateful_sharded_model.step(request_id, mx.array(input_data)))
return output_data, "", output_data.size == 1 and output_data.item() == self.tokenizer.eos_token_id
- async def reset_shard(self, shard: Shard):
- await self.ensure_shard(shard)
- self.stateful_sharded_model.reset()
-
async def ensure_shard(self, shard: Shard):
if self.shard == shard:
return
diff --git a/exo/inference/mlx/sharded_model.py b/exo/inference/mlx/sharded_model.py
index 43e26d17..98c02d96 100644
--- a/exo/inference/mlx/sharded_model.py
+++ b/exo/inference/mlx/sharded_model.py
@@ -11,10 +11,11 @@ class StatefulShardedModel:
def __init__(self, shard: Shard, model: nn.Module):
self.shard = shard
self.model = model
- self.reset()
+ self.request_cache: Dict[str, Tuple[str, KVCache]] = {}
def step(
self,
+ request_id: str,
x,
temp: float = 0.0,
top_p: float = 1.0,
@@ -38,7 +39,9 @@ class StatefulShardedModel:
y = x
- output = self.model(y[None] if self.shard.is_first_layer() else y, cache=self.cache)
+ if request_id not in self.request_cache:
+ self.init_cache(request_id)
+ output = self.model(y[None] if self.shard.is_first_layer() else y, cache=self.request_cache[request_id])
if self.shard.is_last_layer():
logits = output[:, -1, :]
@@ -56,10 +59,10 @@ class StatefulShardedModel:
) -> Generator[Tuple[mx.array, mx.array], None, None]:
return self.step(x, temp, top_p, logit_bias)
- def reset(self):
+ def init_cache(self, request_id: str):
kv_heads = (
[self.model.n_kv_heads] * len(self.model.layers)
if isinstance(self.model.n_kv_heads, int)
else self.model.n_kv_heads
)
- self.cache = [KVCache(self.model.head_dim, n) for n in kv_heads]
+ self.request_cache[request_id] = [KVCache(self.model.head_dim, n) for n in kv_heads]
diff --git a/exo/inference/test_inference_engine.py b/exo/inference/test_inference_engine.py
index cca8b66b..735ba045 100644
--- a/exo/inference/test_inference_engine.py
+++ b/exo/inference/test_inference_engine.py
@@ -8,18 +8,13 @@ import numpy as np
# An inference engine should work the same for any number of Shards, as long as the Shards are continuous.
async def test_inference_engine(inference_engine_1: InferenceEngine, inference_engine_2: InferenceEngine, model_id: str):
prompt = "In a single word only, what is the last name of the current president of the USA?"
- resp_full, inference_state_full, _ = await inference_engine_1.infer_prompt(shard=Shard(model_id=model_id, start_layer=0, end_layer=31, n_layers=32), prompt=prompt)
- next_resp_full, next_inference_state_full, _ = await inference_engine_1.infer_tensor(shard=Shard(model_id=model_id, start_layer=0, end_layer=31, n_layers=32), input_data=resp_full, inference_state=inference_state_full)
+ resp_full, inference_state_full, _ = await inference_engine_1.infer_prompt("A", shard=Shard(model_id=model_id, start_layer=0, end_layer=31, n_layers=32), prompt=prompt)
+ next_resp_full, next_inference_state_full, _ = await inference_engine_1.infer_tensor("A", shard=Shard(model_id=model_id, start_layer=0, end_layer=31, n_layers=32), input_data=resp_full, inference_state=inference_state_full)
- await inference_engine_1.reset_shard(shard=Shard(model_id=model_id, start_layer=0, end_layer=30, n_layers=32))
- resp1, inference_state_1, _ = await inference_engine_1.infer_prompt(shard=Shard(model_id=model_id, start_layer=0, end_layer=30, n_layers=32), prompt=prompt)
-
- await inference_engine_2.reset_shard(shard=Shard(model_id=model_id, start_layer=31, end_layer=31, n_layers=32))
- resp2, inference_state_2, _ = await inference_engine_2.infer_tensor(shard=Shard(model_id=model_id, start_layer=31, end_layer=31, n_layers=32), input_data=resp1, inference_state=inference_state_1)
-
- # don't reset the second time
- resp3, inference_state_3, _ = await inference_engine_1.infer_tensor(shard=Shard(model_id=model_id, start_layer=0, end_layer=30, n_layers=32), input_data=resp2, inference_state=inference_state_2)
- resp4, inference_state_4, _ = await inference_engine_2.infer_tensor(shard=Shard(model_id=model_id, start_layer=31, end_layer=31, n_layers=32), input_data=resp3, inference_state=inference_state_3)
+ resp1, inference_state_1, _ = await inference_engine_1.infer_prompt("B", shard=Shard(model_id=model_id, start_layer=0, end_layer=30, n_layers=32), prompt=prompt)
+ resp2, inference_state_2, _ = await inference_engine_2.infer_tensor("B", shard=Shard(model_id=model_id, start_layer=31, end_layer=31, n_layers=32), input_data=resp1, inference_state=inference_state_1)
+ resp3, inference_state_3, _ = await inference_engine_1.infer_tensor("B", shard=Shard(model_id=model_id, start_layer=0, end_layer=30, n_layers=32), input_data=resp2, inference_state=inference_state_2)
+ resp4, inference_state_4, _ = await inference_engine_2.infer_tensor("B", shard=Shard(model_id=model_id, start_layer=31, end_layer=31, n_layers=32), input_data=resp3, inference_state=inference_state_3)
assert np.array_equal(resp_full, resp2)
assert np.array_equal(next_resp_full, resp4)
diff --git a/exo/inference/tinygrad/inference.py b/exo/inference/tinygrad/inference.py
index 4cf1d1b1..a240e5d4 100644
--- a/exo/inference/tinygrad/inference.py
+++ b/exo/inference/tinygrad/inference.py
@@ -143,7 +143,8 @@ class TinygradDynamicShardInferenceEngine(InferenceEngine):
def __init__(self):
self.shard = None
- async def infer_prompt(self, shard: Shard, prompt: str, inference_state: Optional[str] = None) -> (np.ndarray, str, bool):
+ async def infer_prompt(self, request_id: str, shard: Shard, prompt: str, inference_state: Optional[str] = None) -> (np.ndarray, str, bool):
+ # TODO: we need to refactor models/llamaa to handle per-request-kv-cache. right now it's shared between requests.
await self.ensure_shard(shard)
start_pos = json.loads(inference_state).get("start_pos", 0) if inference_state else 0
@@ -157,7 +158,7 @@ class TinygradDynamicShardInferenceEngine(InferenceEngine):
return output_data, json.dumps({"start_pos": start_pos}), output_data.size == 1 and output_data.item() in self.tokenizer.stop_tokens
- async def infer_tensor(self, shard: Shard, input_data: np.ndarray, inference_state: Optional[str] = None) -> (np.ndarray, str, bool):
+ async def infer_tensor(self, request_id: str, shard: Shard, input_data: np.ndarray, inference_state: Optional[str] = None) -> (np.ndarray, str, bool):
await self.ensure_shard(shard)
start_pos = json.loads(inference_state).get("start_pos", 0) if inference_state else 0
@@ -167,11 +168,6 @@ class TinygradDynamicShardInferenceEngine(InferenceEngine):
return output_data, json.dumps({"start_pos": start_pos}), output_data.size == 1 and output_data.item() in self.tokenizer.stop_tokens
- async def reset_shard(self, shard: Shard):
- await self.ensure_shard(shard)
-
- self.model.reset()
-
async def ensure_shard(self, shard: Shard):
if self.shard == shard:
return
diff --git a/exo/networking/grpc/grpc_peer_handle.py b/exo/networking/grpc/grpc_peer_handle.py
index 9e274be7..2570628e 100644
--- a/exo/networking/grpc/grpc_peer_handle.py
+++ b/exo/networking/grpc/grpc_peer_handle.py
@@ -74,10 +74,6 @@ class GRPCPeerHandle(PeerHandle):
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 reset_shard(self, shard: Shard) -> None:
- 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, visited: set[str], max_depth: int) -> Topology:
request = node_service_pb2.CollectTopologyRequest(visited=visited, max_depth=max_depth)
response = await self.stub.CollectTopology(request)
@@ -90,10 +86,6 @@ class GRPCPeerHandle(PeerHandle):
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)
-
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)
diff --git a/exo/networking/grpc/grpc_server.py b/exo/networking/grpc/grpc_server.py
index 41904708..0d9e7782 100644
--- a/exo/networking/grpc/grpc_server.py
+++ b/exo/networking/grpc/grpc_server.py
@@ -60,12 +60,6 @@ class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):
tensor_data = result[0].tobytes() if result[0] is not None else None
return node_service_pb2.InferenceResult(tensor=node_service_pb2.Tensor(tensor_data=tensor_data, shape=result[0].shape, dtype=str(result[0].dtype)), is_finished=result[1]) if result[0] is not None else node_service_pb2.InferenceResult(is_finished=result[1])
- async def ResetShard(self, request, context):
- shard = Shard(model_id=request.shard.model_id, start_layer=request.shard.start_layer, end_layer=request.shard.end_layer, n_layers=request.shard.n_layers)
- if DEBUG >= 2: print(f"Received ResetShard request: {shard}")
- await self.node.reset_shard(shard)
- return node_service_pb2.Empty()
-
async def CollectTopology(self, request, context):
max_depth = request.max_depth
visited = set(request.visited)
@@ -75,14 +69,6 @@ class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):
if DEBUG >= 2: print(f"CollectTopology {max_depth=} {visited=} {nodes=} {peer_graph=}")
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
- if DEBUG >= 2: print(f"Received GlobalReset request: {base_shard=} {visited=} {max_depth=}")
- await self.node.global_reset(base_shard, visited, max_depth)
- return node_service_pb2.Empty()
-
async def SendResult(self, request, context):
request_id = request.request_id
result = request.result
diff --git a/exo/networking/grpc/node_service.proto b/exo/networking/grpc/node_service.proto
index c0c7c224..d76430ca 100644
--- a/exo/networking/grpc/node_service.proto
+++ b/exo/networking/grpc/node_service.proto
@@ -5,10 +5,8 @@ package node_service;
service NodeService {
rpc SendPrompt (PromptRequest) returns (Tensor) {}
rpc SendTensor (TensorRequest) returns (Tensor) {}
- rpc ResetShard (ResetShardRequest) returns (Empty) {}
rpc GetInferenceResult (GetInferenceResultRequest) returns (InferenceResult) {}
rpc CollectTopology (CollectTopologyRequest) returns (Topology) {}
- rpc GlobalReset (GlobalResetRequest) returns (Empty) {}
rpc SendResult (SendResultRequest) returns (Empty) {}
rpc SendOpaqueStatus (SendOpaqueStatusRequest) returns (Empty) {}
}
@@ -49,21 +47,11 @@ message Tensor {
string dtype = 3;
}
-message ResetShardRequest {
- Shard shard = 1;
-}
-
message CollectTopologyRequest {
repeated string visited = 1;
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 765c2ea2..66e516c7 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\"\x9d\x01\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\x12\x1c\n\x0finference_state\x18\x04 \x01(\tH\x01\x88\x01\x01\x42\r\n\x0b_request_idB\x12\n\x10_inference_state\"\xb3\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\x12\x1c\n\x0finference_state\x18\x04 \x01(\tH\x01\x88\x01\x01\x42\r\n\x0b_request_idB\x12\n\x10_inference_state\"/\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\"7\n\x0b\x44\x65viceFlops\x12\x0c\n\x04\x66p32\x18\x01 \x01(\x02\x12\x0c\n\x04\x66p16\x18\x02 \x01(\x02\x12\x0c\n\x04int8\x18\x03 \x01(\x02\"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\"\x07\n\x05\x45mpty2\xec\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\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\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\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\"\x9d\x01\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\x12\x1c\n\x0finference_state\x18\x04 \x01(\tH\x01\x88\x01\x01\x42\r\n\x0b_request_idB\x12\n\x10_inference_state\"\xb3\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\x12\x1c\n\x0finference_state\x18\x04 \x01(\tH\x01\x88\x01\x01\x42\r\n\x0b_request_idB\x12\n\x10_inference_state\"/\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\"\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\"7\n\x0b\x44\x65viceFlops\x12\x0c\n\x04\x66p32\x18\x01 \x01(\x02\x12\x0c\n\x04\x66p16\x18\x02 \x01(\x02\x12\x0c\n\x04int8\x18\x03 \x01(\x02\"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\"\x07\n\x05\x45mpty2\xde\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^\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\x62\x06proto3')
_globals = globals()
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
@@ -37,30 +37,26 @@ if not _descriptor._USE_C_DESCRIPTORS:
_globals['_INFERENCERESULT']._serialized_end=604
_globals['_TENSOR']._serialized_start=606
_globals['_TENSOR']._serialized_end=665
- _globals['_RESETSHARDREQUEST']._serialized_start=667
- _globals['_RESETSHARDREQUEST']._serialized_end=722
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=724
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=784
- _globals['_GLOBALRESETREQUEST']._serialized_start=786
- _globals['_GLOBALRESETREQUEST']._serialized_end=883
- _globals['_TOPOLOGY']._serialized_start=886
- _globals['_TOPOLOGY']._serialized_end=1156
- _globals['_TOPOLOGY_NODESENTRY']._serialized_start=1007
- _globals['_TOPOLOGY_NODESENTRY']._serialized_end=1085
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=1087
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=1156
- _globals['_PEERS']._serialized_start=1158
- _globals['_PEERS']._serialized_end=1183
- _globals['_DEVICEFLOPS']._serialized_start=1185
- _globals['_DEVICEFLOPS']._serialized_end=1240
- _globals['_DEVICECAPABILITIES']._serialized_start=1242
- _globals['_DEVICECAPABILITIES']._serialized_end=1349
- _globals['_SENDRESULTREQUEST']._serialized_start=1351
- _globals['_SENDRESULTREQUEST']._serialized_end=1427
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_start=1429
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_end=1490
- _globals['_EMPTY']._serialized_start=1492
- _globals['_EMPTY']._serialized_end=1499
- _globals['_NODESERVICE']._serialized_start=1502
- _globals['_NODESERVICE']._serialized_end=2122
+ _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=667
+ _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=727
+ _globals['_TOPOLOGY']._serialized_start=730
+ _globals['_TOPOLOGY']._serialized_end=1000
+ _globals['_TOPOLOGY_NODESENTRY']._serialized_start=851
+ _globals['_TOPOLOGY_NODESENTRY']._serialized_end=929
+ _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=931
+ _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=1000
+ _globals['_PEERS']._serialized_start=1002
+ _globals['_PEERS']._serialized_end=1027
+ _globals['_DEVICEFLOPS']._serialized_start=1029
+ _globals['_DEVICEFLOPS']._serialized_end=1084
+ _globals['_DEVICECAPABILITIES']._serialized_start=1086
+ _globals['_DEVICECAPABILITIES']._serialized_end=1193
+ _globals['_SENDRESULTREQUEST']._serialized_start=1195
+ _globals['_SENDRESULTREQUEST']._serialized_end=1271
+ _globals['_SENDOPAQUESTATUSREQUEST']._serialized_start=1273
+ _globals['_SENDOPAQUESTATUSREQUEST']._serialized_end=1334
+ _globals['_EMPTY']._serialized_start=1336
+ _globals['_EMPTY']._serialized_end=1343
+ _globals['_NODESERVICE']._serialized_start=1346
+ _globals['_NODESERVICE']._serialized_end=1824
# @@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 920a2a3e..6bf04a92 100644
--- a/exo/networking/grpc/node_service_pb2_grpc.py
+++ b/exo/networking/grpc/node_service_pb2_grpc.py
@@ -49,11 +49,6 @@ class NodeServiceStub(object):
request_serializer=node__service__pb2.TensorRequest.SerializeToString,
response_deserializer=node__service__pb2.Tensor.FromString,
_registered_method=True)
- self.ResetShard = channel.unary_unary(
- '/node_service.NodeService/ResetShard',
- request_serializer=node__service__pb2.ResetShardRequest.SerializeToString,
- response_deserializer=node__service__pb2.Empty.FromString,
- _registered_method=True)
self.GetInferenceResult = channel.unary_unary(
'/node_service.NodeService/GetInferenceResult',
request_serializer=node__service__pb2.GetInferenceResultRequest.SerializeToString,
@@ -64,11 +59,6 @@ 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)
self.SendResult = channel.unary_unary(
'/node_service.NodeService/SendResult',
request_serializer=node__service__pb2.SendResultRequest.SerializeToString,
@@ -96,12 +86,6 @@ class NodeServiceServicer(object):
context.set_details('Method not implemented!')
raise NotImplementedError('Method not implemented!')
- def ResetShard(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 GetInferenceResult(self, request, context):
"""Missing associated documentation comment in .proto file."""
context.set_code(grpc.StatusCode.UNIMPLEMENTED)
@@ -114,12 +98,6 @@ 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 SendResult(self, request, context):
"""Missing associated documentation comment in .proto file."""
context.set_code(grpc.StatusCode.UNIMPLEMENTED)
@@ -145,11 +123,6 @@ def add_NodeServiceServicer_to_server(servicer, server):
request_deserializer=node__service__pb2.TensorRequest.FromString,
response_serializer=node__service__pb2.Tensor.SerializeToString,
),
- 'ResetShard': grpc.unary_unary_rpc_method_handler(
- servicer.ResetShard,
- request_deserializer=node__service__pb2.ResetShardRequest.FromString,
- response_serializer=node__service__pb2.Empty.SerializeToString,
- ),
'GetInferenceResult': grpc.unary_unary_rpc_method_handler(
servicer.GetInferenceResult,
request_deserializer=node__service__pb2.GetInferenceResultRequest.FromString,
@@ -160,11 +133,6 @@ 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,
- ),
'SendResult': grpc.unary_unary_rpc_method_handler(
servicer.SendResult,
request_deserializer=node__service__pb2.SendResultRequest.FromString,
@@ -240,33 +208,6 @@ class NodeService(object):
metadata,
_registered_method=True)
- @staticmethod
- def ResetShard(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/ResetShard',
- node__service__pb2.ResetShardRequest.SerializeToString,
- node__service__pb2.Empty.FromString,
- options,
- channel_credentials,
- insecure,
- call_credentials,
- compression,
- wait_for_ready,
- timeout,
- metadata,
- _registered_method=True)
-
@staticmethod
def GetInferenceResult(request,
target,
@@ -321,33 +262,6 @@ class NodeService(object):
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)
-
@staticmethod
def SendResult(request,
target,
diff --git a/exo/networking/peer_handle.py b/exo/networking/peer_handle.py
index f4d478ca..1196d547 100644
--- a/exo/networking/peer_handle.py
+++ b/exo/networking/peer_handle.py
@@ -38,18 +38,10 @@ class PeerHandle(ABC):
async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
pass
- @abstractmethod
- 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
-
@abstractmethod
async def send_result(self, request_id: str, result: List[int], is_finished: bool) -> None:
pass
diff --git a/exo/orchestration/node.py b/exo/orchestration/node.py
index 1caf6618..485185bb 100644
--- a/exo/orchestration/node.py
+++ b/exo/orchestration/node.py
@@ -22,10 +22,6 @@ class Node(ABC):
async def process_tensor(self, shard: Shard, tensor: np.ndarray, request_id: Optional[str] = None, inference_state: Optional[str] = None) -> Optional[np.ndarray]:
pass
- @abstractmethod
- async def reset_shard(self, shard: Shard) -> None:
- pass
-
@abstractmethod
async def get_inference_result(self, request_id: str) -> Tuple[Optional[np.ndarray], bool]:
pass
@@ -34,10 +30,6 @@ class Node(ABC):
async def collect_topology(self, visited: set[str] = set(), max_depth: int = 2) -> Topology:
pass
- @abstractmethod
- async def global_reset(self, base_shard: Shard, visited: set[str] = set(), max_depth: int = 2) -> None:
- pass
-
@property
@abstractmethod
def current_topology(self) -> Topology:
diff --git a/exo/orchestration/standard_node.py b/exo/orchestration/standard_node.py
index 67af8e71..0e08ec9f 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -79,7 +79,7 @@ class StandardNode(Node):
await self.forward_to_next_shard(shard, prompt, request_id)
return
- result, inference_state, is_finished = await self.inference_engine.infer_prompt(shard, prompt, inference_state=inference_state)
+ result, inference_state, is_finished = await self.inference_engine.infer_prompt(request_id, shard, prompt, inference_state=inference_state)
is_finished = is_finished or len(self.buffered_token_output[request_id][0]) >= self.max_generate_tokens
if is_finished:
self.buffered_token_output[request_id] = (self.buffered_token_output[request_id][0], True)
@@ -115,7 +115,7 @@ class StandardNode(Node):
try:
if DEBUG >= 1: print(f"[{request_id}] process_tensor: {tensor.size=} {tensor.shape=}")
- result, inference_state, is_finished = await self.inference_engine.infer_tensor(shard, tensor, inference_state=inference_state)
+ result, inference_state, is_finished = await self.inference_engine.infer_tensor(request_id, shard, tensor, inference_state=inference_state)
is_finished = is_finished or len(self.buffered_token_output[request_id][0]) >= self.max_generate_tokens
if is_finished:
self.buffered_token_output[request_id] = (self.buffered_token_output[request_id][0], True)
@@ -178,12 +178,6 @@ class StandardNode(Node):
raise ValueError(f"No current partition found for node: {self.id}")
return shards[current_partition_index]
- async def reset_shard(self, base_shard: Shard) -> None:
- # Implement shard reset logic
- if DEBUG >= 2: print(f"Resetting shard: {base_shard}")
- self.buffered_token_output = {}
- await self.inference_engine.reset_shard(self.get_current_shard(base_shard))
-
async def update_peers(self, wait_for_peers: int = 0) -> None:
self.peers = await self.discovery.discover_peers(wait_for_peers)
if DEBUG >= 2: print(f"Starting with the following peers: {self.peers}")
@@ -245,31 +239,6 @@ class StandardNode(Node):
if self.topology_viz: self.topology_viz.update_visualization(self.current_topology, self.partitioning_strategy.partition(self.current_topology))
return next_topology
- # 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:
- shard = self.get_current_shard(base_shard)
- await self.reset_shard(shard)
-
- if DEBUG >= 2: print(f"Global reset {base_shard=} {max_depth=} {visited=}")
-
- prev_visited = visited.copy()
- visited.update(p.id() for p in self.peers)
-
- 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}")
-
@property
def on_token(self) -> AsyncCallbackSystem[str, Tuple[str, List[int], bool]]:
return self._on_token
← dd8c5d63 add support for mistral nemo and mistral large
·
back to Exo
·
init 803a4421 →