From a89f76d01abc807a6db2f6774141aac2d2e79db4 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Thu, 12 Sep 2024 20:49:18 -0500 Subject: [PATCH 01/19] WIP --- grpc/README.md | 1 + grpc/client.py | 25 ++++ grpc/data_helper.py | 16 +++ grpc/implemented_sample.py | 39 +++++++ grpc/sample.proto | 61 ++++++++++ grpc/sample_pb2.py | 62 ++++++++++ grpc/sample_pb2_grpc.py | 226 +++++++++++++++++++++++++++++++++++++ 7 files changed, 430 insertions(+) create mode 100644 grpc/README.md create mode 100644 grpc/client.py create mode 100644 grpc/data_helper.py create mode 100644 grpc/implemented_sample.py create mode 100644 grpc/sample.proto create mode 100644 grpc/sample_pb2.py create mode 100644 grpc/sample_pb2_grpc.py diff --git a/grpc/README.md b/grpc/README.md new file mode 100644 index 0000000..dd6eb59 --- /dev/null +++ b/grpc/README.md @@ -0,0 +1 @@ +`python -m grpc.tools.protoc --proto_path=. --python_out=. --grpc_python_out=. sample.proto` diff --git a/grpc/client.py b/grpc/client.py new file mode 100644 index 0000000..370e026 --- /dev/null +++ b/grpc/client.py @@ -0,0 +1,25 @@ +import grpc +import sample_pb2 +import sample_pb2_grpc +import data_helper + +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) + +connect_result = ps_stub.Connect(sample_pb2.EmptyRequest()) + +successful, client_index = connect_result.successful, connect_result.clientIndex +print(successful, client_index) + +join_result = ps_stub.TryToJoinIteration(sample_pb2.TryToJoinIterationRequest(clientIndex = client_index)) +print(join_result.successful) + +grads_and_seeds = ps_stub.PullGradsAndSeeds(sample_pb2.PullGradsAndSeedsRequest(clientIndex = client_index)) + +print(data_helper.protobuf_to_py_list_of_list_of_list_of_floats(grads_and_seeds.grads), data_helper.protobuf_to_py_list_of_list_of_ints(grads_and_seeds.seeds)) + +ps_stub.SubmitIteration(sample_pb2.SubmitIterationRequest(local_update_result = sample_pb2.ListOfListOfFloats(data=[sample_pb2.ListOfFloats(data=[0.123])]))) + +print('success') \ No newline at end of file diff --git a/grpc/data_helper.py b/grpc/data_helper.py new file mode 100644 index 0000000..cfdf773 --- /dev/null +++ b/grpc/data_helper.py @@ -0,0 +1,16 @@ +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] diff --git a/grpc/implemented_sample.py b/grpc/implemented_sample.py new file mode 100644 index 0000000..fbd633f --- /dev/null +++ b/grpc/implemented_sample.py @@ -0,0 +1,39 @@ +import grpc +import data_helper +from concurrent import futures +import sample_pb2_grpc +import sample_pb2 + + +class SampleServer(sample_pb2_grpc.SampleServerServicer): + def Connect(self, request, context): + return sample_pb2.ConnectResponse( + successful = True, + clientIndex = 0 + ) + + def TryToJoinIteration(self, request, context): + return sample_pb2.TryToJoinIterationResponse( + successful = True + ) + + def PullGradsAndSeeds(self, request, context): + seeds = sample_pb2.ListOfListOfInts(data=[sample_pb2.ListOfInts(data=[1])]) + grads = sample_pb2.ListOfListOfListOfFloats(data=[sample_pb2.ListOfListOfFloats(data=[sample_pb2.ListOfFloats(data=[0.3])])]) + return sample_pb2.PullGradsAndSeedsResponse(seeds=seeds, grads=grads) + + def SubmitIteration(self, request, context): + print(data_helper.protobuf_to_py_list_of_list_of_floats(request.local_update_result)) + return sample_pb2.EmptyResponse() + +def serve(rpc_master_port, rpc_num_workers): + server = grpc.server(futures.ThreadPoolExecutor(max_workers=rpc_num_workers)) + sample_pb2_grpc.add_SampleServerServicer_to_server(SampleServer(), server) + server.add_insecure_port(f"localhost:{rpc_master_port}") + print(f"Parameter server starting on [::]:{rpc_master_port}") + server.start() + server.wait_for_termination() + + +if __name__ == "__main__": + serve(4242, 8) diff --git a/grpc/sample.proto b/grpc/sample.proto new file mode 100644 index 0000000..7e4d61d --- /dev/null +++ b/grpc/sample.proto @@ -0,0 +1,61 @@ +syntax = "proto3"; + +package sample_server; + +service SampleServer { + rpc Connect (EmptyRequest) returns (ConnectResponse) {} + rpc TryToJoinIteration (TryToJoinIterationRequest) returns (TryToJoinIterationResponse) {} + rpc PullGradsAndSeeds (PullGradsAndSeedsRequest) returns (PullGradsAndSeedsResponse) {} + rpc SubmitIteration (SubmitIterationRequest) returns (EmptyRequest) {} +} + +message EmptyRequest {} + +message EmptyResponse {} + +message ConnectResponse { + bool successful = 1; + int32 clientIndex = 2; +} + +message TryToJoinIterationRequest { + int32 clientIndex = 1; +} + +message TryToJoinIterationResponse{ + bool successful = 1; +} + +message PullGradsAndSeedsRequest { + int32 clientIndex = 1; +} + +message PullGradsAndSeedsResponse { + ListOfListOfInts seeds = 1; + ListOfListOfListOfFloats grads = 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; +} + + +message SubmitIterationRequest { + ListOfListOfFloats local_update_result = 1; +} diff --git a/grpc/sample_pb2.py b/grpc/sample_pb2.py new file mode 100644 index 0000000..925e907 --- /dev/null +++ b/grpc/sample_pb2.py @@ -0,0 +1,62 @@ +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: sample.proto +# Protobuf Python Version: 5.27.2 +"""Generated protocol buffer code.""" +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import runtime_version as _runtime_version +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder +_runtime_version.ValidateProtobufRuntimeVersion( + _runtime_version.Domain.PUBLIC, + 5, + 27, + 2, + '', + 'sample.proto' +) +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + + + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0csample.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\"0\n\x19TryToJoinIterationRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"0\n\x1aTryToJoinIterationResponse\x12\x12\n\nsuccessful\x18\x01 \x01(\x08\"/\n\x18PullGradsAndSeedsRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"\x83\x01\n\x19PullGradsAndSeedsResponse\x12.\n\x05seeds\x18\x01 \x01(\x0b\x32\x1f.sample_server.ListOfListOfInts\x12\x36\n\x05grads\x18\x02 \x01(\x0b\x32\'.sample_server.ListOfListOfListOfFloats\"\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.ListOfListOfFloats\"X\n\x16SubmitIterationRequest\x12>\n\x13local_update_result\x18\x01 \x01(\x0b\x32!.sample_server.ListOfListOfFloats2\x88\x03\n\x0cSampleServer\x12H\n\x07\x43onnect\x12\x1b.sample_server.EmptyRequest\x1a\x1e.sample_server.ConnectResponse\"\x00\x12k\n\x12TryToJoinIteration\x12(.sample_server.TryToJoinIterationRequest\x1a).sample_server.TryToJoinIterationResponse\"\x00\x12h\n\x11PullGradsAndSeeds\x12\'.sample_server.PullGradsAndSeedsRequest\x1a(.sample_server.PullGradsAndSeedsResponse\"\x00\x12W\n\x0fSubmitIteration\x12%.sample_server.SubmitIterationRequest\x1a\x1b.sample_server.EmptyRequest\"\x00\x62\x06proto3') + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'sample_pb2', _globals) +if not _descriptor._USE_C_DESCRIPTORS: + DESCRIPTOR._loaded_options = None + _globals['_EMPTYREQUEST']._serialized_start=31 + _globals['_EMPTYREQUEST']._serialized_end=45 + _globals['_EMPTYRESPONSE']._serialized_start=47 + _globals['_EMPTYRESPONSE']._serialized_end=62 + _globals['_CONNECTRESPONSE']._serialized_start=64 + _globals['_CONNECTRESPONSE']._serialized_end=122 + _globals['_TRYTOJOINITERATIONREQUEST']._serialized_start=124 + _globals['_TRYTOJOINITERATIONREQUEST']._serialized_end=172 + _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_start=174 + _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_end=222 + _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_start=224 + _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_end=271 + _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_start=274 + _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_end=405 + _globals['_LISTOFINTS']._serialized_start=407 + _globals['_LISTOFINTS']._serialized_end=433 + _globals['_LISTOFLISTOFINTS']._serialized_start=435 + _globals['_LISTOFLISTOFINTS']._serialized_end=494 + _globals['_LISTOFFLOATS']._serialized_start=496 + _globals['_LISTOFFLOATS']._serialized_end=524 + _globals['_LISTOFLISTOFFLOATS']._serialized_start=526 + _globals['_LISTOFLISTOFFLOATS']._serialized_end=589 + _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_start=591 + _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_end=666 + _globals['_SUBMITITERATIONREQUEST']._serialized_start=668 + _globals['_SUBMITITERATIONREQUEST']._serialized_end=756 + _globals['_SAMPLESERVER']._serialized_start=759 + _globals['_SAMPLESERVER']._serialized_end=1151 +# @@protoc_insertion_point(module_scope) diff --git a/grpc/sample_pb2_grpc.py b/grpc/sample_pb2_grpc.py new file mode 100644 index 0000000..8d9c9e5 --- /dev/null +++ b/grpc/sample_pb2_grpc.py @@ -0,0 +1,226 @@ +# Generated by the gRPC Python protocol compiler plugin. DO NOT EDIT! +"""Client and server classes corresponding to protobuf-defined services.""" +import grpc +import warnings + +import sample_pb2 as sample__pb2 + +GRPC_GENERATED_VERSION = '1.66.1' +GRPC_VERSION = grpc.__version__ +_version_not_supported = False + +try: + from grpc._utilities import first_version_is_lower + _version_not_supported = first_version_is_lower(GRPC_VERSION, GRPC_GENERATED_VERSION) +except ImportError: + _version_not_supported = True + +if _version_not_supported: + raise RuntimeError( + f'The grpc package installed is at version {GRPC_VERSION},' + + f' but the generated code in sample_pb2_grpc.py depends on' + + f' grpcio>={GRPC_GENERATED_VERSION}.' + + f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}' + + f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.' + ) + + +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=sample__pb2.EmptyRequest.SerializeToString, + response_deserializer=sample__pb2.ConnectResponse.FromString, + _registered_method=True) + self.TryToJoinIteration = channel.unary_unary( + '/sample_server.SampleServer/TryToJoinIteration', + request_serializer=sample__pb2.TryToJoinIterationRequest.SerializeToString, + response_deserializer=sample__pb2.TryToJoinIterationResponse.FromString, + _registered_method=True) + self.PullGradsAndSeeds = channel.unary_unary( + '/sample_server.SampleServer/PullGradsAndSeeds', + request_serializer=sample__pb2.PullGradsAndSeedsRequest.SerializeToString, + response_deserializer=sample__pb2.PullGradsAndSeedsResponse.FromString, + _registered_method=True) + self.SubmitIteration = channel.unary_unary( + '/sample_server.SampleServer/SubmitIteration', + request_serializer=sample__pb2.SubmitIterationRequest.SerializeToString, + response_deserializer=sample__pb2.EmptyRequest.FromString, + _registered_method=True) + + +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 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 PullGradsAndSeeds(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 add_SampleServerServicer_to_server(servicer, server): + rpc_method_handlers = { + 'Connect': grpc.unary_unary_rpc_method_handler( + servicer.Connect, + request_deserializer=sample__pb2.EmptyRequest.FromString, + response_serializer=sample__pb2.ConnectResponse.SerializeToString, + ), + 'TryToJoinIteration': grpc.unary_unary_rpc_method_handler( + servicer.TryToJoinIteration, + request_deserializer=sample__pb2.TryToJoinIterationRequest.FromString, + response_serializer=sample__pb2.TryToJoinIterationResponse.SerializeToString, + ), + 'PullGradsAndSeeds': grpc.unary_unary_rpc_method_handler( + servicer.PullGradsAndSeeds, + request_deserializer=sample__pb2.PullGradsAndSeedsRequest.FromString, + response_serializer=sample__pb2.PullGradsAndSeedsResponse.SerializeToString, + ), + 'SubmitIteration': grpc.unary_unary_rpc_method_handler( + servicer.SubmitIteration, + request_deserializer=sample__pb2.SubmitIterationRequest.FromString, + response_serializer=sample__pb2.EmptyRequest.SerializeToString, + ), + } + generic_handler = grpc.method_handlers_generic_handler( + 'sample_server.SampleServer', rpc_method_handlers) + server.add_generic_rpc_handlers((generic_handler,)) + server.add_registered_method_handlers('sample_server.SampleServer', rpc_method_handlers) + + + # 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', + sample__pb2.EmptyRequest.SerializeToString, + sample__pb2.ConnectResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @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', + sample__pb2.TryToJoinIterationRequest.SerializeToString, + sample__pb2.TryToJoinIterationResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @staticmethod + def PullGradsAndSeeds(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/PullGradsAndSeeds', + sample__pb2.PullGradsAndSeedsRequest.SerializeToString, + sample__pb2.PullGradsAndSeedsResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @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', + sample__pb2.SubmitIterationRequest.SerializeToString, + sample__pb2.EmptyRequest.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) From b1b70492965cecae7e0644529972e32792946ff1 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Fri, 13 Sep 2024 08:09:22 -0500 Subject: [PATCH 02/19] WIP --- cezo_grpc/README.md | 1 + cezo_grpc/__init__.py | 0 {grpc => cezo_grpc}/data_helper.py | 2 +- {grpc => cezo_grpc}/sample.proto | 0 {grpc => cezo_grpc}/sample_pb2.py | 0 {grpc => cezo_grpc}/sample_pb2_grpc.py | 0 grpc/README.md | 1 - grpc/client.py => grpc_client.py | 8 +- grpc/implemented_sample.py => grpc_server.py | 7 +- sample_pb2.py | 62 +++++ sample_pb2_grpc.py | 226 +++++++++++++++++++ 11 files changed, 299 insertions(+), 8 deletions(-) create mode 100644 cezo_grpc/README.md create mode 100644 cezo_grpc/__init__.py rename {grpc => cezo_grpc}/data_helper.py (95%) rename {grpc => cezo_grpc}/sample.proto (100%) rename {grpc => cezo_grpc}/sample_pb2.py (100%) rename {grpc => cezo_grpc}/sample_pb2_grpc.py (100%) delete mode 100644 grpc/README.md rename grpc/client.py => grpc_client.py (87%) rename grpc/implemented_sample.py => grpc_server.py (92%) create mode 100644 sample_pb2.py create mode 100644 sample_pb2_grpc.py diff --git a/cezo_grpc/README.md b/cezo_grpc/README.md new file mode 100644 index 0000000..ab49ea9 --- /dev/null +++ b/cezo_grpc/README.md @@ -0,0 +1 @@ +`python -m grpc.tools.protoc --proto_path=./cezo_grpc --python_out=./cezo_grpc --grpc_python_out=./cezo_grpc ./cezo_grpc/sample.proto` diff --git a/cezo_grpc/__init__.py b/cezo_grpc/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/grpc/data_helper.py b/cezo_grpc/data_helper.py similarity index 95% rename from grpc/data_helper.py rename to cezo_grpc/data_helper.py index cfdf773..88e6b9d 100644 --- a/grpc/data_helper.py +++ b/cezo_grpc/data_helper.py @@ -1,4 +1,4 @@ -import sample_pb2 +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] diff --git a/grpc/sample.proto b/cezo_grpc/sample.proto similarity index 100% rename from grpc/sample.proto rename to cezo_grpc/sample.proto diff --git a/grpc/sample_pb2.py b/cezo_grpc/sample_pb2.py similarity index 100% rename from grpc/sample_pb2.py rename to cezo_grpc/sample_pb2.py diff --git a/grpc/sample_pb2_grpc.py b/cezo_grpc/sample_pb2_grpc.py similarity index 100% rename from grpc/sample_pb2_grpc.py rename to cezo_grpc/sample_pb2_grpc.py diff --git a/grpc/README.md b/grpc/README.md deleted file mode 100644 index dd6eb59..0000000 --- a/grpc/README.md +++ /dev/null @@ -1 +0,0 @@ -`python -m grpc.tools.protoc --proto_path=. --python_out=. --grpc_python_out=. sample.proto` diff --git a/grpc/client.py b/grpc_client.py similarity index 87% rename from grpc/client.py rename to grpc_client.py index 370e026..9f79836 100644 --- a/grpc/client.py +++ b/grpc_client.py @@ -1,7 +1,9 @@ import grpc -import sample_pb2 -import sample_pb2_grpc -import data_helper +from cezo_grpc import sample_pb2 +from cezo_grpc import sample_pb2_grpc +from cezo_grpc import data_helper + +from cezo_fl import client rpc_master_addr = "localhost" rpc_master_port = 4242 diff --git a/grpc/implemented_sample.py b/grpc_server.py similarity index 92% rename from grpc/implemented_sample.py rename to grpc_server.py index fbd633f..b1dc6ce 100644 --- a/grpc/implemented_sample.py +++ b/grpc_server.py @@ -1,8 +1,9 @@ import grpc -import data_helper from concurrent import futures -import sample_pb2_grpc -import sample_pb2 + +from cezo_grpc import data_helper +from cezo_grpc import sample_pb2_grpc +from cezo_grpc import sample_pb2 class SampleServer(sample_pb2_grpc.SampleServerServicer): diff --git a/sample_pb2.py b/sample_pb2.py new file mode 100644 index 0000000..925e907 --- /dev/null +++ b/sample_pb2.py @@ -0,0 +1,62 @@ +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: sample.proto +# Protobuf Python Version: 5.27.2 +"""Generated protocol buffer code.""" +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import runtime_version as _runtime_version +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder +_runtime_version.ValidateProtobufRuntimeVersion( + _runtime_version.Domain.PUBLIC, + 5, + 27, + 2, + '', + 'sample.proto' +) +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + + + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0csample.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\"0\n\x19TryToJoinIterationRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"0\n\x1aTryToJoinIterationResponse\x12\x12\n\nsuccessful\x18\x01 \x01(\x08\"/\n\x18PullGradsAndSeedsRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"\x83\x01\n\x19PullGradsAndSeedsResponse\x12.\n\x05seeds\x18\x01 \x01(\x0b\x32\x1f.sample_server.ListOfListOfInts\x12\x36\n\x05grads\x18\x02 \x01(\x0b\x32\'.sample_server.ListOfListOfListOfFloats\"\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.ListOfListOfFloats\"X\n\x16SubmitIterationRequest\x12>\n\x13local_update_result\x18\x01 \x01(\x0b\x32!.sample_server.ListOfListOfFloats2\x88\x03\n\x0cSampleServer\x12H\n\x07\x43onnect\x12\x1b.sample_server.EmptyRequest\x1a\x1e.sample_server.ConnectResponse\"\x00\x12k\n\x12TryToJoinIteration\x12(.sample_server.TryToJoinIterationRequest\x1a).sample_server.TryToJoinIterationResponse\"\x00\x12h\n\x11PullGradsAndSeeds\x12\'.sample_server.PullGradsAndSeedsRequest\x1a(.sample_server.PullGradsAndSeedsResponse\"\x00\x12W\n\x0fSubmitIteration\x12%.sample_server.SubmitIterationRequest\x1a\x1b.sample_server.EmptyRequest\"\x00\x62\x06proto3') + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'sample_pb2', _globals) +if not _descriptor._USE_C_DESCRIPTORS: + DESCRIPTOR._loaded_options = None + _globals['_EMPTYREQUEST']._serialized_start=31 + _globals['_EMPTYREQUEST']._serialized_end=45 + _globals['_EMPTYRESPONSE']._serialized_start=47 + _globals['_EMPTYRESPONSE']._serialized_end=62 + _globals['_CONNECTRESPONSE']._serialized_start=64 + _globals['_CONNECTRESPONSE']._serialized_end=122 + _globals['_TRYTOJOINITERATIONREQUEST']._serialized_start=124 + _globals['_TRYTOJOINITERATIONREQUEST']._serialized_end=172 + _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_start=174 + _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_end=222 + _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_start=224 + _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_end=271 + _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_start=274 + _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_end=405 + _globals['_LISTOFINTS']._serialized_start=407 + _globals['_LISTOFINTS']._serialized_end=433 + _globals['_LISTOFLISTOFINTS']._serialized_start=435 + _globals['_LISTOFLISTOFINTS']._serialized_end=494 + _globals['_LISTOFFLOATS']._serialized_start=496 + _globals['_LISTOFFLOATS']._serialized_end=524 + _globals['_LISTOFLISTOFFLOATS']._serialized_start=526 + _globals['_LISTOFLISTOFFLOATS']._serialized_end=589 + _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_start=591 + _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_end=666 + _globals['_SUBMITITERATIONREQUEST']._serialized_start=668 + _globals['_SUBMITITERATIONREQUEST']._serialized_end=756 + _globals['_SAMPLESERVER']._serialized_start=759 + _globals['_SAMPLESERVER']._serialized_end=1151 +# @@protoc_insertion_point(module_scope) diff --git a/sample_pb2_grpc.py b/sample_pb2_grpc.py new file mode 100644 index 0000000..8d9c9e5 --- /dev/null +++ b/sample_pb2_grpc.py @@ -0,0 +1,226 @@ +# Generated by the gRPC Python protocol compiler plugin. DO NOT EDIT! +"""Client and server classes corresponding to protobuf-defined services.""" +import grpc +import warnings + +import sample_pb2 as sample__pb2 + +GRPC_GENERATED_VERSION = '1.66.1' +GRPC_VERSION = grpc.__version__ +_version_not_supported = False + +try: + from grpc._utilities import first_version_is_lower + _version_not_supported = first_version_is_lower(GRPC_VERSION, GRPC_GENERATED_VERSION) +except ImportError: + _version_not_supported = True + +if _version_not_supported: + raise RuntimeError( + f'The grpc package installed is at version {GRPC_VERSION},' + + f' but the generated code in sample_pb2_grpc.py depends on' + + f' grpcio>={GRPC_GENERATED_VERSION}.' + + f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}' + + f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.' + ) + + +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=sample__pb2.EmptyRequest.SerializeToString, + response_deserializer=sample__pb2.ConnectResponse.FromString, + _registered_method=True) + self.TryToJoinIteration = channel.unary_unary( + '/sample_server.SampleServer/TryToJoinIteration', + request_serializer=sample__pb2.TryToJoinIterationRequest.SerializeToString, + response_deserializer=sample__pb2.TryToJoinIterationResponse.FromString, + _registered_method=True) + self.PullGradsAndSeeds = channel.unary_unary( + '/sample_server.SampleServer/PullGradsAndSeeds', + request_serializer=sample__pb2.PullGradsAndSeedsRequest.SerializeToString, + response_deserializer=sample__pb2.PullGradsAndSeedsResponse.FromString, + _registered_method=True) + self.SubmitIteration = channel.unary_unary( + '/sample_server.SampleServer/SubmitIteration', + request_serializer=sample__pb2.SubmitIterationRequest.SerializeToString, + response_deserializer=sample__pb2.EmptyRequest.FromString, + _registered_method=True) + + +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 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 PullGradsAndSeeds(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 add_SampleServerServicer_to_server(servicer, server): + rpc_method_handlers = { + 'Connect': grpc.unary_unary_rpc_method_handler( + servicer.Connect, + request_deserializer=sample__pb2.EmptyRequest.FromString, + response_serializer=sample__pb2.ConnectResponse.SerializeToString, + ), + 'TryToJoinIteration': grpc.unary_unary_rpc_method_handler( + servicer.TryToJoinIteration, + request_deserializer=sample__pb2.TryToJoinIterationRequest.FromString, + response_serializer=sample__pb2.TryToJoinIterationResponse.SerializeToString, + ), + 'PullGradsAndSeeds': grpc.unary_unary_rpc_method_handler( + servicer.PullGradsAndSeeds, + request_deserializer=sample__pb2.PullGradsAndSeedsRequest.FromString, + response_serializer=sample__pb2.PullGradsAndSeedsResponse.SerializeToString, + ), + 'SubmitIteration': grpc.unary_unary_rpc_method_handler( + servicer.SubmitIteration, + request_deserializer=sample__pb2.SubmitIterationRequest.FromString, + response_serializer=sample__pb2.EmptyRequest.SerializeToString, + ), + } + generic_handler = grpc.method_handlers_generic_handler( + 'sample_server.SampleServer', rpc_method_handlers) + server.add_generic_rpc_handlers((generic_handler,)) + server.add_registered_method_handlers('sample_server.SampleServer', rpc_method_handlers) + + + # 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', + sample__pb2.EmptyRequest.SerializeToString, + sample__pb2.ConnectResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @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', + sample__pb2.TryToJoinIterationRequest.SerializeToString, + sample__pb2.TryToJoinIterationResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @staticmethod + def PullGradsAndSeeds(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/PullGradsAndSeeds', + sample__pb2.PullGradsAndSeedsRequest.SerializeToString, + sample__pb2.PullGradsAndSeedsResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @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', + sample__pb2.SubmitIterationRequest.SerializeToString, + sample__pb2.EmptyRequest.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) From 259efccc658c962169308b31a040a96301c52c09 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Fri, 13 Sep 2024 13:52:43 -0500 Subject: [PATCH 03/19] run script at root of folder and make sure import works --- cezo_grpc/README.md | 2 +- cezo_grpc/sample_pb2.py | 64 +++++----- cezo_grpc/sample_pb2_grpc.py | 52 ++++---- sample_pb2.py | 62 ---------- sample_pb2_grpc.py | 226 ----------------------------------- 5 files changed, 59 insertions(+), 347 deletions(-) delete mode 100644 sample_pb2.py delete mode 100644 sample_pb2_grpc.py diff --git a/cezo_grpc/README.md b/cezo_grpc/README.md index ab49ea9..c8fc573 100644 --- a/cezo_grpc/README.md +++ b/cezo_grpc/README.md @@ -1 +1 @@ -`python -m grpc.tools.protoc --proto_path=./cezo_grpc --python_out=./cezo_grpc --grpc_python_out=./cezo_grpc ./cezo_grpc/sample.proto` +run this at root `python -m grpc_tools.protoc -Icezo_grpc=./cezo_grpc --python_out=. --grpc_python_out=. ./cezo_grpc/sample.proto` diff --git a/cezo_grpc/sample_pb2.py b/cezo_grpc/sample_pb2.py index 925e907..4088cf4 100644 --- a/cezo_grpc/sample_pb2.py +++ b/cezo_grpc/sample_pb2.py @@ -1,7 +1,7 @@ # -*- coding: utf-8 -*- # Generated by the protocol buffer compiler. DO NOT EDIT! # NO CHECKED-IN PROTOBUF GENCODE -# source: sample.proto +# source: cezo_grpc/sample.proto # Protobuf Python Version: 5.27.2 """Generated protocol buffer code.""" from google.protobuf import descriptor as _descriptor @@ -15,7 +15,7 @@ 27, 2, '', - 'sample.proto' + 'cezo_grpc/sample.proto' ) # @@protoc_insertion_point(imports) @@ -24,39 +24,39 @@ -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0csample.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\"0\n\x19TryToJoinIterationRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"0\n\x1aTryToJoinIterationResponse\x12\x12\n\nsuccessful\x18\x01 \x01(\x08\"/\n\x18PullGradsAndSeedsRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"\x83\x01\n\x19PullGradsAndSeedsResponse\x12.\n\x05seeds\x18\x01 \x01(\x0b\x32\x1f.sample_server.ListOfListOfInts\x12\x36\n\x05grads\x18\x02 \x01(\x0b\x32\'.sample_server.ListOfListOfListOfFloats\"\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.ListOfListOfFloats\"X\n\x16SubmitIterationRequest\x12>\n\x13local_update_result\x18\x01 \x01(\x0b\x32!.sample_server.ListOfListOfFloats2\x88\x03\n\x0cSampleServer\x12H\n\x07\x43onnect\x12\x1b.sample_server.EmptyRequest\x1a\x1e.sample_server.ConnectResponse\"\x00\x12k\n\x12TryToJoinIteration\x12(.sample_server.TryToJoinIterationRequest\x1a).sample_server.TryToJoinIterationResponse\"\x00\x12h\n\x11PullGradsAndSeeds\x12\'.sample_server.PullGradsAndSeedsRequest\x1a(.sample_server.PullGradsAndSeedsResponse\"\x00\x12W\n\x0fSubmitIteration\x12%.sample_server.SubmitIterationRequest\x1a\x1b.sample_server.EmptyRequest\"\x00\x62\x06proto3') +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\"0\n\x19TryToJoinIterationRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"0\n\x1aTryToJoinIterationResponse\x12\x12\n\nsuccessful\x18\x01 \x01(\x08\"/\n\x18PullGradsAndSeedsRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"\x83\x01\n\x19PullGradsAndSeedsResponse\x12.\n\x05seeds\x18\x01 \x01(\x0b\x32\x1f.sample_server.ListOfListOfInts\x12\x36\n\x05grads\x18\x02 \x01(\x0b\x32\'.sample_server.ListOfListOfListOfFloats\"\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.ListOfListOfFloats\"X\n\x16SubmitIterationRequest\x12>\n\x13local_update_result\x18\x01 \x01(\x0b\x32!.sample_server.ListOfListOfFloats2\x88\x03\n\x0cSampleServer\x12H\n\x07\x43onnect\x12\x1b.sample_server.EmptyRequest\x1a\x1e.sample_server.ConnectResponse\"\x00\x12k\n\x12TryToJoinIteration\x12(.sample_server.TryToJoinIterationRequest\x1a).sample_server.TryToJoinIterationResponse\"\x00\x12h\n\x11PullGradsAndSeeds\x12\'.sample_server.PullGradsAndSeedsRequest\x1a(.sample_server.PullGradsAndSeedsResponse\"\x00\x12W\n\x0fSubmitIteration\x12%.sample_server.SubmitIterationRequest\x1a\x1b.sample_server.EmptyRequest\"\x00\x62\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) -_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'sample_pb2', _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'cezo_grpc.sample_pb2', _globals) if not _descriptor._USE_C_DESCRIPTORS: DESCRIPTOR._loaded_options = None - _globals['_EMPTYREQUEST']._serialized_start=31 - _globals['_EMPTYREQUEST']._serialized_end=45 - _globals['_EMPTYRESPONSE']._serialized_start=47 - _globals['_EMPTYRESPONSE']._serialized_end=62 - _globals['_CONNECTRESPONSE']._serialized_start=64 - _globals['_CONNECTRESPONSE']._serialized_end=122 - _globals['_TRYTOJOINITERATIONREQUEST']._serialized_start=124 - _globals['_TRYTOJOINITERATIONREQUEST']._serialized_end=172 - _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_start=174 - _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_end=222 - _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_start=224 - _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_end=271 - _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_start=274 - _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_end=405 - _globals['_LISTOFINTS']._serialized_start=407 - _globals['_LISTOFINTS']._serialized_end=433 - _globals['_LISTOFLISTOFINTS']._serialized_start=435 - _globals['_LISTOFLISTOFINTS']._serialized_end=494 - _globals['_LISTOFFLOATS']._serialized_start=496 - _globals['_LISTOFFLOATS']._serialized_end=524 - _globals['_LISTOFLISTOFFLOATS']._serialized_start=526 - _globals['_LISTOFLISTOFFLOATS']._serialized_end=589 - _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_start=591 - _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_end=666 - _globals['_SUBMITITERATIONREQUEST']._serialized_start=668 - _globals['_SUBMITITERATIONREQUEST']._serialized_end=756 - _globals['_SAMPLESERVER']._serialized_start=759 - _globals['_SAMPLESERVER']._serialized_end=1151 + _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['_TRYTOJOINITERATIONREQUEST']._serialized_start=134 + _globals['_TRYTOJOINITERATIONREQUEST']._serialized_end=182 + _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_start=184 + _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_end=232 + _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_start=234 + _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_end=281 + _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_start=284 + _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_end=415 + _globals['_LISTOFINTS']._serialized_start=417 + _globals['_LISTOFINTS']._serialized_end=443 + _globals['_LISTOFLISTOFINTS']._serialized_start=445 + _globals['_LISTOFLISTOFINTS']._serialized_end=504 + _globals['_LISTOFFLOATS']._serialized_start=506 + _globals['_LISTOFFLOATS']._serialized_end=534 + _globals['_LISTOFLISTOFFLOATS']._serialized_start=536 + _globals['_LISTOFLISTOFFLOATS']._serialized_end=599 + _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_start=601 + _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_end=676 + _globals['_SUBMITITERATIONREQUEST']._serialized_start=678 + _globals['_SUBMITITERATIONREQUEST']._serialized_end=766 + _globals['_SAMPLESERVER']._serialized_start=769 + _globals['_SAMPLESERVER']._serialized_end=1161 # @@protoc_insertion_point(module_scope) diff --git a/cezo_grpc/sample_pb2_grpc.py b/cezo_grpc/sample_pb2_grpc.py index 8d9c9e5..aea20cd 100644 --- a/cezo_grpc/sample_pb2_grpc.py +++ b/cezo_grpc/sample_pb2_grpc.py @@ -3,7 +3,7 @@ import grpc import warnings -import sample_pb2 as sample__pb2 +from cezo_grpc import sample_pb2 as cezo__grpc_dot_sample__pb2 GRPC_GENERATED_VERSION = '1.66.1' GRPC_VERSION = grpc.__version__ @@ -18,7 +18,7 @@ if _version_not_supported: raise RuntimeError( f'The grpc package installed is at version {GRPC_VERSION},' - + f' but the generated code in sample_pb2_grpc.py depends on' + + f' but the generated code in cezo_grpc/sample_pb2_grpc.py depends on' + f' grpcio>={GRPC_GENERATED_VERSION}.' + f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}' + f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.' @@ -36,23 +36,23 @@ def __init__(self, channel): """ self.Connect = channel.unary_unary( '/sample_server.SampleServer/Connect', - request_serializer=sample__pb2.EmptyRequest.SerializeToString, - response_deserializer=sample__pb2.ConnectResponse.FromString, + request_serializer=cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.ConnectResponse.FromString, _registered_method=True) self.TryToJoinIteration = channel.unary_unary( '/sample_server.SampleServer/TryToJoinIteration', - request_serializer=sample__pb2.TryToJoinIterationRequest.SerializeToString, - response_deserializer=sample__pb2.TryToJoinIterationResponse.FromString, + request_serializer=cezo__grpc_dot_sample__pb2.TryToJoinIterationRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.TryToJoinIterationResponse.FromString, _registered_method=True) self.PullGradsAndSeeds = channel.unary_unary( '/sample_server.SampleServer/PullGradsAndSeeds', - request_serializer=sample__pb2.PullGradsAndSeedsRequest.SerializeToString, - response_deserializer=sample__pb2.PullGradsAndSeedsResponse.FromString, + request_serializer=cezo__grpc_dot_sample__pb2.PullGradsAndSeedsRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.PullGradsAndSeedsResponse.FromString, _registered_method=True) self.SubmitIteration = channel.unary_unary( '/sample_server.SampleServer/SubmitIteration', - request_serializer=sample__pb2.SubmitIterationRequest.SerializeToString, - response_deserializer=sample__pb2.EmptyRequest.FromString, + request_serializer=cezo__grpc_dot_sample__pb2.SubmitIterationRequest.SerializeToString, + response_deserializer=cezo__grpc_dot_sample__pb2.EmptyRequest.FromString, _registered_method=True) @@ -88,23 +88,23 @@ def add_SampleServerServicer_to_server(servicer, server): rpc_method_handlers = { 'Connect': grpc.unary_unary_rpc_method_handler( servicer.Connect, - request_deserializer=sample__pb2.EmptyRequest.FromString, - response_serializer=sample__pb2.ConnectResponse.SerializeToString, + request_deserializer=cezo__grpc_dot_sample__pb2.EmptyRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.ConnectResponse.SerializeToString, ), 'TryToJoinIteration': grpc.unary_unary_rpc_method_handler( servicer.TryToJoinIteration, - request_deserializer=sample__pb2.TryToJoinIterationRequest.FromString, - response_serializer=sample__pb2.TryToJoinIterationResponse.SerializeToString, + request_deserializer=cezo__grpc_dot_sample__pb2.TryToJoinIterationRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.TryToJoinIterationResponse.SerializeToString, ), 'PullGradsAndSeeds': grpc.unary_unary_rpc_method_handler( servicer.PullGradsAndSeeds, - request_deserializer=sample__pb2.PullGradsAndSeedsRequest.FromString, - response_serializer=sample__pb2.PullGradsAndSeedsResponse.SerializeToString, + request_deserializer=cezo__grpc_dot_sample__pb2.PullGradsAndSeedsRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.PullGradsAndSeedsResponse.SerializeToString, ), 'SubmitIteration': grpc.unary_unary_rpc_method_handler( servicer.SubmitIteration, - request_deserializer=sample__pb2.SubmitIterationRequest.FromString, - response_serializer=sample__pb2.EmptyRequest.SerializeToString, + request_deserializer=cezo__grpc_dot_sample__pb2.SubmitIterationRequest.FromString, + response_serializer=cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, ), } generic_handler = grpc.method_handlers_generic_handler( @@ -132,8 +132,8 @@ def Connect(request, request, target, '/sample_server.SampleServer/Connect', - sample__pb2.EmptyRequest.SerializeToString, - sample__pb2.ConnectResponse.FromString, + cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.ConnectResponse.FromString, options, channel_credentials, insecure, @@ -159,8 +159,8 @@ def TryToJoinIteration(request, request, target, '/sample_server.SampleServer/TryToJoinIteration', - sample__pb2.TryToJoinIterationRequest.SerializeToString, - sample__pb2.TryToJoinIterationResponse.FromString, + cezo__grpc_dot_sample__pb2.TryToJoinIterationRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.TryToJoinIterationResponse.FromString, options, channel_credentials, insecure, @@ -186,8 +186,8 @@ def PullGradsAndSeeds(request, request, target, '/sample_server.SampleServer/PullGradsAndSeeds', - sample__pb2.PullGradsAndSeedsRequest.SerializeToString, - sample__pb2.PullGradsAndSeedsResponse.FromString, + cezo__grpc_dot_sample__pb2.PullGradsAndSeedsRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.PullGradsAndSeedsResponse.FromString, options, channel_credentials, insecure, @@ -213,8 +213,8 @@ def SubmitIteration(request, request, target, '/sample_server.SampleServer/SubmitIteration', - sample__pb2.SubmitIterationRequest.SerializeToString, - sample__pb2.EmptyRequest.FromString, + cezo__grpc_dot_sample__pb2.SubmitIterationRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.EmptyRequest.FromString, options, channel_credentials, insecure, diff --git a/sample_pb2.py b/sample_pb2.py deleted file mode 100644 index 925e907..0000000 --- a/sample_pb2.py +++ /dev/null @@ -1,62 +0,0 @@ -# -*- coding: utf-8 -*- -# Generated by the protocol buffer compiler. DO NOT EDIT! -# NO CHECKED-IN PROTOBUF GENCODE -# source: sample.proto -# Protobuf Python Version: 5.27.2 -"""Generated protocol buffer code.""" -from google.protobuf import descriptor as _descriptor -from google.protobuf import descriptor_pool as _descriptor_pool -from google.protobuf import runtime_version as _runtime_version -from google.protobuf import symbol_database as _symbol_database -from google.protobuf.internal import builder as _builder -_runtime_version.ValidateProtobufRuntimeVersion( - _runtime_version.Domain.PUBLIC, - 5, - 27, - 2, - '', - 'sample.proto' -) -# @@protoc_insertion_point(imports) - -_sym_db = _symbol_database.Default() - - - - -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0csample.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\"0\n\x19TryToJoinIterationRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"0\n\x1aTryToJoinIterationResponse\x12\x12\n\nsuccessful\x18\x01 \x01(\x08\"/\n\x18PullGradsAndSeedsRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"\x83\x01\n\x19PullGradsAndSeedsResponse\x12.\n\x05seeds\x18\x01 \x01(\x0b\x32\x1f.sample_server.ListOfListOfInts\x12\x36\n\x05grads\x18\x02 \x01(\x0b\x32\'.sample_server.ListOfListOfListOfFloats\"\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.ListOfListOfFloats\"X\n\x16SubmitIterationRequest\x12>\n\x13local_update_result\x18\x01 \x01(\x0b\x32!.sample_server.ListOfListOfFloats2\x88\x03\n\x0cSampleServer\x12H\n\x07\x43onnect\x12\x1b.sample_server.EmptyRequest\x1a\x1e.sample_server.ConnectResponse\"\x00\x12k\n\x12TryToJoinIteration\x12(.sample_server.TryToJoinIterationRequest\x1a).sample_server.TryToJoinIterationResponse\"\x00\x12h\n\x11PullGradsAndSeeds\x12\'.sample_server.PullGradsAndSeedsRequest\x1a(.sample_server.PullGradsAndSeedsResponse\"\x00\x12W\n\x0fSubmitIteration\x12%.sample_server.SubmitIterationRequest\x1a\x1b.sample_server.EmptyRequest\"\x00\x62\x06proto3') - -_globals = globals() -_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) -_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'sample_pb2', _globals) -if not _descriptor._USE_C_DESCRIPTORS: - DESCRIPTOR._loaded_options = None - _globals['_EMPTYREQUEST']._serialized_start=31 - _globals['_EMPTYREQUEST']._serialized_end=45 - _globals['_EMPTYRESPONSE']._serialized_start=47 - _globals['_EMPTYRESPONSE']._serialized_end=62 - _globals['_CONNECTRESPONSE']._serialized_start=64 - _globals['_CONNECTRESPONSE']._serialized_end=122 - _globals['_TRYTOJOINITERATIONREQUEST']._serialized_start=124 - _globals['_TRYTOJOINITERATIONREQUEST']._serialized_end=172 - _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_start=174 - _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_end=222 - _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_start=224 - _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_end=271 - _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_start=274 - _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_end=405 - _globals['_LISTOFINTS']._serialized_start=407 - _globals['_LISTOFINTS']._serialized_end=433 - _globals['_LISTOFLISTOFINTS']._serialized_start=435 - _globals['_LISTOFLISTOFINTS']._serialized_end=494 - _globals['_LISTOFFLOATS']._serialized_start=496 - _globals['_LISTOFFLOATS']._serialized_end=524 - _globals['_LISTOFLISTOFFLOATS']._serialized_start=526 - _globals['_LISTOFLISTOFFLOATS']._serialized_end=589 - _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_start=591 - _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_end=666 - _globals['_SUBMITITERATIONREQUEST']._serialized_start=668 - _globals['_SUBMITITERATIONREQUEST']._serialized_end=756 - _globals['_SAMPLESERVER']._serialized_start=759 - _globals['_SAMPLESERVER']._serialized_end=1151 -# @@protoc_insertion_point(module_scope) diff --git a/sample_pb2_grpc.py b/sample_pb2_grpc.py deleted file mode 100644 index 8d9c9e5..0000000 --- a/sample_pb2_grpc.py +++ /dev/null @@ -1,226 +0,0 @@ -# Generated by the gRPC Python protocol compiler plugin. DO NOT EDIT! -"""Client and server classes corresponding to protobuf-defined services.""" -import grpc -import warnings - -import sample_pb2 as sample__pb2 - -GRPC_GENERATED_VERSION = '1.66.1' -GRPC_VERSION = grpc.__version__ -_version_not_supported = False - -try: - from grpc._utilities import first_version_is_lower - _version_not_supported = first_version_is_lower(GRPC_VERSION, GRPC_GENERATED_VERSION) -except ImportError: - _version_not_supported = True - -if _version_not_supported: - raise RuntimeError( - f'The grpc package installed is at version {GRPC_VERSION},' - + f' but the generated code in sample_pb2_grpc.py depends on' - + f' grpcio>={GRPC_GENERATED_VERSION}.' - + f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}' - + f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.' - ) - - -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=sample__pb2.EmptyRequest.SerializeToString, - response_deserializer=sample__pb2.ConnectResponse.FromString, - _registered_method=True) - self.TryToJoinIteration = channel.unary_unary( - '/sample_server.SampleServer/TryToJoinIteration', - request_serializer=sample__pb2.TryToJoinIterationRequest.SerializeToString, - response_deserializer=sample__pb2.TryToJoinIterationResponse.FromString, - _registered_method=True) - self.PullGradsAndSeeds = channel.unary_unary( - '/sample_server.SampleServer/PullGradsAndSeeds', - request_serializer=sample__pb2.PullGradsAndSeedsRequest.SerializeToString, - response_deserializer=sample__pb2.PullGradsAndSeedsResponse.FromString, - _registered_method=True) - self.SubmitIteration = channel.unary_unary( - '/sample_server.SampleServer/SubmitIteration', - request_serializer=sample__pb2.SubmitIterationRequest.SerializeToString, - response_deserializer=sample__pb2.EmptyRequest.FromString, - _registered_method=True) - - -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 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 PullGradsAndSeeds(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 add_SampleServerServicer_to_server(servicer, server): - rpc_method_handlers = { - 'Connect': grpc.unary_unary_rpc_method_handler( - servicer.Connect, - request_deserializer=sample__pb2.EmptyRequest.FromString, - response_serializer=sample__pb2.ConnectResponse.SerializeToString, - ), - 'TryToJoinIteration': grpc.unary_unary_rpc_method_handler( - servicer.TryToJoinIteration, - request_deserializer=sample__pb2.TryToJoinIterationRequest.FromString, - response_serializer=sample__pb2.TryToJoinIterationResponse.SerializeToString, - ), - 'PullGradsAndSeeds': grpc.unary_unary_rpc_method_handler( - servicer.PullGradsAndSeeds, - request_deserializer=sample__pb2.PullGradsAndSeedsRequest.FromString, - response_serializer=sample__pb2.PullGradsAndSeedsResponse.SerializeToString, - ), - 'SubmitIteration': grpc.unary_unary_rpc_method_handler( - servicer.SubmitIteration, - request_deserializer=sample__pb2.SubmitIterationRequest.FromString, - response_serializer=sample__pb2.EmptyRequest.SerializeToString, - ), - } - generic_handler = grpc.method_handlers_generic_handler( - 'sample_server.SampleServer', rpc_method_handlers) - server.add_generic_rpc_handlers((generic_handler,)) - server.add_registered_method_handlers('sample_server.SampleServer', rpc_method_handlers) - - - # 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', - sample__pb2.EmptyRequest.SerializeToString, - sample__pb2.ConnectResponse.FromString, - options, - channel_credentials, - insecure, - call_credentials, - compression, - wait_for_ready, - timeout, - metadata, - _registered_method=True) - - @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', - sample__pb2.TryToJoinIterationRequest.SerializeToString, - sample__pb2.TryToJoinIterationResponse.FromString, - options, - channel_credentials, - insecure, - call_credentials, - compression, - wait_for_ready, - timeout, - metadata, - _registered_method=True) - - @staticmethod - def PullGradsAndSeeds(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/PullGradsAndSeeds', - sample__pb2.PullGradsAndSeedsRequest.SerializeToString, - sample__pb2.PullGradsAndSeedsResponse.FromString, - options, - channel_credentials, - insecure, - call_credentials, - compression, - wait_for_ready, - timeout, - metadata, - _registered_method=True) - - @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', - sample__pb2.SubmitIterationRequest.SerializeToString, - sample__pb2.EmptyRequest.FromString, - options, - channel_credentials, - insecure, - call_credentials, - compression, - wait_for_ready, - timeout, - metadata, - _registered_method=True) From d72fa78cd805d15d7c90e9954c112da42c03bae5 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sun, 15 Sep 2024 15:34:43 -0500 Subject: [PATCH 04/19] make it work --- cezo_grpc/data_helper.py | 49 +++++++-- cezo_grpc/sample.proto | 23 +++-- cezo_grpc/sample_pb2.py | 44 ++++---- cezo_grpc/sample_pb2_grpc.py | 40 ++++---- grpc_client.py | 116 ++++++++++++++++++--- grpc_server.py | 193 ++++++++++++++++++++++++++++++++--- 6 files changed, 379 insertions(+), 86 deletions(-) diff --git a/cezo_grpc/data_helper.py b/cezo_grpc/data_helper.py index 88e6b9d..97e55cc 100644 --- a/cezo_grpc/data_helper.py +++ b/cezo_grpc/data_helper.py @@ -1,16 +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] + 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] + 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] + 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 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_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 index 7e4d61d..f581053 100644 --- a/cezo_grpc/sample.proto +++ b/cezo_grpc/sample.proto @@ -4,8 +4,8 @@ package sample_server; service SampleServer { rpc Connect (EmptyRequest) returns (ConnectResponse) {} + rpc Disconnect (DisconnectRequest) returns (EmptyResponse) {} rpc TryToJoinIteration (TryToJoinIterationRequest) returns (TryToJoinIterationResponse) {} - rpc PullGradsAndSeeds (PullGradsAndSeedsRequest) returns (PullGradsAndSeedsResponse) {} rpc SubmitIteration (SubmitIterationRequest) returns (EmptyRequest) {} } @@ -18,22 +18,28 @@ message ConnectResponse { 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 PullGradsAndSeedsRequest { +message SubmitIterationRequest { int32 clientIndex = 1; + ListOfListOfFloats gradTensors = 2; + float stepAccuracy = 3; + float stepLoss = 4; } -message PullGradsAndSeedsResponse { - ListOfListOfInts seeds = 1; - ListOfListOfListOfFloats grads = 2; -} message ListOfInts { repeated int32 data = 1; @@ -54,8 +60,3 @@ message ListOfListOfFloats { message ListOfListOfListOfFloats { repeated ListOfListOfFloats data = 1; } - - -message SubmitIterationRequest { - ListOfListOfFloats local_update_result = 1; -} diff --git a/cezo_grpc/sample_pb2.py b/cezo_grpc/sample_pb2.py index 4088cf4..b059528 100644 --- a/cezo_grpc/sample_pb2.py +++ b/cezo_grpc/sample_pb2.py @@ -24,7 +24,7 @@ -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\"0\n\x19TryToJoinIterationRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"0\n\x1aTryToJoinIterationResponse\x12\x12\n\nsuccessful\x18\x01 \x01(\x08\"/\n\x18PullGradsAndSeedsRequest\x12\x13\n\x0b\x63lientIndex\x18\x01 \x01(\x05\"\x83\x01\n\x19PullGradsAndSeedsResponse\x12.\n\x05seeds\x18\x01 \x01(\x0b\x32\x1f.sample_server.ListOfListOfInts\x12\x36\n\x05grads\x18\x02 \x01(\x0b\x32\'.sample_server.ListOfListOfListOfFloats\"\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.ListOfListOfFloats\"X\n\x16SubmitIterationRequest\x12>\n\x13local_update_result\x18\x01 \x01(\x0b\x32!.sample_server.ListOfListOfFloats2\x88\x03\n\x0cSampleServer\x12H\n\x07\x43onnect\x12\x1b.sample_server.EmptyRequest\x1a\x1e.sample_server.ConnectResponse\"\x00\x12k\n\x12TryToJoinIteration\x12(.sample_server.TryToJoinIterationRequest\x1a).sample_server.TryToJoinIterationResponse\"\x00\x12h\n\x11PullGradsAndSeeds\x12\'.sample_server.PullGradsAndSeedsRequest\x1a(.sample_server.PullGradsAndSeedsResponse\"\x00\x12W\n\x0fSubmitIteration\x12%.sample_server.SubmitIterationRequest\x1a\x1b.sample_server.EmptyRequest\"\x00\x62\x06proto3') +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\"\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\xee\x02\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\x12W\n\x0fSubmitIteration\x12%.sample_server.SubmitIterationRequest\x1a\x1b.sample_server.EmptyRequest\"\x00\x62\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) @@ -37,26 +37,24 @@ _globals['_EMPTYRESPONSE']._serialized_end=72 _globals['_CONNECTRESPONSE']._serialized_start=74 _globals['_CONNECTRESPONSE']._serialized_end=132 - _globals['_TRYTOJOINITERATIONREQUEST']._serialized_start=134 - _globals['_TRYTOJOINITERATIONREQUEST']._serialized_end=182 - _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_start=184 - _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_end=232 - _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_start=234 - _globals['_PULLGRADSANDSEEDSREQUEST']._serialized_end=281 - _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_start=284 - _globals['_PULLGRADSANDSEEDSRESPONSE']._serialized_end=415 - _globals['_LISTOFINTS']._serialized_start=417 - _globals['_LISTOFINTS']._serialized_end=443 - _globals['_LISTOFLISTOFINTS']._serialized_start=445 - _globals['_LISTOFLISTOFINTS']._serialized_end=504 - _globals['_LISTOFFLOATS']._serialized_start=506 - _globals['_LISTOFFLOATS']._serialized_end=534 - _globals['_LISTOFLISTOFFLOATS']._serialized_start=536 - _globals['_LISTOFLISTOFFLOATS']._serialized_end=599 - _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_start=601 - _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_end=676 - _globals['_SUBMITITERATIONREQUEST']._serialized_start=678 - _globals['_SUBMITITERATIONREQUEST']._serialized_end=766 - _globals['_SAMPLESERVER']._serialized_start=769 - _globals['_SAMPLESERVER']._serialized_end=1161 + _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['_LISTOFINTS']._serialized_start=584 + _globals['_LISTOFINTS']._serialized_end=610 + _globals['_LISTOFLISTOFINTS']._serialized_start=612 + _globals['_LISTOFLISTOFINTS']._serialized_end=671 + _globals['_LISTOFFLOATS']._serialized_start=673 + _globals['_LISTOFFLOATS']._serialized_end=701 + _globals['_LISTOFLISTOFFLOATS']._serialized_start=703 + _globals['_LISTOFLISTOFFLOATS']._serialized_end=766 + _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_start=768 + _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_end=843 + _globals['_SAMPLESERVER']._serialized_start=846 + _globals['_SAMPLESERVER']._serialized_end=1212 # @@protoc_insertion_point(module_scope) diff --git a/cezo_grpc/sample_pb2_grpc.py b/cezo_grpc/sample_pb2_grpc.py index aea20cd..ce3c3a0 100644 --- a/cezo_grpc/sample_pb2_grpc.py +++ b/cezo_grpc/sample_pb2_grpc.py @@ -39,16 +39,16 @@ def __init__(self, channel): request_serializer=cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, response_deserializer=cezo__grpc_dot_sample__pb2.ConnectResponse.FromString, _registered_method=True) + 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, + _registered_method=True) 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, _registered_method=True) - self.PullGradsAndSeeds = channel.unary_unary( - '/sample_server.SampleServer/PullGradsAndSeeds', - request_serializer=cezo__grpc_dot_sample__pb2.PullGradsAndSeedsRequest.SerializeToString, - response_deserializer=cezo__grpc_dot_sample__pb2.PullGradsAndSeedsResponse.FromString, - _registered_method=True) self.SubmitIteration = channel.unary_unary( '/sample_server.SampleServer/SubmitIteration', request_serializer=cezo__grpc_dot_sample__pb2.SubmitIterationRequest.SerializeToString, @@ -65,13 +65,13 @@ def Connect(self, request, context): context.set_details('Method not implemented!') raise NotImplementedError('Method not implemented!') - def TryToJoinIteration(self, request, context): + 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 PullGradsAndSeeds(self, request, context): + def TryToJoinIteration(self, request, context): """Missing associated documentation comment in .proto file.""" context.set_code(grpc.StatusCode.UNIMPLEMENTED) context.set_details('Method not implemented!') @@ -91,16 +91,16 @@ def add_SampleServerServicer_to_server(servicer, server): 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, ), - 'PullGradsAndSeeds': grpc.unary_unary_rpc_method_handler( - servicer.PullGradsAndSeeds, - request_deserializer=cezo__grpc_dot_sample__pb2.PullGradsAndSeedsRequest.FromString, - response_serializer=cezo__grpc_dot_sample__pb2.PullGradsAndSeedsResponse.SerializeToString, - ), 'SubmitIteration': grpc.unary_unary_rpc_method_handler( servicer.SubmitIteration, request_deserializer=cezo__grpc_dot_sample__pb2.SubmitIterationRequest.FromString, @@ -145,7 +145,7 @@ def Connect(request, _registered_method=True) @staticmethod - def TryToJoinIteration(request, + def Disconnect(request, target, options=(), channel_credentials=None, @@ -158,9 +158,9 @@ def TryToJoinIteration(request, 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, + '/sample_server.SampleServer/Disconnect', + cezo__grpc_dot_sample__pb2.DisconnectRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, options, channel_credentials, insecure, @@ -172,7 +172,7 @@ def TryToJoinIteration(request, _registered_method=True) @staticmethod - def PullGradsAndSeeds(request, + def TryToJoinIteration(request, target, options=(), channel_credentials=None, @@ -185,9 +185,9 @@ def PullGradsAndSeeds(request, return grpc.experimental.unary_unary( request, target, - '/sample_server.SampleServer/PullGradsAndSeeds', - cezo__grpc_dot_sample__pb2.PullGradsAndSeedsRequest.SerializeToString, - cezo__grpc_dot_sample__pb2.PullGradsAndSeedsResponse.FromString, + '/sample_server.SampleServer/TryToJoinIteration', + cezo__grpc_dot_sample__pb2.TryToJoinIterationRequest.SerializeToString, + cezo__grpc_dot_sample__pb2.TryToJoinIterationResponse.FromString, options, channel_credentials, insecure, diff --git a/grpc_client.py b/grpc_client.py index 9f79836..8c54975 100644 --- a/grpc_client.py +++ b/grpc_client.py @@ -1,27 +1,117 @@ +from huggingface_hub.repository import atexit +from numpy import repeat +import torch import grpc +from tqdm import cli from cezo_grpc import sample_pb2 from cezo_grpc import sample_pb2_grpc from cezo_grpc import data_helper +from cezo_fl import fl_helpers from cezo_fl import client +import config +import preprocess +import decomfl_main +import time -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) -connect_result = ps_stub.Connect(sample_pb2.EmptyRequest()) +def setup_client(args, client_index): + device_map, train_loaders, _ = preprocess.preprocess(args) + client_index = 0 + client_name = fl_helpers.get_client_name(client_index) + client_device = device_map[client_name] + ( + client_model, + client_criterion, + client_optimizer, + client_grad_estimator, + client_accuracy_func, + ) = decomfl_main.prepare_settings_underseed(args, client_device) + client_model.to(client_device) -successful, client_index = connect_result.successful, connect_result.clientIndex -print(successful, client_index) + return client.ResetClient( + client_model, + train_loaders[client_index], + client_grad_estimator, + client_optimizer, + client_criterion, + client_accuracy_func, + client_device, + ) -join_result = ps_stub.TryToJoinIteration(sample_pb2.TryToJoinIterationRequest(clientIndex = client_index)) -print(join_result.successful) -grads_and_seeds = ps_stub.PullGradsAndSeeds(sample_pb2.PullGradsAndSeedsRequest(clientIndex = client_index)) +def repeat_every(fn, pass_fn, repeat_interval=1): + while True: + response = fn() + if pass_fn(response): + return response + time.sleep(repeat_interval) -print(data_helper.protobuf_to_py_list_of_list_of_list_of_floats(grads_and_seeds.grads), data_helper.protobuf_to_py_list_of_list_of_ints(grads_and_seeds.seeds)) -ps_stub.SubmitIteration(sample_pb2.SubmitIterationRequest(local_update_result = sample_pb2.ListOfListOfFloats(data=[sample_pb2.ListOfFloats(data=[0.123])]))) +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 -print('success') \ No newline at end of file + +if __name__ == "__main__": + args = config.get_params().parse_args() + if args.dataset == "shakespeare": + args.num_clients = 139 + print(args) + + ps_stub = get_stub() + + connect_result = repeat_every( + lambda: ps_stub.Connect(sample_pb2.EmptyRequest()), lambda x: x.successful + ) + print("connected") + client_index = connect_result.clientIndex + # when program exits, we need to disconnect this client from server + atexit.register( + lambda: ps_stub.Disconnect(sample_pb2.DisconnectRequest(clientIndex=client_index)) + ) + + 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) + ), + 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( + 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) + print("success") diff --git a/grpc_server.py b/grpc_server.py index b1dc6ce..25d4670 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -1,32 +1,201 @@ 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_fl import server +from byzantine.aggregation import mean +from byzantine.attack import no_byz + + +class ServerStatus(Enum): + connecting = "connecting" + training = "training" + aggregating = "aggregating" + + +def find_first(lst: list, check_fn) -> int: + try: + return next(i for i in range(len(lst)) if check_fn(lst)) + except StopIteration: + return -1 class SampleServer(sample_pb2_grpc.SampleServerServicer): - def Connect(self, request, context): - return sample_pb2.ConnectResponse( - successful = True, - clientIndex = 0 + def __init__(self): + self.num_clients = 1 + self.num_sample_clients = 1 + self.local_update_steps = 1 + self.seed_grad_records = server.SeedAndGradientRecords() + self.client_last_updates = [0 for _ in range(self.num_clients)] + + self.status = ServerStatus.connecting + + self.connect_lock = threading.Lock() + self.connected_clients: list[bool] = [False for _ in range(self.num_clients)] + + self.iteration: int = -1 + self.iteration_seeds: list[int] = [] + self.iteration_sampled_clients: list[int] = [] + self.iteration_local_grad_scalar: dict[int, list[torch.Tensor]] = {} + + def _get_next_connect_client_index(self) -> int: + return find_first(self.connected_clients, lambda x: x) + + def change_status(self, new_status: Enum) -> 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_local_grad_scalar = {} + + def try_swtich_from_connecting_to_training(self): + # 1. check for current status == connecting + if self.status is not ServerStatus.connecting: + return + # 2. check for connected_clients vs num_clients + if not all(self.connected_clients): + return + # 3. initialize training data for next iteration + print("swtich from connecting to training") + self.preprare_for_next_iteration() + # 4. change status to training + self.change_status(ServerStatus.training) + + def swtich_to_connecting(self): + if all(self.connected_clients): + return + print("Switch to Connecting") + self.change_status(ServerStatus.connecting) + + def _get_iteration_grad_scalar_list(self): + return [ + self.iteration_local_grad_scalar[client_index] + for client_index in self.iteration_sampled_clients + ] + + def _aggregate_and_update_server_record(self): + local_grad_scalar_list = no_byz(self._get_iteration_grad_scalar_list()) + grad_scalar = mean(self.num_sample_clients, 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. + self.seed_grad_records.remove_too_old(earliest_record_needs=min(self.client_last_updates)) + + def try_switch_from_training_to_aggregating(self): + # 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 all(self._get_iteration_grad_scalar_list()): + return + print("swtich 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 + self.preprare_for_next_iteration() + self.change_status(ServerStatus.training) + + def Connect(self, request, context): + if self.status is not ServerStatus.connecting: + return sample_pb2.ConnectResponse(successful=False, clientIndex=-1) + + self.connect_lock.acquire() + client_index = self._get_next_connect_client_index() + self.connected_clients[client_index] = True + self.try_swtich_from_connecting_to_training() + self.connect_lock.release() + 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 + self.connect_lock.acquire() + 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() + self.connect_lock.release() + + return sample_pb2.EmptyResponse() + def TryToJoinIteration(self, request, context): + client_index = request.clientIndex + if ( + self.status is not ServerStatus.training + or client_index not in self.iteration_sampled_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 + 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 PullGradsAndSeeds(self, request, context): - seeds = sample_pb2.ListOfListOfInts(data=[sample_pb2.ListOfInts(data=[1])]) - grads = sample_pb2.ListOfListOfListOfFloats(data=[sample_pb2.ListOfListOfFloats(data=[sample_pb2.ListOfFloats(data=[0.3])])]) - return sample_pb2.PullGradsAndSeedsResponse(seeds=seeds, grads=grads) + + # def PullGradsAndSeeds(self, request, context): + # seeds = sample_pb2.ListOfListOfInts(data=[sample_pb2.ListOfInts(data=[1])]) + # grads = sample_pb2.ListOfListOfListOfFloats( + # data=[sample_pb2.ListOfListOfFloats(data=[sample_pb2.ListOfFloats(data=[0.3])])] + # ) + # return sample_pb2.PullGradsAndSeedsResponse(seeds=seeds, grads=grads) + + 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): - print(data_helper.protobuf_to_py_list_of_list_of_floats(request.local_update_result)) + print("submit iteration", request.clientIndex) + client_index = request.clientIndex + + if ( + self.status is not ServerStatus.training + or client_index not in self.iteration_sampled_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.try_switch_from_training_to_aggregating() return sample_pb2.EmptyResponse() + def serve(rpc_master_port, rpc_num_workers): server = grpc.server(futures.ThreadPoolExecutor(max_workers=rpc_num_workers)) sample_pb2_grpc.add_SampleServerServicer_to_server(SampleServer(), server) From 8d6650c5af7444c51764a92a6fe76f942ad9a9c8 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sun, 15 Sep 2024 16:06:41 -0500 Subject: [PATCH 05/19] bug fix --- grpc_client.py | 2 +- grpc_server.py | 29 +++++++++++++++++++++++------ 2 files changed, 24 insertions(+), 7 deletions(-) diff --git a/grpc_client.py b/grpc_client.py index 8c54975..f942be7 100644 --- a/grpc_client.py +++ b/grpc_client.py @@ -67,8 +67,8 @@ def get_stub(): connect_result = repeat_every( lambda: ps_stub.Connect(sample_pb2.EmptyRequest()), lambda x: x.successful ) - print("connected") client_index = 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)) diff --git a/grpc_server.py b/grpc_server.py index 25d4670..d7bb304 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -21,15 +21,15 @@ class ServerStatus(Enum): def find_first(lst: list, check_fn) -> int: try: - return next(i for i in range(len(lst)) if check_fn(lst)) + 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): - self.num_clients = 1 - self.num_sample_clients = 1 + self.num_clients = 3 + self.num_sample_clients = 2 self.local_update_steps = 1 self.seed_grad_records = server.SeedAndGradientRecords() self.client_last_updates = [0 for _ in range(self.num_clients)] @@ -45,7 +45,7 @@ def __init__(self): self.iteration_local_grad_scalar: dict[int, list[torch.Tensor]] = {} def _get_next_connect_client_index(self) -> int: - return find_first(self.connected_clients, lambda x: x) + return find_first(self.connected_clients, lambda x: not x) def change_status(self, new_status: Enum) -> None: self.status = new_status @@ -58,6 +58,7 @@ def preprare_for_next_iteration(self): range(self.num_clients), self.num_sample_clients ) self.iteration_local_grad_scalar = {} + print(f"Iteration: {self.iteration}, sampled_clients: {self.iteration_sampled_clients}") def try_swtich_from_connecting_to_training(self): # 1. check for current status == connecting @@ -78,7 +79,16 @@ def swtich_to_connecting(self): print("Switch to Connecting") self.change_status(ServerStatus.connecting) + def _has_iteration_finished(self): + return all( + [ + self.iteration_local_grad_scalar.get(client_index) + for client_index in self.iteration_sampled_clients + ] + ) + def _get_iteration_grad_scalar_list(self): + print("getting _get_iteration_grad_scalar_list") return [ self.iteration_local_grad_scalar[client_index] for client_index in self.iteration_sampled_clients @@ -98,7 +108,7 @@ def try_switch_from_training_to_aggregating(self): if self.status is not ServerStatus.training: return # 2. check if sampled client all have return result - if not all(self._get_iteration_grad_scalar_list()): + if not self._has_iteration_finished(): return print("swtich from training to aggregating") # 3. change status to aggregating @@ -115,6 +125,11 @@ def Connect(self, request, context): self.connect_lock.acquire() 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) + + print(f"try to assign {client_index}") self.connected_clients[client_index] = True self.try_swtich_from_connecting_to_training() self.connect_lock.release() @@ -180,8 +195,8 @@ def _apply_client_local_update_result(self, client_index, local_update_result): self.iteration_local_grad_scalar[client_index] = local_update_result def SubmitIteration(self, request, context): - print("submit iteration", request.clientIndex) client_index = request.clientIndex + print(f"submit iteration from {client_index}") if ( self.status is not ServerStatus.training @@ -189,10 +204,12 @@ def SubmitIteration(self, request, context): ): return sample_pb2.EmptyResponse() + self.connect_lock.acquire() 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.try_switch_from_training_to_aggregating() + self.connect_lock.release() return sample_pb2.EmptyResponse() From 8e4ba584f002abc98034c88720f33828da091d12 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Fri, 20 Sep 2024 17:46:07 -0500 Subject: [PATCH 06/19] a working example --- cezo_fl/server.py | 1 + cezo_grpc/sample.proto | 11 ++- cezo_grpc/sample_pb2.py | 28 +++--- cezo_grpc/sample_pb2_grpc.py | 178 ++++++++++++++++++++++++++++++++++- grpc_client.py | 3 - grpc_eval_client.py | 103 ++++++++++++++++++++ grpc_server.py | 136 +++++++++++++++++++++----- grpc_server_test.py | 24 +++++ 8 files changed, 443 insertions(+), 41 deletions(-) create mode 100644 grpc_eval_client.py create mode 100644 grpc_server_test.py diff --git a/cezo_fl/server.py b/cezo_fl/server.py index 0198d6c..944ba8c 100644 --- a/cezo_fl/server.py +++ b/cezo_fl/server.py @@ -171,6 +171,7 @@ def train_one_step(self, iteration: int) -> tuple[float, float]: if self.server_model: assert self.optim assert self.random_gradient_estimator + print(seeds, grad_scalar) self.server_model.train() self.random_gradient_estimator.update_model_given_seed_and_grad( self.optim, diff --git a/cezo_grpc/sample.proto b/cezo_grpc/sample.proto index f581053..3c1cb68 100644 --- a/cezo_grpc/sample.proto +++ b/cezo_grpc/sample.proto @@ -6,7 +6,12 @@ service SampleServer { rpc Connect (EmptyRequest) returns (ConnectResponse) {} rpc Disconnect (DisconnectRequest) returns (EmptyResponse) {} rpc TryToJoinIteration (TryToJoinIterationRequest) returns (TryToJoinIterationResponse) {} - rpc SubmitIteration (SubmitIterationRequest) returns (EmptyRequest) {} + 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 {} @@ -40,6 +45,10 @@ message SubmitIterationRequest { float stepLoss = 4; } +message SubmitEvaluationRequest { + float evalAccuracy = 1; + float evalLoss = 2; +} message ListOfInts { repeated int32 data = 1; diff --git a/cezo_grpc/sample_pb2.py b/cezo_grpc/sample_pb2.py index b059528..61f7596 100644 --- a/cezo_grpc/sample_pb2.py +++ b/cezo_grpc/sample_pb2.py @@ -24,7 +24,7 @@ -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\"\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\xee\x02\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\x12W\n\x0fSubmitIteration\x12%.sample_server.SubmitIterationRequest\x1a\x1b.sample_server.EmptyRequest\"\x00\x62\x06proto3') +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) @@ -45,16 +45,18 @@ _globals['_TRYTOJOINITERATIONRESPONSE']._serialized_end=438 _globals['_SUBMITITERATIONREQUEST']._serialized_start=441 _globals['_SUBMITITERATIONREQUEST']._serialized_end=582 - _globals['_LISTOFINTS']._serialized_start=584 - _globals['_LISTOFINTS']._serialized_end=610 - _globals['_LISTOFLISTOFINTS']._serialized_start=612 - _globals['_LISTOFLISTOFINTS']._serialized_end=671 - _globals['_LISTOFFLOATS']._serialized_start=673 - _globals['_LISTOFFLOATS']._serialized_end=701 - _globals['_LISTOFLISTOFFLOATS']._serialized_start=703 - _globals['_LISTOFLISTOFFLOATS']._serialized_end=766 - _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_start=768 - _globals['_LISTOFLISTOFLISTOFFLOATS']._serialized_end=843 - _globals['_SAMPLESERVER']._serialized_start=846 - _globals['_SAMPLESERVER']._serialized_end=1212 + _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 index ce3c3a0..17f6889 100644 --- a/cezo_grpc/sample_pb2_grpc.py +++ b/cezo_grpc/sample_pb2_grpc.py @@ -52,7 +52,27 @@ def __init__(self, channel): 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.EmptyRequest.FromString, + response_deserializer=cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, + _registered_method=True) + 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, + _registered_method=True) + 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, + _registered_method=True) + 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, + _registered_method=True) + 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, _registered_method=True) @@ -83,6 +103,30 @@ def SubmitIteration(self, request, context): 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 = { @@ -104,7 +148,27 @@ def add_SampleServerServicer_to_server(servicer, server): '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.EmptyRequest.SerializeToString, + 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( @@ -214,7 +278,115 @@ def SubmitIteration(request, target, '/sample_server.SampleServer/SubmitIteration', cezo__grpc_dot_sample__pb2.SubmitIterationRequest.SerializeToString, - cezo__grpc_dot_sample__pb2.EmptyRequest.FromString, + cezo__grpc_dot_sample__pb2.EmptyResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @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, + _registered_method=True) + + @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, + _registered_method=True) + + @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, + _registered_method=True) + + @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, diff --git a/grpc_client.py b/grpc_client.py index f942be7..9b08c39 100644 --- a/grpc_client.py +++ b/grpc_client.py @@ -1,8 +1,6 @@ from huggingface_hub.repository import atexit -from numpy import repeat import torch import grpc -from tqdm import cli from cezo_grpc import sample_pb2 from cezo_grpc import sample_pb2_grpc from cezo_grpc import data_helper @@ -17,7 +15,6 @@ def setup_client(args, client_index): device_map, train_loaders, _ = preprocess.preprocess(args) - client_index = 0 client_name = fl_helpers.get_client_name(client_index) client_device = device_map[client_name] ( diff --git a/grpc_eval_client.py b/grpc_eval_client.py new file mode 100644 index 0000000..8b1d7c3 --- /dev/null +++ b/grpc_eval_client.py @@ -0,0 +1,103 @@ +from huggingface_hub.repository import atexit +import torch + +from cezo_grpc import sample_pb2 +from cezo_grpc import data_helper + +import grpc_client + +import config +import preprocess +import decomfl_main + +from shared.metrics import Metric + + +def setup_eval_model(args): + device_map, _, test_loader = preprocess.preprocess(args) + device = device_map["server"] + ( + model, + criterion, + optimizer, + grad_estimator, + accuracy_func, + ) = decomfl_main.prepare_settings_underseed(args, device) + + 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 + batch_labels = batch_labels.to(device) + pred = grad_estimator.model_forward(batch_inputs) + eval_loss.update(criterion(pred, batch_labels)) + eval_accuracy.update(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 + + +if __name__ == "__main__": + args = config.get_params().parse_args() + if args.dataset == "shakespeare": + args.num_clients = 139 + print(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 + ] + print(tensor_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) + print("success") diff --git a/grpc_server.py b/grpc_server.py index d7bb304..f744b6d 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -17,6 +17,7 @@ class ServerStatus(Enum): connecting = "connecting" training = "training" aggregating = "aggregating" + evaluating = "evaluating" def find_first(lst: list, check_fn) -> int: @@ -28,17 +29,23 @@ def find_first(lst: list, check_fn) -> int: class SampleServer(sample_pb2_grpc.SampleServerServicer): def __init__(self): - self.num_clients = 3 - self.num_sample_clients = 2 + self.should_eval = True + self.eval_iteration = 5 + + self.num_clients = 1 + self.num_sample_clients = 1 self.local_update_steps = 1 + self.seed_grad_records = server.SeedAndGradientRecords() self.client_last_updates = [0 for _ in range(self.num_clients)] - self.status = ServerStatus.connecting self.connect_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] = [] @@ -60,12 +67,17 @@ def preprare_for_next_iteration(self): self.iteration_local_grad_scalar = {} print(f"Iteration: {self.iteration}, sampled_clients: {self.iteration_sampled_clients}") - def try_swtich_from_connecting_to_training(self): + 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 not all(self.connected_clients): + if self._should_connect(): return # 3. initialize training data for next iteration print("swtich from connecting to training") @@ -73,13 +85,13 @@ def try_swtich_from_connecting_to_training(self): # 4. change status to training self.change_status(ServerStatus.training) - def swtich_to_connecting(self): - if all(self.connected_clients): + 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): + def _has_iteration_finished(self) -> bool: return all( [ self.iteration_local_grad_scalar.get(client_index) @@ -87,23 +99,27 @@ def _has_iteration_finished(self): ] ) - def _get_iteration_grad_scalar_list(self): - print("getting _get_iteration_grad_scalar_list") + 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): + def _aggregate_and_update_server_record(self) -> None: local_grad_scalar_list = no_byz(self._get_iteration_grad_scalar_list()) grad_scalar = mean(self.num_sample_clients, 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. - self.seed_grad_records.remove_too_old(earliest_record_needs=min(self.client_last_updates)) + 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): + def try_switch_from_training_to_aggregating(self) -> None: # 1. check for current status == training if self.status is not ServerStatus.training: return @@ -115,7 +131,23 @@ def try_switch_from_training_to_aggregating(self): self.change_status(ServerStatus.aggregating) # 4. update seed_grad_records self._aggregate_and_update_server_record() - # 5. change to next training + # 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("swtich from aggregating to training") + self.preprare_for_next_iteration() + self.change_status(ServerStatus.training) + + def switch_from_aggregating_to_evaluating(self) -> None: + print("swtich from aggregating to evaluating") + self.change_status(ServerStatus.evaluating) + + def switch_from_evaluating_to_training(self) -> None: + print("swtich from evaluating to training") self.preprare_for_next_iteration() self.change_status(ServerStatus.training) @@ -180,13 +212,6 @@ def TryToJoinIteration(self, request, context): iterationSeeds=data_helper.py_to_protobuf_list_of_ints(self.iteration_seeds), ) - # def PullGradsAndSeeds(self, request, context): - # seeds = sample_pb2.ListOfListOfInts(data=[sample_pb2.ListOfInts(data=[1])]) - # grads = sample_pb2.ListOfListOfListOfFloats( - # data=[sample_pb2.ListOfListOfFloats(data=[sample_pb2.ListOfFloats(data=[0.3])])] - # ) - # return sample_pb2.PullGradsAndSeedsResponse(seeds=seeds, grads=grads) - def _apply_client_local_update_result(self, client_index, local_update_result): if client_index not in self.iteration_sampled_clients: return @@ -212,6 +237,75 @@ def SubmitIteration(self, request, context): self.connect_lock.release() return sample_pb2.EmptyResponse() + def ConnectEval(self, request, context): + if self.status is not ServerStatus.connecting: + return sample_pb2.ConnectResponse(successful=False, clientIndex=-1) + + self.connect_lock.acquire() + if self.eval_client_connected: + print("A eval client is already connected!") + return sample_pb2.ConnectResponse(successful=False, clientIndex=-1) + + print("Eval Client connected") + self.eval_client_connected = True + self.try_swtich_from_connecting_to_training() + self.connect_lock.release() + 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 + self.connect_lock.acquire() + 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() + self.connect_lock.release() + + def TryToEval(self, request, context): + if self.status is not ServerStatus.evaluating: + print("try to eval") + 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. + print("before return try to eval 1", self.eval_client_last_update) + 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) + print("seeds_list, grad_list", seeds_list, grad_list) + self.eval_client_last_update = self.iteration + print("before return try to eval 4") + 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): + print("submit evaluation") + 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(rpc_master_port, rpc_num_workers): server = grpc.server(futures.ThreadPoolExecutor(max_workers=rpc_num_workers)) diff --git a/grpc_server_test.py b/grpc_server_test.py new file mode 100644 index 0000000..cc18b8f --- /dev/null +++ b/grpc_server_test.py @@ -0,0 +1,24 @@ +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(): + server = SampleServer() + pass From ba89ed05ddb4d628b67f1a5954f4617f6216c2c3 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Fri, 20 Sep 2024 17:47:06 -0500 Subject: [PATCH 07/19] do not print --- cezo_fl/server.py | 1 - 1 file changed, 1 deletion(-) diff --git a/cezo_fl/server.py b/cezo_fl/server.py index 944ba8c..0198d6c 100644 --- a/cezo_fl/server.py +++ b/cezo_fl/server.py @@ -171,7 +171,6 @@ def train_one_step(self, iteration: int) -> tuple[float, float]: if self.server_model: assert self.optim assert self.random_gradient_estimator - print(seeds, grad_scalar) self.server_model.train() self.random_gradient_estimator.update_model_given_seed_and_grad( self.optim, From 89b5331a67b55ed58b13ee4bfd1321fa5d039c08 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Fri, 20 Sep 2024 18:12:56 -0500 Subject: [PATCH 08/19] adjust default parameters in server --- grpc_server.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/grpc_server.py b/grpc_server.py index f744b6d..934093a 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -30,10 +30,10 @@ def find_first(lst: list, check_fn) -> int: class SampleServer(sample_pb2_grpc.SampleServerServicer): def __init__(self): self.should_eval = True - self.eval_iteration = 5 + self.eval_iteration = 25 - self.num_clients = 1 - self.num_sample_clients = 1 + self.num_clients = 3 + self.num_sample_clients = 2 self.local_update_steps = 1 self.seed_grad_records = server.SeedAndGradientRecords() From 4e3fe1a37adcb0801e55f14c9ad5ab524e82528d Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sat, 21 Sep 2024 10:04:42 -0500 Subject: [PATCH 09/19] lock improvement --- grpc_server.py | 234 +++++++++++++++++++++++++------------------------ 1 file changed, 118 insertions(+), 116 deletions(-) diff --git a/grpc_server.py b/grpc_server.py index 934093a..de4bc01 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -18,6 +18,8 @@ class ServerStatus(Enum): 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: @@ -40,7 +42,8 @@ def __init__(self): self.client_last_updates = [0 for _ in range(self.num_clients)] self.status = ServerStatus.connecting - self.connect_lock = threading.Lock() + self.lock = threading.Lock() + self.connected_clients: list[bool] = [False for _ in range(self.num_clients)] self.eval_client_connected: bool = False @@ -51,6 +54,12 @@ def __init__(self): self.iteration_sampled_clients: list[int] = [] 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) @@ -152,66 +161,65 @@ def switch_from_evaluating_to_training(self) -> None: self.change_status(ServerStatus.training) def Connect(self, request, context): - if self.status is not ServerStatus.connecting: - return sample_pb2.ConnectResponse(successful=False, clientIndex=-1) + with self.lock: + if self.status is not ServerStatus.connecting: + return sample_pb2.ConnectResponse(successful=False, clientIndex=-1) - self.connect_lock.acquire() - 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) + 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) - print(f"try to assign {client_index}") - self.connected_clients[client_index] = True - self.try_swtich_from_connecting_to_training() - self.connect_lock.release() - return sample_pb2.ConnectResponse(successful=True, 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 - self.connect_lock.acquire() - 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() - self.connect_lock.release() + 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() + return sample_pb2.EmptyResponse() def TryToJoinIteration(self, request, context): - client_index = request.clientIndex - if ( - self.status is not ServerStatus.training - or client_index not in self.iteration_sampled_clients - ): + with self.lock: + client_index = request.clientIndex + if ( + self.status is not ServerStatus.training + or client_index not in self.iteration_sampled_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=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([]), + 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), ) - 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 @@ -220,91 +228,85 @@ def _apply_client_local_update_result(self, client_index, local_update_result): self.iteration_local_grad_scalar[client_index] = local_update_result def SubmitIteration(self, request, context): - 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 - ): + 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 + ): + 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.try_switch_from_training_to_aggregating() return sample_pb2.EmptyResponse() - self.connect_lock.acquire() - 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.try_switch_from_training_to_aggregating() - self.connect_lock.release() - return sample_pb2.EmptyResponse() - def ConnectEval(self, request, context): - if self.status is not ServerStatus.connecting: - return sample_pb2.ConnectResponse(successful=False, clientIndex=-1) + with self.lock: + if self.status is not ServerStatus.connecting: + return sample_pb2.ConnectResponse(successful=False, clientIndex=-1) - self.connect_lock.acquire() - if self.eval_client_connected: - print("A eval client is already connected!") - 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) - print("Eval Client connected") - self.eval_client_connected = True - self.try_swtich_from_connecting_to_training() - self.connect_lock.release() - return sample_pb2.ConnectResponse(successful=True, 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 - self.connect_lock.acquire() - 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 + 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() - self.connect_lock.release() + self.swtich_to_connecting() def TryToEval(self, request, context): - if self.status is not ServerStatus.evaluating: - print("try to eval") + 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=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([]), + 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([]), ) - # 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. - print("before return try to eval 1", self.eval_client_last_update) - 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) - print("seeds_list, grad_list", seeds_list, grad_list) - self.eval_client_last_update = self.iteration - print("before return try to eval 4") - 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): - print("submit evaluation") - 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() + 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() + return sample_pb2.EmptyResponse() def serve(rpc_master_port, rpc_num_workers): From 2e2ed0c11ceb0171a68bbd81e8cbf22be502dbaf Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sun, 22 Sep 2024 17:01:59 -0500 Subject: [PATCH 10/19] use 1 script to run grpc in multi process --- cezo_grpc/README.md | 2 ++ grpc_client.py | 16 +++++++++------- grpc_eval_client.py | 18 ++++++++++-------- grpc_server.py | 23 ++++++++++++++++------- run_grpc.py | 42 ++++++++++++++++++++++++++++++++++++++++++ 5 files changed, 79 insertions(+), 22 deletions(-) create mode 100644 run_grpc.py diff --git a/cezo_grpc/README.md b/cezo_grpc/README.md index c8fc573..81821d9 100644 --- a/cezo_grpc/README.md +++ b/cezo_grpc/README.md @@ -1 +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/grpc_client.py b/grpc_client.py index 9b08c39..e43e3e3 100644 --- a/grpc_client.py +++ b/grpc_client.py @@ -53,12 +53,7 @@ def get_stub(): return ps_stub -if __name__ == "__main__": - args = config.get_params().parse_args() - if args.dataset == "shakespeare": - args.num_clients = 139 - print(args) - +def train_with_args(args): ps_stub = get_stub() connect_result = repeat_every( @@ -111,4 +106,11 @@ def try_to_join_iteration(): ) repeat_every(try_to_join_iteration, lambda x: False) - print("success") + + +if __name__ == "__main__": + args = config.get_params().parse_args() + if args.dataset == "shakespeare": + args.num_clients = 139 + print(args) + train_with_args(args) diff --git a/grpc_eval_client.py b/grpc_eval_client.py index 8b1d7c3..a0d7e13 100644 --- a/grpc_eval_client.py +++ b/grpc_eval_client.py @@ -54,12 +54,7 @@ def eval_model(): return update_model, eval_model, device -if __name__ == "__main__": - args = config.get_params().parse_args() - if args.dataset == "shakespeare": - args.num_clients = 139 - print(args) - +def eval_with_args(args): ps_stub = grpc_client.get_stub() grpc_client.repeat_every( @@ -89,7 +84,6 @@ def try_to_eval(): tensor_grad_list = [ [torch.tensor(v, device=device) for v in vv] for vv in raw_grad_list ] - print(tensor_grad_list) update_model(pull_seeds_list, tensor_grad_list) eval_loss, eval_accuracy = eval_model() @@ -100,4 +94,12 @@ def try_to_eval(): ) grpc_client.repeat_every(try_to_eval, lambda x: False) - print("success") + + +if __name__ == "__main__": + args = config.get_params().parse_args() + if args.dataset == "shakespeare": + args.num_clients = 139 + print(args) + + eval_with_args(args) diff --git a/grpc_server.py b/grpc_server.py index de4bc01..c8c1c72 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -11,6 +11,7 @@ from cezo_fl import server from byzantine.aggregation import mean from byzantine.attack import no_byz +import config class ServerStatus(Enum): @@ -30,13 +31,13 @@ def find_first(lst: list, check_fn) -> int: class SampleServer(sample_pb2_grpc.SampleServerServicer): - def __init__(self): + def __init__(self, num_clients, num_sample_clients, local_update_steps): self.should_eval = True self.eval_iteration = 25 - self.num_clients = 3 - self.num_sample_clients = 2 - self.local_update_steps = 1 + 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)] @@ -309,9 +310,13 @@ def SubmitEvaluation(self, request, context): return sample_pb2.EmptyResponse() -def serve(rpc_master_port, rpc_num_workers): +def serve(args): + rpc_master_port = 4242 + rpc_num_workers = 8 server = grpc.server(futures.ThreadPoolExecutor(max_workers=rpc_num_workers)) - sample_pb2_grpc.add_SampleServerServicer_to_server(SampleServer(), server) + sample_pb2_grpc.add_SampleServerServicer_to_server( + SampleServer(args.num_clients, args.num_sample_clients, args.local_update_steps), server + ) server.add_insecure_port(f"localhost:{rpc_master_port}") print(f"Parameter server starting on [::]:{rpc_master_port}") server.start() @@ -319,4 +324,8 @@ def serve(rpc_master_port, rpc_num_workers): if __name__ == "__main__": - serve(4242, 8) + args = config.get_params().parse_args() + if args.dataset == "shakespeare": + args.num_clients = 139 + print(args) + serve(args) diff --git a/run_grpc.py b/run_grpc.py new file mode 100644 index 0000000..a9ae295 --- /dev/null +++ b/run_grpc.py @@ -0,0 +1,42 @@ +from multiprocessing import Process + +from grpc_client import train_with_args +from grpc_eval_client import eval_with_args +from grpc_server import serve +import config + + +if __name__ == "__main__": + args = config.get_params().parse_args() + if args.dataset == "shakespeare": + args.num_clients = 139 + + Process(target=serve, args=(args,)).start() + + for _ in range(args.num_clients): + Process(target=train_with_args, args=(args,)).start() + + Process(target=eval_with_args, args=(args,)).start() + + +# from multiprocessing import Process +# import os + + +# def info(title): +# print(title) +# print("module name:", __name__) +# print("parent process:", os.getppid()) +# print("process id:", os.getpid()) + + +# def f(name): +# info("function f") +# print("hello", name) + + +# if __name__ == "__main__": +# info("main line") +# p = Process(target=f, args=("bob",)) +# p.start() +# p.join() From 4227ce45715dea74754fcd6d18785770e31e1c01 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sun, 22 Sep 2024 19:06:47 -0500 Subject: [PATCH 11/19] update environment --- README.md | 5 +++-- environment.yml | 3 +++ environment_cuda.yml | 3 +++ 3 files changed, 9 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index e7d76ed..b64da55 100644 --- a/README.md +++ b/README.md @@ -35,6 +35,7 @@ For READMD.md, we will use `environment.yml` whenever a environment file is need 3. Once installation is finished, run `conda activate decomfl` to use the created virtual env. 4. (Optional) If you see something like `conda init before activate`. Run `conda init`, then restart your terminal/powershell. Then repeat step 3. 5. Run any command provided in [Run Experiments](#run-experiments) section. If code works, then congratulations, you have successfully set up the environment for this repo! +6. Update the environemtn if there are some missing dependencies, most recent change was introduced by adding grpc. Try `conda env update --file environment.yml --prune`. The `--prune` option causes conda to remove any dependencies that are no longer required from the environment. ## Run Experiments @@ -44,7 +45,6 @@ For READMD.md, we will use `environment.yml` whenever a environment file is need - **Run DeComFL:** Follow FL routine, split data into chunks and train on different clients. Usage example: `python decomfl_main.py --large-model=opt-125m --dataset=sst2 --iterations=1000 --train-batch-size=32 --test-batch-size=200 --eval-iterations=25 --num-clients=3 --num-sample-clients=2 --local-update-steps=1 --num-pert=5 --lr=1e-5 --mu=1e-3 --grad-estimate-method=rge-forward` - ## Citation ``` @@ -57,7 +57,8 @@ For READMD.md, we will use `environment.yml` whenever a environment file is need ``` ## Our Team -DeComFL is currently contributed and maintained by **Zidong Liu** (ComboCurve), **Bicheng Ying** (Google) and **Zhe Li** (RIT), and advised by Prof. **Haibo Yang** (RIT). + +DeComFL is currently contributed and maintained by **Zidong Liu** (ComboCurve), **Bicheng Ying** (Google) and **Zhe Li** (RIT), and advised by Prof. **Haibo Yang** (RIT).
Image 1 diff --git a/environment.yml b/environment.yml index 2c28742..7fee4fd 100644 --- a/environment.yml +++ b/environment.yml @@ -17,6 +17,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 4d8a0bb..352a9c6 100644 --- a/environment_cuda.yml +++ b/environment_cuda.yml @@ -19,6 +19,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 From ad8b40c8d946146d7120ada97eb292ecd5c80912 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sun, 22 Sep 2024 19:07:06 -0500 Subject: [PATCH 12/19] rebuild protobuf --- cezo_grpc/sample_pb2.py | 16 +--- cezo_grpc/sample_pb2_grpc.py | 165 +++++++---------------------------- 2 files changed, 35 insertions(+), 146 deletions(-) diff --git a/cezo_grpc/sample_pb2.py b/cezo_grpc/sample_pb2.py index 61f7596..00f098f 100644 --- a/cezo_grpc/sample_pb2.py +++ b/cezo_grpc/sample_pb2.py @@ -1,22 +1,12 @@ # -*- coding: utf-8 -*- # Generated by the protocol buffer compiler. DO NOT EDIT! -# NO CHECKED-IN PROTOBUF GENCODE # source: cezo_grpc/sample.proto -# Protobuf Python Version: 5.27.2 +# 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 runtime_version as _runtime_version from google.protobuf import symbol_database as _symbol_database from google.protobuf.internal import builder as _builder -_runtime_version.ValidateProtobufRuntimeVersion( - _runtime_version.Domain.PUBLIC, - 5, - 27, - 2, - '', - 'cezo_grpc/sample.proto' -) # @@protoc_insertion_point(imports) _sym_db = _symbol_database.Default() @@ -29,8 +19,8 @@ _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) _builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'cezo_grpc.sample_pb2', _globals) -if not _descriptor._USE_C_DESCRIPTORS: - DESCRIPTOR._loaded_options = None +if _descriptor._USE_C_DESCRIPTORS == False: + DESCRIPTOR._options = None _globals['_EMPTYREQUEST']._serialized_start=41 _globals['_EMPTYREQUEST']._serialized_end=55 _globals['_EMPTYRESPONSE']._serialized_start=57 diff --git a/cezo_grpc/sample_pb2_grpc.py b/cezo_grpc/sample_pb2_grpc.py index 17f6889..9f40cb0 100644 --- a/cezo_grpc/sample_pb2_grpc.py +++ b/cezo_grpc/sample_pb2_grpc.py @@ -1,29 +1,9 @@ # Generated by the gRPC Python protocol compiler plugin. DO NOT EDIT! """Client and server classes corresponding to protobuf-defined services.""" import grpc -import warnings from cezo_grpc import sample_pb2 as cezo__grpc_dot_sample__pb2 -GRPC_GENERATED_VERSION = '1.66.1' -GRPC_VERSION = grpc.__version__ -_version_not_supported = False - -try: - from grpc._utilities import first_version_is_lower - _version_not_supported = first_version_is_lower(GRPC_VERSION, GRPC_GENERATED_VERSION) -except ImportError: - _version_not_supported = True - -if _version_not_supported: - raise RuntimeError( - f'The grpc package installed is at version {GRPC_VERSION},' - + f' but the generated code in cezo_grpc/sample_pb2_grpc.py depends on' - + f' grpcio>={GRPC_GENERATED_VERSION}.' - + f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}' - + f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.' - ) - class SampleServerStub(object): """Missing associated documentation comment in .proto file.""" @@ -38,42 +18,42 @@ def __init__(self, channel): '/sample_server.SampleServer/Connect', request_serializer=cezo__grpc_dot_sample__pb2.EmptyRequest.SerializeToString, response_deserializer=cezo__grpc_dot_sample__pb2.ConnectResponse.FromString, - _registered_method=True) + ) 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, - _registered_method=True) + ) 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, - _registered_method=True) + ) 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, - _registered_method=True) + ) 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, - _registered_method=True) + ) 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, - _registered_method=True) + ) 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, - _registered_method=True) + ) 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, - _registered_method=True) + ) class SampleServerServicer(object): @@ -174,7 +154,6 @@ def add_SampleServerServicer_to_server(servicer, server): generic_handler = grpc.method_handlers_generic_handler( 'sample_server.SampleServer', rpc_method_handlers) server.add_generic_rpc_handlers((generic_handler,)) - server.add_registered_method_handlers('sample_server.SampleServer', rpc_method_handlers) # This class is part of an EXPERIMENTAL API. @@ -192,21 +171,11 @@ def Connect(request, wait_for_ready=None, timeout=None, metadata=None): - return grpc.experimental.unary_unary( - request, - target, - '/sample_server.SampleServer/Connect', + 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, - _registered_method=True) + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) @staticmethod def Disconnect(request, @@ -219,21 +188,11 @@ def Disconnect(request, wait_for_ready=None, timeout=None, metadata=None): - return grpc.experimental.unary_unary( - request, - target, - '/sample_server.SampleServer/Disconnect', + 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, - _registered_method=True) + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) @staticmethod def TryToJoinIteration(request, @@ -246,21 +205,11 @@ def TryToJoinIteration(request, wait_for_ready=None, timeout=None, metadata=None): - return grpc.experimental.unary_unary( - request, - target, - '/sample_server.SampleServer/TryToJoinIteration', + 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, - _registered_method=True) + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) @staticmethod def SubmitIteration(request, @@ -273,21 +222,11 @@ def SubmitIteration(request, wait_for_ready=None, timeout=None, metadata=None): - return grpc.experimental.unary_unary( - request, - target, - '/sample_server.SampleServer/SubmitIteration', + 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, - _registered_method=True) + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) @staticmethod def ConnectEval(request, @@ -300,21 +239,11 @@ def ConnectEval(request, wait_for_ready=None, timeout=None, metadata=None): - return grpc.experimental.unary_unary( - request, - target, - '/sample_server.SampleServer/ConnectEval', + 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, - _registered_method=True) + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) @staticmethod def DisconnectEval(request, @@ -327,21 +256,11 @@ def DisconnectEval(request, wait_for_ready=None, timeout=None, metadata=None): - return grpc.experimental.unary_unary( - request, - target, - '/sample_server.SampleServer/DisconnectEval', + 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, - _registered_method=True) + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) @staticmethod def TryToEval(request, @@ -354,21 +273,11 @@ def TryToEval(request, wait_for_ready=None, timeout=None, metadata=None): - return grpc.experimental.unary_unary( - request, - target, - '/sample_server.SampleServer/TryToEval', + 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, - _registered_method=True) + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) @staticmethod def SubmitEvaluation(request, @@ -381,18 +290,8 @@ def SubmitEvaluation(request, wait_for_ready=None, timeout=None, metadata=None): - return grpc.experimental.unary_unary( - request, - target, - '/sample_server.SampleServer/SubmitEvaluation', + 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, - _registered_method=True) + options, channel_credentials, + insecure, call_credentials, compression, wait_for_ready, timeout, metadata) From 95dc435d53291fbda2e055d2a5424d5cd30c2852 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sun, 22 Sep 2024 19:19:05 -0500 Subject: [PATCH 13/19] make seperate args for grpc; allow address and port to change --- config.py | 23 +++++++++++++++++++++++ grpc_client.py | 2 +- grpc_eval_client.py | 2 +- grpc_server.py | 12 ++++++------ run_grpc.py | 29 +++++------------------------ 5 files changed, 36 insertions(+), 32 deletions(-) diff --git a/config.py b/config.py index 740a043..e4ba5af 100644 --- a/config.py +++ b/config.py @@ -156,6 +156,29 @@ def get_params(): return parser +def get_params_grpc(): + parser = get_params() + parser.add_argument( + "--rpc_master_addr", + type=str, + default="localhost", + help="Address of the RPC master node (the parameter server).", + ) + parser.add_argument( + "--rpc_master_port", + type=int, + default=4242, + help="Port of the RPC master node (the parameter server).", + ) + parser.add_argument( + "--rpc_num_workers", + type=int, + default=8, + help="Number of workers for training, excluding the parameter server.", + ) + return parser + + def get_args_dict(args): return {key: getattr(args, key) for key in DEFAULTS.keys()} diff --git a/grpc_client.py b/grpc_client.py index e43e3e3..de47265 100644 --- a/grpc_client.py +++ b/grpc_client.py @@ -109,7 +109,7 @@ def try_to_join_iteration(): if __name__ == "__main__": - args = config.get_params().parse_args() + args = config.get_params_grpc().parse_args() if args.dataset == "shakespeare": args.num_clients = 139 print(args) diff --git a/grpc_eval_client.py b/grpc_eval_client.py index a0d7e13..46ab6c6 100644 --- a/grpc_eval_client.py +++ b/grpc_eval_client.py @@ -97,7 +97,7 @@ def try_to_eval(): if __name__ == "__main__": - args = config.get_params().parse_args() + args = config.get_params_grpc().parse_args() if args.dataset == "shakespeare": args.num_clients = 139 print(args) diff --git a/grpc_server.py b/grpc_server.py index c8c1c72..9e6ae11 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -311,20 +311,20 @@ def SubmitEvaluation(self, request, context): def serve(args): - rpc_master_port = 4242 - rpc_num_workers = 8 - server = grpc.server(futures.ThreadPoolExecutor(max_workers=rpc_num_workers)) + 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 ) - server.add_insecure_port(f"localhost:{rpc_master_port}") - print(f"Parameter server starting on [::]:{rpc_master_port}") + port_str = f"{args.rpc_master_addr}:{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 = config.get_params().parse_args() + args = config.get_params_grpc().parse_args() if args.dataset == "shakespeare": args.num_clients = 139 print(args) diff --git a/run_grpc.py b/run_grpc.py index a9ae295..21cfa26 100644 --- a/run_grpc.py +++ b/run_grpc.py @@ -5,38 +5,19 @@ from grpc_server import serve import config +import time if __name__ == "__main__": - args = config.get_params().parse_args() + args = config.get_params_grpc().parse_args() if args.dataset == "shakespeare": args.num_clients = 139 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() - - -# from multiprocessing import Process -# import os - - -# def info(title): -# print(title) -# print("module name:", __name__) -# print("parent process:", os.getppid()) -# print("process id:", os.getpid()) - - -# def f(name): -# info("function f") -# print("hello", name) - - -# if __name__ == "__main__": -# info("main line") -# p = Process(target=f, args=("bob",)) -# p.start() -# p.join() From ed1fa246606577f4ce3da92827f68bb2e492322b Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sat, 5 Oct 2024 09:58:45 -0500 Subject: [PATCH 14/19] make sure grpc still works --- grpc_eval_client.py | 2 +- grpc_server.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/grpc_eval_client.py b/grpc_eval_client.py index 46ab6c6..2f8c61f 100644 --- a/grpc_eval_client.py +++ b/grpc_eval_client.py @@ -10,7 +10,7 @@ import preprocess import decomfl_main -from shared.metrics import Metric +from cezo_fl.util.metrics import Metric def setup_eval_model(args): diff --git a/grpc_server.py b/grpc_server.py index 9e6ae11..6db22e3 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -117,7 +117,7 @@ def _get_iteration_grad_scalar_list(self) -> list[list[torch.Tensor]]: def _aggregate_and_update_server_record(self) -> None: local_grad_scalar_list = no_byz(self._get_iteration_grad_scalar_list()) - grad_scalar = mean(self.num_sample_clients, local_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 From 4560555d14a6ef58ed286c2690c8b256f6fe02cb Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sat, 19 Jul 2025 11:05:56 -0500 Subject: [PATCH 15/19] make it work again --- cezo_grpc/cli_interface.py | 35 +++++++++++++++++++ experiment_helper/cli_parser.py | 19 ++++++++++ grpc_client.py | 54 +++++++++++++++++------------ grpc_eval_client.py | 61 +++++++++++++++++++++------------ grpc_server.py | 9 +++-- run_grpc.py | 7 ++-- test_model.py | 7 ---- 7 files changed, 133 insertions(+), 59 deletions(-) create mode 100644 cezo_grpc/cli_interface.py delete mode 100644 test_model.py 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/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 index de47265..e7ce02b 100644 --- a/grpc_client.py +++ b/grpc_client.py @@ -1,38 +1,52 @@ 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 -import config -import preprocess -import decomfl_main -import time +from experiment_helper import prepare_settings +from experiment_helper.device import use_device +from experiment_helper.data import get_dataloaders -def setup_client(args, client_index): - device_map, train_loaders, _ = preprocess.preprocess(args) + +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, - client_criterion, - client_optimizer, - client_grad_estimator, - client_accuracy_func, - ) = decomfl_main.prepare_settings_underseed(args, client_device) - client_model.to(client_device) + 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, - client_criterion, - client_accuracy_func, + metrics.train_loss, + metrics.train_acc, client_device, ) @@ -53,13 +67,13 @@ def get_stub(): return ps_stub -def train_with_args(args): +def train_with_args(args: cli_interface.CliSetting): ps_stub = get_stub() connect_result = repeat_every( lambda: ps_stub.Connect(sample_pb2.EmptyRequest()), lambda x: x.successful ) - client_index = connect_result.clientIndex + 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( @@ -109,8 +123,6 @@ def try_to_join_iteration(): if __name__ == "__main__": - args = config.get_params_grpc().parse_args() - if args.dataset == "shakespeare": - args.num_clients = 139 + args = cli_interface.CliSetting() print(args) train_with_args(args) diff --git a/grpc_eval_client.py b/grpc_eval_client.py index 2f8c61f..7ac5d87 100644 --- a/grpc_eval_client.py +++ b/grpc_eval_client.py @@ -6,23 +6,42 @@ import grpc_client -import config -import preprocess -import decomfl_main - 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): - device_map, _, test_loader = preprocess.preprocess(args) - device = device_map["server"] - ( - model, - criterion, - optimizer, - grad_estimator, - accuracy_func, - ) = decomfl_main.prepare_settings_underseed(args, device) +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() @@ -42,10 +61,11 @@ def eval_model(): 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 - batch_labels = batch_labels.to(device) - pred = grad_estimator.model_forward(batch_inputs) - eval_loss.update(criterion(pred, batch_labels)) - eval_accuracy.update(accuracy_func(pred, batch_labels)) + 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}%", ) @@ -97,9 +117,6 @@ def try_to_eval(): if __name__ == "__main__": - args = config.get_params_grpc().parse_args() - if args.dataset == "shakespeare": - args.num_clients = 139 + args = cli_interface.CliSetting() print(args) - eval_with_args(args) diff --git a/grpc_server.py b/grpc_server.py index 6db22e3..3462818 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -8,10 +8,11 @@ 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 -import config class ServerStatus(Enum): @@ -310,7 +311,7 @@ def SubmitEvaluation(self, request, context): return sample_pb2.EmptyResponse() -def serve(args): +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( @@ -324,8 +325,6 @@ def serve(args): if __name__ == "__main__": - args = config.get_params_grpc().parse_args() - if args.dataset == "shakespeare": - args.num_clients = 139 + args = cli_interface.CliSetting() print(args) serve(args) diff --git a/run_grpc.py b/run_grpc.py index 21cfa26..6228527 100644 --- a/run_grpc.py +++ b/run_grpc.py @@ -3,14 +3,13 @@ from grpc_client import train_with_args from grpc_eval_client import eval_with_args from grpc_server import serve -import config +from cezo_grpc import cli_interface import time if __name__ == "__main__": - args = config.get_params_grpc().parse_args() - if args.dataset == "shakespeare": - args.num_clients = 139 + args = cli_interface.CliSetting() + print(args) Process(target=serve, 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) From d8e2bc6bc9d5eb7759bd43cabf8d4a5816491f5d Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sat, 19 Jul 2025 11:16:57 -0500 Subject: [PATCH 16/19] fix ruff --- grpc_server_test.py | 1 - ruff.toml | 2 ++ 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/grpc_server_test.py b/grpc_server_test.py index cc18b8f..f45ec86 100644 --- a/grpc_server_test.py +++ b/grpc_server_test.py @@ -20,5 +20,4 @@ def test_get_next_connect_client_index(): def test_preprare_for_next_iteration(): - server = SampleServer() pass 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 From 305538bcbecf77d2cbef927892a583f8b1849428 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Sat, 19 Jul 2025 11:35:25 -0500 Subject: [PATCH 17/19] fix mypy --- grpc_client.py | 9 +++++---- grpc_server.py | 4 ++-- mypy.ini | 3 ++- 3 files changed, 9 insertions(+), 7 deletions(-) diff --git a/grpc_client.py b/grpc_client.py index e7ce02b..ae796a0 100644 --- a/grpc_client.py +++ b/grpc_client.py @@ -71,13 +71,14 @@ def train_with_args(args: cli_interface.CliSetting): ps_stub = get_stub() connect_result = repeat_every( - lambda: ps_stub.Connect(sample_pb2.EmptyRequest()), lambda x: x.successful + 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)) + lambda: ps_stub.Disconnect(sample_pb2.DisconnectRequest(clientIndex=client_index)) # type: ignore[attr-defined] ) with torch.no_grad(): @@ -86,7 +87,7 @@ def train_with_args(args: cli_interface.CliSetting): def try_to_join_iteration(): join_result = repeat_every( lambda: ps_stub.TryToJoinIteration( - sample_pb2.TryToJoinIterationRequest(clientIndex=client_index) + sample_pb2.TryToJoinIterationRequest(clientIndex=client_index) # type: ignore[attr-defined] ), lambda x: x.successful, ) @@ -109,7 +110,7 @@ def try_to_join_iteration(): print("submit result") ps_stub.SubmitIteration( - sample_pb2.SubmitIterationRequest( + 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] diff --git a/grpc_server.py b/grpc_server.py index 3462818..0940365 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -32,7 +32,7 @@ def find_first(lst: list, check_fn) -> int: class SampleServer(sample_pb2_grpc.SampleServerServicer): - def __init__(self, num_clients, num_sample_clients, local_update_steps): + def __init__(self, num_clients: int, num_sample_clients: int, local_update_steps: int): self.should_eval = True self.eval_iteration = 25 @@ -65,7 +65,7 @@ def _get_connect_status(self): def _get_next_connect_client_index(self) -> int: return find_first(self.connected_clients, lambda x: not x) - def change_status(self, new_status: Enum) -> None: + def change_status(self, new_status: ServerStatus) -> None: self.status = new_status def preprare_for_next_iteration(self): 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 From 09b116a38ee8b81fbd32fa875b88acabf9b812d0 Mon Sep 17 00:00:00 2001 From: ZidongLiu Date: Thu, 24 Jul 2025 16:48:08 -0500 Subject: [PATCH 18/19] add extra check for iteration finished clients --- grpc_server.py | 23 +++++++++++++++-------- 1 file changed, 15 insertions(+), 8 deletions(-) diff --git a/grpc_server.py b/grpc_server.py index 0940365..1f36077 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -54,6 +54,7 @@ def __init__(self, num_clients: int, num_sample_clients: int, local_update_steps 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): @@ -75,6 +76,7 @@ def preprare_for_next_iteration(self): 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}") @@ -91,7 +93,7 @@ def try_swtich_from_connecting_to_training(self) -> None: if self._should_connect(): return # 3. initialize training data for next iteration - print("swtich from connecting to training") + print("Switch from connecting to training") self.preprare_for_next_iteration() # 4. change status to training self.change_status(ServerStatus.training) @@ -103,10 +105,12 @@ def swtich_to_connecting(self) -> None: self.change_status(ServerStatus.connecting) def _has_iteration_finished(self) -> bool: - return all( + 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( [ - self.iteration_local_grad_scalar.get(client_index) - for client_index in self.iteration_sampled_clients + sorted_sampled_clients[i] == sorted_finished_clients[i] + for i in range(len(sorted_sampled_clients)) ] ) @@ -137,7 +141,7 @@ def try_switch_from_training_to_aggregating(self) -> None: # 2. check if sampled client all have return result if not self._has_iteration_finished(): return - print("swtich from training to aggregating") + print("Switch from training to aggregating") # 3. change status to aggregating self.change_status(ServerStatus.aggregating) # 4. update seed_grad_records @@ -149,16 +153,16 @@ def try_switch_from_training_to_aggregating(self) -> None: self.switch_from_aggregating_to_training() def switch_from_aggregating_to_training(self) -> None: - print("swtich from aggregating to training") + 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("swtich from aggregating to evaluating") + print("Switch from aggregating to evaluating") self.change_status(ServerStatus.evaluating) def switch_from_evaluating_to_training(self) -> None: - print("swtich from evaluating to training") + print("Switch from evaluating to training") self.preprare_for_next_iteration() self.change_status(ServerStatus.training) @@ -196,6 +200,7 @@ def TryToJoinIteration(self, request, context): 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, @@ -237,12 +242,14 @@ def SubmitIteration(self, request, context): 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() From 954c0e6c8698fb2703ee0ac46d29511475911de7 Mon Sep 17 00:00:00 2001 From: Zidong Liu Date: Thu, 28 Aug 2025 09:51:49 -0400 Subject: [PATCH 19/19] always listen on [::] for server side --- grpc_server.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/grpc_server.py b/grpc_server.py index 1f36077..2933e55 100644 --- a/grpc_server.py +++ b/grpc_server.py @@ -324,7 +324,7 @@ def serve(args: cli_interface.CliSetting): sample_pb2_grpc.add_SampleServerServicer_to_server( SampleServer(args.num_clients, args.num_sample_clients, args.local_update_steps), server ) - port_str = f"{args.rpc_master_addr}:{rpc_master_port}" + port_str = f"[::]:{rpc_master_port}" server.add_insecure_port(port_str) print(f"Parameter server starting on {port_str}") server.start()