From 2d5cac7e7780bf8a8452fc7c5ae0d7ab92009606 Mon Sep 17 00:00:00 2001 From: justinlu Date: Thu, 3 Sep 2026 16:46:05 -0700 Subject: [PATCH] Integrate WeightSynchronizationWorkerService gRPC mode into WeightSynchronizer and RaidenController. PiperOrigin-RevId: 975991956 --- tpu_sync/rpc/BUILD | 7 + tpu_sync/rpc/raiden_controller.py | 190 ++++++++++++++---- tpu_sync/rpc/raiden_controller_test.py | 147 +++++++++++++- tpu_sync/weight_sync/BUILD | 16 +- .../weight_synchronization_worker_service.cc | 7 + .../weight_sync/weight_synchronizer_base.cc | 41 +++- .../weight_sync/weight_synchronizer_base.h | 12 +- .../weight_synchronizer_listener_test.cc | 101 ++++++++++ 8 files changed, 454 insertions(+), 67 deletions(-) diff --git a/tpu_sync/rpc/BUILD b/tpu_sync/rpc/BUILD index ad6eb4a5..c86758e4 100644 --- a/tpu_sync/rpc/BUILD +++ b/tpu_sync/rpc/BUILD @@ -42,6 +42,7 @@ cc_grpc_library( name = "raiden_service_cc_grpc", srcs = [":raiden_service_proto"], generate_mocks = True, + grpc_only = True, visibility = ["//visibility:public"], deps = [":raiden_service_cc_proto"], ) @@ -116,8 +117,11 @@ py_library( deps = [ ":controller_service_py_pb2", ":raiden_service_py_pb2", + ":raiden_service_py_pb2_grpc", "//tpu_sync/api:common", "//tpu_sync/kv_cache:nd_slice_math", + "//tpu_sync/weight_sync:weight_synchronization_worker_service_client_py", + "@com_github_grpc_grpc//src/python/grpcio/grpc:grpcio", "@com_google_absl_py//absl/logging", ], ) @@ -128,6 +132,9 @@ py_test( deps = [ ":raiden_controller", ":raiden_service_py_pb2", + ":raiden_service_py_pb2_grpc", + "//net/grpc/python:use_insecure_channel_for_test", + "@com_github_grpc_grpc//src/python/grpcio/grpc:grpcio", "@com_google_absl_py//absl/testing:absltest", ], ) diff --git a/tpu_sync/rpc/raiden_controller.py b/tpu_sync/rpc/raiden_controller.py index 018b9ad0..7d54a80b 100644 --- a/tpu_sync/rpc/raiden_controller.py +++ b/tpu_sync/rpc/raiden_controller.py @@ -34,6 +34,7 @@ from tpu_sync.kv_cache import nd_slice_math from tpu_sync.rpc import controller_service_pb2 from tpu_sync.rpc import raiden_service_pb2 +from tpu_sync.weight_sync import weight_synchronization_worker_service_client @dataclasses.dataclass @@ -427,12 +428,34 @@ def connect_socket( time.sleep(2.0) +async def _await_grpc_future(fut: Any) -> Any: + """Adapts a grpc.Future to an asyncio awaitable without blocking the event loop.""" + loop = asyncio.get_running_loop() + async_fut = loop.create_future() + + def _done_callback(f): + try: + res = f.result() + loop.call_soon_threadsafe( + lambda: not async_fut.done() and async_fut.set_result(res) + ) + except Exception as e: # pylint: disable=broad-exception-caught + loop.call_soon_threadsafe( + lambda: not async_fut.done() and async_fut.set_exception(e) + ) + + fut.add_done_callback(_done_callback) + return await async_fut + + class WorkerRpcClient: """Distributed RPC Client connecting to Native C++ Control Daemons with Event-Driven resolution. Maintains an asynchronous endpoint catalog that resolves worker network coordinates instantaneously when participating worker tasks self-register, completely eliminating hardcoded active polling loops or arbitrary delays. + Supports both raw TCP socket streams and gRPC + WeightSynchronizationWorkerService. """ def __init__( @@ -441,6 +464,7 @@ def __init__( resolve_timeout: float = 300.0, name_resolver: Optional[NameResolver] = None, proto_module: Optional[Any] = None, + use_grpc: bool = False, ): """Instantiates RPC Client with an optional initial endpoint mapping. @@ -450,6 +474,8 @@ def __init__( task to self-register before raising a Timeout RuntimeError. name_resolver: Interface for resolving remote coordinates (e.g. BNS). proto_module: Optional protobuf module to use for ControlRequest/Response. + use_grpc: If True, uses gRPC WeightSynchronizationWorkerServiceClient + instead of raw TCP socket streams. """ self._endpoints = {} if endpoint_addresses: @@ -459,11 +485,44 @@ def __init__( self._resolve_timeout = resolve_timeout self._name_resolver = name_resolver self._proto_module = proto_module or raiden_service_pb2 + self._use_grpc = use_grpc + self._grpc_clients: dict[ + str, + weight_synchronization_worker_service_client.WeightSynchronizationWorkerServiceClient, + ] = {} + self._grpc_lock = threading.Lock() @property def name_resolver(self) -> Optional[NameResolver]: return self._name_resolver + @property + def use_grpc(self) -> bool: + return self._use_grpc + + def get_grpc_client( + self, addr: str + ) -> ( + weight_synchronization_worker_service_client.WeightSynchronizationWorkerServiceClient + ): + """Returns or creates a cached WeightSynchronizationWorkerServiceClient for addr.""" + if not self._use_grpc: + raise ValueError("WorkerRpcClient is not configured to use gRPC") + resolved_addr = addr + if self._name_resolver: + try: + resolved_addr = self._name_resolver.resolve(addr) + except Exception: # pylint: disable=broad-except + pass + with self._grpc_lock: + client = self._grpc_clients.get(resolved_addr) + if client is None: + client = weight_synchronization_worker_service_client.WeightSynchronizationWorkerServiceClient( + target=resolved_addr + ) + self._grpc_clients[resolved_addr] = client + return client + def register_worker_endpoint( self, worker_name: RaidenId, rpc_address: str ) -> None: @@ -559,13 +618,34 @@ def _send_rpc_sync( finally: sock.close() + async def _send_control_request( + self, addr: str, req: Any, timeout: float = 600.0 + ) -> Any: + """Sends a ControlRequest via gRPC or raw TCP socket, verifying success.""" + if self._use_grpc: + client = self.get_grpc_client(addr) + fut = client.handle_control(req, timeout=timeout) + resp = await _await_grpc_future(fut) + else: + resp_bytes = await self._send_rpc( + addr, req.SerializeToString(), timeout=timeout + ) + resp = self._proto_module.ControlResponse() + resp.ParseFromString(resp_bytes) + + if not resp.success: + raise RuntimeError( + f"Raiden remote native execution failed: {resp.message}" + ) + return resp + async def start_transfer( self, target_id: RaidenId, transfer_plan: TransferPlan, address: Optional[str] = None, ) -> None: - """Connects to remote Worker servicer and dispatches encoded collective transfer commands. + """Connects to remote Worker servicer and dispatches collective transfer commands. Args: target_id: Target participating worker RaidenId. @@ -581,20 +661,16 @@ async def start_transfer( native execution reports failure status. """ try: - payload = self._encode_start_transfer(target_id, transfer_plan) - if not payload: + req = self._build_start_transfer_request(target_id, transfer_plan) + if req is None: return except NotImplementedError: return addrs = [address] if address else await self._resolve_endpoints(target_id) await asyncio.gather( - *[self._send_and_verify(addr, payload) for addr in addrs] + *[self._send_control_request(addr, req) for addr in addrs] ) - async def _send_and_verify(self, addr: str, payload: bytes) -> None: - resp_bytes = await self._send_rpc(addr, payload) - self._verify_response(resp_bytes) - def _raiden_id_to_proto(self, unit: RaidenId) -> Any: return self._proto_module.RaidenIdProto( job_name=unit.job_name, @@ -603,18 +679,10 @@ def _raiden_id_to_proto(self, unit: RaidenId) -> Any: data_replica_idx=unit.data_replica_idx, ) - def _encode_start_transfer( + def _build_start_transfer_request( self, target_id: RaidenId, transfer_plan: TransferPlan - ) -> Optional[bytes]: - """Serializes domain-specific binary Protobuf command for collective transfer kickoff. - - Args: - target_id: Target worker RaidenId coordinate. - transfer_plan: Top-level distributed Collective Transfer execution plan. - - Returns: - Serialized binary bytes payload, or None for no-op execution. - """ + ) -> Optional[Any]: + """Constructs domain-specific Protobuf ControlRequest for collective transfer kickoff.""" if ( target_id not in transfer_plan.src_units and target_id not in transfer_plan.dst_units @@ -773,7 +841,14 @@ def _encode_start_transfer( start_req.shard_push_schedules[shard_idx].CopyFrom(schedule_proto) req.start_transfer_request.CopyFrom(start_req) - return req.SerializeToString() + return req + + def _encode_start_transfer( + self, target_id: RaidenId, transfer_plan: TransferPlan + ) -> Optional[bytes]: + """Serializes domain-specific binary Protobuf command for collective transfer kickoff.""" + req = self._build_start_transfer_request(target_id, transfer_plan) + return req.SerializeToString() if req is not None else None def _verify_response(self, resp_bytes: bytes) -> None: """Validates demarshaled remote response bytes returned from C++ workers.""" @@ -794,25 +869,54 @@ def get_registered_endpoints(self, worker_name: RaidenId) -> list[str]: async def shutdown_workers(self, timeout: float = 10.0) -> None: """Dispatches remote shutdown signaling payloads to all registered worker daemons.""" - payload = self._encode_shutdown() all_addrs = set() for addrs in self._endpoints.values(): all_addrs.update(addrs) - if all_addrs: - await asyncio.gather( - *[ - self._send_rpc(addr, payload, timeout=timeout) - for addr in all_addrs - ], - return_exceptions=True, - ) + if not all_addrs: + return + + req = self._build_shutdown_request() + await asyncio.gather( + *[ + self._send_control_request(addr, req, timeout=timeout) + for addr in all_addrs + ], + return_exceptions=True, + ) + + def _build_shutdown_request(self) -> Any: + """Constructs domain-specific binary command for remote shutdown signaling.""" + return self._proto_module.ControlRequest( + command=self._proto_module.ControlRequest.COMMAND_SHUTDOWN + ) def _encode_shutdown(self) -> bytes: """Serializes domain-specific binary command for remote shutdown signaling.""" + req = self._build_shutdown_request() + return req.SerializeToString() + + async def query_metadata( + self, addr: str, timeout: float = 600.0 + ) -> list[Any]: + """Queries metadata from a remote worker endpoint.""" req = self._proto_module.ControlRequest( - command=self._proto_module.ControlRequest.COMMAND_SHUTDOWN + command=self._proto_module.ControlRequest.COMMAND_GET_METADATA ) - return req.SerializeToString() + resp = await self._send_control_request(addr, req, timeout=timeout) + return list(resp.get_metadata_response.metadata) + + def close(self) -> None: + """Closes all cached gRPC client channels.""" + with self._grpc_lock: + for client in self._grpc_clients.values(): + client.close() + self._grpc_clients.clear() + + def __enter__(self) -> "WorkerRpcClient": + return self + + def __exit__(self, exc_type, exc_val, exc_tb) -> None: + self.close() class WeightSyncWorkerRpcClient(WorkerRpcClient): @@ -1254,6 +1358,7 @@ def __init__( request_registry_ttl_s: float = 600.0, broadcast_k: Optional[int] = None, enable_plan_cache: bool = True, + use_grpc: bool = False, ): """Initializes the RaidenController. @@ -1264,6 +1369,8 @@ def __init__( broadcast_k: Fan-out factor K for tree-based broadcast transfers. enable_plan_cache: Whether to cache transfer planning and resharding schedules across transfer invocations with identical topologies. + use_grpc: Whether to default to WorkerRpcClient in gRPC mode if + worker_rpc_client is not specified. """ self.port = port self.broadcast_k = ( @@ -1294,7 +1401,13 @@ def __init__( if request_registry_ttl_s <= 0: raise ValueError("request_registry_ttl_s must be positive") self._request_registry_ttl_s = request_registry_ttl_s - self.worker_rpc_client = worker_rpc_client or WorkerRpcClient() + use_grpc_effective = use_grpc or ( + os.environ.get("RAIDEN_WEIGHT_SYNC_USE_GRPC", "").lower() + in ("1", "true") + ) + self.worker_rpc_client = worker_rpc_client or WorkerRpcClient( + use_grpc=use_grpc_effective + ) self._registered_variables = {} def register_work_unit( @@ -1509,17 +1622,8 @@ def _resolve_shards(self, unit: RaidenId) -> list[str]: return list(shards) async def _query_remote_metadata(self, addr: str) -> list[Any]: - req = raiden_service_pb2.ControlRequest( - command=raiden_service_pb2.ControlRequest.COMMAND_GET_METADATA - ) - resp_bytes = await self.worker_rpc_client._send_rpc( - addr, req.SerializeToString() - ) - resp = raiden_service_pb2.ControlResponse() - resp.ParseFromString(resp_bytes) - if not resp.success: - raise RuntimeError(f"Failed to query remote metadata: {resp.message}") - return list(resp.get_metadata_response.metadata) + """Queries metadata from a remote controller or worker endpoint.""" + return await self.worker_rpc_client.query_metadata(addr) def _get_local_metadata(self, units: list[RaidenId]) -> list[Any]: with self._lock: diff --git a/tpu_sync/rpc/raiden_controller_test.py b/tpu_sync/rpc/raiden_controller_test.py index f4fa3497..78096c71 100644 --- a/tpu_sync/rpc/raiden_controller_test.py +++ b/tpu_sync/rpc/raiden_controller_test.py @@ -12,13 +12,14 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Tests for Raiden Controller high-level transfer API under rpc/.""" - import asyncio +from concurrent import futures import socket from absl.testing import absltest +import grpc from tpu_sync.rpc import raiden_controller from tpu_sync.rpc import raiden_service_pb2 +from tpu_sync.rpc import raiden_service_pb2_grpc class DummyWorkerRpcClient(raiden_controller.WorkerRpcClient): @@ -2703,5 +2704,147 @@ def test_format_units(self): self.assertEqual(raiden_controller._format_units(b"unit_bytes"), "b'unit_bytes'") +class FakeWeightSyncWorkerServicer( + raiden_service_pb2_grpc.WeightSynchronizationWorkerServiceServicer +): + """Fake gRPC servicer for testing WeightSynchronizationWorkerService.""" + + def __init__(self, succeed: bool = True, failure_message: str = "Error"): + self.requests = [] + self.succeed = succeed + self.failure_message = failure_message + + def HandleControl(self, request, context): + self.requests.append(request) + resp = raiden_service_pb2.ControlResponse() + resp.success = self.succeed + resp.message = "SUCCESS" if self.succeed else self.failure_message + return resp + + +class WorkerRpcClientGrpcModeTest(absltest.TestCase): + + def setUp(self): + super().setUp() + self.servicer = FakeWeightSyncWorkerServicer() + self.server = grpc.server(futures.ThreadPoolExecutor(max_workers=2)) + raiden_service_pb2_grpc.add_WeightSynchronizationWorkerServiceServicer_to_server( + self.servicer, self.server + ) + port = self.server.add_insecure_port("127.0.0.1:0") + self.server.start() + self.server_port = port + self.server_address = f"127.0.0.1:{port}" + + def tearDown(self): + self.server.stop(grace=None) + super().tearDown() + + def test_controller_default_and_grpc_instantiation(self): + controller_grpc = raiden_controller.RaidenController(port=0, use_grpc=True) + self.assertTrue(controller_grpc.worker_rpc_client.use_grpc) + self.assertIsNotNone( + controller_grpc.worker_rpc_client.get_grpc_client("127.0.0.1:8000") + ) + + controller_default = raiden_controller.RaidenController(port=0) + self.assertFalse(controller_default.worker_rpc_client.use_grpc) + with self.assertRaisesRegex( + ValueError, "WorkerRpcClient is not configured to use gRPC" + ): + controller_default.worker_rpc_client.get_grpc_client("127.0.0.1:8000") + + def test_grpc_mode_start_transfer_and_shutdown(self): + client = raiden_controller.WorkerRpcClient(use_grpc=True) + src_unit = raiden_controller.RaidenId("trainer", "0", "weights") + dst_unit = raiden_controller.RaidenId("sampler", "0", "weights") + + client.register_worker_endpoint(src_unit, self.server_address) + + plan = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={dst_unit: ["127.0.0.1:8000"]}, + uuid=12345, + ) + + asyncio.run(client.start_transfer(src_unit, plan)) + self.assertLen(self.servicer.requests, 1) + req = self.servicer.requests[0] + self.assertEqual( + req.command, raiden_service_pb2.ControlRequest.COMMAND_START_TRANSFER + ) + self.assertEqual(req.start_transfer_request.uuid, 12345) + self.assertTrue(req.start_transfer_request.is_sender) + + asyncio.run(client.shutdown_workers()) + self.assertLen(self.servicer.requests, 2) + shutdown_req = self.servicer.requests[1] + self.assertEqual( + shutdown_req.command, raiden_service_pb2.ControlRequest.COMMAND_SHUTDOWN + ) + + client.close() + + def test_grpc_mode_error_handling(self): + self.servicer.succeed = False + self.servicer.failure_message = "Native execution failure" + + client = raiden_controller.WorkerRpcClient(use_grpc=True) + src_unit = raiden_controller.RaidenId("trainer", "0", "weights") + dst_unit = raiden_controller.RaidenId("sampler", "0", "weights") + client.register_worker_endpoint(src_unit, self.server_address) + + plan = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={dst_unit: ["127.0.0.1:8000"]}, + uuid=12345, + ) + + with self.assertRaisesRegex(RuntimeError, "Native execution failure"): + asyncio.run(client.start_transfer(src_unit, plan)) + + client.close() + + def test_controller_e2e_with_grpc_mode(self): + controller = raiden_controller.RaidenController(port=0, use_grpc=True) + + src_unit = raiden_controller.RaidenId("trainer", "0", "weights") + dst_unit = raiden_controller.RaidenId("sampler", "0", "weights") + + controller.register_work_unit( + src_unit, + ["127.0.0.1:8000"], + control_plane_rpc_address=self.server_address, + ) + controller.register_work_unit( + dst_unit, + ["127.0.0.1:8001"], + control_plane_rpc_address=self.server_address, + ) + + fut = controller.start_transfer( + src_units=[src_unit], + dst_units=[dst_unit], + req_id="grpc_e2e_req", + ) + asyncio.run(fut.wait()) + + # Verify both destination and source workers received start_transfer + self.assertGreaterEqual(len(self.servicer.requests), 2) + commands = [r.command for r in self.servicer.requests] + self.assertTrue( + all( + c == raiden_service_pb2.ControlRequest.COMMAND_START_TRANSFER + for c in commands + ) + ) + + controller.worker_rpc_client.close() + + if __name__ == "__main__": absltest.main() diff --git a/tpu_sync/weight_sync/BUILD b/tpu_sync/weight_sync/BUILD index 2ee3c0a5..72884e2b 100644 --- a/tpu_sync/weight_sync/BUILD +++ b/tpu_sync/weight_sync/BUILD @@ -41,10 +41,12 @@ cc_library( cc_library( name = "weight_synchronizer_base", srcs = [ + "weight_synchronization_worker_service.cc", "weight_synchronizer_base.cc", "weight_synchronizer_listener.cc", ], hdrs = [ + "weight_synchronization_worker_service.h", "weight_synchronizer_base.h", "weight_synchronizer_listener.h", ], @@ -63,8 +65,10 @@ cc_library( "//tpu_sync/core:raw_transfer_core", "//tpu_sync/core:status_macros", "//tpu_sync/core:xla_raw_transfer_headers", + "//tpu_sync/rpc:raiden_service_cc_grpc", "//tpu_sync/rpc:raiden_service_cc_proto", "//tpu_sync/transport:buffer_push_task", + "@com_github_grpc_grpc//:grpc++", "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/container:flat_hash_set", @@ -115,6 +119,7 @@ cc_test( ], features = ["-use_header_modules"], deps = [ + ":weight_synchronization_worker_service_client", ":weight_synchronizer_base", "//tpu_sync/rpc:raiden_service_cc_proto", "@com_google_googletest//:gtest", @@ -141,24 +146,13 @@ cc_test( cc_library( name = "weight_synchronization_worker_service", - srcs = ["weight_synchronization_worker_service.cc"], hdrs = ["weight_synchronization_worker_service.h"], - copts = [ - "-fno-strict-aliasing", - "-fexceptions", - ], - features = ["-use_header_modules"], visibility = ["//visibility:public"], deps = [ ":weight_synchronizer_base", "//tpu_sync/rpc:raiden_service_cc_grpc", "//tpu_sync/rpc:raiden_service_cc_proto", "@com_github_grpc_grpc//:grpc++", - "@com_google_absl//absl/container:flat_hash_map", - "@com_google_absl//absl/log", - "@com_google_absl//absl/status", - "@com_google_absl//absl/strings", - "@com_google_absl//absl/time", ], ) diff --git a/tpu_sync/weight_sync/weight_synchronization_worker_service.cc b/tpu_sync/weight_sync/weight_synchronization_worker_service.cc index 1d593c19..3dd60758 100644 --- a/tpu_sync/weight_sync/weight_synchronization_worker_service.cc +++ b/tpu_sync/weight_sync/weight_synchronization_worker_service.cc @@ -140,6 +140,13 @@ grpc::Status WeightSynchronizationWorkerServiceImpl::HandleControl( tpu_sync::rpc::ControlRequest::COMMAND_SHUTDOWN) { LOG(INFO) << "gRPC WeightSynchronizationWorkerService received SHUTDOWN " "command. Initiating clean exit."; + if (engine_) { + if (engine_->control_delegate()) { + engine_->control_delegate()->DrainPendingH2d(); + } else { + engine_->DrainPendingH2d(); + } + } if (shutdown_callback_) { std::thread([cb = shutdown_callback_]() { absl::SleepFor(absl::Milliseconds(50)); diff --git a/tpu_sync/weight_sync/weight_synchronizer_base.cc b/tpu_sync/weight_sync/weight_synchronizer_base.cc index 3b07631f..6bfe9773 100644 --- a/tpu_sync/weight_sync/weight_synchronizer_base.cc +++ b/tpu_sync/weight_sync/weight_synchronizer_base.cc @@ -56,11 +56,15 @@ #include "tpu_sync/rpc/raiden_service.pb.h" #include "tpu_sync/transport/buffer_push_task.h" #include "tpu_sync/weight_sync/tiling_utils.h" +#include "tpu_sync/weight_sync/weight_synchronization_worker_service.h" #include "tpu_sync/weight_sync/weight_synchronizer_listener.h" ABSL_FLAG(size_t, raiden_weight_sync_host_buffer_scratchpad_size, 256 * 1024, "Amount of scratchpad to allocate to host buffers for resharding " "pulls."); +ABSL_FLAG(bool, raiden_weight_sync_use_grpc, false, + "Whether to use gRPC WeightSynchronizationWorkerService instead of " + "raw TCP socket listener."); namespace tpu_raiden { namespace weight_sync { @@ -71,13 +75,15 @@ WeightSynchronizerBase::WeightSynchronizerBase( std::optional> external_host_ptrs, bool unsafe_skip_buffer_lock, int parallelism, std::optional listener_port, std::optional bind_ip, - std::vector layer_names, bool auto_h2d) + std::vector layer_names, bool auto_h2d, bool use_grpc) : tpu_raiden::RaidenManagerBase( layer_buffers.size(), layer_buffers.empty() ? 0 : layer_buffers[0].size(), layer_buffers.empty() ? 0 : layer_buffers[0][0].GetOnDeviceSizeInBytes(), local_port, parallelism, bind_ip), + use_grpc_(use_grpc || absl::GetFlag(FLAGS_raiden_weight_sync_use_grpc) || + (std::getenv("RAIDEN_WEIGHT_SYNC_USE_GRPC") != nullptr)), auto_h2d_(auto_h2d) { if (layer_names.empty()) { layer_names_.reserve(num_layers_); @@ -221,8 +227,13 @@ WeightSynchronizerBase::WeightSynchronizerBase( } if (listener_port) { - listener_ = - std::make_unique(this, *listener_port); + if (use_grpc_) { + grpc_service_ = std::make_unique( + this, *listener_port); + } else { + listener_ = + std::make_unique(this, *listener_port); + } } if (auto_h2d_) { h2d_pool_ = std::make_unique( @@ -237,23 +248,25 @@ WeightSynchronizerBase::WeightSynchronizerBase( std::optional local_port, std::optional host_blocks_to_allocate, int parallelism, std::optional listener_port, std::optional bind_ip, std::vector layer_names, - bool auto_h2d) + bool auto_h2d, bool use_grpc) : WeightSynchronizerBase(num_layers, num_shards, std::vector(num_layers, slice_byte_size), local_port, host_blocks_to_allocate, parallelism, listener_port, bind_ip, std::move(layer_names), - auto_h2d) {} + auto_h2d, use_grpc) {} WeightSynchronizerBase::WeightSynchronizerBase( size_t num_layers, size_t num_shards, std::vector slice_byte_sizes, std::optional local_port, std::optional host_blocks_to_allocate, int parallelism, std::optional listener_port, std::optional bind_ip, std::vector layer_names, - bool auto_h2d) + bool auto_h2d, bool use_grpc) : tpu_raiden::RaidenManagerBase( num_layers, num_shards, slice_byte_sizes.empty() ? 0 : slice_byte_sizes[0], local_port, parallelism, bind_ip), + use_grpc_(use_grpc || absl::GetFlag(FLAGS_raiden_weight_sync_use_grpc) || + (std::getenv("RAIDEN_WEIGHT_SYNC_USE_GRPC") != nullptr)), auto_h2d_(auto_h2d) { if (layer_names.empty()) { layer_names_.reserve(num_layers_); @@ -301,8 +314,13 @@ WeightSynchronizerBase::WeightSynchronizerBase( } if (listener_port) { - listener_ = - std::make_unique(this, *listener_port); + if (use_grpc_) { + grpc_service_ = std::make_unique( + this, *listener_port); + } else { + listener_ = + std::make_unique(this, *listener_port); + } } if (auto_h2d_) { h2d_pool_ = std::make_unique( @@ -313,6 +331,9 @@ WeightSynchronizerBase::WeightSynchronizerBase( } std::optional WeightSynchronizerBase::listener_port() const { + if (grpc_service_) { + return grpc_service_->server_port(); + } if (listener_) { return listener_->listener_port(); } @@ -320,6 +341,9 @@ std::optional WeightSynchronizerBase::listener_port() const { } bool WeightSynchronizerBase::is_listener_active() const { + if (grpc_service_) { + return grpc_service_->is_active(); + } if (listener_) { return listener_->is_active(); } @@ -347,6 +371,7 @@ WeightSynchronizerBase::get_local_endpoints() const { WeightSynchronizerBase::~WeightSynchronizerBase() { StopTransportServer(); + grpc_service_.reset(); listener_.reset(); h2d_pool_.reset(); push_pool_.reset(); diff --git a/tpu_sync/weight_sync/weight_synchronizer_base.h b/tpu_sync/weight_sync/weight_synchronizer_base.h index 45312f4b..6bbe953c 100644 --- a/tpu_sync/weight_sync/weight_synchronizer_base.h +++ b/tpu_sync/weight_sync/weight_synchronizer_base.h @@ -79,6 +79,7 @@ struct WeightSyncMetrics { }; class WeightSynchronizerListener; +class WeightSynchronizationWorkerService; class WeightSynchronizerControlDelegate { public: @@ -109,7 +110,8 @@ class WeightSynchronizerBase : public tpu_raiden::RaidenManagerBase { bool unsafe_skip_buffer_lock = false, int parallelism = 1, std::optional listener_port = std::nullopt, std::optional bind_ip = std::nullopt, - std::vector layer_names = {}, bool auto_h2d = false); + std::vector layer_names = {}, bool auto_h2d = false, + bool use_grpc = false); // CPU-only constructor for remote workers and mock E2E testing WeightSynchronizerBase( @@ -118,7 +120,8 @@ class WeightSynchronizerBase : public tpu_raiden::RaidenManagerBase { std::optional host_blocks_to_allocate = std::nullopt, int parallelism = 1, std::optional listener_port = std::nullopt, std::optional bind_ip = std::nullopt, - std::vector layer_names = {}, bool auto_h2d = false); + std::vector layer_names = {}, bool auto_h2d = false, + bool use_grpc = false); // CPU-only constructor for remote workers and mock E2E testing supporting // heterogeneous slice sizes and custom layer names. @@ -129,7 +132,8 @@ class WeightSynchronizerBase : public tpu_raiden::RaidenManagerBase { std::optional host_blocks_to_allocate = std::nullopt, int parallelism = 1, std::optional listener_port = std::nullopt, std::optional bind_ip = std::nullopt, - std::vector layer_names = {}, bool auto_h2d = false); + std::vector layer_names = {}, bool auto_h2d = false, + bool use_grpc = false); std::optional listener_port() const; bool is_listener_active() const; @@ -287,6 +291,8 @@ class WeightSynchronizerBase : public tpu_raiden::RaidenManagerBase { protected: std::unique_ptr listener_; + std::unique_ptr grpc_service_; + bool use_grpc_ = false; const PJRT_Api* c_api_ = nullptr; const PJRT_RawBuffer_Extension* extension_ = nullptr; size_t physical_size_ = 0; diff --git a/tpu_sync/weight_sync/weight_synchronizer_listener_test.cc b/tpu_sync/weight_sync/weight_synchronizer_listener_test.cc index a410127c..124cd2e3 100644 --- a/tpu_sync/weight_sync/weight_synchronizer_listener_test.cc +++ b/tpu_sync/weight_sync/weight_synchronizer_listener_test.cc @@ -27,6 +27,7 @@ #include #include "tpu_sync/rpc/raiden_service.pb.h" +#include "tpu_sync/weight_sync/weight_synchronization_worker_service_client.h" #include "tpu_sync/weight_sync/weight_synchronizer_base.h" namespace tpu_raiden { @@ -288,6 +289,106 @@ TEST(WeightSynchronizerListenerTest, PushWeightsReshardedSuccess) { close(sock); } +TEST(WeightSynchronizerListenerTest, GrpcModePushWeightsReshardedSuccess) { + WeightSynchronizerBase src_engine( + /*num_layers=*/1, /*num_shards=*/4, /*slice_byte_size=*/16, + /*local_port=*/0, /*host_blocks_to_allocate=*/std::nullopt, + /*parallelism=*/1, /*listener_port=*/0, /*bind_ip=*/std::nullopt, + /*layer_names=*/{}, /*auto_h2d=*/false, /*use_grpc=*/true); + ASSERT_TRUE(src_engine.listener_port().has_value()); + EXPECT_TRUE(src_engine.is_listener_active()); + + WeightSynchronizerBase dst_engine( + /*num_layers=*/1, /*num_shards=*/4, /*slice_byte_size=*/16, + /*local_port=*/0, /*host_blocks_to_allocate=*/1, + /*parallelism=*/1, /*listener_port=*/0, /*bind_ip=*/std::nullopt, + /*layer_names=*/{}, /*auto_h2d=*/false, /*use_grpc=*/true); + ASSERT_TRUE(dst_engine.local_port().has_value()); + ASSERT_TRUE(dst_engine.listener_port().has_value()); + EXPECT_TRUE(dst_engine.is_listener_active()); + + std::string dst_peer = + "127.0.0.1:" + std::to_string(*dst_engine.local_port()); + + // Populate source buffers + std::vector> src_data = { + {0, 1, 2, 3, 8, 9, 10, 11, 16, 17, 18, 19, 24, 25, 26, 27}, + {4, 5, 6, 7, 12, 13, 14, 15, 20, 21, 22, 23, 28, 29, 30, 31}, + {32, 33, 34, 35, 40, 41, 42, 43, 48, 49, 50, 51, 56, 57, 58, 59}, + {36, 37, 38, 39, 44, 45, 46, 47, 52, 53, 54, 55, 60, 61, 62, 63}, + }; + + for (size_t i = 0; i < 4; ++i) { + uint8_t* ptr = src_engine.GetHostPointer(0, i); + ASSERT_NE(ptr, nullptr); + std::memcpy(ptr, src_data[i].data(), 16); + } + + uint64_t test_uuid = 112233; + + WeightSynchronizationWorkerServiceClient dst_client( + "localhost:" + std::to_string(*dst_engine.listener_port())); + WeightSynchronizationWorkerServiceClient src_client( + "localhost:" + std::to_string(*src_engine.listener_port())); + + // 1. Arm destination receiver using gRPC client + StartTransferRequest dst_start_req; + dst_start_req.set_is_sender(false); + dst_start_req.set_uuid(test_uuid); + dst_start_req.set_expected_block_count(8); + + auto dst_resp_or = dst_client.StartTransfer(dst_start_req).Await(); + ASSERT_TRUE(dst_resp_or.ok()); + EXPECT_TRUE(dst_resp_or->success()); + + // 2. Dispatch sender push schedules using gRPC client + StartTransferRequest src_start_req; + src_start_req.set_is_sender(true); + src_start_req.set_uuid(test_uuid); + + auto& push_schedules = *src_start_req.mutable_shard_push_schedules(); + + // S0 -> D0 + ShardPushScheduleProto s0_sched; + for (int r = 0; r < 4; ++r) { + ShardPushEntryProto* e = s0_sched.add_entries(); + e->set_dst_peer(dst_peer); + e->set_dst_shard_idx(0); + e->set_src_offset_bytes(r * 4); + e->set_dst_offset_bytes(r * 2); + e->set_size_bytes(2); + e->set_layer_idx(0); + } + push_schedules[0] = s0_sched; + + // S2 -> D0 + ShardPushScheduleProto s2_sched; + for (int r = 0; r < 4; ++r) { + ShardPushEntryProto* e = s2_sched.add_entries(); + e->set_dst_peer(dst_peer); + e->set_dst_shard_idx(0); + e->set_src_offset_bytes(r * 4); + e->set_dst_offset_bytes(8 + r * 2); + e->set_size_bytes(2); + e->set_layer_idx(0); + } + push_schedules[2] = s2_sched; + + auto src_resp_or = src_client.StartTransfer(src_start_req).Await(); + ASSERT_TRUE(src_resp_or.ok()); + EXPECT_TRUE(src_resp_or->success()); + + // Verify Destination Shard 0 host memory + uint8_t* dst_ptr = dst_engine.GetHostPointer(0, 0); + ASSERT_NE(dst_ptr, nullptr); + + std::vector expected_d0 = {0, 1, 8, 9, 16, 17, 24, 25, + 32, 33, 40, 41, 48, 49, 56, 57}; + for (size_t k = 0; k < 16; ++k) { + EXPECT_EQ(dst_ptr[k], expected_d0[k]) << "Mismatch at byte " << k; + } +} + } // namespace } // namespace weight_sync } // namespace tpu_raiden