← back to Exo
Correct loss propagation so we can see the actual loss instead of just the requestor shard's loss
9283f6d7bd192f80ccfdaa3d41c02728d5cd3947 · 2024-12-06 00:50:52 -0800 · Nel Nibcord
Files touched
M exo/api/chatgpt_api.pyM exo/inference/mlx/sharded_inference_engine.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/orchestration/standard_node.pyM setup.py
Diff
commit 9283f6d7bd192f80ccfdaa3d41c02728d5cd3947
Author: Nel Nibcord <blindcrone@tuta.io>
Date: Fri Dec 6 00:50:52 2024 -0800
Correct loss propagation so we can see the actual loss instead of just the requestor shard's loss
---
exo/api/chatgpt_api.py | 1 -
exo/inference/mlx/sharded_inference_engine.py | 4 +-
exo/networking/grpc/grpc_peer_handle.py | 14 ++---
exo/networking/grpc/grpc_server.py | 12 +++--
exo/networking/grpc/node_service.proto | 8 +--
exo/networking/grpc/node_service_pb2.py | 74 +++++++++++++--------------
exo/networking/grpc/node_service_pb2_grpc.py | 6 +--
exo/orchestration/standard_node.py | 41 +++++----------
setup.py | 4 +-
9 files changed, 76 insertions(+), 88 deletions(-)
diff --git a/exo/api/chatgpt_api.py b/exo/api/chatgpt_api.py
index b7447d67..7aa8e595 100644
--- a/exo/api/chatgpt_api.py
+++ b/exo/api/chatgpt_api.py
@@ -153,7 +153,6 @@ class ChatGPTAPI:
def __init__(self, node: Node, inference_engine_classname: str, response_timeout: int = 90, on_chat_completion_request: Callable[[str, ChatCompletionRequest, str], None] = None, default_model: Optional[str] = None):
self.node = node
self.inference_engine_classname = inference_engine_classname
- print(self.inference_engine_classname)
self.response_timeout = response_timeout
self.on_chat_completion_request = on_chat_completion_request
self.app = web.Application(client_max_size=100*1024*1024) # 100MB to support image upload
diff --git a/exo/inference/mlx/sharded_inference_engine.py b/exo/inference/mlx/sharded_inference_engine.py
index d954f292..e5138e82 100644
--- a/exo/inference/mlx/sharded_inference_engine.py
+++ b/exo/inference/mlx/sharded_inference_engine.py
@@ -99,7 +99,7 @@ class MLXDynamicShardInferenceEngine(InferenceEngine):
l = mx.array(lengths)
score = await loop.run_in_executor(self.executor, self.session['loss'], self.model, x, y, l)
#print(f"evaluate out -> {score}")
- return np.array(score)
+ return score
async def ensure_train(self, shard: Shard, loss: str, opt=optim.SGD, lr=1e-5, trainable_layers=['input_layernorm', 'gate_proj']):
await self.ensure_shard(shard)
@@ -134,7 +134,7 @@ class MLXDynamicShardInferenceEngine(InferenceEngine):
layers = [{k: v["weight"] for k,v in l.items() if 'weight' in v} for l in gradients if l]
#print(layers[0])
- return np.array(score).reshape(1, -1), np.array(layers[0]['input_layernorm'])
+ return score, np.array(layers[0]['input_layernorm'])
async def ensure_shard(self, shard: Shard):
if self.shard == shard:
diff --git a/exo/networking/grpc/grpc_peer_handle.py b/exo/networking/grpc/grpc_peer_handle.py
index df733d86..35b6968a 100644
--- a/exo/networking/grpc/grpc_peer_handle.py
+++ b/exo/networking/grpc/grpc_peer_handle.py
@@ -118,16 +118,16 @@ class GRPCPeerHandle(PeerHandle):
example=node_service_pb2.Tensor(tensor_data=example.tobytes(), shape=example.shape, dtype=str(example.dtype)),
target=node_service_pb2.Tensor(tensor_data=target.tobytes(), shape=target.shape, dtype=str(target.dtype)),
length=node_service_pb2.Tensor(tensor_data=length.tobytes(), shape=length.shape, dtype=str(length.dtype)),
- train = train,
+ train=train,
request_id=request_id,
)
response = await self.stub.SendExample(request)
-
- if not response.tensor_data or not response.shape or not response.dtype:
- return None
-
- out = np.frombuffer(response.tensor_data, dtype=np.dtype(response.dtype)).reshape(response.shape)
- return out
+ loss = response.loss
+ if train and not shard.is_first_layer():
+ grads = np.frombuffer(response.grads.tensor_data, dtype=np.dtype(response.grads.dtype)).reshape(response.grads.shape)
+ return loss, grads
+ else:
+ return loss
async def send_loss(self, shard: Shard, tensor: np.ndarray, request_id: Optional[str] = None) -> Optional[np.array]:
request = node_service_pb2.TensorRequest(
diff --git a/exo/networking/grpc/grpc_server.py b/exo/networking/grpc/grpc_server.py
index 1fba89e9..c6a56685 100644
--- a/exo/networking/grpc/grpc_server.py
+++ b/exo/networking/grpc/grpc_server.py
@@ -83,10 +83,14 @@ class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):
train = request.train
request_id = request.request_id
- result = await self.node.process_example(shard, example, target, length, train, request_id)
- if DEBUG >= 5: print(f"SendTensor tensor {shard=} {example=} {target=} {length=} {request_id=} result: {result}")
- tensor_data = result.tobytes()
- return node_service_pb2.Tensor(tensor_data=tensor_data, shape=result.shape, dtype=str(result.dtype))
+ if train and not shard.is_first_layer():
+ loss, grad = await self.node.process_example(shard, example, target, length, train, request_id)
+ tensor_data = grad.tobytes()
+ grad_tensor = node_service_pb2.Tensor(tensor_data=tensor_data, shape=grad.shape, dtype=str(grad.dtype))
+ return node_service_pb2.Loss(loss=loss, grads=grad_tensor)
+ else:
+ loss = await self.node.process_example(shard, example, target, length, train, request_id)
+ return node_service_pb2.Loss(loss=loss, grads=None)
async def CollectTopology(self, request, context):
max_depth = request.max_depth
diff --git a/exo/networking/grpc/node_service.proto b/exo/networking/grpc/node_service.proto
index 1bbfc96f..83f1f03c 100644
--- a/exo/networking/grpc/node_service.proto
+++ b/exo/networking/grpc/node_service.proto
@@ -5,7 +5,7 @@ package node_service;
service NodeService {
rpc SendPrompt (PromptRequest) returns (Tensor) {}
rpc SendTensor (TensorRequest) returns (Tensor) {}
- rpc SendExample (ExampleRequest) returns (Tensor) {}
+ rpc SendExample (ExampleRequest) returns (Loss) {}
rpc GetInferenceResult (GetInferenceResultRequest) returns (InferenceResult) {}
rpc CollectTopology (CollectTopologyRequest) returns (Topology) {}
rpc SendResult (SendResultRequest) returns (Empty) {}
@@ -41,9 +41,9 @@ message ExampleRequest {
optional string request_id = 6;
}
-message LossRequest {
- Tensor loss = 1;
- Tensor grads = 2;
+message Loss {
+ float loss = 1;
+ optional Tensor grads = 2;
}
message GetInferenceResultRequest {
diff --git a/exo/networking/grpc/node_service_pb2.py b/exo/networking/grpc/node_service_pb2.py
index 46f3496a..3e34e365 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\"V\n\x0bLossRequest\x12\"\n\x04loss\x18\x01 \x01(\x0b\x32\x14.node_service.Tensor\x12#\n\x05grads\x18\x02 \x01(\x0b\x32\x14.node_service.Tensor\"/\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\xf9\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\x43\n\x0bSendExample\x12\x1c.node_service.ExampleRequest\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\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\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')
_globals = globals()
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
@@ -43,40 +43,40 @@ if not _descriptor._USE_C_DESCRIPTORS:
_globals['_TENSORREQUEST']._serialized_end=360
_globals['_EXAMPLEREQUEST']._serialized_start=363
_globals['_EXAMPLEREQUEST']._serialized_end=585
- _globals['_LOSSREQUEST']._serialized_start=587
- _globals['_LOSSREQUEST']._serialized_end=673
- _globals['_GETINFERENCERESULTREQUEST']._serialized_start=675
- _globals['_GETINFERENCERESULTREQUEST']._serialized_end=722
- _globals['_INFERENCERESULT']._serialized_start=724
- _globals['_INFERENCERESULT']._serialized_end=816
- _globals['_TENSOR']._serialized_start=818
- _globals['_TENSOR']._serialized_end=877
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=879
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=939
- _globals['_TOPOLOGY']._serialized_start=942
- _globals['_TOPOLOGY']._serialized_end=1222
- _globals['_TOPOLOGY_NODESENTRY']._serialized_start=1063
- _globals['_TOPOLOGY_NODESENTRY']._serialized_end=1141
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=1143
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=1222
- _globals['_PEERCONNECTION']._serialized_start=1224
- _globals['_PEERCONNECTION']._serialized_end=1297
- _globals['_PEERCONNECTIONS']._serialized_start=1299
- _globals['_PEERCONNECTIONS']._serialized_end=1367
- _globals['_DEVICEFLOPS']._serialized_start=1369
- _globals['_DEVICEFLOPS']._serialized_end=1424
- _globals['_DEVICECAPABILITIES']._serialized_start=1426
- _globals['_DEVICECAPABILITIES']._serialized_end=1533
- _globals['_SENDRESULTREQUEST']._serialized_start=1535
- _globals['_SENDRESULTREQUEST']._serialized_end=1611
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_start=1613
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_end=1674
- _globals['_HEALTHCHECKREQUEST']._serialized_start=1676
- _globals['_HEALTHCHECKREQUEST']._serialized_end=1696
- _globals['_HEALTHCHECKRESPONSE']._serialized_start=1698
- _globals['_HEALTHCHECKRESPONSE']._serialized_end=1739
- _globals['_EMPTY']._serialized_start=1741
- _globals['_EMPTY']._serialized_end=1748
- _globals['_NODESERVICE']._serialized_start=1751
- _globals['_NODESERVICE']._serialized_end=2384
+ _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
# @@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 71856d35..9cf336ce 100644
--- a/exo/networking/grpc/node_service_pb2_grpc.py
+++ b/exo/networking/grpc/node_service_pb2_grpc.py
@@ -47,7 +47,7 @@ class NodeServiceStub(object):
self.SendExample = channel.unary_unary(
'/node_service.NodeService/SendExample',
request_serializer=node__service__pb2.ExampleRequest.SerializeToString,
- response_deserializer=node__service__pb2.Tensor.FromString,
+ response_deserializer=node__service__pb2.Loss.FromString,
_registered_method=True)
self.GetInferenceResult = channel.unary_unary(
'/node_service.NodeService/GetInferenceResult',
@@ -143,7 +143,7 @@ def add_NodeServiceServicer_to_server(servicer, server):
'SendExample': grpc.unary_unary_rpc_method_handler(
servicer.SendExample,
request_deserializer=node__service__pb2.ExampleRequest.FromString,
- response_serializer=node__service__pb2.Tensor.SerializeToString,
+ response_serializer=node__service__pb2.Loss.SerializeToString,
),
'GetInferenceResult': grpc.unary_unary_rpc_method_handler(
servicer.GetInferenceResult,
@@ -251,7 +251,7 @@ class NodeService(object):
target,
'/node_service.NodeService/SendExample',
node__service__pb2.ExampleRequest.SerializeToString,
- node__service__pb2.Tensor.FromString,
+ node__service__pb2.Loss.FromString,
options,
channel_credentials,
insecure,
diff --git a/exo/orchestration/standard_node.py b/exo/orchestration/standard_node.py
index 57b475a3..b0b62874 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -122,7 +122,6 @@ class StandardNode(Node):
await self.inference_engine.ensure_shard(shard)
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
- self.trigger_on_token_callbacks(request_id, self.buffered_token_output[request_id][0], is_finished)
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)
@@ -211,13 +210,14 @@ class StandardNode(Node):
):
shard = self.get_current_shard(base_shard)
if shard.is_first_layer():
- resp = await self.process_example(shard, example, target, length, train, request_id)
+ loss = await self.process_example(shard, example, target, length, train, request_id)
+ return loss
else:
if request_id is None:
request_id = str(uuid.uuid4())
self.outstanding_requests[request_id] = "waiting"
- resp = await self.forward_example(shard, example, target, length, train, request_id, 0)
- return resp
+ loss = await self.forward_example(shard, example, target, length, train, request_id, 0)
+ return loss
async def coordinate_save(
self,
@@ -279,7 +279,6 @@ class StandardNode(Node):
"shard": shard.to_dict(),
"request_id": request_id,
"elapsed_time_ns": elapsed_time_ns,
- "result_size": resp.size if resp is not None else 0,
}),
)
)
@@ -308,11 +307,14 @@ class StandardNode(Node):
self.outstanding_requests[request_id] = "preprocessing"
step = await self.inference_engine.infer_tensor(request_id, shard, example)
self.outstanding_requests[request_id] = "waiting"
- backgrad = await self.forward_example(shard, step, target, length, train, request_id, self.get_partition_index(offset = 1))
+ loss, backgrad = await self.forward_example(shard, step, target, length, train, request_id, self.get_partition_index(offset = 1))
self.outstanding_requests[request_id] = "training"
- loss, grad = await self.inference_engine.train(request_id, shard, example, backgrad, length, loss="back_gradient")
+ partial_loss, grad = await self.inference_engine.train(request_id, shard, example, backgrad, length, loss="back_gradient")
self.outstanding_requests.pop(request_id)
- return loss.reshape(1, -1) if shard.is_first_layer() else grad
+ if shard.is_first_layer():
+ return loss
+ else:
+ return loss, grad
else:
if shard.is_last_layer():
self.outstanding_requests[request_id] = "evaluating"
@@ -323,7 +325,7 @@ class StandardNode(Node):
self.outstanding_requests[request_id] = "waiting"
loss = await self.forward_example(shard, step, target, length, train, request_id, self.get_partition_index(offset = 1))
self.outstanding_requests.pop(request_id)
- return loss.reshape(1, -1)
+ return loss
except Exception as e:
self.outstanding_requests.pop(request_id)
print(f"Error processing example for shard {shard}: {e}")
@@ -413,25 +415,8 @@ class StandardNode(Node):
if not target_peer:
raise ValueError(f"peer for {target_index} not found")
if DEBUG >= 1: print(f"sending example to {target_peer.id()}: {step} => {target} ({length})")
- ret = await target_peer.send_example(target_shard, step, target, length, request_id=request_id, train=train)
- return ret
-
- async def forward_loss(
- self,
- base_shard: Shard,
- loss: np.ndarray,
- request_id: str,
- target_index: int,
- ) -> None:
- if DEBUG >= 1: print(f"target partition index: {target_index}")
- target_id = self.partitioning_strategy.partition(self.topology)[target_index].node_id
- target_shard = self.get_current_shard(base_shard, target_index)
- if DEBUG >= 2: print(f"computed target from: {base_shard} {target_index}, {self.topology}. target shard: {target_shard}")
- target_peer = next((p for p in self.peers if p.id() == target_id), None)
- if not target_peer:
- raise ValueError(f"peer for {target_index} not found")
- if DEBUG >= 1: print(f"sending tensor to {target_peer.id()}: {loss}")
- await target_peer.send_loss(target_shard, step, target, length, request_id=request_id)
+ resp = await target_peer.send_example(target_shard, step, target, length, request_id=request_id, train=train)
+ return resp
async def forward_prompt(
self,
diff --git a/setup.py b/setup.py
index ead654f1..6b59a9ac 100644
--- a/setup.py
+++ b/setup.py
@@ -35,8 +35,8 @@ extras_require = {
"yapf==0.40.2",
],
"apple_silicon": [
- "mlx",
- "mlx-lm",
+ "mlx==0.20.0",
+ "mlx-lm==0.19.3",
],
}
← 9eadee31 Basic model saving
·
back to Exo
·
Made models save properly 0d3abfca →