From 1806a039782f586ebfade3d294ecd5abf6d33875 Mon Sep 17 00:00:00 2001 From: Andy Twigg Date: Fri, 25 Sep 2026 23:59:42 -0700 Subject: [PATCH] Add offline symbolic and parallel weight-sync plan computation. Previously, `ReshardPlanner` computed the entire resharding transfer schedule sequentially on the centralized controller at step 0 after all workers registered their dynamic `ip:port` endpoints. For large topologies and multi-hundred-billion-parameter models, this centralized calculation delayed step 0 execution. This change introduces two complementary capabilities to eliminate step 0 planning overhead: - Offline symbolic planning and parallel worker loading: `ReshardPlanner.compute_offline_schedule` and `save_offline_plan` compute the N-D resharding schedule ahead of time using deterministic `raiden_symbolic://` endpoint placeholders and serialize per-worker `ControlRequest` protobufs (`.pb`) alongside `offline_schedule.json`. At runtime, `load_offline_worker_plan`, `load_offline_worker_plans_parallel`, and `load_offline_schedule` load the precomputed plans in parallel and bind symbolic placeholders to live `ip:port` endpoints in `O(unique_plans)`. Configurable via `offline_plan_path` / `RAIDEN_OFFLINE_PLAN_PATH` and `save_offline_plan_path` / `RAIDEN_SAVE_OFFLINE_PLAN_PATH`. - Parallel per-worker schedule computation: `compute_transfer_schedule_from_metadata` now accepts `target_src_units` and `parallel_worker_planning`, allowing each source worker's schedule slice to be computed independently or across a parallel worker pool (`compute_schedules_in_parallel_workers`) and merged via `merge_worker_schedules`. Configurable via `parallel_worker_planning` / `RAIDEN_PARALLEL_WORKER_PLANNING` or `RAIDEN_PLANNING_MODE=parallel_worker`. PiperOrigin-RevId: 988717008 --- tpu_sync/api/weight_synchronizer_manager.py | 82 ++ tpu_sync/rpc/raiden_controller.py | 208 +++- tpu_sync/rpc/raiden_controller_test.py | 237 +++++ tpu_sync/weight_sync/manager/BUILD | 1 + .../weight_sync/manager/controller_types.py | 133 ++- .../weight_sync/manager/reshard_planner.py | 970 ++++++++++++++++++ .../manager/reshard_planner_test.py | 367 +++++++ 7 files changed, 1987 insertions(+), 11 deletions(-) diff --git a/tpu_sync/api/weight_synchronizer_manager.py b/tpu_sync/api/weight_synchronizer_manager.py index 90d9a77fd..2d5c2eebf 100644 --- a/tpu_sync/api/weight_synchronizer_manager.py +++ b/tpu_sync/api/weight_synchronizer_manager.py @@ -104,6 +104,8 @@ def __init__( broadcast_k: Optional[int] = None, enable_plan_cache: bool = True, auto_start_server: bool = False, + parallel_worker_planning: Optional[bool] = None, + offline_plan_path: Optional[str] = None, ): """Initializes the WeightSynchronizerManager. @@ -118,6 +120,13 @@ def __init__( schedules across transfer invocations with identical topologies. auto_start_server: Whether to automatically spawn the background TCP servicer loop on initialization. + parallel_worker_planning: If True, computes each source worker's slice of + the transfer schedule independently in parallel across worker units. + Defaults to RAIDEN_PARALLEL_WORKER_PLANNING or RAIDEN_PLANNING_MODE env + vars when None. + offline_plan_path: Optional path to a precomputed offline symbolic plan + directory or file. Defaults to RAIDEN_OFFLINE_PLAN_PATH env var when + None. """ self._controller = raiden_controller.RaidenController( port=port, @@ -125,6 +134,8 @@ def __init__( request_registry_ttl_s=request_registry_ttl_s, broadcast_k=broadcast_k, enable_plan_cache=enable_plan_cache, + parallel_worker_planning=parallel_worker_planning, + offline_plan_path=offline_plan_path, ) self._server: Optional[raiden_controller.RaidenControllerServer] = None if auto_start_server: @@ -210,6 +221,70 @@ def get_plan(self, req_id: str) -> Optional[raiden_controller.TransferPlan]: """Returns the generated TransferPlan for a given transfer request ID.""" return self._controller.get_plan(req_id) + def save_offline_plan( + self, + path: str, + src_units: list[raiden_controller.RaidenId], + dst_units: list[raiden_controller.RaidenId], + group_size: int = 1, + skip_tiling: Optional[dict[int, bool]] = None, + ) -> Any: + """Computes an offline symbolic resharding plan and saves it to `path`.""" + return self._controller.save_offline_plan( + path=path, + src_units=src_units, + dst_units=dst_units, + group_size=group_size, + skip_tiling=skip_tiling, + ) + + def load_offline_schedule( + self, + path: str, + bind_registered_endpoints: bool = True, + ) -> Any: + """Loads an offline schedule from `path` and optionally binds registered endpoints.""" + return self._controller.load_offline_schedule( + path=path, + bind_registered_endpoints=bind_registered_endpoints, + ) + + def load_offline_worker_plan( + self, + path: str, + unit: raiden_controller.RaidenId, + bind_registered_endpoints: bool = True, + req_id: Optional[str] = None, + uuid: Optional[int] = None, + ) -> Any: + """Loads a single worker's offline `ControlRequest` plan and binds live endpoints.""" + return self._controller.load_offline_worker_plan( + path=path, + unit=unit, + bind_registered_endpoints=bind_registered_endpoints, + req_id=req_id, + uuid=uuid, + ) + + def load_offline_worker_plans_parallel( + self, + path: str, + units: Sequence[raiden_controller.RaidenId], + bind_registered_endpoints: bool = True, + req_id: Optional[str] = None, + uuid: Optional[int] = None, + max_workers: Optional[int] = None, + ) -> dict[raiden_controller.RaidenId, Any]: + """Loads per-worker offline `ControlRequest` plans in parallel and binds endpoints.""" + return self._controller.load_offline_worker_plans_parallel( + path=path, + units=units, + bind_registered_endpoints=bind_registered_endpoints, + req_id=req_id, + uuid=uuid, + max_workers=max_workers, + ) + def start_transfer( self, src_units: list[raiden_controller.RaidenId], @@ -233,6 +308,8 @@ def start_transfer( skip_tiling: Optional[dict[int, bool]] = None, group_size: int = 1, use_cached_plan: bool = False, + parallel_worker_planning: Optional[bool] = None, + offline_plan_path: Optional[str] = None, ) -> str: """Initiates an asynchronous distributed transfer. @@ -258,6 +335,9 @@ def start_transfer( skip_tiling: Optional per-device tiling bypass settings. group_size: Variable group size for broadcast pipeline. use_cached_plan: Whether to reuse precomputed plans. + parallel_worker_planning: Optional override to compute per-worker schedule + slices in parallel across source workers. + offline_plan_path: Optional path to a precomputed offline symbolic plan. Returns: Request ID string of the started transfer task. @@ -284,6 +364,8 @@ def start_transfer( skip_tiling=skip_tiling, group_size=group_size, use_cached_plan=use_cached_plan, + parallel_worker_planning=parallel_worker_planning, + offline_plan_path=offline_plan_path, ) def get_transfer_status(self, req_id: str) -> int: diff --git a/tpu_sync/rpc/raiden_controller.py b/tpu_sync/rpc/raiden_controller.py index fd5b27c2b..45d1d6a0a 100644 --- a/tpu_sync/rpc/raiden_controller.py +++ b/tpu_sync/rpc/raiden_controller.py @@ -54,6 +54,16 @@ _proto_to_nd_slice = controller_types.proto_to_nd_slice _raiden_id_from_proto = controller_types.raiden_id_from_proto _raiden_id_to_proto = controller_types.raiden_id_to_proto +SYMBOLIC_ENDPOINT_PREFIX = controller_types.SYMBOLIC_ENDPOINT_PREFIX +make_symbolic_endpoint = controller_types.make_symbolic_endpoint +make_symbolic_shards = controller_types.make_symbolic_shards +is_symbolic_endpoint = controller_types.is_symbolic_endpoint +parse_symbolic_endpoint = controller_types.parse_symbolic_endpoint +resolve_symbolic_endpoint = controller_types.resolve_symbolic_endpoint +bind_symbolic_endpoints_in_proto = ( + controller_types.bind_symbolic_endpoints_in_proto +) +unit_filename_stem = controller_types.unit_filename_stem HostDescriptor = job_entity.HostDescriptor HostGroup = job_entity.HostGroup @@ -115,6 +125,8 @@ def __init__( request_registry_ttl_s: float = 600.0, broadcast_k: Optional[int] = None, enable_plan_cache: bool = True, + parallel_worker_planning: Optional[bool] = None, + offline_plan_path: Optional[str] = None, ): """Initializes the RaidenController. @@ -126,6 +138,13 @@ 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. + parallel_worker_planning: If True, computes each source worker's slice of + the transfer schedule independently in parallel across worker units. + Defaults to RAIDEN_PARALLEL_WORKER_PLANNING or RAIDEN_PLANNING_MODE env + vars when None. + offline_plan_path: Optional path to a precomputed offline symbolic plan + directory or file. Defaults to RAIDEN_OFFLINE_PLAN_PATH env var when + None. """ self.port = port self.broadcast_k = ( @@ -134,6 +153,27 @@ def __init__( else int(os.environ.get("RAIDEN_BROADCAST_K", "64")) ) self.enable_plan_cache = enable_plan_cache + if parallel_worker_planning is None: + env_val = ( + os.environ.get("RAIDEN_PARALLEL_WORKER_PLANNING", "").strip().lower() + ) + mode_val = os.environ.get("RAIDEN_PLANNING_MODE", "").strip().lower() + parallel_worker_planning = env_val in ( + "1", + "true", + "yes", + ) or mode_val in ( + "parallel", + "parallel_worker", + "parallel_workers", + "decentralized", + ) + self.parallel_worker_planning = bool(parallel_worker_planning) + self.offline_plan_path = ( + offline_plan_path + if offline_plan_path is not None + else (os.environ.get("RAIDEN_OFFLINE_PLAN_PATH", "").strip() or None) + ) self._plan_cache: dict[Any, _CachedTransferSchedule] = {} self._active_transfers: dict[str, TransferPlan] = {} self._active_tasks: dict[str, RaidenFuture] = {} @@ -657,6 +697,8 @@ async def warmup_transfer_plan( skip_tiling: Optional[dict[int, bool]] = None, dst_controller_address: Optional[str] = None, src_controller_address: Optional[str] = None, + parallel_worker_planning: Optional[bool] = None, + offline_plan_path: Optional[str] = None, ) -> _CachedTransferSchedule: """Precomputes and caches transfer schedule outside of the critical path.""" if group_size <= 0: @@ -681,6 +723,8 @@ async def warmup_transfer_plan( skip_tiling=skip_tiling, req_id="warmup", uuid=str(random.randint(1, 2**63 - 1)), + parallel_worker_planning=parallel_worker_planning, + offline_plan_path=offline_plan_path, ) raw_schedules = schedule.direct_schedules or schedule.computed_schedules if raw_schedules: @@ -706,6 +750,9 @@ async def _compute_transfer_schedule( ] = None, req_id: str = "warmup", uuid: Any = "", + target_src_units: Optional[Sequence[RaidenId]] = None, + parallel_worker_planning: Optional[bool] = None, + offline_plan_path: Optional[str] = None, ) -> _CachedTransferSchedule: """Computes transfer schedule math via ReshardPlanner.""" t_start = time.perf_counter() @@ -722,12 +769,95 @@ async def _compute_transfer_schedule( dst_metadata = self._get_local_metadata(dst_units) worker_endpoints = self.get_entity_rpc_addresses() + effective_offline_path = ( + offline_plan_path + if offline_plan_path is not None + else self.offline_plan_path + ) + effective_parallel = ( + parallel_worker_planning + if parallel_worker_planning is not None + else self.parallel_worker_planning + ) + + if effective_offline_path and not shard_push_schedules: + live_shards = dict(self._registered_shards) + for item in dst_metadata: + live_shards[_raiden_id_from_proto(item.unit)] = list(item.shards) + schedule = self._planner.load_offline_schedule( + effective_offline_path, + registered_shards=live_shards, + worker_endpoints=worker_endpoints, + entities=self._entities, + ) + else: + schedule = self._planner.compute_transfer_schedule_from_metadata( + src_units=src_units, + dst_units=dst_units, + dst_metadata=dst_metadata, + entities=self._entities, + registered_variables=self._registered_variables, + registered_global_shapes=self._registered_global_shapes, + registered_mesh_shapes=self._registered_mesh_shapes, + registered_mesh_axes=self._registered_mesh_axes, + registered_host_subgrids=self._registered_host_subgrids, + registered_layouts=self._registered_layouts, + registered_itemsizes=self._registered_itemsizes, + registered_shards=self._registered_shards, + computed_phys_meshes=self._computed_phys_meshes, + worker_endpoints=worker_endpoints, + broadcast_k=self.broadcast_k, + lock=self._lock, + group_size=group_size, + skip_tiling=skip_tiling, + shard_push_schedules=shard_push_schedules, + req_id=req_id, + uuid=uuid, + target_src_units=target_src_units, + parallel_worker_planning=effective_parallel, + ) + common.record_histogram( + "weight_sync_schedule_generation_time_ms", + (time.perf_counter() - t_start) * 1000.0, + ) + return schedule - schedule = self._planner.compute_transfer_schedule_from_metadata( + def save_offline_plan( + self, + path: str, + src_units: list[RaidenId], + dst_units: list[RaidenId], + group_size: int = 1, + skip_tiling: Optional[dict[int, bool]] = None, + ) -> _CachedTransferSchedule: + """Computes an offline symbolic resharding plan and saves it to `path`.""" + all_units = list(dict.fromkeys([*src_units, *dst_units])) + symbolic_shards: dict[RaidenId, list[str]] = {} + with self._lock: + for u in all_units: + num_shards = len(self._registered_shards.get(u, ())) + if num_shards <= 0: + subgrid = self._registered_host_subgrids.get( + u + ) or self._registered_mesh_shapes.get(u) + if subgrid: + num_shards = 1 + for d in subgrid: + num_shards *= d + else: + num_shards = 1 + symbolic_shards[u] = make_symbolic_shards(u, num_shards) + dst_metadata = [] + for u in dst_units: + meta = self._metadata_proto_locked(u) + del meta.shards[:] + meta.shards.extend(symbolic_shards[u]) + dst_metadata.append(meta) + + schedule = self._planner.compute_offline_schedule( src_units=src_units, dst_units=dst_units, dst_metadata=dst_metadata, - entities=self._entities, registered_variables=self._registered_variables, registered_global_shapes=self._registered_global_shapes, registered_mesh_shapes=self._registered_mesh_shapes, @@ -735,22 +865,76 @@ async def _compute_transfer_schedule( registered_host_subgrids=self._registered_host_subgrids, registered_layouts=self._registered_layouts, registered_itemsizes=self._registered_itemsizes, - registered_shards=self._registered_shards, + registered_shards=symbolic_shards, computed_phys_meshes=self._computed_phys_meshes, - worker_endpoints=worker_endpoints, broadcast_k=self.broadcast_k, lock=self._lock, group_size=group_size, skip_tiling=skip_tiling, - shard_push_schedules=shard_push_schedules, + ) + self._planner.save_offline_plan( + schedule, path, src_units=src_units, dst_units=dst_units + ) + return schedule + + def load_offline_schedule( + self, + path: str, + bind_registered_endpoints: bool = True, + ) -> _CachedTransferSchedule: + """Loads an offline schedule from `path` and optionally binds registered endpoints.""" + return self._planner.load_offline_schedule( + path=path, + registered_shards=( + self._registered_shards if bind_registered_endpoints else None + ), + worker_endpoints=( + self.get_entity_rpc_addresses() + if bind_registered_endpoints + else None + ), + entities=self._entities if bind_registered_endpoints else None, + ) + + def load_offline_worker_plan( + self, + path: str, + unit: RaidenId, + bind_registered_endpoints: bool = True, + req_id: Optional[str] = None, + uuid: Optional[int] = None, + ) -> Any: + """Loads a single worker's offline `ControlRequest` plan and binds live endpoints.""" + return self._planner.load_offline_worker_plan( + path=path, + unit=unit, + registered_shards=( + self._registered_shards if bind_registered_endpoints else None + ), req_id=req_id, uuid=uuid, ) - common.record_histogram( - "weight_sync_schedule_generation_time_ms", - (time.perf_counter() - t_start) * 1000.0, + + def load_offline_worker_plans_parallel( + self, + path: str, + units: Sequence[RaidenId], + bind_registered_endpoints: bool = True, + req_id: Optional[str] = None, + uuid: Optional[int] = None, + max_workers: Optional[int] = None, + ) -> dict[RaidenId, Any]: + """Loads per-worker offline `ControlRequest` plans in parallel and binds endpoints.""" + return self._planner.load_offline_worker_plans_parallel( + path=path, + units=units, + registered_shards=( + self._registered_shards if bind_registered_endpoints else None + ), + req_id=req_id, + uuid=uuid, + max_workers=max_workers, ) - return schedule def _metadata_proto_locked(self, unit: RaidenId) -> Any: """Builds an owned registration proto while `_lock` is held.""" @@ -761,7 +945,7 @@ def _metadata_proto_locked(self, unit: RaidenId) -> Any: data_name=unit.data_name, data_replica_idx=unit.data_replica_idx, ), - shards=self._registered_shards[unit], + shards=self._registered_shards.get(unit, ()), control_plane_rpc_address=self.worker_endpoints.get(unit, ""), itemsize=self._registered_itemsizes.get(unit, 0), layout_fingerprint=self._registered_layout_fingerprints.get(unit, ""), @@ -938,6 +1122,8 @@ def start_transfer( skip_tiling: Optional[dict[int, bool]] = None, group_size: int = 1, use_cached_plan: bool = True, + parallel_worker_planning: Optional[bool] = None, + offline_plan_path: Optional[str] = None, ) -> RaidenFuture: """Generates a transfer plan for the requested entities and dispatches it.""" if group_size <= 0: @@ -1170,6 +1356,8 @@ async def _execute_transfer_inner() -> None: shard_push_schedules=shard_push_schedules, req_id=req_id, uuid=uuid, + parallel_worker_planning=parallel_worker_planning, + offline_plan_path=offline_plan_path, ) if ( self.enable_plan_cache diff --git a/tpu_sync/rpc/raiden_controller_test.py b/tpu_sync/rpc/raiden_controller_test.py index f3bef96de..ecae905bb 100644 --- a/tpu_sync/rpc/raiden_controller_test.py +++ b/tpu_sync/rpc/raiden_controller_test.py @@ -5814,6 +5814,243 @@ def counting_get_global_indices(unit, *args, **kwargs): finally: controller.worker_rpc_client.close() + def test_controller_offline_plan_save_and_load_transfer(self): + """Verifies offline plan generation, parallel worker plan loading, and step-0 transfer without online reshard math.""" + offline_controller = raiden_controller.RaidenController( + port=0, enable_plan_cache=False + ) + src_units = [ + raiden_controller.RaidenId("trainer", str(i), "weights", 0) + for i in range(2) + ] + dst_units = [ + raiden_controller.RaidenId("sampler", str(j), "weights", 0) + for j in range(2) + ] + src_vars = [ + raiden_service_pb2.VariableMetadataProto( + name="layer_0_w", + shape=[64, 64], + mesh_shape=[2, 4], + layout=[1, 0], + item_size=2, + layer_idx=0, + sharding_spec=["fsdp", "tp"], + ), + raiden_service_pb2.VariableMetadataProto( + name="layer_1_w", + shape=[64, 64], + mesh_shape=[2, 4], + layout=[1, 0], + item_size=2, + layer_idx=1, + sharding_spec=["fsdp", "tp"], + ), + ] + dst_vars = [ + raiden_service_pb2.VariableMetadataProto( + name="layer_0_w", + shape=[64, 64], + mesh_shape=[1, 4], + layout=[1, 0], + item_size=2, + layer_idx=0, + sharding_spec=["fsdp", "tp"], + ), + raiden_service_pb2.VariableMetadataProto( + name="layer_1_w", + shape=[64, 64], + mesh_shape=[1, 4], + layout=[1, 0], + item_size=2, + layer_idx=1, + sharding_spec=["fsdp", "tp"], + ), + ] + try: + # Register with placeholder offline addresses + for i, u in enumerate(src_units): + offline_controller.register_work_unit( + u, + [f"0.0.0.{i + 1}:{8000 + d}" for d in range(4)], + control_plane_rpc_address=f"0.0.0.{i + 1}:9000", + variables=src_vars, + mesh_shape=[2, 4], + mesh_axes=["fsdp", "tp"], + ) + for j, u in enumerate(dst_units): + offline_controller.register_work_unit( + u, + [f"0.0.1.{j + 1}:{8000 + d}" for d in range(4)], + control_plane_rpc_address=f"0.0.1.{j + 1}:9000", + variables=dst_vars, + mesh_shape=[1, 4], + mesh_axes=["fsdp", "tp"], + ) + + plan_dir = self.create_tempdir().full_path + offline_controller.save_offline_plan(plan_dir, src_units, dst_units) + finally: + offline_controller.worker_rpc_client.close() + + # Start a runtime controller configured with offline_plan_path and live IPs. + client = RecordingWorkerRpcClient() + runtime_controller = raiden_controller.RaidenController( + port=0, + worker_rpc_client=client, + offline_plan_path=plan_dir, + ) + try: + for i, u in enumerate(src_units): + runtime_controller.register_work_unit( + u, + [f"10.20.0.{i + 1}:{18000 + d}" for d in range(4)], + control_plane_rpc_address=f"10.20.0.{i + 1}:19000", + variables=src_vars, + mesh_shape=[2, 4], + mesh_axes=["fsdp", "tp"], + ) + for j, u in enumerate(dst_units): + runtime_controller.register_work_unit( + u, + [f"10.30.0.{j + 1}:{28000 + d}" for d in range(4)], + control_plane_rpc_address=f"10.30.0.{j + 1}:29000", + variables=dst_vars, + mesh_shape=[1, 4], + mesh_axes=["fsdp", "tp"], + ) + + # Verify parallel worker loading binds the new live 10.30.0.x endpoints + worker_plans = runtime_controller.load_offline_worker_plans_parallel( + plan_dir, src_units, req_id="offline_req", uuid=999 + ) + self.assertLen(worker_plans, 2) + for u in src_units: + for ( + _, + sched_proto, + ) in worker_plans[ + u + ].start_transfer_request.shard_push_schedules.items(): + for entry in sched_proto.entries: + self.assertTrue(entry.dst_peer.startswith("10.30.0.")) + + # Verify start_transfer uses the offline plan without calling + # compute_transfer_schedule_from_metadata. + with mock.patch.object( + runtime_controller._planner, + "compute_transfer_schedule_from_metadata", + side_effect=AssertionError("Should not run online planning math!"), + ): + fut = runtime_controller.start_transfer( + src_units=src_units, + dst_units=dst_units, + use_block_chunks=True, + req_id="step0_offline", + ) + asyncio.run(fut.wait()) + + plan = runtime_controller.get_plan("step0_offline") + self.assertIsNotNone(plan) + for u in src_units: + for _, entries in plan.shard_push_schedules[u].items(): + for entry in entries: + self.assertTrue(entry[0].startswith("10.30.0.")) + finally: + runtime_controller.worker_rpc_client.close() + + def test_controller_parallel_worker_planning_transfer(self): + """Verifies RaidenController with parallel_worker_planning=True produces identical plans to sequential planning.""" + src_units = [ + raiden_controller.RaidenId("trainer", str(i), "weights", 0) + for i in range(2) + ] + dst_units = [ + raiden_controller.RaidenId("sampler", str(j), "weights", 0) + for j in range(2) + ] + src_vars = [ + raiden_service_pb2.VariableMetadataProto( + name="layer_0_w", + shape=[64, 64], + mesh_shape=[2, 4], + layout=[1, 0], + item_size=2, + layer_idx=0, + sharding_spec=["fsdp", "tp"], + ), + ] + dst_vars = [ + raiden_service_pb2.VariableMetadataProto( + name="layer_0_w", + shape=[64, 64], + mesh_shape=[1, 4], + layout=[1, 0], + item_size=2, + layer_idx=0, + sharding_spec=["fsdp", "tp"], + ), + ] + + plans = [] + for parallel_flag in (False, True): + client = RecordingWorkerRpcClient() + ctrl = raiden_controller.RaidenController( + port=0, + worker_rpc_client=client, + enable_plan_cache=False, + parallel_worker_planning=parallel_flag, + ) + try: + for i, u in enumerate(src_units): + ctrl.register_work_unit( + u, + [f"10.0.0.{i + 1}:{8000 + d}" for d in range(4)], + control_plane_rpc_address=f"10.0.0.{i + 1}:9000", + variables=src_vars, + mesh_shape=[2, 4], + mesh_axes=["fsdp", "tp"], + ) + for j, u in enumerate(dst_units): + ctrl.register_work_unit( + u, + [f"10.0.1.{j + 1}:{8000 + d}" for d in range(4)], + control_plane_rpc_address=f"10.0.1.{j + 1}:9000", + variables=dst_vars, + mesh_shape=[1, 4], + mesh_axes=["fsdp", "tp"], + ) + fut = ctrl.start_transfer( + src_units=src_units, + dst_units=dst_units, + use_block_chunks=True, + req_id="req_cmp", + uuid=12345, + ) + asyncio.run(fut.wait()) + plans.append(ctrl.get_plan("req_cmp")) + finally: + ctrl.worker_rpc_client.close() + + self.assertLen(plans, 2) + seq_plan = plans[0] + par_plan = plans[1] + self.assertEqual( + seq_plan.expected_block_count, par_plan.expected_block_count + ) + self.assertEqual(seq_plan.dst_endpoint_counts, par_plan.dst_endpoint_counts) + for u in src_units: + self.assertEqual( + {k: list(v) for k, v in seq_plan.shard_push_schedules[u].items()}, + {k: list(v) for k, v in par_plan.shard_push_schedules[u].items()}, + ) + + with mock.patch.dict( + os.environ, {"RAIDEN_PLANNING_MODE": "parallel"}, clear=False + ): + env_ctrl = raiden_controller.RaidenController(port=0) + self.assertTrue(env_ctrl.parallel_worker_planning) + if __name__ == "__main__": absltest.main() diff --git a/tpu_sync/weight_sync/manager/BUILD b/tpu_sync/weight_sync/manager/BUILD index 6966bcfe5..d7ce20566 100644 --- a/tpu_sync/weight_sync/manager/BUILD +++ b/tpu_sync/weight_sync/manager/BUILD @@ -80,6 +80,7 @@ py_library( ":job_entity", "//tpu_sync/api:common", "//tpu_sync/kv_cache:nd_slice_math", + "//tpu_sync/rpc:raiden_service_py_pb2", "@com_google_absl_py//absl/logging", ], ) diff --git a/tpu_sync/weight_sync/manager/controller_types.py b/tpu_sync/weight_sync/manager/controller_types.py index d75ca5190..f99fe424c 100644 --- a/tpu_sync/weight_sync/manager/controller_types.py +++ b/tpu_sync/weight_sync/manager/controller_types.py @@ -18,7 +18,7 @@ import dataclasses import enum import threading -from typing import Any, Callable, Mapping, Optional, Protocol +from typing import Any, Callable, Mapping, Optional, Protocol, Sequence from tpu_sync.api.common import RaidenId from tpu_sync.rpc import raiden_service_pb2 @@ -507,6 +507,130 @@ def _entity_key_from_unit(unit: RaidenId) -> RaidenId: return RaidenId(job_name=job_name) +SYMBOLIC_ENDPOINT_PREFIX = "raiden_symbolic://" + + +def _make_symbolic_endpoint(unit: RaidenId, shard_idx: int) -> str: + """Builds a deterministic symbolic endpoint string for an offline shard.""" + return ( + f"{SYMBOLIC_ENDPOINT_PREFIX}{unit.job_name}/{unit.job_replica_id}/" + f"{unit.data_name}/{unit.data_replica_idx}:{int(shard_idx)}" + ) + + +def _make_symbolic_shards(unit: RaidenId, num_shards: int) -> list[str]: + """Builds a list of symbolic shard endpoints for `unit`.""" + return [_make_symbolic_endpoint(unit, i) for i in range(max(1, num_shards))] + + +def _is_symbolic_endpoint(endpoint: str) -> bool: + """Returns True if `endpoint` is an offline symbolic shard placeholder.""" + return isinstance(endpoint, str) and endpoint.startswith( + SYMBOLIC_ENDPOINT_PREFIX + ) + + +def _parse_symbolic_endpoint(endpoint: str) -> Optional[tuple[RaidenId, int]]: + """Parses a symbolic endpoint into (RaidenId, shard_idx), or None if not symbolic.""" + if not _is_symbolic_endpoint(endpoint): + return None + body = endpoint[len(SYMBOLIC_ENDPOINT_PREFIX) :] + if ":" not in body: + return None + unit_part, shard_str = body.rsplit(":", 1) + parts = unit_part.split("/") + if len(parts) != 4: + return None + try: + shard_idx = int(shard_str) + data_rep_idx = int(parts[3]) + except ValueError: + return None + return ( + RaidenId( + job_name=parts[0], + job_replica_id=parts[1], + data_name=parts[2], + data_replica_idx=data_rep_idx, + ), + shard_idx, + ) + + +def _resolve_symbolic_endpoint( + endpoint: str, + live_data_addresses: Mapping[RaidenId, Sequence[str]], +) -> str: + """Resolves a symbolic endpoint string against `live_data_addresses`.""" + parsed = _parse_symbolic_endpoint(endpoint) + if parsed is None: + return endpoint + unit, shard_idx = parsed + shards = live_data_addresses.get(unit) + if not shards: + return endpoint + if 0 <= shard_idx < len(shards): + return shards[shard_idx] + return shards[0] + + +def _bind_symbolic_endpoints_in_proto( + req: Any, + live_data_addresses: Mapping[RaidenId, Sequence[str]], +) -> Any: + """Binds symbolic endpoints in a ControlRequest or StartTransferRequest proto in-place.""" + if not live_data_addresses: + return req + if hasattr(req, "peers") and req.peers: + resolved_peers = [ + _resolve_symbolic_endpoint(p, live_data_addresses) for p in req.peers + ] + del req.peers[:] + req.peers.extend(resolved_peers) + + start_req = ( + req.start_transfer_request + if hasattr(req, "start_transfer_request") + else req + ) + if hasattr(start_req, "shard_push_schedules"): + for _, schedule_proto in start_req.shard_push_schedules.items(): + for entry in schedule_proto.entries: + if entry.dst_peers: + resolved_list = [] + for p in entry.dst_peers: + rp = _resolve_symbolic_endpoint(p, live_data_addresses) + if rp not in resolved_list: + resolved_list.append(rp) + del entry.dst_peers[:] + entry.dst_peers.extend(resolved_list) + if resolved_list: + entry.dst_peer = resolved_list[0] + elif entry.dst_peer: + entry.dst_peer = _resolve_symbolic_endpoint( + entry.dst_peer, live_data_addresses + ) + return req + + +def _unit_filename_stem(unit: RaidenId) -> str: + """Returns a filesystem-safe filename stem for a work unit's offline plan.""" + + def _clean(s: str) -> str: + return ( + str(s) + .replace("/", "_") + .replace(":", "_") + .replace("\\", "_") + .replace(" ", "_") + ) + + return ( + f"plan_{_clean(unit.job_name)}_{_clean(unit.job_replica_id)}_" + f"{_clean(unit.data_name)}_{int(unit.data_replica_idx)}" + ) + + VariableMetadata = _VariableMetadata CachedTransferSchedule = _CachedTransferSchedule PlanReferencedShardSchedule = _PlanReferencedShardSchedule @@ -520,3 +644,10 @@ def _entity_key_from_unit(unit: RaidenId) -> RaidenId: coerce_pool_spec_proto = _coerce_pool_spec_proto format_unit = _format_unit format_units = _format_units +make_symbolic_endpoint = _make_symbolic_endpoint +make_symbolic_shards = _make_symbolic_shards +is_symbolic_endpoint = _is_symbolic_endpoint +parse_symbolic_endpoint = _parse_symbolic_endpoint +resolve_symbolic_endpoint = _resolve_symbolic_endpoint +bind_symbolic_endpoints_in_proto = _bind_symbolic_endpoints_in_proto +unit_filename_stem = _unit_filename_stem diff --git a/tpu_sync/weight_sync/manager/reshard_planner.py b/tpu_sync/weight_sync/manager/reshard_planner.py index ebc717dc0..76d664630 100644 --- a/tpu_sync/weight_sync/manager/reshard_planner.py +++ b/tpu_sync/weight_sync/manager/reshard_planner.py @@ -14,7 +14,11 @@ """Resharding plan and N-D slice math for RaidenController.""" +import concurrent.futures +import dataclasses +import json import math +import os import sys import threading from typing import Any, Mapping, Optional, Sequence @@ -23,6 +27,7 @@ from tpu_sync.api.common import RaidenId from tpu_sync.kv_cache import nd_slice_math +from tpu_sync.rpc import raiden_service_pb2 from tpu_sync.weight_sync.manager import broadcast_engine from tpu_sync.weight_sync.manager import controller_types from tpu_sync.weight_sync.manager import job_entity @@ -32,11 +37,20 @@ _CachedTransferSchedule = controller_types.CachedTransferSchedule _PlanReferencedShardSchedule = controller_types.PlanReferencedShardSchedule _VariableMetadata = controller_types.VariableMetadata +_bind_symbolic_endpoints_in_proto = ( + controller_types.bind_symbolic_endpoints_in_proto +) _extract_host_ip = controller_types.extract_host_ip _format_units = controller_types.format_units +_is_symbolic_endpoint = controller_types.is_symbolic_endpoint _is_variable_spec_identical = controller_types.is_variable_spec_identical +_make_symbolic_endpoint = controller_types.make_symbolic_endpoint +_make_symbolic_shards = controller_types.make_symbolic_shards +_parse_symbolic_endpoint = controller_types.parse_symbolic_endpoint _proto_to_nd_slice = controller_types.proto_to_nd_slice _raiden_id_from_proto = controller_types.raiden_id_from_proto +_resolve_symbolic_endpoint = controller_types.resolve_symbolic_endpoint +_unit_filename_stem = controller_types.unit_filename_stem def to_physical(logical_shape, logical_mesh_shape, minor_to_major): @@ -725,11 +739,62 @@ def compute_transfer_schedule_from_metadata( ] = None, req_id: str = "warmup", uuid: Any = "", + target_src_units: Optional[Sequence[RaidenId]] = None, + parallel_worker_planning: Optional[bool] = None, ) -> _CachedTransferSchedule: """Computes transfer schedule math and returns a _CachedTransferSchedule.""" if group_size <= 0: raise ValueError("group_size must be positive") + if parallel_worker_planning is None: + env_parallel = ( + os.environ.get("RAIDEN_PARALLEL_WORKER_PLANNING", "").strip().lower() + ) + env_mode = os.environ.get("RAIDEN_PLANNING_MODE", "").strip().lower() + parallel_worker_planning = env_parallel in ( + "1", + "true", + "yes", + ) or env_mode in ( + "parallel", + "parallel_worker", + "parallel_workers", + "decentralized", + ) + + if ( + parallel_worker_planning + and target_src_units is None + and not shard_push_schedules + and len(src_units) > 1 + ): + return cls.compute_schedules_in_parallel_workers( + src_units=src_units, + dst_units=dst_units, + dst_metadata=dst_metadata, + entities=entities, + registered_variables=registered_variables, + registered_global_shapes=registered_global_shapes, + registered_mesh_shapes=registered_mesh_shapes, + registered_mesh_axes=registered_mesh_axes, + registered_host_subgrids=registered_host_subgrids, + registered_layouts=registered_layouts, + registered_itemsizes=registered_itemsizes, + registered_shards=registered_shards, + computed_phys_meshes=computed_phys_meshes, + worker_endpoints=worker_endpoints, + broadcast_k=broadcast_k, + lock=lock, + group_size=group_size, + skip_tiling=skip_tiling, + req_id=req_id, + uuid=uuid, + ) + + target_src_set = ( + set(target_src_units) if target_src_units is not None else None + ) + computed_schedules = {} computed_slices = {} data_address_to_unit = {} @@ -1195,6 +1260,9 @@ def _get_or_compute_slices( unit_layers_tuple_by_pid, ) = classified + if target_src_set is not None and src_unit not in target_src_set: + continue + # Calculate the plan ONLY ONCE per unique `plan_id` on this src_unit. # Any subsequent variable sharing the same `plan_id` already points to # `plan_id` via `unit_var_to_plan_id` and skips calculation completely. @@ -1809,3 +1877,905 @@ def _get_or_compute_slices( variable_plans=variable_plans, variable_to_plan_id=variable_to_plan_id, ) + + @classmethod + def merge_worker_schedules( + cls, + worker_schedules: Sequence[_CachedTransferSchedule], + src_units: Sequence[RaidenId], + dst_units: Sequence[RaidenId], + broadcast_k: int = 64, + group_size: int = 1, + ) -> _CachedTransferSchedule: + """Merges independently computed per-worker schedules into a unified schedule.""" + del cls + if not worker_schedules: + return _CachedTransferSchedule( + computed_schedules={}, + direct_schedules={}, + broadcast_groups={}, + local_skip_tiling={}, + expected_block_count=0, + dst_unit_layer_counts={}, + data_address_to_unit={}, + direct_dsts=[], + rpc_addresses={}, + data_addresses={u: [] for u in dst_units}, + ) + + computed_schedules: dict[Any, Any] = {} + variable_plans: dict[Any, Any] = {} + variable_to_plan_id: dict[Any, Any] = {} + local_skip_tiling: dict[int, bool] = {} + data_address_to_unit: dict[str, Any] = {} + rpc_addresses: dict[Any, str] = {} + data_addresses: dict[Any, list[str]] = {u: [] for u in dst_units} + is_weight_sync = False + + can_fast_path_direct = len(set(dst_units)) <= max(1, broadcast_k) + direct_schedules: dict[Any, Any] = {} + direct_dsts: list[Any] = [] + direct_dsts_set: set[Any] = set() + dst_unit_counts: dict[Any, int] = {} + dst_unit_layer_counts: dict[Any, dict[int, int]] = {} + dst_endpoint_counts: dict[str, int] = {} + dst_endpoint_layer_counts: dict[str, dict[int, int]] = {} + + sched_by_src: dict[RaidenId, _CachedTransferSchedule] = {} + for w_sched in worker_schedules: + for u in w_sched.computed_schedules: + sched_by_src[u] = w_sched + for u in w_sched.variable_plans: + if u not in sched_by_src: + sched_by_src[u] = w_sched + + ordered_scheds = list( + { + id(sched_by_src[u]): sched_by_src[u] + for u in src_units + if u in sched_by_src + }.values() + ) or list({id(s): s for s in worker_schedules}.values()) + + for w_sched in ordered_scheds: + computed_schedules.update(w_sched.computed_schedules) + variable_plans.update(w_sched.variable_plans) + variable_to_plan_id.update(w_sched.variable_to_plan_id) + local_skip_tiling.update(w_sched.local_skip_tiling) + data_address_to_unit.update(w_sched.data_address_to_unit) + rpc_addresses.update(w_sched.rpc_addresses) + for k, v in w_sched.data_addresses.items(): + if v: + data_addresses[k] = list(v) + if w_sched.is_weight_sync: + is_weight_sync = True + + if can_fast_path_direct and not w_sched.broadcast_groups: + direct_schedules.update(w_sched.direct_schedules) + for d_u in w_sched.direct_dsts: + if d_u not in direct_dsts_set: + direct_dsts_set.add(d_u) + direct_dsts.append(d_u) + for d_u, cnt in w_sched.dst_unit_counts.items(): + dst_unit_counts[d_u] = dst_unit_counts.get(d_u, 0) + cnt + for d_u, l_map in w_sched.dst_unit_layer_counts.items(): + target_l_map = dst_unit_layer_counts.setdefault(d_u, {}) + for l_idx, cnt in l_map.items(): + target_l_map[l_idx] = target_l_map.get(l_idx, 0) + cnt + for d_h, cnt in w_sched.dst_endpoint_counts.items(): + dst_endpoint_counts[d_h] = dst_endpoint_counts.get(d_h, 0) + cnt + for d_h, l_map in w_sched.dst_endpoint_layer_counts.items(): + target_h_map = dst_endpoint_layer_counts.setdefault(d_h, {}) + for l_idx, cnt in l_map.items(): + target_h_map[l_idx] = target_h_map.get(l_idx, 0) + cnt + + if can_fast_path_direct and all( + not w.broadcast_groups for w in ordered_scheds + ): + broadcast_groups = {} + expected_block_count = ( + max(dst_unit_counts.values()) if dst_unit_counts else 0 + ) + else: + groups = {} + for src_unit in src_units: + schedules = computed_schedules.get(src_unit, {}) + for shard_idx, entries in schedules.items(): + for entry in entries: + ( + dst_peer, + dst_shard_idx, + dst_block_offset, + src_block_offset, + size, + src_block_id, + dst_block_id, + src_stride, + dst_stride, + count, + layer_idx, + pool_group, + ) = entry + dst_unit = data_address_to_unit.get(dst_peer) + if not dst_unit: + continue + key = ( + src_unit, + shard_idx, + src_block_id, + src_block_offset, + size, + src_stride, + count, + layer_idx, + pool_group, + ) + val = ( + dst_unit, + dst_peer, + dst_shard_idx, + dst_block_id, + dst_block_offset, + dst_stride, + ) + groups.setdefault(key, []).append(val) + + direct_schedules, broadcast_groups = ( + BroadcastEngine.partition_direct_and_broadcast_groups( + groups, broadcast_k, group_size + ) + ) + dst_unit_counts = {} + dst_unit_layer_counts = {} + dst_endpoint_counts = {} + dst_endpoint_layer_counts = {} + expected_block_count = 0 + direct_dsts = [] + if direct_schedules: + for _, schedules in direct_schedules.items(): + for _, entries in schedules.items(): + for entry in entries: + dst_peer = entry[0] + dst_unit = data_address_to_unit.get(dst_peer) + if dst_unit: + if dst_unit not in direct_dsts: + direct_dsts.append(dst_unit) + layer_idx = entry[10] if len(entry) > 10 else 0 + dst_unit_counts[dst_unit] = dst_unit_counts.get(dst_unit, 0) + 1 + dst_unit_layer_counts.setdefault(dst_unit, {})[layer_idx] = ( + dst_unit_layer_counts.get(dst_unit, {}).get(layer_idx, 0) + + 1 + ) + dst_host = _extract_host_ip(dst_peer) + if dst_host: + dst_endpoint_counts[dst_host] = ( + dst_endpoint_counts.get(dst_host, 0) + 1 + ) + dst_endpoint_layer_counts.setdefault(dst_host, {})[ + layer_idx + ] = ( + dst_endpoint_layer_counts.get(dst_host, {}).get( + layer_idx, 0 + ) + + 1 + ) + if dst_unit_counts: + expected_block_count = max(dst_unit_counts.values()) + + return _CachedTransferSchedule( + computed_schedules=computed_schedules, + direct_schedules=direct_schedules, + broadcast_groups=broadcast_groups, + local_skip_tiling=local_skip_tiling, + expected_block_count=expected_block_count, + dst_unit_layer_counts=dst_unit_layer_counts, + data_address_to_unit=data_address_to_unit, + direct_dsts=direct_dsts, + rpc_addresses=rpc_addresses, + data_addresses=data_addresses, + dst_unit_counts=dst_unit_counts, + dst_endpoint_counts=dst_endpoint_counts, + dst_endpoint_layer_counts=dst_endpoint_layer_counts, + is_weight_sync=is_weight_sync, + variable_plans=variable_plans, + variable_to_plan_id=variable_to_plan_id, + ) + + @classmethod + def compute_schedules_in_parallel_workers( + cls, + src_units: list[RaidenId], + dst_units: list[RaidenId], + dst_metadata: list[Any], + entities: Mapping[RaidenId, JobEntity], + registered_variables: Mapping[RaidenId, list[Any]], + registered_global_shapes: Mapping[RaidenId, list[int]], + registered_mesh_shapes: Mapping[RaidenId, list[int]], + registered_mesh_axes: Mapping[RaidenId, list[str]], + registered_host_subgrids: Mapping[RaidenId, list[int]], + registered_layouts: Mapping[RaidenId, list[int]], + registered_itemsizes: Mapping[RaidenId, int], + registered_shards: Mapping[RaidenId, list[str]], + computed_phys_meshes: dict[RaidenId, list[int]], + worker_endpoints: dict[RaidenId, str], + broadcast_k: int, + lock: threading.Lock, + group_size: int = 1, + skip_tiling: Optional[dict[int, bool]] = None, + req_id: str = "warmup", + uuid: Any = "", + max_workers: Optional[int] = None, + ) -> _CachedTransferSchedule: + """Computes each source worker's schedule in parallel and merges the results.""" + + def _compute_for_unit(unit: RaidenId) -> _CachedTransferSchedule: + local_phys_meshes: dict[RaidenId, list[int]] = {} + sched = cls.compute_transfer_schedule_from_metadata( + src_units=src_units, + dst_units=dst_units, + dst_metadata=dst_metadata, + entities=entities, + registered_variables=registered_variables, + registered_global_shapes=registered_global_shapes, + registered_mesh_shapes=registered_mesh_shapes, + registered_mesh_axes=registered_mesh_axes, + registered_host_subgrids=registered_host_subgrids, + registered_layouts=registered_layouts, + registered_itemsizes=registered_itemsizes, + registered_shards=registered_shards, + computed_phys_meshes=local_phys_meshes, + worker_endpoints=worker_endpoints, + broadcast_k=broadcast_k, + lock=lock, + group_size=group_size, + skip_tiling=skip_tiling, + req_id=req_id, + uuid=uuid, + target_src_units=[unit], + parallel_worker_planning=False, + ) + with lock: + computed_phys_meshes.update(local_phys_meshes) + return sched + + num_threads = max_workers or min(32, max(1, len(src_units))) + with concurrent.futures.ThreadPoolExecutor(max_workers=num_threads) as pool: + worker_schedules = list(pool.map(_compute_for_unit, src_units)) + + return cls.merge_worker_schedules( + worker_schedules=worker_schedules, + src_units=src_units, + dst_units=dst_units, + broadcast_k=broadcast_k, + group_size=group_size, + ) + + @classmethod + def compute_offline_schedule( + cls, + src_units: list[RaidenId], + dst_units: list[RaidenId], + src_variables: Optional[Mapping[RaidenId, list[Any]]] = None, + dst_variables: Optional[Mapping[RaidenId, list[Any]]] = None, + num_src_shards_per_unit: int = 8, + num_dst_shards_per_unit: int = 8, + src_mesh_shapes: Optional[Mapping[RaidenId, Sequence[int]]] = None, + src_mesh_axes: Optional[Mapping[RaidenId, Sequence[str]]] = None, + src_host_subgrids: Optional[Mapping[RaidenId, Sequence[int]]] = None, + dst_mesh_shapes: Optional[Mapping[RaidenId, Sequence[int]]] = None, + dst_mesh_axes: Optional[Mapping[RaidenId, Sequence[str]]] = None, + dst_host_subgrids: Optional[Mapping[RaidenId, Sequence[int]]] = None, + num_shards_by_unit: Optional[Mapping[RaidenId, int]] = None, + broadcast_k: int = 64, + group_size: int = 1, + skip_tiling: Optional[dict[int, bool]] = None, + target_src_units: Optional[Sequence[RaidenId]] = None, + parallel_worker_planning: bool = False, + **legacy_kwargs: Any, + ) -> _CachedTransferSchedule: + """Computes a symbolic transfer schedule offline without live worker IP:ports.""" + lock = legacy_kwargs.get("lock") or threading.Lock() + if src_variables is None: + src_variables = legacy_kwargs.get("registered_variables", {}) + if src_mesh_shapes is None: + src_mesh_shapes = legacy_kwargs.get("registered_mesh_shapes") + if src_mesh_axes is None: + src_mesh_axes = legacy_kwargs.get("registered_mesh_axes") + if src_host_subgrids is None: + src_host_subgrids = legacy_kwargs.get("registered_host_subgrids") + legacy_shards = legacy_kwargs.get("registered_shards", {}) + legacy_dst_metadata = legacy_kwargs.get("dst_metadata") + + entities = {} + registered_shards = {} + registered_mesh_shapes = {} + registered_mesh_axes = {} + registered_host_subgrids = {} + worker_endpoints = {} + + for u in src_units: + if num_shards_by_unit and u in num_shards_by_unit: + n_shards = num_shards_by_unit[u] + elif u in legacy_shards and legacy_shards[u]: + n_shards = len(legacy_shards[u]) + else: + n_shards = num_src_shards_per_unit + shards = _make_symbolic_shards(u, n_shards) + registered_shards[u] = shards + entities[u] = JobEntity(unit=u, shards=shards) + if src_mesh_shapes and u in src_mesh_shapes: + registered_mesh_shapes[u] = list(src_mesh_shapes[u]) + if src_mesh_axes and u in src_mesh_axes: + registered_mesh_axes[u] = list(src_mesh_axes[u]) + if src_host_subgrids and u in src_host_subgrids: + registered_host_subgrids[u] = list(src_host_subgrids[u]) + worker_endpoints[u] = _make_symbolic_endpoint(u, 0) + + dst_metadata = [] + if legacy_dst_metadata is not None and dst_variables is None: + for item in legacy_dst_metadata: + u = controller_types.raiden_id_from_proto(item.unit) + if num_shards_by_unit and u in num_shards_by_unit: + n_shards = num_shards_by_unit[u] + elif item.shards: + n_shards = len(item.shards) + else: + n_shards = num_dst_shards_per_unit + sym_shards = _make_symbolic_shards(u, n_shards) + registered_shards[u] = sym_shards + meta = raiden_service_pb2.RegisterWorkUnitRequest() + meta.CopyFrom(item) + del meta.shards[:] + meta.shards.extend(sym_shards) + meta.control_plane_rpc_address = _make_symbolic_endpoint(u, 0) + dst_metadata.append(meta) + else: + dst_vars_map = dst_variables or {} + proto_vars_cache = {} + for u in dst_units: + if num_shards_by_unit and u in num_shards_by_unit: + n_shards = num_shards_by_unit[u] + elif u in legacy_shards and legacy_shards[u]: + n_shards = len(legacy_shards[u]) + else: + n_shards = num_dst_shards_per_unit + shards = _make_symbolic_shards(u, n_shards) + registered_shards[u] = shards + meta = raiden_service_pb2.RegisterWorkUnitRequest( + unit=raiden_service_pb2.RaidenIdProto( + job_name=u.job_name, + job_replica_id=str(u.job_replica_id), + data_name=u.data_name, + data_replica_idx=u.data_replica_idx, + ), + control_plane_rpc_address=_make_symbolic_endpoint(u, 0), + ) + meta.shards.extend(shards) + if dst_mesh_shapes and u in dst_mesh_shapes: + meta.mesh_shape.extend(dst_mesh_shapes[u]) + if dst_mesh_axes and u in dst_mesh_axes: + meta.mesh_axes.extend(dst_mesh_axes[u]) + if dst_host_subgrids and u in dst_host_subgrids: + meta.host_subgrid.extend(dst_host_subgrids[u]) + u_vars = dst_vars_map.get(u, []) + proto_vars = proto_vars_cache.get(id(u_vars)) + if proto_vars is None: + tmpl = raiden_service_pb2.RegisterWorkUnitRequest() + for v in u_vars: + vp = tmpl.variables.add() + vp.CopyFrom( + controller_types.coerce_variable_proto(v, raiden_service_pb2) + ) + proto_vars = list(tmpl.variables) + proto_vars_cache[id(u_vars)] = proto_vars + meta.variables.extend(proto_vars) + dst_metadata.append(meta) + + return cls.compute_transfer_schedule_from_metadata( + src_units=src_units, + dst_units=dst_units, + dst_metadata=dst_metadata, + entities=entities, + registered_variables=src_variables, + registered_global_shapes=legacy_kwargs.get( + "registered_global_shapes", {} + ), + registered_mesh_shapes=registered_mesh_shapes, + registered_mesh_axes=registered_mesh_axes, + registered_host_subgrids=registered_host_subgrids, + registered_layouts=legacy_kwargs.get("registered_layouts", {}), + registered_itemsizes=legacy_kwargs.get("registered_itemsizes", {}), + registered_shards=registered_shards, + computed_phys_meshes={}, + worker_endpoints=worker_endpoints, + broadcast_k=broadcast_k, + lock=lock, + group_size=group_size, + skip_tiling=skip_tiling, + req_id="offline", + uuid=0, + target_src_units=target_src_units, + parallel_worker_planning=parallel_worker_planning, + ) + + @classmethod + def bind_symbolic_schedule( + cls, + schedule: _CachedTransferSchedule, + live_data_addresses: Mapping[RaidenId, Sequence[str]], + live_rpc_addresses: Optional[Mapping[RaidenId, str]] = None, + ) -> _CachedTransferSchedule: + """Binds symbolic endpoints in an offline _CachedTransferSchedule to live IP:ports.""" + del cls + if not live_data_addresses: + return schedule + + bound_variable_plans: dict[Any, dict[int, dict[int, list[Any]]]] = {} + bound_computed_schedules: dict[Any, Any] = {} + + for src_unit, unit_plans_by_id in schedule.variable_plans.items(): + new_unit_plans_by_id: dict[int, dict[int, list[Any]]] = {} + new_unit_shard_plans_by_id: dict[int, dict[int, list[Any]]] = {} + for pid, tmpl_for_var in unit_plans_by_id.items(): + new_tmpl_for_var: dict[int, list[Any]] = {} + for local_src_idx, entries in tmpl_for_var.items(): + bound_entries = [ + (_resolve_symbolic_endpoint(e[0], live_data_addresses), *e[1:]) + for e in entries + ] + new_tmpl_for_var[local_src_idx] = bound_entries + new_unit_shard_plans_by_id.setdefault(local_src_idx, {})[ + pid + ] = bound_entries + new_unit_plans_by_id[pid] = new_tmpl_for_var + bound_variable_plans[src_unit] = new_unit_plans_by_id + + orig_unit_scheds = schedule.computed_schedules.get(src_unit, {}) + unit_var_to_pid = schedule.variable_to_plan_id.get(src_unit, {}) + new_unit_scheds = {} + for local_src_idx, orig_sched in orig_unit_scheds.items(): + ordered_vars = getattr(orig_sched, "_ordered_vars", None) + if local_src_idx in new_unit_shard_plans_by_id: + new_unit_scheds[local_src_idx] = _PlanReferencedShardSchedule( + new_unit_shard_plans_by_id[local_src_idx], + unit_var_to_pid, + ordered_vars, + ) + elif isinstance(orig_sched, (list, tuple)): + new_unit_scheds[local_src_idx] = [ + (_resolve_symbolic_endpoint(e[0], live_data_addresses), *e[1:]) + for e in orig_sched + ] + if new_unit_scheds: + bound_computed_schedules[src_unit] = new_unit_scheds + + for src_unit, orig_unit_scheds in schedule.computed_schedules.items(): + if src_unit not in bound_computed_schedules: + new_unit_scheds = {} + for local_src_idx, orig_sched in orig_unit_scheds.items(): + new_unit_scheds[local_src_idx] = [ + (_resolve_symbolic_endpoint(e[0], live_data_addresses), *e[1:]) + for e in orig_sched + ] + bound_computed_schedules[src_unit] = new_unit_scheds + + if not schedule.broadcast_groups: + bound_direct_schedules = { + u: {s_idx: entries for s_idx, entries in scheds.items() if entries} + for u, scheds in bound_computed_schedules.items() + if any(scheds.values()) + } + bound_broadcast_groups = {} + else: + bound_direct_schedules = {} + for src_unit, scheds in schedule.direct_schedules.items(): + if ( + src_unit in bound_computed_schedules + and scheds == schedule.computed_schedules.get(src_unit) + ): + bound_direct_schedules[src_unit] = bound_computed_schedules[src_unit] + else: + bound_direct_schedules[src_unit] = { + s_idx: [ + ( + _resolve_symbolic_endpoint(e[0], live_data_addresses), + *e[1:], + ) + for e in entries + ] + for s_idx, entries in scheds.items() + } + bound_broadcast_groups = {} + for group_key, k_and_t_list in schedule.broadcast_groups.items(): + bound_list = [] + for k, targets in k_and_t_list: + bound_targets = [ + ( + t[0], + _resolve_symbolic_endpoint(t[1], live_data_addresses), + *t[2:], + ) + for t in targets + ] + bound_list.append((k, bound_targets)) + bound_broadcast_groups[group_key] = bound_list + + bound_data_addresses = dict(schedule.data_addresses) + bound_data_address_to_unit = {} + for u, shards in bound_data_addresses.items(): + if u in live_data_addresses and live_data_addresses[u]: + bound_shards = list(live_data_addresses[u]) + else: + bound_shards = [ + _resolve_symbolic_endpoint(s, live_data_addresses) for s in shards + ] + bound_data_addresses[u] = bound_shards + for s in bound_shards: + bound_data_address_to_unit[s] = u + for u, shards in live_data_addresses.items(): + if u not in bound_data_addresses and shards: + bound_data_addresses[u] = list(shards) + for s in shards: + bound_data_address_to_unit[s] = u + + bound_dst_endpoint_counts: dict[str, int] = {} + for sym_host, cnt in schedule.dst_endpoint_counts.items(): + parsed = _parse_symbolic_endpoint(f"{sym_host}:0") + if parsed is not None and parsed[0] in live_data_addresses: + live_shards = live_data_addresses[parsed[0]] + live_host = ( + _extract_host_ip(live_shards[0]) if live_shards else sym_host + ) + else: + live_host = sym_host + bound_dst_endpoint_counts[live_host] = ( + bound_dst_endpoint_counts.get(live_host, 0) + cnt + ) + + bound_dst_endpoint_layer_counts: dict[str, dict[int, int]] = {} + for sym_host, l_map in schedule.dst_endpoint_layer_counts.items(): + parsed = _parse_symbolic_endpoint(f"{sym_host}:0") + if parsed is not None and parsed[0] in live_data_addresses: + live_shards = live_data_addresses[parsed[0]] + live_host = ( + _extract_host_ip(live_shards[0]) if live_shards else sym_host + ) + else: + live_host = sym_host + target_map = bound_dst_endpoint_layer_counts.setdefault(live_host, {}) + for l_idx, cnt in l_map.items(): + target_map[l_idx] = target_map.get(l_idx, 0) + cnt + + bound_rpc_addresses = dict(schedule.rpc_addresses) + if live_rpc_addresses: + bound_rpc_addresses.update(live_rpc_addresses) + + return _CachedTransferSchedule( + computed_schedules=bound_computed_schedules, + direct_schedules=bound_direct_schedules, + broadcast_groups=bound_broadcast_groups, + local_skip_tiling=dict(schedule.local_skip_tiling), + expected_block_count=schedule.expected_block_count, + dst_unit_layer_counts={ + u: dict(m) for u, m in schedule.dst_unit_layer_counts.items() + }, + data_address_to_unit=bound_data_address_to_unit, + direct_dsts=list(schedule.direct_dsts), + rpc_addresses=bound_rpc_addresses, + data_addresses=bound_data_addresses, + dst_unit_counts=dict(schedule.dst_unit_counts), + dst_endpoint_counts=bound_dst_endpoint_counts, + dst_endpoint_layer_counts=bound_dst_endpoint_layer_counts, + is_weight_sync=schedule.is_weight_sync, + variable_plans=bound_variable_plans, + variable_to_plan_id={ + u: dict(m) for u, m in schedule.variable_to_plan_id.items() + }, + ) + + @classmethod + def _to_json_serializable(cls, obj: Any) -> Any: + """Recursively converts a schedule structure to JSON-compatible primitives.""" + if isinstance(obj, RaidenId): + return { + "__raiden_id__": [ + obj.job_name, + str(obj.job_replica_id), + obj.data_name, + int(obj.data_replica_idx), + ] + } + if isinstance(obj, _PlanReferencedShardSchedule): + ordered_vars = getattr(obj, "_ordered_vars", []) + return { + "__plan_ref__": [ + [int(l_idx), int(pid)] for l_idx, pid in ordered_vars + ] + } + if isinstance(obj, _CachedTransferSchedule): + skip_fields = { + "sender_push_schedule_protos", + "cached_serialized_payloads", + } + return { + "__cached_schedule__": { + f.name: cls._to_json_serializable(getattr(obj, f.name)) + for f in dataclasses.fields(obj) + if f.name not in skip_fields + } + } + if isinstance(obj, tuple): + return {"__tuple__": [cls._to_json_serializable(x) for x in obj]} + if isinstance(obj, list): + return [cls._to_json_serializable(x) for x in obj] + if isinstance(obj, dict): + if all(isinstance(k, str) for k in obj.keys()): + return {k: cls._to_json_serializable(v) for k, v in obj.items()} + return { + "__dict__": [ + [ + cls._to_json_serializable(k), + cls._to_json_serializable(v), + ] + for k, v in obj.items() + ] + } + return obj + + @classmethod + def _from_json_serializable(cls, obj: Any) -> Any: + """Recursively reconstructs a schedule structure from JSON primitives.""" + if isinstance(obj, list): + return [cls._from_json_serializable(x) for x in obj] + if isinstance(obj, dict): + if "__raiden_id__" in obj: + parts = obj["__raiden_id__"] + return RaidenId( + str(parts[0]), str(parts[1]), str(parts[2]), int(parts[3]) + ) + if "__tuple__" in obj: + return tuple(cls._from_json_serializable(x) for x in obj["__tuple__"]) + if "__dict__" in obj: + return { + cls._from_json_serializable(k): cls._from_json_serializable(v) + for k, v in obj["__dict__"] + } + if "__plan_ref__" in obj: + return {"__plan_ref__": [tuple(pair) for pair in obj["__plan_ref__"]]} + if "__cached_schedule__" in obj: + fields_dict = { + k: cls._from_json_serializable(v) + for k, v in obj["__cached_schedule__"].items() + } + var_plans = fields_dict.get("variable_plans", {}) + var_to_pid = fields_dict.get("variable_to_plan_id", {}) + for src_unit, unit_plans_by_id in var_plans.items(): + shard_plans_by_id: dict[int, dict[int, list[Any]]] = {} + for pid, tmpl_for_var in unit_plans_by_id.items(): + for local_src_idx, entries in tmpl_for_var.items(): + shard_plans_by_id.setdefault(local_src_idx, {})[pid] = entries + unit_var_to_pid = var_to_pid.get(src_unit, {}) + for sched_key in ("computed_schedules", "direct_schedules"): + unit_scheds = fields_dict.get(sched_key, {}).get(src_unit) + if not unit_scheds: + continue + for local_src_idx, val in list(unit_scheds.items()): + if isinstance(val, dict) and "__plan_ref__" in val: + unit_scheds[local_src_idx] = _PlanReferencedShardSchedule( + shard_plans_by_id.get(local_src_idx, {}), + unit_var_to_pid, + val["__plan_ref__"], + ) + return _CachedTransferSchedule(**fields_dict) + return {k: cls._from_json_serializable(v) for k, v in obj.items()} + return obj + + @classmethod + def save_offline_plan( + cls, + schedule: _CachedTransferSchedule, + path: str, + src_units: Optional[Sequence[RaidenId]] = None, + dst_units: Optional[Sequence[RaidenId]] = None, + dst_mem_type: int = controller_types.RaidenMemoryType.DRAM, + parallelism: int = 1, + ) -> dict[Any, str]: + """Saves per-worker ControlRequest protobufs and schedule bundle to `path`.""" + os.makedirs(path, exist_ok=True) + if src_units is not None: + src_list = list(src_units) + else: + src_keys = ( + schedule.computed_schedules.keys() or schedule.direct_schedules.keys() + ) + src_list = list(src_keys) + src_set = set(src_list) + if dst_units is not None: + dst_list = list(dst_units) + else: + inferred_dsts = [u for u in schedule.data_addresses if u not in src_set] + dst_list = inferred_dsts or list(schedule.dst_unit_counts.keys()) + + raw_schedules = schedule.direct_schedules or schedule.computed_schedules + transfer_plan = controller_types.TransferPlan( + src_units=src_list, + dst_units=dst_list, + plan=None, + shard_push_schedules=raw_schedules, + worker_rpc_addresses=dict(schedule.rpc_addresses), + worker_data_addresses=dict(schedule.data_addresses), + uuid=0, + dst_mem_type=dst_mem_type, + use_block_chunks=True, + is_sender=True, + expected_block_count=schedule.expected_block_count, + dst_expected_layer_chunk_counts=schedule.dst_unit_layer_counts, + dst_expected_block_counts=schedule.dst_unit_counts, + dst_endpoint_counts=schedule.dst_endpoint_counts, + dst_endpoint_layer_counts=schedule.dst_endpoint_layer_counts, + src_schedule_keys={u: i for i, u in enumerate(src_list)}, + req_id="offline", + skip_d2h=False, + skip_tiling=schedule.local_skip_tiling, + parallelism=parallelism, + is_weight_sync=schedule.is_weight_sync, + variable_plans=schedule.variable_plans, + variable_to_plan_id=schedule.variable_to_plan_id, + ) + + written_files: dict[Any, str] = {} + for u in src_list: + ent = JobEntity( + unit=u, + shards=schedule.data_addresses.get(u) or _make_symbolic_shards(u, 1), + weight_sync_mode=schedule.is_weight_sync, + ) + payload = ent.encode_start_transfer(transfer_plan, unit=u) + if payload: + file_path = os.path.join(path, f"{_unit_filename_stem(u)}.pb") + with open(file_path, "wb") as f: + f.write(payload) + written_files[u] = file_path + written_files[controller_types.format_unit(u)] = file_path + + for u in dst_list: + ent = JobEntity( + unit=u, + shards=schedule.data_addresses.get(u) or _make_symbolic_shards(u, 1), + weight_sync_mode=schedule.is_weight_sync, + ) + payload = ent.encode_start_transfer(transfer_plan, unit=u) + if payload: + file_path = os.path.join(path, f"{_unit_filename_stem(u)}.pb") + with open(file_path, "wb") as f: + f.write(payload) + written_files[u] = file_path + written_files[controller_types.format_unit(u)] = file_path + + bundle_path = os.path.join(path, "offline_schedule.json") + payload_dict = cls._to_json_serializable({ + "schedule": schedule, + "src_units": src_list, + "dst_units": dst_list, + }) + with open(bundle_path, "w", encoding="utf-8") as f: + json.dump(payload_dict, f) + written_files["__schedule_bundle__"] = bundle_path + return written_files + + @classmethod + def load_offline_schedule( + cls, + path: str, + live_data_addresses: Optional[Mapping[RaidenId, Sequence[str]]] = None, + live_rpc_addresses: Optional[Mapping[RaidenId, str]] = None, + registered_shards: Optional[Mapping[RaidenId, Sequence[str]]] = None, + worker_endpoints: Optional[Mapping[RaidenId, str]] = None, + entities: Optional[Mapping[RaidenId, Any]] = None, + ) -> _CachedTransferSchedule: + """Loads an offline _CachedTransferSchedule from `path` and binds live endpoints.""" + del entities + bundle_path = ( + os.path.join(path, "offline_schedule.json") + if os.path.isdir(path) + else path + ) + with open(bundle_path, "r", encoding="utf-8") as f: + raw_bundle = json.load(f) + bundle = cls._from_json_serializable(raw_bundle) + schedule = bundle["schedule"] if isinstance(bundle, dict) else bundle + effective_data = ( + live_data_addresses + if live_data_addresses is not None + else registered_shards + ) + effective_rpc = ( + live_rpc_addresses + if live_rpc_addresses is not None + else worker_endpoints + ) + if effective_data: + return cls.bind_symbolic_schedule( + schedule, + live_data_addresses=effective_data, + live_rpc_addresses=effective_rpc, + ) + return schedule + + @classmethod + def load_offline_worker_plan( + cls, + path: str, + unit: RaidenId, + live_data_addresses: Optional[Mapping[RaidenId, Sequence[str]]] = None, + registered_shards: Optional[Mapping[RaidenId, Sequence[str]]] = None, + uuid: Optional[int] = None, + req_id: Optional[str] = None, + skip_d2h: Optional[bool] = None, + ) -> raiden_service_pb2.ControlRequest: + """Loads a single worker's precomputed ControlRequest proto and binds live endpoints.""" + del cls + file_path = ( + os.path.join(path, f"{_unit_filename_stem(unit)}.pb") + if os.path.isdir(path) + else path + ) + with open(file_path, "rb") as f: + raw_bytes = f.read() + req = raiden_service_pb2.ControlRequest() + req.ParseFromString(raw_bytes) + effective_data = ( + live_data_addresses + if live_data_addresses is not None + else registered_shards + ) + if effective_data: + _bind_symbolic_endpoints_in_proto(req, effective_data) + if uuid is not None: + req.start_transfer_request.uuid = int(uuid) + if req_id is not None: + req.start_transfer_request.req_id = str(req_id) + if skip_d2h is not None: + req.start_transfer_request.skip_d2h = bool(skip_d2h) + return req + + @classmethod + def load_offline_worker_plans_parallel( + cls, + path: str, + units: Sequence[RaidenId], + live_data_addresses: Optional[Mapping[RaidenId, Sequence[str]]] = None, + registered_shards: Optional[Mapping[RaidenId, Sequence[str]]] = None, + uuid: Optional[int] = None, + req_id: Optional[str] = None, + skip_d2h: Optional[bool] = None, + max_workers: Optional[int] = None, + ) -> dict[RaidenId, raiden_service_pb2.ControlRequest]: + """Loads and binds precomputed ControlRequest plans for `units` in parallel.""" + unit_list = list(units) + if not unit_list: + return {} + effective_data = ( + live_data_addresses + if live_data_addresses is not None + else registered_shards + ) + + def _load_one( + u: RaidenId, + ) -> tuple[RaidenId, raiden_service_pb2.ControlRequest]: + return ( + u, + cls.load_offline_worker_plan( + path=path, + unit=u, + live_data_addresses=effective_data, + uuid=uuid, + req_id=req_id, + skip_d2h=skip_d2h, + ), + ) + + num_threads = max_workers or min(32, max(1, len(unit_list))) + with concurrent.futures.ThreadPoolExecutor(max_workers=num_threads) as pool: + return dict(pool.map(_load_one, unit_list)) diff --git a/tpu_sync/weight_sync/manager/reshard_planner_test.py b/tpu_sync/weight_sync/manager/reshard_planner_test.py index 1c624ee94..69ea804a3 100644 --- a/tpu_sync/weight_sync/manager/reshard_planner_test.py +++ b/tpu_sync/weight_sync/manager/reshard_planner_test.py @@ -14,9 +14,11 @@ """Tests and benchmarks for ReshardPlanner schedule computation and deduplication.""" +import os import threading import timeit from typing import Any +from unittest import mock from absl import logging from absl.testing import absltest @@ -518,6 +520,371 @@ def test_deterministic_circular_shift_latin_square_scheduling(self): self.assertGreater(multi_dest_vars_tested, 0) + def test_parallel_worker_planning_exact_equivalence(self): + """Verifies Approach 2 (parallel per-worker planning) matches centralized planning bit-for-bit.""" + src_units = [RaidenId("trainer", str(i), "weights", 0) for i in range(4)] + dst_units = [ + RaidenId("rollout_0", "0", "weights", 0), + RaidenId("rollout_0", "1", "weights", 0), + RaidenId("rollout_1", "0", "weights", 0), + RaidenId("rollout_1", "1", "weights", 0), + ] + src_vars = _build_qwen3_397b_variables( + num_layers=4, src_fsdp=4, is_src=True + ) + dst_vars = _build_qwen3_397b_variables( + num_layers=4, src_fsdp=4, is_src=False + ) + + central_kwargs = self._build_planner_inputs( + src_vars_by_unit={u: src_vars for u in src_units}, + dst_vars_by_unit={u: dst_vars for u in dst_units}, + src_phys_mesh=[1, 1, 4, 4, 2], + src_mesh_axes=["data", "stage", "fsdp", "context", "expert"], + src_host_subgrid=[1, 1, 1, 4, 2], + dst_phys_mesh=[2, 8], + dst_mesh_axes=["x", "y"], + dst_host_subgrid=[1, 8], + ) + central_sched = ( + reshard_planner.ReshardPlanner.compute_transfer_schedule_from_metadata( + **central_kwargs, + parallel_worker_planning=False, + ) + ) + + # 1. Test independent single-worker slice computation + # (`target_src_units=[u]`) + worker_slices = [] + for u in src_units: + worker_kwargs = self._build_planner_inputs( + src_vars_by_unit={su: src_vars for su in src_units}, + dst_vars_by_unit={du: dst_vars for du in dst_units}, + src_phys_mesh=[1, 1, 4, 4, 2], + src_mesh_axes=["data", "stage", "fsdp", "context", "expert"], + src_host_subgrid=[1, 1, 1, 4, 2], + dst_phys_mesh=[2, 8], + dst_mesh_axes=["x", "y"], + dst_host_subgrid=[1, 8], + ) + w_sched = reshard_planner.ReshardPlanner.compute_transfer_schedule_from_metadata( + **worker_kwargs, + target_src_units=[u], + parallel_worker_planning=False, + ) + self.assertEqual(list(w_sched.computed_schedules.keys()), [u]) + self.assertEqual( + w_sched.variable_to_plan_id[u], + central_sched.variable_to_plan_id[u], + ) + self.assertEqual( + w_sched.variable_plans[u], + central_sched.variable_plans[u], + ) + ent = central_kwargs["entities"][u] + w_protos = ent.build_sender_push_schedule_protos( + w_sched.direct_schedules[u] + ) + c_protos = ent.build_sender_push_schedule_protos( + central_sched.direct_schedules[u] + ) + self.assertEqual( + { + k: p.SerializeToString(deterministic=True) + for k, p in w_protos.items() + }, + { + k: p.SerializeToString(deterministic=True) + for k, p in c_protos.items() + }, + ) + worker_slices.append(w_sched) + + # 2. Test merging independent worker slices + parallel worker execution + parallel_kwargs = self._build_planner_inputs( + src_vars_by_unit={u: src_vars for u in src_units}, + dst_vars_by_unit={u: dst_vars for u in dst_units}, + src_phys_mesh=[1, 1, 4, 4, 2], + src_mesh_axes=["data", "stage", "fsdp", "context", "expert"], + src_host_subgrid=[1, 1, 1, 4, 2], + dst_phys_mesh=[2, 8], + dst_mesh_axes=["x", "y"], + dst_host_subgrid=[1, 8], + ) + parallel_sched = ( + reshard_planner.ReshardPlanner.compute_transfer_schedule_from_metadata( + **parallel_kwargs, + parallel_worker_planning=True, + ) + ) + + self.assertEqual( + parallel_sched.expected_block_count, + central_sched.expected_block_count, + ) + self.assertEqual( + parallel_sched.dst_unit_counts, central_sched.dst_unit_counts + ) + self.assertEqual( + parallel_sched.dst_unit_layer_counts, + central_sched.dst_unit_layer_counts, + ) + self.assertEqual( + parallel_sched.dst_endpoint_counts, + central_sched.dst_endpoint_counts, + ) + self.assertEqual( + parallel_sched.dst_endpoint_layer_counts, + central_sched.dst_endpoint_layer_counts, + ) + for u in src_units: + ent = central_kwargs["entities"][u] + p_protos = ent.build_sender_push_schedule_protos( + parallel_sched.direct_schedules[u] + ) + c_protos = ent.build_sender_push_schedule_protos( + central_sched.direct_schedules[u] + ) + self.assertEqual( + { + k: p.SerializeToString(deterministic=True) + for k, p in p_protos.items() + }, + { + k: p.SerializeToString(deterministic=True) + for k, p in c_protos.items() + }, + ) + for shard_idx in range(8): + self.assertEqual( + list(parallel_sched.computed_schedules[u][shard_idx]), + list(central_sched.computed_schedules[u][shard_idx]), + ) + + def test_parallel_worker_planning_via_env_vars(self): + """Verifies RAIDEN_PARALLEL_WORKER_PLANNING and RAIDEN_PLANNING_MODE env vars trigger parallel worker planning.""" + src_units = [RaidenId("trainer", str(i), "weights", 0) for i in range(2)] + dst_units = [RaidenId("rollout_0", "0", "weights", 0)] + src_vars = _build_qwen3_397b_variables( + num_layers=1, src_fsdp=2, is_src=True + ) + dst_vars = _build_qwen3_397b_variables( + num_layers=1, src_fsdp=2, is_src=False + ) + + for env_dict in ( + {"RAIDEN_PARALLEL_WORKER_PLANNING": "1"}, + {"RAIDEN_PLANNING_MODE": "parallel"}, + {"RAIDEN_PLANNING_MODE": "parallel_worker"}, + {"RAIDEN_PLANNING_MODE": "decentralized"}, + ): + with mock.patch.dict(os.environ, env_dict, clear=False): + with mock.patch.object( + reshard_planner.ReshardPlanner, + "compute_schedules_in_parallel_workers", + wraps=reshard_planner.ReshardPlanner.compute_schedules_in_parallel_workers, + ) as spy_parallel: + kwargs = self._build_planner_inputs( + src_vars_by_unit={u: src_vars for u in src_units}, + dst_vars_by_unit={u: dst_vars for u in dst_units}, + src_phys_mesh=[1, 1, 2, 4, 2], + src_mesh_axes=["data", "stage", "fsdp", "context", "expert"], + src_host_subgrid=[1, 1, 1, 4, 2], + dst_phys_mesh=[1, 8], + dst_mesh_axes=["x", "y"], + dst_host_subgrid=[1, 8], + ) + sched = reshard_planner.ReshardPlanner.compute_transfer_schedule_from_metadata( + **kwargs + ) + self.assertEqual(spy_parallel.call_count, 1) + self.assertLen(sched.computed_schedules, 2) + + # Verify merge_worker_schedules deduplicates multi-unit schedules + merged_batched = ( + reshard_planner.ReshardPlanner.merge_worker_schedules( + [sched], + src_units=src_units, + dst_units=dst_units, + ) + ) + self.assertEqual( + merged_batched.expected_block_count, sched.expected_block_count + ) + self.assertEqual( + merged_batched.dst_unit_counts, sched.dst_unit_counts + ) + + def test_offline_symbolic_plan_save_and_parallel_worker_load(self): + """Verifies Approach 1: offline symbolic schedule computation, serialization, parallel worker load, and late endpoint binding.""" + src_units = [RaidenId("trainer", str(i), "weights", 0) for i in range(4)] + dst_units = [ + RaidenId("rollout_0", "0", "weights", 0), + RaidenId("rollout_0", "1", "weights", 0), + RaidenId("rollout_1", "0", "weights", 0), + RaidenId("rollout_1", "1", "weights", 0), + ] + src_vars = _build_qwen3_397b_variables( + num_layers=4, src_fsdp=4, is_src=True + ) + dst_vars = _build_qwen3_397b_variables( + num_layers=4, src_fsdp=4, is_src=False + ) + + online_kwargs = self._build_planner_inputs( + src_vars_by_unit={u: src_vars for u in src_units}, + dst_vars_by_unit={u: dst_vars for u in dst_units}, + src_phys_mesh=[1, 1, 4, 4, 2], + src_mesh_axes=["data", "stage", "fsdp", "context", "expert"], + src_host_subgrid=[1, 1, 1, 4, 2], + dst_phys_mesh=[2, 8], + dst_mesh_axes=["x", "y"], + dst_host_subgrid=[1, 8], + ) + online_sched = ( + reshard_planner.ReshardPlanner.compute_transfer_schedule_from_metadata( + **online_kwargs + ) + ) + + # Compute offline symbolic schedule without any real IPs/ports + offline_kwargs = self._build_planner_inputs( + src_vars_by_unit={u: src_vars for u in src_units}, + dst_vars_by_unit={u: dst_vars for u in dst_units}, + src_phys_mesh=[1, 1, 4, 4, 2], + src_mesh_axes=["data", "stage", "fsdp", "context", "expert"], + src_host_subgrid=[1, 1, 1, 4, 2], + dst_phys_mesh=[2, 8], + dst_mesh_axes=["x", "y"], + dst_host_subgrid=[1, 8], + ) + live_registered_shards = dict(offline_kwargs["registered_shards"]) + for item in offline_kwargs["dst_metadata"]: + u = controller_types.raiden_id_from_proto(item.unit) + live_registered_shards[u] = list(item.shards) + live_worker_endpoints = dict(offline_kwargs["worker_endpoints"]) + live_entities = dict(offline_kwargs["entities"]) + + # Strip live endpoints from offline_kwargs to prove offline computation + # needs no live IPs. + offline_kwargs["registered_shards"] = { + u: [f"0.0.0.0:{d}" for d in range(8)] for u in [*src_units, *dst_units] + } + offline_kwargs["worker_endpoints"] = {} + offline_kwargs["entities"] = {} + + symbolic_sched = reshard_planner.ReshardPlanner.compute_offline_schedule( + **offline_kwargs + ) + # Verify symbolic endpoints are present before binding + first_entry = symbolic_sched.direct_schedules[src_units[0]][0][0] + self.assertTrue(controller_types.is_symbolic_endpoint(first_entry[0])) + + plan_dir = self.create_tempdir().full_path + saved_files = reshard_planner.ReshardPlanner.save_offline_plan( + symbolic_sched, plan_dir, src_units=src_units, dst_units=dst_units + ) + self.assertIn("__schedule_bundle__", saved_files) + for u in src_units: + self.assertIn(controller_types.format_unit(u), saved_files) + self.assertTrue( + os.path.exists(saved_files[controller_types.format_unit(u)]) + ) + + # 1. Parallel worker load + late endpoint binding of per-worker + # ControlRequest protobufs. + worker_reqs = ( + reshard_planner.ReshardPlanner.load_offline_worker_plans_parallel( + plan_dir, + units=src_units, + registered_shards=live_registered_shards, + req_id="step_0_sync", + uuid=424242, + ) + ) + self.assertLen(worker_reqs, len(src_units)) + for u in src_units: + req = worker_reqs[u] + self.assertEqual(req.start_transfer_request.req_id, "step_0_sync") + self.assertEqual(req.start_transfer_request.uuid, 424242) + ent = live_entities[u] + online_protos = ent.build_sender_push_schedule_protos( + online_sched.direct_schedules[u] + ) + self.assertEqual( + { + k: p.SerializeToString(deterministic=True) + for k, p in ( + req.start_transfer_request.shard_push_schedules.items() + ) + }, + { + k: p.SerializeToString(deterministic=True) + for k, p in online_protos.items() + }, + ) + + # 2. Full schedule load + late endpoint binding for controller + bound_sched = reshard_planner.ReshardPlanner.load_offline_schedule( + plan_dir, + registered_shards=live_registered_shards, + worker_endpoints=live_worker_endpoints, + entities=live_entities, + ) + self.assertEqual( + bound_sched.expected_block_count, online_sched.expected_block_count + ) + self.assertEqual(bound_sched.dst_unit_counts, online_sched.dst_unit_counts) + self.assertEqual( + bound_sched.dst_endpoint_counts, online_sched.dst_endpoint_counts + ) + self.assertEqual( + bound_sched.dst_endpoint_layer_counts, + online_sched.dst_endpoint_layer_counts, + ) + for u in src_units: + ent = live_entities[u] + bound_protos = ent.build_sender_push_schedule_protos( + bound_sched.direct_schedules[u] + ) + online_protos = ent.build_sender_push_schedule_protos( + online_sched.direct_schedules[u] + ) + self.assertEqual( + { + k: p.SerializeToString(deterministic=True) + for k, p in bound_protos.items() + }, + { + k: p.SerializeToString(deterministic=True) + for k, p in online_protos.items() + }, + ) + + # 3. Re-bind the exact same offline plan to a brand new set of runtime + # IPs/ports. + new_runtime_shards = { + u: [f"192.168.50.{idx + 1}:{15000 + d}" for d in range(8)] + for idx, u in enumerate([*src_units, *dst_units]) + } + rebound_reqs = ( + reshard_planner.ReshardPlanner.load_offline_worker_plans_parallel( + plan_dir, + units=src_units, + registered_shards=new_runtime_shards, + ) + ) + for u in src_units: + for ( + _, + sched_proto, + ) in rebound_reqs[u].start_transfer_request.shard_push_schedules.items(): + for entry in sched_proto.entries: + self.assertFalse( + controller_types.is_symbolic_endpoint(entry.dst_peer) + ) + self.assertTrue(entry.dst_peer.startswith("192.168.50.")) + if __name__ == "__main__": absltest.main()