← back to Exo
Naive network-propagated loss implementation on MLX
75c8650f1f8d6a60f5847694d8bcd1a68aba900c · 2024-12-06 00:50:23 -0800 · Nel Nibcord
Files touched
M exo/inference/mlx/losses.pyM exo/inference/mlx/sharded_inference_engine.pyM exo/main.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.py
Diff
commit 75c8650f1f8d6a60f5847694d8bcd1a68aba900c
Author: Nel Nibcord <blindcrone@tuta.io>
Date: Fri Dec 6 00:50:23 2024 -0800
Naive network-propagated loss implementation on MLX
---
exo/inference/mlx/losses.py | 11 ++++
exo/inference/mlx/sharded_inference_engine.py | 26 ++++++----
exo/main.py | 4 +-
exo/networking/grpc/grpc_server.py | 15 ------
exo/networking/grpc/node_service.proto | 6 ++-
exo/networking/grpc/node_service_pb2.py | 72 ++++++++++++++-------------
exo/networking/grpc/node_service_pb2_grpc.py | 43 ----------------
exo/orchestration/standard_node.py | 24 ++++-----
8 files changed, 83 insertions(+), 118 deletions(-)
diff --git a/exo/inference/mlx/losses.py b/exo/inference/mlx/losses.py
index 72f94c41..258d09e3 100644
--- a/exo/inference/mlx/losses.py
+++ b/exo/inference/mlx/losses.py
@@ -12,3 +12,14 @@ def length_masked_ce_loss(model, inputs, targets, lengths):
loss = ce.sum() / length_mask.sum()
return loss
+#Naive intermediate layer loss, where we replace the targets with gradients and just multiply the output by the gradients to derive the loss. This is naive and may warrant some further iteration, but will do the job for now
+def back_gradient_loss(model, inputs, gradients, shard_proportion):
+ out = model(inputs)
+ logits = out[:, -1, :]
+ loss = (logits * gradients).mean()
+ return loss
+
+loss_fns = {
+ "back_gradient": back_gradient_loss,
+ "length_masked_ce": length_masked_ce_loss,
+}
diff --git a/exo/inference/mlx/sharded_inference_engine.py b/exo/inference/mlx/sharded_inference_engine.py
index 42c3db6a..00c07f9f 100644
--- a/exo/inference/mlx/sharded_inference_engine.py
+++ b/exo/inference/mlx/sharded_inference_engine.py
@@ -6,7 +6,7 @@ import mlx.optimizers as optim
from ..inference_engine import InferenceEngine
from .stateful_model import StatefulModel
from .sharded_utils import load_shard, get_image_from_str
-from .losses import length_masked_ce_loss
+from .losses import loss_fns
from ..shard import Shard
from typing import Dict, Optional, Tuple
from exo.download.shard_download import ShardDownloader
@@ -64,33 +64,39 @@ class MLXDynamicShardInferenceEngine(InferenceEngine):
#print(f"infer_tensor out -> {output_data}")
return output_data
- async def evaluate(self, request_id: str, shard: Shard, inputs, targets, lengths, loss=length_masked_ce_loss):
+ async def evaluate(self, request_id: str, shard: Shard, inputs, targets, lengths, loss: str = "length_masked_ce"):
await self.ensure_shard(shard)
- await self.ensure_session('loss', lambda: loss)
+ await self.ensure_session('loss', lambda: loss_fns[loss])
await self.ensure_session('task', lambda: ('eval', self.model.eval()))
#print(f"evaluate in <- {inputs}")
x = mx.array(inputs).astype(mx.int64) if self.shard.is_first_layer() else mx.array(inputs)
- y = mx.array(targets).astype(mx.int64)
+ y = mx.array(targets)
l = mx.array(lengths)
score = await asyncio.get_running_loop().run_in_executor(self.executor, self.session['loss'], self.model, x, y, l)
#print(f"evaluate out -> {score}")
return np.array(score)
+
+ async def update_model(self, grad, lval):
+ await self.ensure_shard(shard)
+ self.session['opt'].update(self.model, grad)
+ mx.eval(self.model.parameters(), self.session['opt'].state, lval)
- async def train(self, request_id: str, shard: Shard, inputs, targets, lengths, loss=length_masked_ce_loss, opt=optim.Adam, lr=1e-5):
+ async def train(self, request_id: str, shard: Shard, inputs, targets, lengths, loss: str = "length_masked_ce", opt=optim.Adam, lr=1e-5):
await self.ensure_shard(shard)
- await self.ensure_session('loss', lambda: loss)
+ await self.ensure_session('loss', lambda: loss_fns[loss])
await self.ensure_session('LVaG', lambda: nn.value_and_grad(self.model, self.session['loss']))
await self.ensure_session('opt', lambda: opt(lr))
await self.ensure_session('task', lambda: ('train', self.model.train()))
x = mx.array(inputs).astype(mx.int64) if self.shard.is_first_layer() else mx.array(inputs)
- y = mx.array(targets).astype(mx.int64)
+ y = mx.array(targets)
l = mx.array(lengths)
loop = asyncio.get_running_loop()
- loss, grad = await loop.run_in_executor(self.executor, self.session['LVaG'], self.model, x, y, l)
- await loop.run_in_executor(self.executor, lambda: self.session['opt'].update(self.model, grad))
+ score, grad = await loop.run_in_executor(self.executor, self.session['LVaG'], self.model, x, y, l)
+ loop.run_in_executor(self.executor, self.update_model, grad, score)
+ layers = [{k: v["weight"].shape for k,v in l.items() if 'weight' in v} for l in grad['model']['model']['layers'] if l]
- return np.array(loss), np.array(grad)
+ return np.array(score).reshape(inputs.shape[0], -1), np.array(layers[0]['input_layernorm']).reshape(inputs.shape[0], -1)
async def ensure_shard(self, shard: Shard):
if self.shard == shard:
diff --git a/exo/main.py b/exo/main.py
index eb6f9c6e..25ec0fe5 100644
--- a/exo/main.py
+++ b/exo/main.py
@@ -225,7 +225,7 @@ async def eval_model_cli(node: Node, inference_engine: InferenceEngine, model_na
tokenizer = await resolve_tokenizer(get_repo(shard.model_id, inference_class))
train, val, test = dataloader(lambda i: tokenizer.encode(i))
dataset = test
- print(f"Evaluating {len(dataset)} examples with batch_size {batch_size}")
+ print(f"Evaluating {len(test)} examples with batch_size {batch_size}")
losses = []
tokens = []
for batch in tqdm(iterate_batches(test, batch_size), total=len(dataset) // batch_size):
@@ -243,7 +243,7 @@ async def train_model_cli(node: Node, inference_engine: InferenceEngine, model_n
return
tokenizer = await resolve_tokenizer(get_repo(shard.model_id, inference_class))
train, val, test = dataloader(lambda i: tokenizer.encode(i))
- print(f"Training on {len(val)} examples with batch_size {batch_size}")
+ print(f"Training on {len(train)} examples with batch_size {batch_size}")
for epoch in range(iters):
losses = []
tokens = []
diff --git a/exo/networking/grpc/grpc_server.py b/exo/networking/grpc/grpc_server.py
index e99c3575..1fba89e9 100644
--- a/exo/networking/grpc/grpc_server.py
+++ b/exo/networking/grpc/grpc_server.py
@@ -87,21 +87,6 @@ class GRPCServer(node_service_pb2_grpc.NodeServiceServicer):
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))
-
- async def SendLoss(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,
- )
- loss = np.frombuffer(request.loss.tensor_data, dtype=np.dtype(request.loss.dtype)).reshape(request.loss.shape)
- request_id = request.request_id
-
- if shard.is_first_layer():
- asyncself.node.backward_loss(shard, loss, request_id)
- if DEBUG >= 5: print(f"SendTensor tensor {shard=} {example=} {target=} {length=} {request_id=} result: {result}")
- return node_service_pb2.Tensor(tensor_data=tensor_data, shape=result.shape, dtype=str(result.dtype)) if result is not None else node_service_pb2.Tensor()
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 57ab3515..1bbfc96f 100644
--- a/exo/networking/grpc/node_service.proto
+++ b/exo/networking/grpc/node_service.proto
@@ -5,7 +5,6 @@ package node_service;
service NodeService {
rpc SendPrompt (PromptRequest) returns (Tensor) {}
rpc SendTensor (TensorRequest) returns (Tensor) {}
- rpc SendLoss (TensorRequest) returns (Empty) {}
rpc SendExample (ExampleRequest) returns (Tensor) {}
rpc GetInferenceResult (GetInferenceResultRequest) returns (InferenceResult) {}
rpc CollectTopology (CollectTopologyRequest) returns (Topology) {}
@@ -41,6 +40,11 @@ message ExampleRequest {
bool train = 5;
optional string request_id = 6;
}
+
+message LossRequest {
+ Tensor loss = 1;
+ Tensor grads = 2;
+}
message GetInferenceResultRequest {
string request_id = 1;
diff --git a/exo/networking/grpc/node_service_pb2.py b/exo/networking/grpc/node_service_pb2.py
index 6e1f6d8b..46f3496a 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\"/\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\xb9\x05\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\x08SendLoss\x12\x1b.node_service.TensorRequest\x1a\x13.node_service.Empty\"\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\"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')
_globals = globals()
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
@@ -43,38 +43,40 @@ if not _descriptor._USE_C_DESCRIPTORS:
_globals['_TENSORREQUEST']._serialized_end=360
_globals['_EXAMPLEREQUEST']._serialized_start=363
_globals['_EXAMPLEREQUEST']._serialized_end=585
- _globals['_GETINFERENCERESULTREQUEST']._serialized_start=587
- _globals['_GETINFERENCERESULTREQUEST']._serialized_end=634
- _globals['_INFERENCERESULT']._serialized_start=636
- _globals['_INFERENCERESULT']._serialized_end=728
- _globals['_TENSOR']._serialized_start=730
- _globals['_TENSOR']._serialized_end=789
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_start=791
- _globals['_COLLECTTOPOLOGYREQUEST']._serialized_end=851
- _globals['_TOPOLOGY']._serialized_start=854
- _globals['_TOPOLOGY']._serialized_end=1134
- _globals['_TOPOLOGY_NODESENTRY']._serialized_start=975
- _globals['_TOPOLOGY_NODESENTRY']._serialized_end=1053
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_start=1055
- _globals['_TOPOLOGY_PEERGRAPHENTRY']._serialized_end=1134
- _globals['_PEERCONNECTION']._serialized_start=1136
- _globals['_PEERCONNECTION']._serialized_end=1209
- _globals['_PEERCONNECTIONS']._serialized_start=1211
- _globals['_PEERCONNECTIONS']._serialized_end=1279
- _globals['_DEVICEFLOPS']._serialized_start=1281
- _globals['_DEVICEFLOPS']._serialized_end=1336
- _globals['_DEVICECAPABILITIES']._serialized_start=1338
- _globals['_DEVICECAPABILITIES']._serialized_end=1445
- _globals['_SENDRESULTREQUEST']._serialized_start=1447
- _globals['_SENDRESULTREQUEST']._serialized_end=1523
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_start=1525
- _globals['_SENDOPAQUESTATUSREQUEST']._serialized_end=1586
- _globals['_HEALTHCHECKREQUEST']._serialized_start=1588
- _globals['_HEALTHCHECKREQUEST']._serialized_end=1608
- _globals['_HEALTHCHECKRESPONSE']._serialized_start=1610
- _globals['_HEALTHCHECKRESPONSE']._serialized_end=1651
- _globals['_EMPTY']._serialized_start=1653
- _globals['_EMPTY']._serialized_end=1660
- _globals['_NODESERVICE']._serialized_start=1663
- _globals['_NODESERVICE']._serialized_end=2360
+ _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
# @@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 ce21706c..71856d35 100644
--- a/exo/networking/grpc/node_service_pb2_grpc.py
+++ b/exo/networking/grpc/node_service_pb2_grpc.py
@@ -44,11 +44,6 @@ class NodeServiceStub(object):
request_serializer=node__service__pb2.TensorRequest.SerializeToString,
response_deserializer=node__service__pb2.Tensor.FromString,
_registered_method=True)
- self.SendLoss = channel.unary_unary(
- '/node_service.NodeService/SendLoss',
- request_serializer=node__service__pb2.TensorRequest.SerializeToString,
- response_deserializer=node__service__pb2.Empty.FromString,
- _registered_method=True)
self.SendExample = channel.unary_unary(
'/node_service.NodeService/SendExample',
request_serializer=node__service__pb2.ExampleRequest.SerializeToString,
@@ -96,12 +91,6 @@ class NodeServiceServicer(object):
context.set_details('Method not implemented!')
raise NotImplementedError('Method not implemented!')
- def SendLoss(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 SendExample(self, request, context):
"""Missing associated documentation comment in .proto file."""
context.set_code(grpc.StatusCode.UNIMPLEMENTED)
@@ -151,11 +140,6 @@ def add_NodeServiceServicer_to_server(servicer, server):
request_deserializer=node__service__pb2.TensorRequest.FromString,
response_serializer=node__service__pb2.Tensor.SerializeToString,
),
- 'SendLoss': grpc.unary_unary_rpc_method_handler(
- servicer.SendLoss,
- request_deserializer=node__service__pb2.TensorRequest.FromString,
- response_serializer=node__service__pb2.Empty.SerializeToString,
- ),
'SendExample': grpc.unary_unary_rpc_method_handler(
servicer.SendExample,
request_deserializer=node__service__pb2.ExampleRequest.FromString,
@@ -251,33 +235,6 @@ class NodeService(object):
metadata,
_registered_method=True)
- @staticmethod
- def SendLoss(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/SendLoss',
- node__service__pb2.TensorRequest.SerializeToString,
- node__service__pb2.Empty.FromString,
- options,
- channel_credentials,
- insecure,
- call_credentials,
- compression,
- wait_for_ready,
- timeout,
- metadata,
- _registered_method=True)
-
@staticmethod
def SendExample(request,
target,
diff --git a/exo/orchestration/standard_node.py b/exo/orchestration/standard_node.py
index fbea74da..0aa2ae97 100644
--- a/exo/orchestration/standard_node.py
+++ b/exo/orchestration/standard_node.py
@@ -269,24 +269,24 @@ class StandardNode(Node):
if request_id is None:
request_id = str(uuid.uuid4())
shard = self.get_current_shard(base_shard)
-
if DEBUG >= 1: print(f"[{request_id}] process_example: {example.shape=}")
try:
- if shard.is_last_layer():
- if train:
+ target = target.astype(int)
+ if train:
+ if shard.is_last_layer():
loss, grad = await self.inference_engine.train(request_id, shard, example, target, length)
- return loss.reshape(example.shape[0], -1) if shard.is_first_layer() else grad
else:
- loss = await self.inference_engine.evaluate(request_id, shard, example, target, length)
- return loss.reshape(example.shape[0], -1)
+ step = await self.inference_engine.infer_tensor(request_id, shard, example)
+ backgrad = await self.forward_example(shard, step, target, length, train, request_id, self.get_partition_index(offset = 1))
+ loss, grad = await self.inference_engine.train(request_id, shard, example, backgrad, length, loss="back_gradient")
+ return loss.reshape(example.shape[0], -1) if shard.is_first_layer() else grad
else:
- step = await self.inference_engine.infer_tensor(request_id, shard, example)
- result = await self.forward_example(shard, step, target, length, train, request_id, self.get_partition_index(offset = 1))
- if train:
- forward = self.get_current_shard(self.get_partition_index(offset = 1))
- return result
+ if shard.is_last_layer():
+ loss = await self.inference_engine.evaluate(request_id, shard, example, target, length)
else:
- return result.reshape(example.shape[0], -1)
+ step = await self.inference_engine.infer_tensor(request_id, shard, example)
+ loss = await self.forward_example(shard, step, target, length, train, request_id, self.get_partition_index(offset = 1))
+ return loss.reshape(example.shape[0], -1)
except Exception as e:
print(f"Error processing example for shard {shard}: {e}")
traceback.print_exc()
← 83685682 WIP: Training works on mlx
·
back to Exo
·
Okay we should probably await the update 3e869051 →