Skip to content
Draft

grpc #56

Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions cezo_grpc/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
run this at root `python -m grpc_tools.protoc -Icezo_grpc=./cezo_grpc --python_out=. --grpc_python_out=. ./cezo_grpc/sample.proto`

exaple to run large llm `python run_grpc.py --dataset=sst2 --eval-iterations=25 --large-model=opt-125m --model-dtype=float16 --seed=365 --iterations=2000 --train-batch-size=32 --test-batch-size=64 --num-clients=3 --num-sample-clients=2 --local-update-steps=1 --num-pert=5 --lr=0.000005 --momentum=0 --grad-estimate-method=rge-forward --mu=0.001`
Empty file added cezo_grpc/__init__.py
Empty file.
35 changes: 35 additions & 0 deletions cezo_grpc/cli_interface.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
from experiment_helper.cli_parser import (
GeneralSetting,
DeviceSetting,
DataSetting,
ModelSetting,
OptimizerSetting,
FederatedLearningSetting,
RGESetting,
ByzantineSetting,
GRPCSetting,
)


class CliSetting(
GeneralSetting,
DeviceSetting,
DataSetting,
ModelSetting,
OptimizerSetting,
FederatedLearningSetting,
RGESetting,
ByzantineSetting,
GRPCSetting,
):
"""
This is a replacement for regular argparse module.
We used a third party library pydantic_setting to make command line interface easier to manage.
Example:
if __name__ == "__main__":
args = CliSetting()

args will have all parameters defined by all components.
"""

pass
51 changes: 51 additions & 0 deletions cezo_grpc/data_helper.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
from cezo_grpc import sample_pb2


def protobuf_to_py_list_of_ints(message: sample_pb2.ListOfInts) -> list[int]:
return [v for v in message.data]


def protobuf_to_py_list_of_list_of_ints(message: sample_pb2.ListOfListOfInts) -> list[list[int]]:
return [protobuf_to_py_list_of_ints(v) for v in message.data]


def protobuf_to_py_list_of_floats(message: sample_pb2.ListOfFloats) -> list[float]:
return [v for v in message.data]


def protobuf_to_py_list_of_list_of_floats(
message: sample_pb2.ListOfListOfFloats,
) -> list[list[float]]:
return [protobuf_to_py_list_of_floats(v) for v in message.data]


def protobuf_to_py_list_of_list_of_list_of_floats(
message: sample_pb2.ListOfListOfListOfFloats,
) -> list[list[list[float]]]:
return [protobuf_to_py_list_of_list_of_floats(v) for v in message.data]


def py_to_protobuf_list_of_ints(py_data: list[int]) -> sample_pb2.ListOfInts:
return sample_pb2.ListOfInts(data=py_data)


def py_to_protobuf_list_of_list_of_ints(py_data: list[list[int]]) -> sample_pb2.ListOfListOfInts:
return sample_pb2.ListOfListOfInts(data=[py_to_protobuf_list_of_ints(v) for v in py_data])


def py_to_protobuf_list_of_floats(py_data: list[float]) -> sample_pb2.ListOfFloats:
return sample_pb2.ListOfFloats(data=py_data)


def py_to_protobuf_list_of_list_of_floats(
py_data: list[list[float]],
) -> sample_pb2.ListOfListOfFloats:
return sample_pb2.ListOfListOfFloats(data=[py_to_protobuf_list_of_floats(v) for v in py_data])


def py_to_protobuf_list_of_list_of_list_of_floats(
py_data: list[list[list[float]]],
) -> sample_pb2.ListOfListOfListOfFloats:
return sample_pb2.ListOfListOfListOfFloats(
data=[py_to_protobuf_list_of_list_of_floats(v) for v in py_data]
)
71 changes: 71 additions & 0 deletions cezo_grpc/sample.proto
Original file line number Diff line number Diff line change
@@ -0,0 +1,71 @@
syntax = "proto3";

package sample_server;

service SampleServer {
rpc Connect (EmptyRequest) returns (ConnectResponse) {}
rpc Disconnect (DisconnectRequest) returns (EmptyResponse) {}
rpc TryToJoinIteration (TryToJoinIterationRequest) returns (TryToJoinIterationResponse) {}
rpc SubmitIteration (SubmitIterationRequest) returns (EmptyResponse) {}

rpc ConnectEval (EmptyRequest) returns (ConnectResponse) {}
rpc DisconnectEval (EmptyRequest) returns (EmptyResponse) {}
rpc TryToEval (EmptyRequest) returns (TryToJoinIterationResponse) {}
rpc SubmitEvaluation (SubmitEvaluationRequest) returns (EmptyResponse) {}
}

message EmptyRequest {}

message EmptyResponse {}

message ConnectResponse {
bool successful = 1;
int32 clientIndex = 2;
}

message DisconnectRequest {
int32 clientIndex = 1;
}

message TryToJoinIterationRequest {
int32 clientIndex = 1;
}

message TryToJoinIterationResponse{
bool successful = 1;
ListOfListOfInts pullSeeds = 2;
ListOfListOfListOfFloats pullGrads = 3;
ListOfInts iterationSeeds = 4;
}

message SubmitIterationRequest {
int32 clientIndex = 1;
ListOfListOfFloats gradTensors = 2;
float stepAccuracy = 3;
float stepLoss = 4;
}

message SubmitEvaluationRequest {
float evalAccuracy = 1;
float evalLoss = 2;
}

message ListOfInts {
repeated int32 data = 1;
}

message ListOfListOfInts {
repeated ListOfInts data = 1;
}

message ListOfFloats {
repeated float data = 1;
}

message ListOfListOfFloats {
repeated ListOfFloats data = 1;
}

message ListOfListOfListOfFloats {
repeated ListOfListOfFloats data = 1;
}
52 changes: 52 additions & 0 deletions cezo_grpc/sample_pb2.py

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading
Loading