diff --git a/cezo_grpc/README.md b/cezo_grpc/README.md new file mode 100644 index 0000000..81821d9 --- /dev/null +++ b/cezo_grpc/README.md @@ -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` diff --git a/cezo_grpc/__init__.py b/cezo_grpc/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/cezo_grpc/cli_interface.py b/cezo_grpc/cli_interface.py new file mode 100644 index 0000000..1ce8f68 --- /dev/null +++ b/cezo_grpc/cli_interface.py @@ -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 diff --git a/cezo_grpc/data_helper.py b/cezo_grpc/data_helper.py new file mode 100644 index 0000000..97e55cc --- /dev/null +++ b/cezo_grpc/data_helper.py @@ -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] + ) diff --git a/cezo_grpc/sample.proto b/cezo_grpc/sample.proto new file mode 100644 index 0000000..3c1cb68 --- /dev/null +++ b/cezo_grpc/sample.proto @@ -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; +} diff --git a/cezo_grpc/sample_pb2.py b/cezo_grpc/sample_pb2.py new file mode 100644 index 0000000..00f098f --- /dev/null +++ b/cezo_grpc/sample_pb2.py @@ -0,0 +1,52 @@ +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# source: cezo_grpc/sample.proto +# Protobuf Python Version: 4.25.1 +"""Generated protocol buffer code.""" +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + + + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x16\x63\x65zo_grpc/sample.proto\x12\rsample_server\"\x0e\n\x0c\x45mptyRequest\"\x0f\n\rEmptyResponse\":\n\x0f\x43onnectResponse\x12\x12\n\nsuccessful\x18\x01 \x01(\x08\x12\x13\n\x0b\x63lientIndex\x18\x02 \x01(\x05\"(\n\x11\x44isconnectRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"0\n\x19TryToJoinIterationRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"\xd3\x01\n\x1aTryToJoinIterationResponse\x12\x12\n\nsuccessful\x18\x01 \x01(\x08\x12\x32\n\tpullSeeds\x18\x02 \x01(\x0b\x32\x1f.sample_server.ListOfListOfInts\x12:\n\tpullGrads\x18\x03 \x01(\x0b\x32\'.sample_server.ListOfListOfListOfFloats\x12\x31\n\x0eiterationSeeds\x18\x04 \x01(\x0b\x32\x19.sample_server.ListOfInts\"\x8d\x01\n\x16SubmitIterationRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\x12\x36\n\x0bgradTensors\x18\x02 \x01(\x0b\x32!.sample_server.ListOfListOfFloats\x12\x14\n\x0cstepAccuracy\x18\x03 \x01(\x02\x12\x10\n\x08stepLoss\x18\x04 \x01(\x02\"A\n\x17SubmitEvaluationRequest\x12\x14\n\x0c\x65valAccuracy\x18\x01 \x01(\x02\x12\x10\n\x08\x65valLoss\x18\x02 \x01(\x02\"\x1a\n\nListOfInts\x12\x0c\n\x04\x64\x61ta\x18\x01 \x03(\x05\";\n\x10ListOfListOfInts\x12\'\n\x04\x64\x61ta\x18\x01 \x03(\x0b\x32\x19.sample_server.ListOfInts\"\x1c\n\x0cListOfFloats\x12\x0c\n\x04\x64\x61ta\x18\x01 \x03(\x02\"?\n\x12ListOfListOfFloats\x12)\n\x04\x64\x61ta\x18\x01 \x03(\x0b\x32\x1b.sample_server.ListOfFloats\"K\n\x18ListOfListOfListOfFloats\x12/\n\x04\x64\x61ta\x18\x01 \x03(\x0b\x32!.sample_server.ListOfListOfFloats2\xbf\x05\n\x0cSampleServer\x12H\n\x07\x43onnect\x12\x1b.sample_server.EmptyRequest\x1a\x1e.sample_server.ConnectResponse\"\x00\x12N\n\nDisconnect\x12 .sample_server.DisconnectRequest\x1a\x1c.sample_server.EmptyResponse\"\x00\x12k\n\x12TryToJoinIteration\x12(.sample_server.TryToJoinIterationRequest\x1a).sample_server.TryToJoinIterationResponse\"\x00\x12X\n\x0fSubmitIteration\x12%.sample_server.SubmitIterationRequest\x1a\x1c.sample_server.EmptyResponse\"\x00\x12L\n\x0b\x43onnectEval\x12\x1b.sample_server.EmptyRequest\x1a\x1e.sample_server.ConnectResponse\"\x00\x12M\n\x0e\x44isconnectEval\x12\x1b.sample_server.EmptyRequest\x1a\x1c.sample_server.EmptyResponse\"\x00\x12U\n\tTryToEval\x12\x1b.sample_server.EmptyRequest\x1a).sample_server.TryToJoinIterationResponse\"\x00\x12Z\n\x10SubmitEvaluation\x12&.sample_server.SubmitEvaluationRequest\x1a\x1c.sample_server.EmptyResponse\"\x00\x62\x06proto3') + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'cezo_grpc.sample_pb2', _globals) +if _descriptor._USE_C_DESCRIPTORS == False: + DESCRIPTOR._options = None + _globals['_EMPTYREQUEST']._serialized_start=41 + _globals['_EMPTYREQUEST']._serialized_end=55 + _globals['_EMPTYRESPONSE']._serialized_start=57 + _globals['_EMPTYRESPONSE']._serialized_end=72 + _globals['_CONNECTRESPONSE']._serialized_start=74 + _globals['_CONNECTRESPONSE']._serialized_end=132 + _globals['_DISCONNECTREQUEST']._serialized_start=134 + _globals['_DISCONNECTREQUEST']._serialized_end=174 + _globals['_TRYTOJOINITERATIONREQUEST']._serialized_start=176 + _globals['_TRYTOJOINITERATIONREQUEST']._serialized_end=224 + _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_start=227 + _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_end=438 + _globals['_SUBMITITERATIONREQUEST']._serialized_start=441 + _globals['_SUBMITITERATIONREQUEST']._serialized_end=582 + _globals['_SUBMITEVALUATIONREQUEST']._serialized_start=584 + _globals['_SUBMITEVALUATIONREQUEST']._serialized_end=649 + _globals['_LISTOFINTS']._serialized_start=651 + _globals['_LISTOFINTS']._serialized_end=677 + _globals['_LISTOFLISTOFINTS']._serialized_start=679 + _globals['_LISTOFLISTOFINTS']._serialized_end=738 + _globals['_LISTOFFLOATS']._serialized_start=740 + _globals['_LISTOFFLOATS']._serialized_end=768 + _globals['_LISTOFLISTOFFLOATS']._serialized_start=770 + _globals['_LISTOFLISTOFFLOATS']._serialized_end=833 + _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_start=835 + _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_end=910 + _globals['_SAMPLESERVER']._serialized_start=913 + _globals['_SAMPLESERVER']._serialized_end=1616 +# @@protoc_insertion_point(module_scope) diff --git a/cezo_grpc/sample_pb2_grpc.py b/cezo_grpc/sample_pb2_grpc.py new file mode 100644 index 0000000..9f40cb0 --- /dev/null +++ b/cezo_grpc/sample_pb2_grpc.py @@ -0,0 +1,297 @@ +# Generated by the gRPC Python protocol compiler plugin. DO NOT EDIT! +"""Client and server classes corresponding to protobuf-defined services.""" +import grpc + +from cezo_grpc import sample_pb2 as cezo__grpc_dot_sample__pb2 + + +class SampleServerStub(object): + """Missing associated documentation comment in .proto file.""" + + def __init__(self, channel): + """Constructor. + + Args: + channel: A grpc.Channel. + """ + self.Connect = channel.unary_unary( + '/sample_server.SampleServer/Connect', + request_serializer=cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.ConnectResponse.FromString, + ) + self.Disconnect = channel.unary_unary( + '/sample_server.SampleServer/Disconnect', + request_serializer=cezo__grpc_dot_sample__pb2.DisconnectRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, + ) + self.TryToJoinIteration = channel.unary_unary( + '/sample_server.SampleServer/TryToJoinIteration', + request_serializer=cezo__grpc_dot_sample__pb2.TryToJoinIterationRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.TryToJoinIterationResponse.FromString, + ) + self.SubmitIteration = channel.unary_unary( + '/sample_server.SampleServer/SubmitIteration', + request_serializer=cezo__grpc_dot_sample__pb2.SubmitIterationRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, + ) + self.ConnectEval = channel.unary_unary( + '/sample_server.SampleServer/ConnectEval', + request_serializer=cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.ConnectResponse.FromString, + ) + self.DisconnectEval = channel.unary_unary( + '/sample_server.SampleServer/DisconnectEval', + request_serializer=cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, + ) + self.TryToEval = channel.unary_unary( + '/sample_server.SampleServer/TryToEval', + request_serializer=cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.TryToJoinIterationResponse.FromString, + ) + self.SubmitEvaluation = channel.unary_unary( + '/sample_server.SampleServer/SubmitEvaluation', + request_serializer=cezo__grpc_dot_sample__pb2.SubmitEvaluationRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, + ) + + +class SampleServerServicer(object): + """Missing associated documentation comment in .proto file.""" + + def Connect(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 Disconnect(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 TryToJoinIteration(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 SubmitIteration(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 ConnectEval(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 DisconnectEval(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 TryToEval(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 SubmitEvaluation(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 add_SampleServerServicer_to_server(servicer, server): + rpc_method_handlers = { + 'Connect': grpc.unary_unary_rpc_method_handler( + servicer.Connect, + request_deserializer=cezo__grpc_dot_sample__pb2.EmptyRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.ConnectResponse.SerializeToString, + ), + 'Disconnect': grpc.unary_unary_rpc_method_handler( + servicer.Disconnect, + request_deserializer=cezo__grpc_dot_sample__pb2.DisconnectRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.EmptyResponse.SerializeToString, + ), + 'TryToJoinIteration': grpc.unary_unary_rpc_method_handler( + servicer.TryToJoinIteration, + request_deserializer=cezo__grpc_dot_sample__pb2.TryToJoinIterationRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.TryToJoinIterationResponse.SerializeToString, + ), + 'SubmitIteration': grpc.unary_unary_rpc_method_handler( + servicer.SubmitIteration, + request_deserializer=cezo__grpc_dot_sample__pb2.SubmitIterationRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.EmptyResponse.SerializeToString, + ), + 'ConnectEval': grpc.unary_unary_rpc_method_handler( + servicer.ConnectEval, + request_deserializer=cezo__grpc_dot_sample__pb2.EmptyRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.ConnectResponse.SerializeToString, + ), + 'DisconnectEval': grpc.unary_unary_rpc_method_handler( + servicer.DisconnectEval, + request_deserializer=cezo__grpc_dot_sample__pb2.EmptyRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.EmptyResponse.SerializeToString, + ), + 'TryToEval': grpc.unary_unary_rpc_method_handler( + servicer.TryToEval, + request_deserializer=cezo__grpc_dot_sample__pb2.EmptyRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.TryToJoinIterationResponse.SerializeToString, + ), + 'SubmitEvaluation': grpc.unary_unary_rpc_method_handler( + servicer.SubmitEvaluation, + request_deserializer=cezo__grpc_dot_sample__pb2.SubmitEvaluationRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.EmptyResponse.SerializeToString, + ), + } + generic_handler = grpc.method_handlers_generic_handler( + 'sample_server.SampleServer', rpc_method_handlers) + server.add_generic_rpc_handlers((generic_handler,)) + + + # This class is part of an EXPERIMENTAL API. +class SampleServer(object): + """Missing associated documentation comment in .proto file.""" + + @staticmethod + def Connect(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, '/sample_server.SampleServer/Connect', + cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.ConnectResponse.FromString, + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) + + @staticmethod + def Disconnect(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, '/sample_server.SampleServer/Disconnect', + cezo__grpc_dot_sample__pb2.DisconnectRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) + + @staticmethod + def TryToJoinIteration(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, '/sample_server.SampleServer/TryToJoinIteration', + cezo__grpc_dot_sample__pb2.TryToJoinIterationRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.TryToJoinIterationResponse.FromString, + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) + + @staticmethod + def SubmitIteration(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, '/sample_server.SampleServer/SubmitIteration', + cezo__grpc_dot_sample__pb2.SubmitIterationRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) + + @staticmethod + def ConnectEval(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, '/sample_server.SampleServer/ConnectEval', + cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.ConnectResponse.FromString, + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) + + @staticmethod + def DisconnectEval(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, '/sample_server.SampleServer/DisconnectEval', + cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) + + @staticmethod + def TryToEval(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, '/sample_server.SampleServer/TryToEval', + cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.TryToJoinIterationResponse.FromString, + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) + + @staticmethod + def SubmitEvaluation(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, '/sample_server.SampleServer/SubmitEvaluation', + cezo__grpc_dot_sample__pb2.SubmitEvaluationRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) diff --git a/environment.yml b/environment.yml index 376c05a..8116d24 100644 --- a/environment.yml +++ b/environment.yml @@ -16,6 +16,9 @@ dependencies: - pytorch::pytorch=2.3 - pytorch::torchvision=0.18 + - grpcio-tools=1.62.2 + - conda-forge::grpcio=1.62.2 + - numpy=1.23 - tensorboardx=2.6 - tqdm=4.66 diff --git a/environment_cuda.yml b/environment_cuda.yml index 808fa7e..a6dc561 100644 --- a/environment_cuda.yml +++ b/environment_cuda.yml @@ -18,6 +18,9 @@ dependencies: - pytorch::pytorch=2.3 - pytorch::torchvision=0.18 + - grpcio-tools=1.62.2 + - conda-forge::grpcio=1.62.2 + - numpy=1.23 - tensorboardx=2.6 - tqdm=4.66 diff --git a/experiment_helper/cli_parser.py b/experiment_helper/cli_parser.py index 0b21834..964b9e9 100644 --- a/experiment_helper/cli_parser.py +++ b/experiment_helper/cli_parser.py @@ -197,3 +197,22 @@ class FOFLSetting(FrozenSetting): @cached_property def fo_fl_setting(self) -> "FOFLSetting": return FOFLSetting() + + +class GRPCSetting(FrozenSetting): + # gRPC + rpc_master_addr: str = Field( + default="localhost", + validation_alias=AliasChoices("rpc-master-addr"), + description="Address of the RPC master node (the parameter server).", + ) + rpc_master_port: int = Field( + default=4242, + validation_alias=AliasChoices("rpc-master-port"), + description="Port of the RPC master node (the parameter server).", + ) + rpc_num_workers: int = Field(default=8, validation_alias=AliasChoices("rpc-num-workers")) + + @cached_property + def grpc_setting(self) -> "GRPCSetting": + return GRPCSetting() diff --git a/grpc_client.py b/grpc_client.py new file mode 100644 index 0000000..ae796a0 --- /dev/null +++ b/grpc_client.py @@ -0,0 +1,129 @@ +from huggingface_hub.repository import atexit +import torch +import grpc +import time + +from cezo_grpc import sample_pb2 +from cezo_grpc import sample_pb2_grpc +from cezo_grpc import data_helper +from cezo_grpc import cli_interface + +from cezo_fl import fl_helpers +from cezo_fl import client + +from experiment_helper import prepare_settings +from experiment_helper.device import use_device +from experiment_helper.data import get_dataloaders + + +def setup_client(args: cli_interface.CliSetting, client_index: int): + device_map = use_device(args.device_setting, args.num_clients) + train_loaders, _ = get_dataloaders( + args.data_setting, args.num_clients, args.seed, args.get_hf_model_name() + ) + model_inferences, metrics = prepare_settings.get_model_inferences_and_metrics( + args.dataset, args.model_setting + ) + client_name = fl_helpers.get_client_name(client_index) + client_device = device_map[client_name] + client_model = prepare_settings.get_model( + dataset=args.dataset, model_setting=args.model_setting, seed=args.seed + ).to(client_device) + client_optimizer = prepare_settings.get_optimizer( + model=client_model, dataset=args.dataset, optimizer_setting=args.optimizer_setting + ) + client_grad_estimator = prepare_settings.get_gradient_estimator( + model=client_model, + device=client_device, + rge_setting=args.rge_setting, + model_setting=args.model_setting, + ) + + return client.ResetClient( + client_model, + model_inferences.train_inference, + train_loaders[client_index], + client_grad_estimator, + client_optimizer, + metrics.train_loss, + metrics.train_acc, + client_device, + ) + + +def repeat_every(fn, pass_fn, repeat_interval=1): + while True: + response = fn() + if pass_fn(response): + return response + time.sleep(repeat_interval) + + +def get_stub(): + rpc_master_addr = "localhost" + rpc_master_port = 4242 + channel = grpc.insecure_channel(f"{rpc_master_addr}:{rpc_master_port}") + ps_stub = sample_pb2_grpc.SampleServerStub(channel) + return ps_stub + + +def train_with_args(args: cli_interface.CliSetting): + ps_stub = get_stub() + + connect_result = repeat_every( + lambda: ps_stub.Connect(sample_pb2.EmptyRequest()), # type: ignore[attr-defined] + lambda x: x.successful, + ) + client_index: int = connect_result.clientIndex + print(f"connected as client: {client_index}") + # when program exits, we need to disconnect this client from server + atexit.register( + lambda: ps_stub.Disconnect(sample_pb2.DisconnectRequest(clientIndex=client_index)) # type: ignore[attr-defined] + ) + + with torch.no_grad(): + client_instance = setup_client(args, client_index) + + def try_to_join_iteration(): + join_result = repeat_every( + lambda: ps_stub.TryToJoinIteration( + sample_pb2.TryToJoinIterationRequest(clientIndex=client_index) # type: ignore[attr-defined] + ), + lambda x: x.successful, + ) + + print("join iteration") + pull_seeds_list = data_helper.protobuf_to_py_list_of_list_of_ints(join_result.pullSeeds) + raw_grad_list = data_helper.protobuf_to_py_list_of_list_of_list_of_floats( + join_result.pullGrads + ) + tensor_grad_list = [ + [torch.tensor(v, device=client_instance.device) for v in vv] for vv in raw_grad_list + ] + iteration_seeds = data_helper.protobuf_to_py_list_of_ints(join_result.iterationSeeds) + + # step 2: client pull to update its model to latest + client_instance.pull_model(pull_seeds_list, tensor_grad_list) + + # step 3: client local update and get its result + client_local_update_result = client_instance.local_update(seeds=iteration_seeds) + + print("submit result") + ps_stub.SubmitIteration( + sample_pb2.SubmitIterationRequest( # type: ignore[attr-defined] + clientIndex=client_index, + gradTensors=data_helper.py_to_protobuf_list_of_list_of_floats( + [t.tolist() for t in client_local_update_result.grad_tensors] + ), + stepAccuracy=client_local_update_result.step_accuracy, + stepLoss=client_local_update_result.step_loss, + ) + ) + + repeat_every(try_to_join_iteration, lambda x: False) + + +if __name__ == "__main__": + args = cli_interface.CliSetting() + print(args) + train_with_args(args) diff --git a/grpc_eval_client.py b/grpc_eval_client.py new file mode 100644 index 0000000..7ac5d87 --- /dev/null +++ b/grpc_eval_client.py @@ -0,0 +1,122 @@ +from huggingface_hub.repository import atexit +import torch + +from cezo_grpc import sample_pb2 +from cezo_grpc import data_helper + +import grpc_client + +from cezo_fl.util.metrics import Metric + +from cezo_grpc import cli_interface + +from experiment_helper import prepare_settings +from experiment_helper.device import use_device +from experiment_helper.data import get_dataloaders + + +def setup_eval_model(args: cli_interface.CliSetting): + device_map = use_device(args.device_setting, args.num_clients) + device_name = "server" + device = device_map[device_name] + + _, test_loader = get_dataloaders( + args.data_setting, args.num_clients, args.seed, args.get_hf_model_name() + ) + model_inferences, metrics = prepare_settings.get_model_inferences_and_metrics( + args.dataset, args.model_setting + ) + server_model_inference = model_inferences.test_inference + server_criterion = metrics.test_loss + server_accuracy_func = metrics.test_acc + + model = prepare_settings.get_model( + dataset=args.dataset, model_setting=args.model_setting, seed=args.seed + ).to(device) + optimizer = prepare_settings.get_optimizer( + model=model, dataset=args.dataset, optimizer_setting=args.optimizer_setting + ) + grad_estimator = prepare_settings.get_gradient_estimator( + model=model, + device=device, + rge_setting=args.rge_setting, + model_setting=args.model_setting, + ) + + def update_model(seeds_list, grad_scalar_list): + model.train() + for iteration_seeds, iteration_grad_sclar in zip(seeds_list, grad_scalar_list): + grad_estimator.update_model_given_seed_and_grad( + optimizer, + iteration_seeds, + iteration_grad_sclar, + ) + + def eval_model(): + model.eval() + eval_loss = Metric("Eval loss") + eval_accuracy = Metric("Eval accuracy") + with torch.no_grad(): + for _, (batch_inputs, batch_labels) in enumerate(test_loader): + if device != torch.device("cpu") or grad_estimator.torch_dtype != torch.float32: + batch_inputs = batch_inputs.to(device, grad_estimator.torch_dtype) + # NOTE: label does not convert to dtype + if isinstance(batch_labels, torch.Tensor): + batch_labels = batch_labels.to(device) + pred = server_model_inference(model, batch_inputs) + eval_loss.update(server_criterion(pred, batch_labels)) + eval_accuracy.update(server_accuracy_func(pred, batch_labels)) + print( + f"Eval Loss:{eval_loss.avg:.4f}, " f"Accuracy:{eval_accuracy.avg * 100:.2f}%", + ) + return eval_loss.avg, eval_accuracy.avg + + return update_model, eval_model, device + + +def eval_with_args(args): + ps_stub = grpc_client.get_stub() + + grpc_client.repeat_every( + lambda: ps_stub.ConnectEval(sample_pb2.EmptyRequest()), lambda x: x.successful + ) + + print("Connected as Evaluation Client") + # when program exits, we need to disconnect this client from server + atexit.register(lambda: ps_stub.DisconnectEval(sample_pb2.EmptyRequest())) + + with torch.no_grad(): + update_model, eval_model, device = setup_eval_model(args) + + def try_to_eval(): + try_to_eval_result = grpc_client.repeat_every( + lambda: ps_stub.TryToEval(sample_pb2.EmptyRequest()), + lambda x: x.successful, + ) + + print("Start update eval model") + pull_seeds_list = data_helper.protobuf_to_py_list_of_list_of_ints( + try_to_eval_result.pullSeeds + ) + raw_grad_list = data_helper.protobuf_to_py_list_of_list_of_list_of_floats( + try_to_eval_result.pullGrads + ) + tensor_grad_list = [ + [torch.tensor(v, device=device) for v in vv] for vv in raw_grad_list + ] + + update_model(pull_seeds_list, tensor_grad_list) + eval_loss, eval_accuracy = eval_model() + + print("Submit Eval Result") + ps_stub.SubmitEvaluation( + sample_pb2.SubmitEvaluationRequest(evalAccuracy=eval_accuracy, evalLoss=eval_loss) + ) + + grpc_client.repeat_every(try_to_eval, lambda x: False) + + +if __name__ == "__main__": + args = cli_interface.CliSetting() + print(args) + eval_with_args(args) diff --git a/grpc_server.py b/grpc_server.py new file mode 100644 index 0000000..2933e55 --- /dev/null +++ b/grpc_server.py @@ -0,0 +1,337 @@ +import grpc +import threading +from concurrent import futures +import random +from enum import Enum +import torch + +from cezo_grpc import data_helper +from cezo_grpc import sample_pb2_grpc +from cezo_grpc import sample_pb2 +from cezo_grpc import cli_interface + +from cezo_fl import server +from byzantine.aggregation import mean +from byzantine.attack import no_byz + + +class ServerStatus(Enum): + connecting = "connecting" + training = "training" + aggregating = "aggregating" + evaluating = "evaluating" + # TODO: implement finishing logic when training iteration reaches the target + finishing = "finishing" + + +def find_first(lst: list, check_fn) -> int: + try: + return next(i for i, v in enumerate(lst) if check_fn(v)) + except StopIteration: + return -1 + + +class SampleServer(sample_pb2_grpc.SampleServerServicer): + def __init__(self, num_clients: int, num_sample_clients: int, local_update_steps: int): + self.should_eval = True + self.eval_iteration = 25 + + self.num_clients = num_clients + self.num_sample_clients = num_sample_clients + self.local_update_steps = local_update_steps + + self.seed_grad_records = server.SeedAndGradientRecords() + self.client_last_updates = [0 for _ in range(self.num_clients)] + self.status = ServerStatus.connecting + + self.lock = threading.Lock() + + self.connected_clients: list[bool] = [False for _ in range(self.num_clients)] + + self.eval_client_connected: bool = False + self.eval_client_last_update: int = 0 + + self.iteration: int = -1 + self.iteration_seeds: list[int] = [] + self.iteration_sampled_clients: list[int] = [] + self.iteration_finished_clients: set[int] = set() + self.iteration_local_grad_scalar: dict[int, list[torch.Tensor]] = {} + + def _get_connect_status(self): + if self.should_eval: + return f"train clients: {self.connected_clients}, eval client: {self.eval_client_connected}" + else: + return f"train clients: {self.connected_clients}" + + def _get_next_connect_client_index(self) -> int: + return find_first(self.connected_clients, lambda x: not x) + + def change_status(self, new_status: ServerStatus) -> None: + self.status = new_status + + def preprare_for_next_iteration(self): + print(f"finish iteration {self.iteration}, starting next iteration") + self.iteration += 1 + self.iteration_seeds = [random.randint(0, 1000000) for _ in range(self.local_update_steps)] + self.iteration_sampled_clients = random.sample( + range(self.num_clients), self.num_sample_clients + ) + self.iteration_finished_clients = set() + self.iteration_local_grad_scalar = {} + print(f"Iteration: {self.iteration}, sampled_clients: {self.iteration_sampled_clients}") + + def _should_connect(self) -> bool: + training_not_all_connected = not all(self.connected_clients) + should_eval_and_eval_not_connected = self.should_eval and not self.eval_client_connected + return training_not_all_connected or should_eval_and_eval_not_connected + + def try_swtich_from_connecting_to_training(self) -> None: + # 1. check for current status == connecting + if self.status is not ServerStatus.connecting: + return + # 2. check for connected_clients vs num_clients + if self._should_connect(): + return + # 3. initialize training data for next iteration + print("Switch from connecting to training") + self.preprare_for_next_iteration() + # 4. change status to training + self.change_status(ServerStatus.training) + + def swtich_to_connecting(self) -> None: + if not self._should_connect(): + return + print("Switch to Connecting") + self.change_status(ServerStatus.connecting) + + def _has_iteration_finished(self) -> bool: + sorted_sampled_clients = sorted(self.iteration_sampled_clients) + sorted_finished_clients = sorted(self.iteration_finished_clients) + return len(sorted_sampled_clients) == len(sorted_finished_clients) and all( + [ + sorted_sampled_clients[i] == sorted_finished_clients[i] + for i in range(len(sorted_sampled_clients)) + ] + ) + + def _get_iteration_grad_scalar_list(self) -> list[list[torch.Tensor]]: + return [ + self.iteration_local_grad_scalar[client_index] + for client_index in self.iteration_sampled_clients + ] + + def _aggregate_and_update_server_record(self) -> None: + local_grad_scalar_list = no_byz(self._get_iteration_grad_scalar_list()) + grad_scalar = mean(local_grad_scalar_list) + + self.seed_grad_records.add_records(seeds=self.iteration_seeds, grad=grad_scalar) + # Optional: optimize the memory. Remove is exclusive, i.e., the min last updates + # information is still kept. + if self.should_eval: + last_update_iterations = self.client_last_updates + [self.eval_client_last_update] + else: + last_update_iterations = self.client_last_updates + + self.seed_grad_records.remove_too_old(earliest_record_needs=min(last_update_iterations)) + + def try_switch_from_training_to_aggregating(self) -> None: + # 1. check for current status == training + if self.status is not ServerStatus.training: + return + # 2. check if sampled client all have return result + if not self._has_iteration_finished(): + return + print("Switch from training to aggregating") + # 3. change status to aggregating + self.change_status(ServerStatus.aggregating) + # 4. update seed_grad_records + self._aggregate_and_update_server_record() + # 5. change to next training or evaluating + if self.should_eval and self.iteration % self.eval_iteration == 0: + self.switch_from_aggregating_to_evaluating() + else: + self.switch_from_aggregating_to_training() + + def switch_from_aggregating_to_training(self) -> None: + print("Switch from aggregating to training") + self.preprare_for_next_iteration() + self.change_status(ServerStatus.training) + + def switch_from_aggregating_to_evaluating(self) -> None: + print("Switch from aggregating to evaluating") + self.change_status(ServerStatus.evaluating) + + def switch_from_evaluating_to_training(self) -> None: + print("Switch from evaluating to training") + self.preprare_for_next_iteration() + self.change_status(ServerStatus.training) + + def Connect(self, request, context): + with self.lock: + if self.status is not ServerStatus.connecting: + return sample_pb2.ConnectResponse(successful=False, clientIndex=-1) + + client_index = self._get_next_connect_client_index() + if client_index == -1: + print("All clients slot are engaged, decline this connect request") + return sample_pb2.ConnectResponse(successful=False, clientIndex=client_index) + + self.connected_clients[client_index] = True + print(f"Just assigned {client_index}. {self._get_connect_status()}") + self.try_swtich_from_connecting_to_training() + return sample_pb2.ConnectResponse(successful=True, clientIndex=client_index) + + def Disconnect(self, request, context): + # this can happen at any status, we force status to be connecting when this method is called + # TODO: depending on status, need to abort current iteration. May need to + with self.lock: + client_index = request.clientIndex + print(f"client {client_index} disconnected, waiting for new client to connect") + self.connected_clients[client_index] = False + # next connected client will need to update from 0's iteration + self.client_last_updates[client_index] = 0 + self.swtich_to_connecting() + + return sample_pb2.EmptyResponse() + + def TryToJoinIteration(self, request, context): + with self.lock: + client_index = request.clientIndex + if ( + self.status is not ServerStatus.training + or client_index not in self.iteration_sampled_clients + or client_index in self.iteration_finished_clients + ): + return sample_pb2.TryToJoinIterationResponse( + successful=False, + pullSeeds=data_helper.py_to_protobuf_list_of_list_of_ints([]), + pullGrads=data_helper.py_to_protobuf_list_of_list_of_list_of_floats([]), + iterationSeeds=data_helper.py_to_protobuf_list_of_ints([]), + ) + + last_update_iter = self.client_last_updates[client_index] + # The seed and grad in last_update_iter is fetched as well + # Note at that iteration, we just reset the client model so that iteration + # information is needed as well. + seeds_list = self.seed_grad_records.fetch_seed_records(last_update_iter) + grad_list = self.seed_grad_records.fetch_grad_records(last_update_iter) + + self.client_last_updates[client_index] = self.iteration + + return sample_pb2.TryToJoinIterationResponse( + successful=True, + pullSeeds=data_helper.py_to_protobuf_list_of_list_of_ints(seeds_list), + pullGrads=data_helper.py_to_protobuf_list_of_list_of_list_of_floats( + [[ts.tolist() for ts in vv] for vv in grad_list] + ), + iterationSeeds=data_helper.py_to_protobuf_list_of_ints(self.iteration_seeds), + ) + + def _apply_client_local_update_result(self, client_index, local_update_result): + if client_index not in self.iteration_sampled_clients: + return + + if client_index not in self.iteration_local_grad_scalar: + self.iteration_local_grad_scalar[client_index] = local_update_result + + def SubmitIteration(self, request, context): + with self.lock: + client_index = request.clientIndex + print(f"submit iteration from {client_index}") + + if ( + self.status is not ServerStatus.training + or client_index not in self.iteration_sampled_clients + or client_index in self.iteration_finished_clients + ): + return sample_pb2.EmptyResponse() + + raw_grad_list = data_helper.protobuf_to_py_list_of_list_of_floats(request.gradTensors) + grad_tensors = [torch.tensor(v) for v in raw_grad_list] + self._apply_client_local_update_result(client_index, grad_tensors) + self.iteration_finished_clients.add(client_index) + self.try_switch_from_training_to_aggregating() + return sample_pb2.EmptyResponse() + + def ConnectEval(self, request, context): + with self.lock: + if self.status is not ServerStatus.connecting: + return sample_pb2.ConnectResponse(successful=False, clientIndex=-1) + + if self.eval_client_connected: + print("A eval client is already connected!") + return sample_pb2.ConnectResponse(successful=False, clientIndex=-1) + + self.eval_client_connected = True + print(f"Eval Client connected. {self._get_connect_status()}") + self.try_swtich_from_connecting_to_training() + return sample_pb2.ConnectResponse(successful=True, clientIndex=-1) + + def DisconnectEval(self, request, context): + # this can happen at any status, we force status to be connecting when this method is called + # TODO: depending on status, need to abort current iteration. May need to + with self.lock: + print("Eval client disconnected, waiting for new client to connect") + self.eval_client_connected = False + # next connected client will need to update from 0's iteration + self.eval_client_last_update = 0 + + self.swtich_to_connecting() + + def TryToEval(self, request, context): + with self.lock: + if self.status is not ServerStatus.evaluating: + return sample_pb2.TryToJoinIterationResponse( + successful=False, + pullSeeds=data_helper.py_to_protobuf_list_of_list_of_ints([]), + pullGrads=data_helper.py_to_protobuf_list_of_list_of_list_of_floats([]), + iterationSeeds=data_helper.py_to_protobuf_list_of_ints([]), + ) + + # The seed and grad in last_update_iter is fetched as well + # Note at that iteration, we just reset the client model so that iteration + # information is needed as well. + seeds_list = self.seed_grad_records.fetch_seed_records(self.eval_client_last_update) + grad_list = self.seed_grad_records.fetch_grad_records(self.eval_client_last_update) + self.eval_client_last_update = self.iteration + return sample_pb2.TryToJoinIterationResponse( + successful=True, + pullSeeds=data_helper.py_to_protobuf_list_of_list_of_ints(seeds_list), + pullGrads=data_helper.py_to_protobuf_list_of_list_of_list_of_floats( + [[ts.tolist() for ts in vv] for vv in grad_list] + ), + iterationSeeds=data_helper.py_to_protobuf_list_of_ints([]), + ) + + def SubmitEvaluation(self, request, context): + with self.lock: + if self.status is not ServerStatus.evaluating: + return sample_pb2.EmptyResponse() + + eval_loss, eval_accuracy = request.evalLoss, request.evalAccuracy + print( + f"\nEvaluation(Iteration {self.iteration}): ", + f"Eval Loss:{eval_loss:.4f}, " f"Accuracy:{eval_accuracy * 100:.2f}%", + ) + self.switch_from_evaluating_to_training() + + return sample_pb2.EmptyResponse() + + +def serve(args: cli_interface.CliSetting): + rpc_master_port = args.rpc_master_port + server = grpc.server(futures.ThreadPoolExecutor(max_workers=args.rpc_num_workers)) + sample_pb2_grpc.add_SampleServerServicer_to_server( + SampleServer(args.num_clients, args.num_sample_clients, args.local_update_steps), server + ) + port_str = f"[::]:{rpc_master_port}" + server.add_insecure_port(port_str) + print(f"Parameter server starting on {port_str}") + server.start() + server.wait_for_termination() + + +if __name__ == "__main__": + args = cli_interface.CliSetting() + print(args) + serve(args) diff --git a/grpc_server_test.py b/grpc_server_test.py new file mode 100644 index 0000000..f45ec86 --- /dev/null +++ b/grpc_server_test.py @@ -0,0 +1,23 @@ +from grpc_server import SampleServer + + +def test_get_next_connect_client_index(): + server = SampleServer() + server.connected_clients = [False, False, False] + assert server._get_next_connect_client_index() == 0 + + server.connected_clients = [True, False, False] + assert server._get_next_connect_client_index() == 1 + + server.connected_clients = [True, True, False] + assert server._get_next_connect_client_index() == 2 + + server.connected_clients = [False, True, False] + assert server._get_next_connect_client_index() == 0 + + server.connected_clients = [True, True, True] + assert server._get_next_connect_client_index() == -1 + + +def test_preprare_for_next_iteration(): + pass diff --git a/mypy.ini b/mypy.ini index 63be554..2d225f3 100644 --- a/mypy.ini +++ b/mypy.ini @@ -4,4 +4,5 @@ warn_return_any = True warn_unused_configs = True ignore_missing_imports = True -# disallow_untyped_defs = True, comment for now, need to resolve another 150 errors, try to enable this after initial PR \ No newline at end of file +# disallow_untyped_defs = True, comment for now, need to resolve another 150 errors, try to enable this after initial PR +exclude = (?x)(^cezo_grpc/.*) \ No newline at end of file diff --git a/ruff.toml b/ruff.toml index 074434e..5b4d02b 100644 --- a/ruff.toml +++ b/ruff.toml @@ -27,6 +27,8 @@ exclude = [ "site-packages", "venv", "*.ipynb", + "cezo_grpc/sample_pb2_grpc.py", + "cezo_grpc/sample_pb2.py" ] line-length = 100 diff --git a/run_grpc.py b/run_grpc.py new file mode 100644 index 0000000..6228527 --- /dev/null +++ b/run_grpc.py @@ -0,0 +1,22 @@ +from multiprocessing import Process + +from grpc_client import train_with_args +from grpc_eval_client import eval_with_args +from grpc_server import serve +from cezo_grpc import cli_interface + +import time + +if __name__ == "__main__": + args = cli_interface.CliSetting() + print(args) + + Process(target=serve, args=(args,)).start() + + # HACK: sleep 2 seconds to make sure server can spin up first before clients try to connect + time.sleep(2) + + for _ in range(args.num_clients): + Process(target=train_with_args, args=(args,)).start() + + Process(target=eval_with_args, args=(args,)).start() diff --git a/test_model.py b/test_model.py deleted file mode 100644 index bb8eeba..0000000 --- a/test_model.py +++ /dev/null @@ -1,7 +0,0 @@ -import torch -from transformers import AutoModelForCausalLM - -hf_model_name = "deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B" -torch_dtype = torch.float32 - -model = AutoModelForCausalLM.from_pretrained(hf_model_name, torch_dtype=torch_dtype)