[object Object]

← 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

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 →