diff --git a/tpu_sync/rpc/raiden_controller.py b/tpu_sync/rpc/raiden_controller.py index 63eb8e00f..b99e7f1ef 100644 --- a/tpu_sync/rpc/raiden_controller.py +++ b/tpu_sync/rpc/raiden_controller.py @@ -111,6 +111,9 @@ class _CachedTransferSchedule: sender_push_schedule_protos: dict[Any, dict[int, Any]] = dataclasses.field( default_factory=dict ) + cached_serialized_payloads: dict[Any, bytes] = dataclasses.field( + default_factory=dict + ) def to_physical(logical_shape, logical_mesh_shape, minor_to_major): @@ -533,6 +536,12 @@ class TransferPlan: sender_push_schedule_protos: dict[RaidenId, dict[int, Any]] = ( dataclasses.field(default_factory=dict, repr=False, compare=False) ) + cached_serialized_payloads: dict[Any, bytes] = dataclasses.field( + default_factory=dict, repr=False, compare=False + ) + endpoint_to_shards: dict[Any, Any] = dataclasses.field( + default_factory=dict, repr=False, compare=False + ) def _coerce_pool_spec_proto(pool: Any) -> Any: @@ -732,6 +741,138 @@ def _send_rpc_sync( message_type=self._proto_module.ControlRequest.DESCRIPTOR.full_name, ) + def _get_worker_owned_shards( + self, + target_id: RaidenId, + transfer_plan: TransferPlan, + address: Optional[str], + ) -> Optional[set[int]]: + """Determines shard indices owned by a sender worker endpoint.""" + if not address: + return None + + addr_clean = address.strip() + if ( + hasattr(transfer_plan, "endpoint_to_shards") + and transfer_plan.endpoint_to_shards + ): + if (target_id, addr_clean) in transfer_plan.endpoint_to_shards: + return set(transfer_plan.endpoint_to_shards[(target_id, addr_clean)]) + if addr_clean in transfer_plan.endpoint_to_shards: + return set(transfer_plan.endpoint_to_shards[addr_clean]) + + endpoints = self._endpoints.get(target_id, []) + if not endpoints and hasattr(transfer_plan, "worker_rpc_addresses"): + rpc_addr = transfer_plan.worker_rpc_addresses.get(target_id, "") + if rpc_addr: + if isinstance(rpc_addr, (list, tuple)): + endpoints = [str(a).strip() for a in rpc_addr if str(a).strip()] + else: + endpoints = [a.strip() for a in str(rpc_addr).split(",") if a.strip()] + + if not endpoints or len(endpoints) <= 1: + return None + + norm_endpoints = [] + for e in endpoints: + clean_e = e.strip() + if clean_e and clean_e not in norm_endpoints: + norm_endpoints.append(clean_e) + + if addr_clean not in norm_endpoints: + return None + worker_idx = norm_endpoints.index(addr_clean) + num_workers = len(norm_endpoints) + + # Determine total number of shards for target_id + num_shards = 0 + data_shards = getattr(transfer_plan, "worker_data_addresses", {}).get( + target_id, [] + ) + if data_shards: + num_shards = len(data_shards) + if num_shards == 0: + cached_protos = getattr( + transfer_plan, "sender_push_schedule_protos", None + ) + if ( + cached_protos + and target_id in cached_protos + and cached_protos[target_id] + ): + num_shards = max( + len(cached_protos[target_id]), + max(cached_protos[target_id].keys()) + 1, + ) + if num_shards == 0: + push_schedules = getattr(transfer_plan, "shard_push_schedules", {}).get( + target_id, {} + ) + if push_schedules: + num_shards = max( + len(push_schedules), + max(push_schedules.keys()) + 1, + ) + + if num_shards <= 1: + return None + + if num_shards < num_workers: + return {worker_idx} if worker_idx < num_shards else set() + + start_shard = (worker_idx * num_shards) // num_workers + end_shard = ((worker_idx + 1) * num_shards) // num_workers + return set(range(start_shard, end_shard)) + + def _is_payload_invariant_across_addrs( + self, + target_id: RaidenId, + transfer_plan: TransferPlan, + addrs: list[str], + ) -> bool: + """Returns True if payload is identical across all worker addresses.""" + if len(addrs) <= 1: + return True + + is_sender = target_id in getattr( + transfer_plan, "src_units", [] + ) and getattr(transfer_plan, "is_sender", False) + if is_sender: + cached_protos = getattr( + transfer_plan, "sender_push_schedule_protos", None + ) + has_protos = bool(cached_protos and target_id in cached_protos) + push_schedules = getattr(transfer_plan, "shard_push_schedules", {}).get( + target_id + ) + if not has_protos and not push_schedules: + return True + + # Check if schedule slicing actually differentiates the addresses. + # When sending the full schedule (e.g. slicing returns None) or + # when all addresses share identical owned shards, payload is invariant. + first_owned = self._get_worker_owned_shards( + target_id, transfer_plan, addrs[0] + ) + for addr in addrs[1:]: + if ( + self._get_worker_owned_shards(target_id, transfer_plan, addr) + != first_owned + ): + return False + return True + + # Receiver: check if endpoint specialization is active + dst_counts = getattr(transfer_plan, "dst_endpoint_counts", None) + dst_layer_counts = getattr(transfer_plan, "dst_endpoint_layer_counts", None) + if dst_counts or dst_layer_counts: + return False + if self.include_receiver_push_schedules(transfer_plan) and getattr( + transfer_plan, "shard_push_schedules", None + ): + return False + return True + async def start_transfer( self, target_id: RaidenId, @@ -759,20 +900,34 @@ async def start_transfer( addrs = await self._resolve_endpoints(target_id) coros = [] - for addr in addrs: + if self._is_payload_invariant_across_addrs(target_id, transfer_plan, addrs): try: + spec_addr = addrs[0] if addrs else None try: - spec_addr = addr if len(addrs) > 1 else None payload = self._encode_start_transfer( target_id, transfer_plan, address=spec_addr ) except TypeError: payload = self._encode_start_transfer(target_id, transfer_plan) - if not payload: - continue except NotImplementedError: - continue - coros.append(self._send_and_verify(addr, payload)) + payload = None + if payload: + for addr in addrs: + coros.append(self._send_and_verify(addr, payload)) + else: + for addr in addrs: + try: + try: + payload = self._encode_start_transfer( + target_id, transfer_plan, address=addr + ) + except TypeError: + payload = self._encode_start_transfer(target_id, transfer_plan) + if not payload: + continue + except NotImplementedError: + continue + coros.append(self._send_and_verify(addr, payload)) if coros: await asyncio.gather(*coros) @@ -812,6 +967,69 @@ def _encode_start_transfer( ): return None + payload_cache = getattr(transfer_plan, "cached_serialized_payloads", None) + uuid_val = getattr(transfer_plan, "uuid", None) + req_id_val = getattr(transfer_plan, "req_id", None) + skip_d2h_val = bool(getattr(transfer_plan, "skip_d2h", False)) + is_sender = target_id in transfer_plan.src_units and transfer_plan.is_sender + is_ws = getattr(transfer_plan, "is_weight_sync", False) + ep_count = len(self._endpoints.get(target_id, [])) + include_recv_sched = self.include_receiver_push_schedules(transfer_plan) + cache_key = ( + target_id, + address, + uuid_val, + req_id_val, + skip_d2h_val, + is_sender, + is_ws, + ep_count, + include_recv_sched, + ) + steady_key = ( + target_id, + address, + uuid_val, + skip_d2h_val, + is_sender, + is_ws, + ep_count, + include_recv_sched, + ) + template_key = ( + "__template__", + target_id, + address, + is_sender, + is_ws, + int(transfer_plan.dst_mem_type), + bool(transfer_plan.use_block_chunks), + int(transfer_plan.parallelism or 0), + ep_count, + include_recv_sched, + ) + + if payload_cache is not None: + if cache_key in payload_cache: + return payload_cache[cache_key] + # When uuid and skip_d2h are invariant across steps (e.g. uuid == 0 or + # repeated uuid), req_id is unused by C++ WeightSynchronizer so the exact + # serialized payload bytes can be returned directly. + if is_sender and is_ws and steady_key in payload_cache: + return payload_cache[steady_key] + # When uuid or skip_d2h changed across steps, reuse the pre-populated + # ControlRequest proto template and update only scalar step fields (uuid, + # req_id, skip_d2h) without rebuilding or copying push schedules. + if is_sender and is_ws and template_key in payload_cache: + cached_req = payload_cache[template_key] + cached_req.start_transfer_request.uuid = int(uuid_val or 0) + cached_req.start_transfer_request.req_id = str(req_id_val or "") + cached_req.start_transfer_request.skip_d2h = skip_d2h_val + serialized_bytes = cached_req.SerializeToString() + payload_cache[cache_key] = serialized_bytes + payload_cache[steady_key] = serialized_bytes + return serialized_bytes + peers = [] for dst in transfer_plan.dst_units: dst_coords = transfer_plan.worker_data_addresses.get(dst) @@ -857,16 +1075,13 @@ def _encode_start_transfer( dst_units=[ self._raiden_id_to_proto(u) for u in transfer_plan.dst_units ], - uuid=transfer_plan.uuid, is_sender=is_sender, dst_mem_type=int(transfer_plan.dst_mem_type), use_block_chunks=transfer_plan.use_block_chunks, expected_block_count=expected_block_count, - req_id=transfer_plan.req_id, transfer_pool_indices=transfer_plan.transfer_pool_indices, pool_dtype_tags=transfer_plan.pool_dtype_tags, parallelism=transfer_plan.parallelism, - skip_d2h=transfer_plan.skip_d2h, ) for layer_idx, skip in transfer_plan.skip_tiling.items(): start_req.skip_tiling[layer_idx] = skip @@ -886,9 +1101,15 @@ def _encode_start_transfer( ) group_proto.order_rank = int(group.get("order_rank", 0)) - if transfer_plan.shard_push_schedules: + cached_protos = getattr( + transfer_plan, "sender_push_schedule_protos", None + ) + if transfer_plan.shard_push_schedules or cached_protos: if not is_sender: - if self.include_receiver_push_schedules(transfer_plan): + if ( + transfer_plan.shard_push_schedules + and self.include_receiver_push_schedules(transfer_plan) + ): # Receiver path: send FILTERED plan, only containing entries for this # receiver target_endpoints = transfer_plan.worker_data_addresses.get( @@ -916,9 +1137,7 @@ def _encode_start_transfer( key_idx = src_base * num_src_shards + shard_idx schedule_proto = self._proto_module.ShardPushScheduleProto() raw_entries = ( - schedule.entries - if hasattr(schedule, "entries") - else schedule + schedule.entries if hasattr(schedule, "entries") else schedule ) target_endpoints_set = set(target_endpoints) for entry_item in raw_entries: @@ -983,12 +1202,19 @@ def _encode_start_transfer( start_req.shard_push_schedules[key_idx].CopyFrom(schedule_proto) else: # Sender path: reuse cached pre-built ShardPushScheduleProto if present + owned_shards = None + if address: + owned_shards = self._get_worker_owned_shards( + target_id, transfer_plan, address + ) + cached_protos = getattr( transfer_plan, "sender_push_schedule_protos", None ) if cached_protos is not None and target_id in cached_protos: for shard_idx, schedule_proto in cached_protos[target_id].items(): - start_req.shard_push_schedules[shard_idx].CopyFrom(schedule_proto) + if owned_shards is None or shard_idx in owned_shards: + start_req.shard_push_schedules[shard_idx].CopyFrom(schedule_proto) else: push_schedules = transfer_plan.shard_push_schedules.get(target_id) if push_schedules: @@ -996,17 +1222,29 @@ def _encode_start_transfer( push_schedules ) for shard_idx, schedule_proto in target_protos.items(): - start_req.shard_push_schedules[shard_idx].CopyFrom(schedule_proto) + if owned_shards is None or shard_idx in owned_shards: + start_req.shard_push_schedules[shard_idx].CopyFrom( + schedule_proto + ) if cached_protos is not None: cached_protos[target_id] = target_protos + start_req.uuid = int(uuid_val or 0) + start_req.req_id = str(req_id_val or "") + start_req.skip_d2h = skip_d2h_val req.start_transfer_request.CopyFrom(start_req) - return req.SerializeToString() + serialized_bytes = req.SerializeToString() + if payload_cache is not None: + payload_cache[cache_key] = serialized_bytes + if is_sender and is_ws: + payload_cache[steady_key] = serialized_bytes + payload_cache[template_key] = req + return serialized_bytes def build_sender_push_schedule_protos( self, push_schedules: dict[int, list[Any]] ) -> dict[int, Any]: - """Builds ShardPushScheduleProto objects for each shard from raw schedule tuples.""" + """Builds ShardPushScheduleProto objects for shards from schedule tuples.""" target_protos = {} for shard_idx, entries in push_schedules.items(): schedule_proto = self._proto_module.ShardPushScheduleProto() @@ -1029,14 +1267,10 @@ def build_sender_push_schedule_protos( dst_stride = entry_item.dst_stride_bytes count = entry_item.count layer_idx = ( - entry_item.layer_idx - if entry_item.HasField("layer_idx") - else 0 + entry_item.layer_idx if entry_item.HasField("layer_idx") else 0 ) pool_group = ( - entry_item.pool_group - if entry_item.HasField("pool_group") - else 0 + entry_item.pool_group if entry_item.HasField("pool_group") else 0 ) else: ( @@ -1932,8 +2166,12 @@ def register_work_unit( else: if endpoints: cached_sched.rpc_addresses[unit] = ",".join(endpoints) + cached_sched.cached_serialized_payloads.clear() + cached_sched.sender_push_schedule_protos.clear() if unit in cached_sched.data_addresses: cached_sched.data_addresses[unit] = list(normalized_shards) + cached_sched.cached_serialized_payloads.clear() + cached_sched.sender_push_schedule_protos.clear() for k in keys_to_clear: self._plan_cache.pop(k, None) @@ -2546,6 +2784,10 @@ async def _compute_transfer_schedule( unit = _raiden_id_from_proto(meta.unit) if unit in data_addresses: data_addresses[unit] = list(meta.shards) + for unit in src_units: + with self._lock: + if unit in self._registered_shards: + data_addresses[unit] = list(self._registered_shards[unit]) # Group flat entries into slices for broadcast groups = {} @@ -2749,8 +2991,7 @@ def _get_local_metadata(self, units: list[RaidenId]) -> list[Any]: @classmethod def _metadata_by_unit( - cls, - metadata: typing.Sequence[Any], units: typing.Sequence[RaidenId] + cls, metadata: typing.Sequence[Any], units: typing.Sequence[RaidenId] ) -> dict[RaidenId, Any]: """Selects exact requested metadata and rejects duplicate identities.""" requested = set(units) @@ -3187,6 +3428,13 @@ async def _execute_transfer() -> None: is_weight_sync=cached_schedule.is_weight_sync, sender_push_schedule_protos=( cached_schedule.sender_push_schedule_protos + if not broadcast_groups + else {} + ), + cached_serialized_payloads=( + cached_schedule.cached_serialized_payloads + if not broadcast_groups + else {} ), ) with self._lock: @@ -3224,6 +3472,9 @@ async def _execute_transfer() -> None: sender_push_schedule_protos=( cached_schedule.sender_push_schedule_protos ), + cached_serialized_payloads=( + cached_schedule.cached_serialized_payloads + ), ) # 1. Arm direct schedule receivers diff --git a/tpu_sync/rpc/raiden_controller_test.py b/tpu_sync/rpc/raiden_controller_test.py index 1582a1c43..1e9d735f6 100644 --- a/tpu_sync/rpc/raiden_controller_test.py +++ b/tpu_sync/rpc/raiden_controller_test.py @@ -4502,7 +4502,10 @@ def test_ep_multi_host_dst_endpoint_counts_matches_push_tasks(self): dst_endpoint_layer_counts=cached.dst_endpoint_layer_counts, is_weight_sync=True, ) - for dst_unit, host in [(dst_unit_0, "10.0.1.1"), (dst_unit_1, "10.0.1.2")]: + for dst_unit, host in [ + (dst_unit_0, "10.0.1.1"), + (dst_unit_1, "10.0.1.2"), + ]: encoded = ws_client._encode_start_transfer( dst_unit, plan, address=f"{host}:9000" ) @@ -4651,5 +4654,565 @@ def test_encode_start_transfer_receiver_filtering_with_dst_peers(self): client.close() +class SenderScheduleSlicingAndPayloadCachingTest(absltest.TestCase): + + def test_per_worker_schedule_slicing_pathways_multinuma(self): + """Verifies that in Pathways multi-NUMA mode (2 workers per host, sharing IP with different ports), + + each worker receives ONLY its own local shards in ShardPushScheduleProto. + """ + client = raiden_controller.WorkerRpcClient() + try: + src_unit = raiden_controller.RaidenId("trainer", "0", "weights", 0) + dst_unit = raiden_controller.RaidenId("rollout", "0", "weights", 0) + + # 2 hosts, 2 workers per host = 4 endpoints, 8 shards (2 shards per worker). + endpoints = [ + "10.0.0.1:9000", + "10.0.0.1:9001", + "10.0.0.2:9000", + "10.0.0.2:9001", + ] + shards = [ + "10.0.0.1:8000", + "10.0.0.1:8001", + "10.0.0.1:8002", + "10.0.0.1:8003", + "10.0.0.2:8000", + "10.0.0.2:8001", + "10.0.0.2:8002", + "10.0.0.2:8003", + ] + for ep in endpoints: + client.register_worker_endpoint(src_unit, ep) + + # Build 8 shard push schedules + shard_schedules = {} + for shard_idx in range(8): + sched_proto = raiden_service_pb2.ShardPushScheduleProto() + entry = sched_proto.entries.add() + entry.dst_peer = "10.0.1.1:8000" + entry.dst_shard_idx = shard_idx + entry.size_bytes = 4096 + shard_schedules[shard_idx] = sched_proto + + plan = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + shard_push_schedules={src_unit: {}}, + worker_data_addresses={ + src_unit: shards, + dst_unit: ["10.0.1.1:8000"], + }, + is_sender=True, + is_weight_sync=True, + sender_push_schedule_protos={src_unit: shard_schedules}, + ) + + # Worker 0 (host 1 port 9000): should own shards [0, 1] + encoded0 = client._encode_start_transfer( + src_unit, plan, address="10.0.0.1:9000" + ) + req0 = raiden_service_pb2.ControlRequest() + req0.ParseFromString(encoded0) + self.assertEqual( + set(req0.start_transfer_request.shard_push_schedules.keys()), {0, 1} + ) + + # Worker 1 (host 1 port 9001): should own shards [2, 3] + encoded1 = client._encode_start_transfer( + src_unit, plan, address="10.0.0.1:9001" + ) + req1 = raiden_service_pb2.ControlRequest() + req1.ParseFromString(encoded1) + self.assertEqual( + set(req1.start_transfer_request.shard_push_schedules.keys()), {2, 3} + ) + + # Worker 2 (host 2 port 9000): should own shards [4, 5] + encoded2 = client._encode_start_transfer( + src_unit, plan, address="10.0.0.2:9000" + ) + req2 = raiden_service_pb2.ControlRequest() + req2.ParseFromString(encoded2) + self.assertEqual( + set(req2.start_transfer_request.shard_push_schedules.keys()), {4, 5} + ) + + # Worker 3 (host 2 port 9001): should own shards [6, 7] + encoded3 = client._encode_start_transfer( + src_unit, plan, address="10.0.0.2:9001" + ) + req3 = raiden_service_pb2.ControlRequest() + req3.ParseFromString(encoded3) + self.assertEqual( + set(req3.start_transfer_request.shard_push_schedules.keys()), {6, 7} + ) + finally: + client.close() + + def test_serialize_once_when_payload_invariant(self): + """Verifies that when payload is invariant across workers, _encode_start_transfer is called once.""" + + class CountingWorkerRpcClient(raiden_controller.WorkerRpcClient): + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.encode_count = 0 + self.dispatched = [] + + def _encode_start_transfer(self, target_id, transfer_plan, address=None): + self.encode_count += 1 + return super()._encode_start_transfer( + target_id, transfer_plan, address=address + ) + + async def _send_rpc(self, addr, payload, timeout=600.0): + self.dispatched.append((addr, payload)) + resp = raiden_service_pb2.ControlResponse(success=True) + return resp.SerializeToString() + + client = CountingWorkerRpcClient() + try: + src_unit = raiden_controller.RaidenId("trainer", "0", "weights", 0) + dst_unit = raiden_controller.RaidenId("rollout", "0", "weights", 0) + + # 4 worker endpoints on receiver (no endpoint specialization) + for i in range(4): + client.register_worker_endpoint(dst_unit, f"10.0.0.{i+1}:9000") + + plan = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={ + dst_unit: ["10.0.0.1:8000"], + }, + is_sender=False, + expected_block_count=10, + is_weight_sync=True, + ) + + asyncio.run(client.start_transfer(dst_unit, plan)) + # Invariant across 4 workers: encode MUST be called only 1 time + self.assertEqual(client.encode_count, 1) + self.assertEqual(len(client.dispatched), 4) + # All 4 workers must receive the exact same payload bytes + for _, payload in client.dispatched: + self.assertEqual(payload, client.dispatched[0][1]) + finally: + client.close() + + def test_steady_state_payload_caching_step1_reuses_bytes(self): + """Verifies that steady-state step 1+ reuses cached serialized bytes without re-serializing.""" + client = raiden_controller.WorkerRpcClient() + try: + src_unit = raiden_controller.RaidenId("trainer", "0", "weights", 0) + dst_unit = raiden_controller.RaidenId("rollout", "0", "weights", 0) + + client.register_worker_endpoint(src_unit, "10.0.0.1:9000") + client.register_worker_endpoint(src_unit, "10.0.0.1:9001") + + sched0 = raiden_service_pb2.ShardPushScheduleProto() + sched0.entries.add( + dst_peer="10.0.1.1:8000", dst_shard_idx=0, size_bytes=1024 + ) + sched1 = raiden_service_pb2.ShardPushScheduleProto() + sched1.entries.add( + dst_peer="10.0.1.1:8000", dst_shard_idx=1, size_bytes=1024 + ) + + plan = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={ + src_unit: ["10.0.0.1:8000", "10.0.0.1:8001"], + dst_unit: ["10.0.1.1:8000"], + }, + is_sender=True, + is_weight_sync=True, + uuid=42, + req_id="steady_step_0", + sender_push_schedule_protos={src_unit: {0: sched0, 1: sched1}}, + ) + + # Step 0: Initial serialization populates plan.cached_serialized_payloads + bytes_w0 = client._encode_start_transfer( + src_unit, plan, address="10.0.0.1:9000" + ) + bytes_w1 = client._encode_start_transfer( + src_unit, plan, address="10.0.0.1:9001" + ) + self.assertTrue(len(plan.cached_serialized_payloads) > 0) + + # Clear sender_push_schedule_protos to prove Step 1 reuses cached bytes + plan.sender_push_schedule_protos.clear() + + bytes_w0_step1 = client._encode_start_transfer( + src_unit, plan, address="10.0.0.1:9000" + ) + bytes_w1_step1 = client._encode_start_transfer( + src_unit, plan, address="10.0.0.1:9001" + ) + + self.assertEqual(bytes_w0, bytes_w0_step1) + self.assertEqual(bytes_w1, bytes_w1_step1) + finally: + client.close() + + def test_serialize_once_when_payload_invariant_sender_full_schedule(self): + """Verifies that when sender sends full schedule, encode runs once.""" + + class CountingWorkerRpcClient(raiden_controller.WorkerRpcClient): + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.encode_count = 0 + self.dispatched = [] + + def _encode_start_transfer(self, target_id, transfer_plan, address=None): + self.encode_count += 1 + return super()._encode_start_transfer( + target_id, transfer_plan, address=address + ) + + async def _send_rpc(self, addr, payload, timeout=600.0): + self.dispatched.append((addr, payload)) + resp = raiden_service_pb2.ControlResponse(success=True) + return resp.SerializeToString() + + client = CountingWorkerRpcClient() + try: + src_unit = raiden_controller.RaidenId("trainer", "0", "weights", 0) + dst_unit = raiden_controller.RaidenId("rollout", "0", "weights", 0) + + # 4 addresses, but only 1 endpoint registered on target_id (cannot slice, + # sends full schedule). + client.register_worker_endpoint(src_unit, "10.0.0.1:9000") + + sched = raiden_service_pb2.ShardPushScheduleProto() + sched.entries.add( + dst_peer="10.0.1.1:8000", dst_shard_idx=0, size_bytes=1024 + ) + + plan = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={ + src_unit: ["10.0.0.1:8000"], + dst_unit: ["10.0.1.1:8000"], + }, + is_sender=True, + is_weight_sync=True, + sender_push_schedule_protos={src_unit: {0: sched}}, + ) + + # When address="10.0.0.1:9000, 10.0.0.2:9000" but slicing is inactive, + # payload is invariant. + asyncio.run( + client.start_transfer( + src_unit, plan, address="10.0.0.1:9000, 10.0.0.2:9000" + ) + ) + self.assertEqual(client.encode_count, 1) + self.assertEqual(len(client.dispatched), 2) + self.assertEqual(client.dispatched[0][1], client.dispatched[1][1]) + finally: + client.close() + + def test_steady_state_payload_caching_step1_reuses_bytes_across_plans(self): + """Verifies that across distinct plans, cached bytes are reused.""" + client = raiden_controller.WorkerRpcClient() + try: + src_unit = raiden_controller.RaidenId("trainer", "0", "weights", 0) + dst_unit = raiden_controller.RaidenId("rollout", "0", "weights", 0) + + client.register_worker_endpoint(src_unit, "10.0.0.1:9000") + client.register_worker_endpoint(src_unit, "10.0.0.1:9001") + + sched0 = raiden_service_pb2.ShardPushScheduleProto() + sched0.entries.add( + dst_peer="10.0.1.1:8000", dst_shard_idx=0, size_bytes=1024 + ) + sched1 = raiden_service_pb2.ShardPushScheduleProto() + sched1.entries.add( + dst_peer="10.0.1.1:8000", dst_shard_idx=1, size_bytes=1024 + ) + + shared_payload_cache = {} + plan_step0 = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={ + src_unit: ["10.0.0.1:8000", "10.0.0.1:8001"], + dst_unit: ["10.0.1.1:8000"], + }, + is_sender=True, + is_weight_sync=True, + skip_d2h=True, + uuid=100, + req_id="step_0", + sender_push_schedule_protos={src_unit: {0: sched0, 1: sched1}}, + cached_serialized_payloads=shared_payload_cache, + ) + + bytes_w0_step0 = client._encode_start_transfer( + src_unit, plan_step0, address="10.0.0.1:9000" + ) + bytes_w1_step0 = client._encode_start_transfer( + src_unit, plan_step0, address="10.0.0.1:9001" + ) + self.assertNotEmpty(shared_payload_cache) + + # Step 1: Newly instantiated TransferPlan with different req_id and uuid + plan_step1 = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={ + src_unit: ["10.0.0.1:8000", "10.0.0.1:8001"], + dst_unit: ["10.0.1.1:8000"], + }, + is_sender=True, + is_weight_sync=True, + skip_d2h=True, + uuid=101, + req_id="step_1", + # Empty: would fail/empty if re-serialized + sender_push_schedule_protos={}, + cached_serialized_payloads=shared_payload_cache, + ) + + bytes_w0_step1 = client._encode_start_transfer( + src_unit, plan_step1, address="10.0.0.1:9000" + ) + bytes_w1_step1 = client._encode_start_transfer( + src_unit, plan_step1, address="10.0.0.1:9001" + ) + + req_w0_step0 = raiden_service_pb2.ControlRequest() + req_w0_step0.ParseFromString(bytes_w0_step0) + req_w0_step1 = raiden_service_pb2.ControlRequest() + req_w0_step1.ParseFromString(bytes_w0_step1) + self.assertEqual(req_w0_step1.start_transfer_request.uuid, 101) + self.assertEqual(req_w0_step1.start_transfer_request.req_id, "step_1") + self.assertEqual( + req_w0_step1.start_transfer_request.shard_push_schedules, + req_w0_step0.start_transfer_request.shard_push_schedules, + ) + + req_w1_step0 = raiden_service_pb2.ControlRequest() + req_w1_step0.ParseFromString(bytes_w1_step0) + req_w1_step1 = raiden_service_pb2.ControlRequest() + req_w1_step1.ParseFromString(bytes_w1_step1) + self.assertEqual(req_w1_step1.start_transfer_request.uuid, 101) + self.assertEqual( + req_w1_step1.start_transfer_request.shard_push_schedules, + req_w1_step0.start_transfer_request.shard_push_schedules, + ) + + # Step 2: When uuid and skip_d2h match Step 0 (only req_id differs), + # exact serialized bytes are returned without re-encoding. + plan_step2 = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={ + src_unit: ["10.0.0.1:8000", "10.0.0.1:8001"], + dst_unit: ["10.0.1.1:8000"], + }, + is_sender=True, + is_weight_sync=True, + skip_d2h=True, + uuid=100, + req_id="step_2", + sender_push_schedule_protos={}, + cached_serialized_payloads=shared_payload_cache, + ) + bytes_w0_step2 = client._encode_start_transfer( + src_unit, plan_step2, address="10.0.0.1:9000" + ) + bytes_w1_step2 = client._encode_start_transfer( + src_unit, plan_step2, address="10.0.0.1:9001" + ) + self.assertEqual(bytes_w0_step0, bytes_w0_step2) + self.assertEqual(bytes_w1_step0, bytes_w1_step2) + finally: + client.close() + + def test_per_worker_schedule_slicing_uneven_shards(self): + """Verifies that uneven shard-to-worker division accounts for all shards.""" + client = raiden_controller.WorkerRpcClient() + try: + src_unit = raiden_controller.RaidenId("trainer", "0", "weights", 0) + dst_unit = raiden_controller.RaidenId("rollout", "0", "weights", 0) + + # 2 workers, 5 shards + client.register_worker_endpoint(src_unit, "10.0.0.1:9000") + client.register_worker_endpoint(src_unit, "10.0.0.1:9001") + + schedules = {} + for s in range(5): + sched = raiden_service_pb2.ShardPushScheduleProto() + sched.entries.add( + dst_peer="10.0.1.1:8000", dst_shard_idx=s, size_bytes=1024 + ) + schedules[s] = sched + + plan = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={ + src_unit: [f"10.0.0.1:800{s}" for s in range(5)], + dst_unit: ["10.0.1.1:8000"], + }, + is_sender=True, + is_weight_sync=True, + sender_push_schedule_protos={src_unit: schedules}, + ) + + owned0 = client._get_worker_owned_shards( + src_unit, plan, address="10.0.0.1:9000" + ) + owned1 = client._get_worker_owned_shards( + src_unit, plan, address="10.0.0.1:9001" + ) + + self.assertEqual(owned0, {0, 1}) + self.assertEqual(owned1, {2, 3, 4}) + self.assertEqual(owned0 | owned1, set(range(5))) + self.assertEqual(len(owned0 & owned1), 0) + finally: + client.close() + + def test_per_worker_schedule_slicing_more_workers_than_shards(self): + """Verifies slicing when there are more workers than shards.""" + client = raiden_controller.WorkerRpcClient() + try: + src_unit = raiden_controller.RaidenId("trainer", "0", "weights", 0) + dst_unit = raiden_controller.RaidenId("rollout", "0", "weights", 0) + + # 4 workers, 2 shards + for i in range(4): + client.register_worker_endpoint(src_unit, f"10.0.0.1:900{i}") + + plan = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={ + src_unit: ["10.0.0.1:8000", "10.0.0.1:8001"], + dst_unit: ["10.0.1.1:8000"], + }, + is_sender=True, + is_weight_sync=True, + ) + + owned0 = client._get_worker_owned_shards( + src_unit, plan, address="10.0.0.1:9000" + ) + owned1 = client._get_worker_owned_shards( + src_unit, plan, address="10.0.0.1:9001" + ) + owned2 = client._get_worker_owned_shards( + src_unit, plan, address="10.0.0.1:9002" + ) + owned3 = client._get_worker_owned_shards( + src_unit, plan, address="10.0.0.1:9003" + ) + + self.assertEqual(owned0, {0}) + self.assertEqual(owned1, {1}) + self.assertEqual(owned2, set()) + self.assertEqual(owned3, set()) + finally: + client.close() + + def test_single_address_dispatch_preserves_slicing(self): + """Verifies that start_transfer with a single address preserves slicing.""" + + class SingleAddrClient(raiden_controller.WorkerRpcClient): + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.dispatched_reqs = [] + + async def _send_rpc(self, addr, payload, timeout=600.0): + req = raiden_service_pb2.ControlRequest() + req.ParseFromString(payload) + self.dispatched_reqs.append((addr, req)) + resp = raiden_service_pb2.ControlResponse(success=True) + return resp.SerializeToString() + + client = SingleAddrClient() + try: + src_unit = raiden_controller.RaidenId("trainer", "0", "weights", 0) + dst_unit = raiden_controller.RaidenId("rollout", "0", "weights", 0) + + # 2 workers registered on src_unit + client.register_worker_endpoint(src_unit, "10.0.0.1:9000") + client.register_worker_endpoint(src_unit, "10.0.0.1:9001") + + s0 = raiden_service_pb2.ShardPushScheduleProto() + s0.entries.add(dst_peer="10.0.1.1:8000", dst_shard_idx=0, size_bytes=1024) + s1 = raiden_service_pb2.ShardPushScheduleProto() + s1.entries.add(dst_peer="10.0.1.1:8000", dst_shard_idx=1, size_bytes=1024) + + plan = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + worker_data_addresses={ + src_unit: ["10.0.0.1:8000", "10.0.0.1:8001"], + dst_unit: ["10.0.1.1:8000"], + }, + is_sender=True, + is_weight_sync=True, + sender_push_schedule_protos={src_unit: {0: s0, 1: s1}}, + ) + + # Dispatch to single address directly + asyncio.run( + client.start_transfer(src_unit, plan, address="10.0.0.1:9001") + ) + self.assertEqual(len(client.dispatched_reqs), 1) + addr, req = client.dispatched_reqs[0] + self.assertEqual(addr, "10.0.0.1:9001") + # Worker 1 must only receive shard 1, not shard 0 + self.assertEqual( + set(req.start_transfer_request.shard_push_schedules.keys()), {1} + ) + finally: + client.close() + + def test_endpoint_to_shards_direct_address_key(self): + """Verifies endpoint_to_shards when keyed directly by address string.""" + client = raiden_controller.WorkerRpcClient() + try: + src_unit = raiden_controller.RaidenId("trainer", "0", "weights", 0) + dst_unit = raiden_controller.RaidenId("rollout", "0", "weights", 0) + + plan = raiden_controller.TransferPlan( + src_units=[src_unit], + dst_units=[dst_unit], + plan={}, + is_sender=True, + is_weight_sync=True, + endpoint_to_shards={"10.0.0.1:9000": [3, 7]}, + ) + + owned = client._get_worker_owned_shards( + src_unit, plan, address="10.0.0.1:9000" + ) + self.assertEqual(owned, {3, 7}) + finally: + client.close() + + if __name__ == "__main__": absltest.main() diff --git a/tpu_sync/weight_sync/weight_synchronizer_base.cc b/tpu_sync/weight_sync/weight_synchronizer_base.cc index da3fd6247..a58c1bb93 100644 --- a/tpu_sync/weight_sync/weight_synchronizer_base.cc +++ b/tpu_sync/weight_sync/weight_synchronizer_base.cc @@ -751,6 +751,31 @@ absl::Status WeightSynchronizerBase::PushWeightsReshardedLocal( } } } + std::vector sorted_sliced_keys; + if (schedules.size() == num_shards_ && num_shards_ > 0) { + bool is_default_zero_based = true; + bool any_direct_match = false; + for (size_t i = 0; i < num_shards_; ++i) { + int64_t g_sh = global_shard_index(i); + int64_t l_sh = local_shard_index(i); + if (g_sh != static_cast(i) || l_sh != static_cast(i)) { + is_default_zero_based = false; + break; + } + if (schedules.contains(static_cast(l_sh)) || + schedules.contains(static_cast(g_sh))) { + any_direct_match = true; + break; + } + } + if (is_default_zero_based && !any_direct_match) { + sorted_sliced_keys.reserve(schedules.size()); + for (const auto& [sched_key, _] : schedules) { + sorted_sliced_keys.push_back(sched_key); + } + std::sort(sorted_sliced_keys.begin(), sorted_sliced_keys.end()); + } + } for (size_t i = 0; i < num_shards_; ++i) { int64_t global_shard = global_shard_index(i); int64_t local_shard = local_shard_index(i); @@ -767,6 +792,9 @@ absl::Status WeightSynchronizerBase::PushWeightsReshardedLocal( it = schedules.find(static_cast(global_shard)); } } + if (it == schedules.end() && i < sorted_sliced_keys.size()) { + it = schedules.find(sorted_sliced_keys[i]); + } if (it == schedules.end()) { continue; } diff --git a/tpu_sync/weight_sync/weight_synchronizer_test.cc b/tpu_sync/weight_sync/weight_synchronizer_test.cc index cb24f27cc..ce79749c7 100644 --- a/tpu_sync/weight_sync/weight_synchronizer_test.cc +++ b/tpu_sync/weight_sync/weight_synchronizer_test.cc @@ -1782,6 +1782,80 @@ TEST_F(WeightSynchronizerTest, } } +TEST_F(WeightSynchronizerTest, + PreSlicedWorkerScheduleWithDefaultZeroBasedIndicesAndMultiStepUuid) { + const size_t num_layers = 1; + const size_t num_shards = 4; + const size_t slice_byte_size = 256; + + // Source host is Worker 1 (owning cluster shards 4..7), but initialized with + // default 0-based local/global shard indices {0, 1, 2, 3} (e.g. PyTorch or + // default JAX WeightSynchronizer). + auto ws_source = std::make_unique( + num_layers, num_shards, slice_byte_size, + /*local_port=*/0, /*host_blocks_to_allocate=*/1); + auto ws_dest = std::make_unique( + num_layers, num_shards, slice_byte_size, + /*local_port=*/0, /*host_blocks_to_allocate=*/1); + + ASSERT_TRUE(ws_source->local_port().has_value()); + ASSERT_TRUE(ws_dest->local_port().has_value()); + std::string dest_peer = "localhost:" + std::to_string(*ws_dest->local_port()); + + tpu_sync::rpc::StartTransferRequest request; + request.set_skip_d2h(true); + + // Controller sends ONLY the pre-sliced schedule keys {4, 5, 6, 7} owned by + // Worker 1. + auto* schedules = request.mutable_shard_push_schedules(); + for (size_t s = 0; s < num_shards; ++s) { + size_t sliced_key = 4 + s; + auto* entry = (*schedules)[sliced_key].add_entries(); + entry->set_dst_peer(dest_peer); + entry->set_dst_shard_idx(s); + entry->set_src_offset_bytes(0); + entry->set_dst_offset_bytes(0); + entry->set_size_bytes(slice_byte_size); + entry->set_count(1); + entry->set_layer_idx(0); + } + + // Execute two consecutive steady-state steps (uuid=60001, then uuid=60002) + // to verify both pre-sliced key mapping and multi-step uuid synchronization. + const uint8_t step_patterns[2][4] = { + {0x11, 0x22, 0x33, 0x44}, + {0x55, 0x66, 0x77, 0x88}, + }; + const uint64_t step_uuids[2] = {60001, 60002}; + + for (int step = 0; step < 2; ++step) { + for (size_t s = 0; s < num_shards; ++s) { + uint8_t* src_ptr = + const_cast(ws_source->GetHostBufferPtr(0, s)); + uint8_t* dst_ptr = const_cast(ws_dest->GetHostBufferPtr(0, s)); + ASSERT_NE(src_ptr, nullptr); + ASSERT_NE(dst_ptr, nullptr); + std::memset(src_ptr, step_patterns[step][s], slice_byte_size); + std::memset(dst_ptr, 0x00, slice_byte_size); + } + + request.set_uuid(step_uuids[step]); + ASSERT_OK(ws_dest->RegisterExpectedChunks(request.uuid(), num_shards)); + absl::Status status = ws_source->PushWeightsResharded(request); + EXPECT_TRUE(status.ok()) << status.message(); + ASSERT_OK(ws_dest->WaitForTransferCompletion(request.uuid())); + + for (size_t s = 0; s < num_shards; ++s) { + const uint8_t* dst_ptr = ws_dest->GetHostBufferPtr(0, s); + ASSERT_NE(dst_ptr, nullptr); + for (size_t b = 0; b < slice_byte_size; ++b) { + EXPECT_EQ(dst_ptr[b], step_patterns[step][s]) + << "Step " << step << " mismatch at slot " << s << " byte " << b; + } + } + } +} + } // namespace } // namespace weight_sync } // namespace tpu_raiden