From e807b6f7a5554fcf9b14dbedf6131fcb3bc56fa3 Mon Sep 17 00:00:00 2001 From: fhzhang Date: Tue, 22 Sep 2026 10:57:58 -0700 Subject: [PATCH] Add TestOnlyRateLimiter in tpu_raiden transport to emulate physical NIC egress/ingress bandwidth constraints in CPU loopback integration tests. PiperOrigin-RevId: 986103579 --- tpu_sync/api/jax/weight_synchronizer.py | 61 +- tpu_sync/api/jax/weight_synchronizer_test.py | 49 ++ tpu_sync/core/BUILD | 1 + tpu_sync/core/raiden_manager_base.cc | 18 + tpu_sync/core/raiden_manager_base.h | 9 + tpu_sync/frameworks/jax/BUILD | 2 + .../frameworks/jax/tpu_raiden_jax_module.cc | 62 +- .../frameworks/jax/weight_synchronizer.cc | 134 ++++ tpu_sync/frameworks/jax/weight_synchronizer.h | 42 ++ tpu_sync/transport/block_transport.h | 7 + tpu_sync/transport/lib/BUILD | 24 + .../transport/lib/raw_buffer_transport.cc | 36 + tpu_sync/transport/lib/raw_buffer_transport.h | 14 + .../transport/lib/test_only_rate_limiter.cc | 48 ++ .../transport/lib/test_only_rate_limiter.h | 50 ++ .../lib/test_only_rate_limiter_test.cc | 67 ++ tpu_sync/weight_sync/BUILD | 18 + .../weight_sync_fanout_perf_test.py | 633 ++++++++++++++++++ .../weight_sync/weight_synchronizer_base.cc | 8 + .../weight_sync/weight_synchronizer_base.h | 5 + 20 files changed, 1283 insertions(+), 5 deletions(-) create mode 100644 tpu_sync/transport/lib/test_only_rate_limiter.cc create mode 100644 tpu_sync/transport/lib/test_only_rate_limiter.h create mode 100644 tpu_sync/transport/lib/test_only_rate_limiter_test.cc create mode 100644 tpu_sync/weight_sync/weight_sync_fanout_perf_test.py diff --git a/tpu_sync/api/jax/weight_synchronizer.py b/tpu_sync/api/jax/weight_synchronizer.py index b4e99b4d8..31c162ce9 100644 --- a/tpu_sync/api/jax/weight_synchronizer.py +++ b/tpu_sync/api/jax/weight_synchronizer.py @@ -14,7 +14,7 @@ """High-performance JAX Weight Synchronizer for RL Trainer-Inference Pipelines.""" -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, Union import jax @@ -95,6 +95,55 @@ def __init__( global_shard_indices, ) + @classmethod + def test_only_create_cpu_instance( + cls, + num_layers: int, + num_shards: int, + slice_byte_size: Union[int, List[int]], + local_port: Optional[int] = None, + parallelism: int = 1, + listener_port: Optional[int] = None, + bind_ip: Optional[str] = "127.0.0.1", + auto_h2d: bool = False, + global_shard_indices: Optional[List[int]] = None, + test_only_simulated_egress_gbps: float = 0.0, + test_only_simulated_ingress_gbps: float = 0.0, + ) -> "WeightSynchronizer": + """Instantiates a CPU-only WeightSynchronizer allocating host DRAM without TPU devices.""" + instance = cls.__new__(cls) + instance._impl = ( + _weight_synchronizer.WeightSynchronizer.test_only_create_cpu_instance( + num_layers, + num_shards, + slice_byte_size, + local_port, + parallelism, + listener_port, + bind_ip, + auto_h2d, + global_shard_indices, + test_only_simulated_egress_gbps, + test_only_simulated_ingress_gbps, + ) + ) + instance._global_shard_indices = list(global_shard_indices or []) + instance._has_explicit_global_shard_indices = ( + global_shard_indices is not None + ) + return instance + + def test_only_set_bandwidth_limit( + self, + test_only_simulated_egress_gbps: float = 0.0, + test_only_simulated_ingress_gbps: float = 0.0, + ) -> None: + """Sets the simulated transport bandwidth limit in Gbps (for testing only).""" + self._impl.test_only_set_bandwidth_limit( + test_only_simulated_egress_gbps, + test_only_simulated_ingress_gbps, + ) + def d2h(self) -> None: """Triggers asynchronous Device-to-Host (D2H) copy of current weights to Host buffer.""" self._impl.D2h() @@ -155,16 +204,22 @@ def get_local_endpoints(self) -> List[Dict[str, Any]]: @property def local_port(self) -> Optional[int]: """Returns the active local port assigned to the transceiving sockets server.""" + if self._impl is None: + return None return self._impl.local_port @property def listener_port(self) -> Optional[int]: """Returns the active local port assigned to the C++ Listener.""" + if self._impl is None: + return None return self._impl.listener_port @property def is_listener_active(self) -> bool: """Returns whether the native C++ Listener is actively running.""" + if self._impl is None: + return False return self._impl.is_listener_active @property @@ -239,3 +294,7 @@ def get_metrics(self) -> dict[str, float | int]: def reset_metrics(self) -> None: """Resets all recorded internal metrics.""" self._impl.reset_metrics() + + def shutdown(self) -> None: + """Releases and shuts down the underlying C++ synchronizer instance.""" + self._impl = None diff --git a/tpu_sync/api/jax/weight_synchronizer_test.py b/tpu_sync/api/jax/weight_synchronizer_test.py index d9dff8735..8f09f0a9e 100644 --- a/tpu_sync/api/jax/weight_synchronizer_test.py +++ b/tpu_sync/api/jax/weight_synchronizer_test.py @@ -117,6 +117,55 @@ def test_push_synchronization(self): for arr in dst2_arrs: np.testing.assert_array_equal(np.asarray(arr), 5.0) + def test_create_cpu_instance(self): + ws = WeightSynchronizer.test_only_create_cpu_instance( + num_layers=2, + num_shards=1, + slice_byte_size=1024, + local_port=0, + listener_port=0, + bind_ip="127.0.0.1", + ) + self.assertIsNotNone(ws.local_port) + self.assertIsNotNone(ws.listener_port) + self.assertEqual(ws.num_layers, 2) + self.assertEqual(ws.num_shards, 1) + self.assertEqual(ws.slice_byte_size, 1024) + + buf = ws.get_host_buffer(layer_idx=0, shard_idx=0) + self.assertGreaterEqual(len(buf), 1024) + buf[:10] = 42 + self.assertEqual(buf[0], 42) + + ws.shutdown() + self.assertIsNone(ws.local_port) + + def test_create_cpu_instance_heterogeneous(self): + ws = WeightSynchronizer.test_only_create_cpu_instance( + num_layers=2, + num_shards=1, + slice_byte_size=[512, 2048], + local_port=0, + listener_port=0, + bind_ip="127.0.0.1", + ) + self.assertIsNotNone(ws.local_port) + self.assertIsNotNone(ws.listener_port) + self.assertEqual(ws.num_layers, 2) + self.assertEqual(ws.num_shards, 1) + + buf0 = ws.get_host_buffer(layer_idx=0, shard_idx=0) + buf1 = ws.get_host_buffer(layer_idx=1, shard_idx=0) + self.assertGreaterEqual(len(buf0), 512) + self.assertGreaterEqual(len(buf1), 2048) + buf0[:4] = 11 + buf1[:4] = 22 + self.assertEqual(buf0[0], 11) + self.assertEqual(buf1[0], 22) + + ws.shutdown() + self.assertIsNone(ws.local_port) + def test_wait_for_transfer_completion_api_exists(self): arrs = [ jax.device_put(jnp.zeros(self.shape, dtype=self.dtype), self.sharding) diff --git a/tpu_sync/core/BUILD b/tpu_sync/core/BUILD index f8fdbc00e..1cce20628 100644 --- a/tpu_sync/core/BUILD +++ b/tpu_sync/core/BUILD @@ -130,6 +130,7 @@ cc_library( "//tpu_sync/transport:block_transport", "//tpu_sync/transport:block_transport_delegate", "//tpu_sync/transport:buffer_push_task", + "//tpu_sync/transport/lib:test_only_rate_limiter", "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/log", diff --git a/tpu_sync/core/raiden_manager_base.cc b/tpu_sync/core/raiden_manager_base.cc index 34fe50ff3..000a5a42c 100644 --- a/tpu_sync/core/raiden_manager_base.cc +++ b/tpu_sync/core/raiden_manager_base.cc @@ -39,6 +39,7 @@ #include "tpu_sync/core/tpu_utils.h" #include "tpu_sync/transport/block_transport.h" #include "tpu_sync/transport/buffer_push_task.h" +#include "tpu_sync/transport/lib/test_only_rate_limiter.h" namespace tpu_raiden { @@ -104,6 +105,18 @@ void RaidenManagerBase::StopTransportServer() { } } +void RaidenManagerBase::SetTestOnlyRateLimiters( + std::shared_ptr egress, + std::shared_ptr ingress) { + absl::MutexLock lock(server_init_mu_); + test_only_egress_rate_limiter_ = std::move(egress); + test_only_ingress_rate_limiter_ = std::move(ingress); + if (server_ != nullptr) { + server_->SetTestOnlyRateLimiters(test_only_egress_rate_limiter_, + test_only_ingress_rate_limiter_); + } +} + std::vector RaidenManagerBase::GetHostNics() const { return GetLocalHostNicAddresses(); } @@ -176,6 +189,11 @@ RaidenManagerBase::InitTransportServer() { server_ = std::make_unique( this, local_port_cfg_, local_ips_, parallelism_); + if (test_only_egress_rate_limiter_ != nullptr || + test_only_ingress_rate_limiter_ != nullptr) { + server_->SetTestOnlyRateLimiters(test_only_egress_rate_limiter_, + test_only_ingress_rate_limiter_); + } return server_.get(); } diff --git a/tpu_sync/core/raiden_manager_base.h b/tpu_sync/core/raiden_manager_base.h index 8e83db6c3..4fa09d8b9 100644 --- a/tpu_sync/core/raiden_manager_base.h +++ b/tpu_sync/core/raiden_manager_base.h @@ -33,6 +33,7 @@ #include "tpu_sync/transport/block_transport.h" #include "tpu_sync/transport/block_transport_delegate.h" #include "tpu_sync/transport/buffer_push_task.h" +#include "tpu_sync/transport/lib/test_only_rate_limiter.h" namespace tpu_raiden { @@ -96,6 +97,10 @@ class RaidenManagerBase : public tpu_raiden::transport::BlockTransportDelegate { virtual void ForgetPushProgress(uint64_t uuid); + void SetTestOnlyRateLimiters( + std::shared_ptr egress, + std::shared_ptr ingress); + // Stops and joins the underlying raw transport server if active. void StopTransportServer(); @@ -163,6 +168,10 @@ class RaidenManagerBase : public tpu_raiden::transport::BlockTransportDelegate { mutable absl::Mutex server_init_mu_; std::unique_ptr server_ ABSL_GUARDED_BY(server_init_mu_); + std::shared_ptr + test_only_egress_rate_limiter_ ABSL_GUARDED_BY(server_init_mu_); + std::shared_ptr + test_only_ingress_rate_limiter_ ABSL_GUARDED_BY(server_init_mu_); std::vector layers_; diff --git a/tpu_sync/frameworks/jax/BUILD b/tpu_sync/frameworks/jax/BUILD index 47cc1ee6b..eb03f09c7 100644 --- a/tpu_sync/frameworks/jax/BUILD +++ b/tpu_sync/frameworks/jax/BUILD @@ -316,6 +316,7 @@ cc_library( "//tpu_sync/core:raw_transfer_core", "//tpu_sync/core:tpu_utils", "//tpu_sync/rpc:raiden_service_cc_proto", + "//tpu_sync/transport/lib:test_only_rate_limiter", "//tpu_sync/weight_sync:weight_synchronizer_base", "@com_google_absl//absl/container:btree", "@com_google_absl//absl/container:flat_hash_map", @@ -484,6 +485,7 @@ cc_library( "//tpu_sync/core:raw_transfer_core", "//tpu_sync/core:tpu_utils", "//tpu_sync/rpc:raiden_service_cc_proto", + "//tpu_sync/transport/lib:test_only_rate_limiter", "//tpu_sync/weight_sync:weight_synchronizer_base", "@com_google_absl//absl/container:btree", "@com_google_absl//absl/container:flat_hash_map", diff --git a/tpu_sync/frameworks/jax/tpu_raiden_jax_module.cc b/tpu_sync/frameworks/jax/tpu_raiden_jax_module.cc index 04b56fc2c..5b3204d7c 100644 --- a/tpu_sync/frameworks/jax/tpu_raiden_jax_module.cc +++ b/tpu_sync/frameworks/jax/tpu_raiden_jax_module.cc @@ -33,6 +33,7 @@ #include // IWYU pragma: keep #include // IWYU pragma: keep #include // IWYU pragma: keep +#include // IWYU pragma: keep #include // IWYU pragma: keep #include "xla/pjrt/status_casters.h" #include "tpu_sync/core/raiden_future.h" @@ -339,6 +340,58 @@ NB_MODULE(_tpu_raiden_jax, m) { nb::arg("bind_ip") = nb::none(), nb::arg("auto_h2d") = false, nb::arg("global_shard_indices") = nb::none()) + .def("test_only_set_bandwidth_limit", + &WeightSynchronizer::test_only_set_bandwidth_limit, + nb::arg("test_only_simulated_egress_gbps") = 0.0, + nb::arg("test_only_simulated_ingress_gbps") = 0.0) + + .def_static( + "test_only_create_cpu_instance", + [](size_t num_layers, size_t num_shards, size_t slice_byte_size, + std::optional local_port, int parallelism, + std::optional listener_port, + std::optional bind_ip, bool auto_h2d, + std::optional> global_shard_indices, + double test_only_simulated_egress_gbps, + double test_only_simulated_ingress_gbps) { + return WeightSynchronizer::test_only_create_cpu_instance( + num_layers, num_shards, slice_byte_size, local_port, + parallelism, listener_port, bind_ip, auto_h2d, + global_shard_indices, test_only_simulated_egress_gbps, + test_only_simulated_ingress_gbps); + }, + nb::arg("num_layers"), nb::arg("num_shards"), + nb::arg("slice_byte_size"), nb::arg("local_port") = nb::none(), + nb::arg("parallelism") = 1, nb::arg("listener_port") = nb::none(), + nb::arg("bind_ip") = nb::none(), nb::arg("auto_h2d") = false, + nb::arg("global_shard_indices") = nb::none(), + nb::arg("test_only_simulated_egress_gbps") = 0.0, + nb::arg("test_only_simulated_ingress_gbps") = 0.0) + + .def_static( + "test_only_create_cpu_instance", + [](size_t num_layers, size_t num_shards, + std::vector slice_byte_sizes, + std::optional local_port, int parallelism, + std::optional listener_port, + std::optional bind_ip, bool auto_h2d, + std::optional> global_shard_indices, + double test_only_simulated_egress_gbps, + double test_only_simulated_ingress_gbps) { + return WeightSynchronizer::test_only_create_cpu_instance( + num_layers, num_shards, std::move(slice_byte_sizes), local_port, + parallelism, listener_port, bind_ip, auto_h2d, + global_shard_indices, test_only_simulated_egress_gbps, + test_only_simulated_ingress_gbps); + }, + nb::arg("num_layers"), nb::arg("num_shards"), + nb::arg("slice_byte_size"), nb::arg("local_port") = nb::none(), + nb::arg("parallelism") = 1, nb::arg("listener_port") = nb::none(), + nb::arg("bind_ip") = nb::none(), nb::arg("auto_h2d") = false, + nb::arg("global_shard_indices") = nb::none(), + nb::arg("test_only_simulated_egress_gbps") = 0.0, + nb::arg("test_only_simulated_ingress_gbps") = 0.0) + .def( "D2h", [](WeightSynchronizer& self) { @@ -413,12 +466,13 @@ NB_MODULE(_tpu_raiden_jax, m) { if (!ptr) { throw std::runtime_error("Invalid layer or shard index"); } - size_t size = self.slice_byte_size() + 256 * 1024; + size_t size = self.GetHostBufferSize(layer_idx, shard_idx); + if (size == 0) { + size = self.slice_byte_size() + 256 * 1024; + } size_t shape[1] = {size}; return nb::ndarray( - const_cast(ptr), 1, shape, - nb::handle() /* view only, no ownership copy */ - ); + const_cast(ptr), 1, shape, nb::find(&self)); }, nb::arg("layer_idx") = 0, nb::arg("shard_idx") = 0) .def_prop_ro("local_port", &WeightSynchronizer::local_port) diff --git a/tpu_sync/frameworks/jax/weight_synchronizer.cc b/tpu_sync/frameworks/jax/weight_synchronizer.cc index f41dfb0fe..e3eecc58b 100644 --- a/tpu_sync/frameworks/jax/weight_synchronizer.cc +++ b/tpu_sync/frameworks/jax/weight_synchronizer.cc @@ -42,6 +42,7 @@ #include "tpu_sync/core/raw_transfer_core.h" #include "tpu_sync/core/tpu_utils.h" #include "tpu_sync/rpc/raiden_service.pb.h" +#include "tpu_sync/transport/lib/test_only_rate_limiter.h" #include "tpu_sync/weight_sync/weight_synchronizer_base.h" #ifndef WITHOUT_PYTHON @@ -151,6 +152,39 @@ NumaAwareWeightSynchronizer::NumaAwareWeightSynchronizer( } } +NumaAwareWeightSynchronizer::NumaAwareWeightSynchronizer( + size_t num_layers, size_t num_shards, std::vector slice_byte_sizes, + std::optional local_port, int parallelism, + std::optional listener_port, std::optional bind_ip, + bool auto_h2d, std::optional> global_shard_indices) + : total_num_shards_(num_shards), + num_layers_(num_layers), + slice_byte_size_(slice_byte_sizes.empty() ? 0 : slice_byte_sizes[0]), + global_shard_indices_( + global_shard_indices.value_or(std::vector{})) { + auto sub = std::make_unique( + num_layers, num_shards, slice_byte_sizes, local_port, + /*host_blocks_to_allocate=*/std::nullopt, parallelism, listener_port, + bind_ip, /*layer_names=*/std::vector{}, auto_h2d); + sub_synchronizers_.push_back(std::move(sub)); + global_shard_to_submanager_.resize(total_num_shards_); + submanager_to_global_shards_.resize(1); + submanager_to_local_shards_.resize(1); + for (size_t i = 0; i < total_num_shards_; ++i) { + global_shard_to_submanager_[i] = {0, static_cast(i)}; + int64_t gidx = (i < global_shard_indices_.size()) ? global_shard_indices_[i] + : static_cast(i); + submanager_to_global_shards_[0].push_back(gidx); + submanager_to_local_shards_[0].push_back(static_cast(i)); + } + if (!sub_synchronizers_.empty() && sub_synchronizers_[0]) { + sub_synchronizers_[0]->SetGlobalShardIndices( + submanager_to_global_shards_[0]); + sub_synchronizers_[0]->SetLocalShardIndices(submanager_to_local_shards_[0]); + sub_synchronizers_[0]->SetControlDelegate(this); + } +} + NumaAwareWeightSynchronizer::NumaAwareWeightSynchronizer( std::vector> sub_synchronizers) { @@ -425,6 +459,17 @@ const uint8_t* NumaAwareWeightSynchronizer::GetHostBufferPtr( return sub_synchronizers_[sub_idx]->GetHostBufferPtr(layer_idx, local_shard); } +size_t NumaAwareWeightSynchronizer::GetHostBufferSize(size_t layer_idx, + size_t shard_idx) const { + if (shard_idx >= global_shard_to_submanager_.size()) return 0; + auto [sub_idx, local_shard] = global_shard_to_submanager_[shard_idx]; + if (sub_idx < 0 || sub_idx >= static_cast(sub_synchronizers_.size()) || + !sub_synchronizers_[sub_idx]) { + return 0; + } + return sub_synchronizers_[sub_idx]->GetHostSize(layer_idx, local_shard); +} + absl::StatusOr NumaAwareWeightSynchronizer::D2h( uint64_t uuid) { if (sub_synchronizers_.empty()) return raiden::PjRtCopyFuture(); @@ -833,6 +878,33 @@ void NumaAwareWeightSynchronizer::SetSubmanagerShardsForTesting( } } +void NumaAwareWeightSynchronizer::SetTestOnlyRateLimiters( + double test_only_simulated_egress_gbps, + double test_only_simulated_ingress_gbps) { + std::shared_ptr egress_limiter = nullptr; + if (test_only_simulated_egress_gbps > 0.0) { + const uint64_t egress_bytes_per_sec = + static_cast(test_only_simulated_egress_gbps * 1e9 / 8.0); + egress_limiter = std::make_shared( + egress_bytes_per_sec); + } + + std::shared_ptr ingress_limiter = + nullptr; + if (test_only_simulated_ingress_gbps > 0.0) { + const uint64_t ingress_bytes_per_sec = + static_cast(test_only_simulated_ingress_gbps * 1e9 / 8.0); + ingress_limiter = std::make_shared( + ingress_bytes_per_sec); + } + + for (auto& sub : sub_synchronizers_) { + if (sub) { + sub->SetTestOnlyRateLimiters(egress_limiter, ingress_limiter); + } + } +} + // ============================================================================ // WeightSynchronizer (Top-Level Facade) Implementation // ============================================================================ @@ -863,6 +935,16 @@ WeightSynchronizer::WeightSynchronizer( listener_port, bind_ip, auto_h2d, global_shard_indices); } +WeightSynchronizer::WeightSynchronizer( + size_t num_layers, size_t num_shards, std::vector slice_byte_sizes, + std::optional local_port, int parallelism, + std::optional listener_port, std::optional bind_ip, + bool auto_h2d, std::optional> global_shard_indices) { + numa_manager_ = std::make_unique( + num_layers, num_shards, std::move(slice_byte_sizes), local_port, + parallelism, listener_port, bind_ip, auto_h2d, global_shard_indices); +} + WeightSynchronizer::WeightSynchronizer( std::vector> sub_synchronizers) { @@ -903,6 +985,11 @@ const uint8_t* WeightSynchronizer::GetHostBufferPtr(size_t layer_idx, return numa_manager_->GetHostBufferPtr(layer_idx, shard_idx); } +size_t WeightSynchronizer::GetHostBufferSize(size_t layer_idx, + size_t shard_idx) const { + return numa_manager_->GetHostBufferSize(layer_idx, shard_idx); +} + std::optional WeightSynchronizer::local_port() const { return numa_manager_->local_port(); } @@ -936,5 +1023,52 @@ size_t WeightSynchronizer::slice_byte_size() const { return numa_manager_->slice_byte_size(); } +void WeightSynchronizer::test_only_set_bandwidth_limit( + double test_only_simulated_egress_gbps, + double test_only_simulated_ingress_gbps) { + if (numa_manager_) { + numa_manager_->SetTestOnlyRateLimiters(test_only_simulated_egress_gbps, + test_only_simulated_ingress_gbps); + } +} + +std::unique_ptr +WeightSynchronizer::test_only_create_cpu_instance( + size_t num_layers, size_t num_shards, size_t slice_byte_size, + std::optional local_port, int parallelism, + std::optional listener_port, std::optional bind_ip, + bool auto_h2d, std::optional> global_shard_indices, + double test_only_simulated_egress_gbps, + double test_only_simulated_ingress_gbps) { + auto ws = std::make_unique( + num_layers, num_shards, slice_byte_size, local_port, parallelism, + listener_port, bind_ip, auto_h2d, global_shard_indices); + if (test_only_simulated_egress_gbps > 0.0 || + test_only_simulated_ingress_gbps > 0.0) { + ws->test_only_set_bandwidth_limit(test_only_simulated_egress_gbps, + test_only_simulated_ingress_gbps); + } + return ws; +} + +std::unique_ptr +WeightSynchronizer::test_only_create_cpu_instance( + size_t num_layers, size_t num_shards, std::vector slice_byte_sizes, + std::optional local_port, int parallelism, + std::optional listener_port, std::optional bind_ip, + bool auto_h2d, std::optional> global_shard_indices, + double test_only_simulated_egress_gbps, + double test_only_simulated_ingress_gbps) { + auto ws = std::make_unique( + num_layers, num_shards, std::move(slice_byte_sizes), local_port, + parallelism, listener_port, bind_ip, auto_h2d, global_shard_indices); + if (test_only_simulated_egress_gbps > 0.0 || + test_only_simulated_ingress_gbps > 0.0) { + ws->test_only_set_bandwidth_limit(test_only_simulated_egress_gbps, + test_only_simulated_ingress_gbps); + } + return ws; +} + } // namespace jax } // namespace tpu_raiden diff --git a/tpu_sync/frameworks/jax/weight_synchronizer.h b/tpu_sync/frameworks/jax/weight_synchronizer.h index 8d03b59f3..554221987 100644 --- a/tpu_sync/frameworks/jax/weight_synchronizer.h +++ b/tpu_sync/frameworks/jax/weight_synchronizer.h @@ -33,6 +33,7 @@ #include #include "tpu_sync/frameworks/jax/jax_utils.h" #endif +#include "tpu_sync/transport/lib/test_only_rate_limiter.h" #include "tpu_sync/weight_sync/weight_synchronizer_base.h" namespace tpu_sync { @@ -74,6 +75,13 @@ class NumaAwareWeightSynchronizer std::optional listener_port = std::nullopt, std::optional bind_ip = std::nullopt, bool auto_h2d = false, std::optional> global_shard_indices = std::nullopt); + NumaAwareWeightSynchronizer( + size_t num_layers, size_t num_shards, + std::vector slice_byte_sizes, + std::optional local_port = std::nullopt, int parallelism = 1, + std::optional listener_port = std::nullopt, + std::optional bind_ip = std::nullopt, bool auto_h2d = false, + std::optional> global_shard_indices = std::nullopt); // Test-only constructor for injecting mock sub-synchronizers explicit NumaAwareWeightSynchronizer( @@ -94,6 +102,7 @@ class NumaAwareWeightSynchronizer std::vector get_local_endpoints() const; const uint8_t* GetHostBufferPtr(size_t layer_idx, size_t shard_idx) const; + size_t GetHostBufferSize(size_t layer_idx, size_t shard_idx) const; absl::StatusOr D2h(uint64_t uuid = 0); absl::StatusOr H2d(uint64_t uuid = 0); @@ -128,6 +137,9 @@ class NumaAwareWeightSynchronizer void SetSubmanagerShardsForTesting( const std::vector>& assignment); + void SetTestOnlyRateLimiters(double test_only_simulated_egress_gbps, + double test_only_simulated_ingress_gbps); + private: void InitSubManagers( const std::vector>& layer_buffers, @@ -180,6 +192,13 @@ class WeightSynchronizer { std::optional listener_port = std::nullopt, std::optional bind_ip = std::nullopt, bool auto_h2d = false, std::optional> global_shard_indices = std::nullopt); + WeightSynchronizer( + size_t num_layers, size_t num_shards, + std::vector slice_byte_sizes, + std::optional local_port = std::nullopt, int parallelism = 1, + std::optional listener_port = std::nullopt, + std::optional bind_ip = std::nullopt, bool auto_h2d = false, + std::optional> global_shard_indices = std::nullopt); // Test-only constructor for injecting mock sub-synchronizers explicit WeightSynchronizer( @@ -202,6 +221,7 @@ class WeightSynchronizer { void ResetMetrics(); const uint8_t* GetHostBufferPtr(size_t layer_idx, size_t shard_idx) const; + size_t GetHostBufferSize(size_t layer_idx, size_t shard_idx) const; std::optional local_port() const; std::optional listener_port() const; bool is_listener_active() const; @@ -213,6 +233,28 @@ class WeightSynchronizer { size_t num_shards() const; size_t slice_byte_size() const; + void test_only_set_bandwidth_limit(double test_only_simulated_egress_gbps, + double test_only_simulated_ingress_gbps); + + static std::unique_ptr test_only_create_cpu_instance( + size_t num_layers, size_t num_shards, size_t slice_byte_size, + std::optional local_port = std::nullopt, int parallelism = 1, + std::optional listener_port = std::nullopt, + std::optional bind_ip = std::nullopt, bool auto_h2d = false, + std::optional> global_shard_indices = std::nullopt, + double test_only_simulated_egress_gbps = 0.0, + double test_only_simulated_ingress_gbps = 0.0); + + static std::unique_ptr test_only_create_cpu_instance( + size_t num_layers, size_t num_shards, + std::vector slice_byte_sizes, + std::optional local_port = std::nullopt, int parallelism = 1, + std::optional listener_port = std::nullopt, + std::optional bind_ip = std::nullopt, bool auto_h2d = false, + std::optional> global_shard_indices = std::nullopt, + double test_only_simulated_egress_gbps = 0.0, + double test_only_simulated_ingress_gbps = 0.0); + private: std::unique_ptr numa_manager_; }; diff --git a/tpu_sync/transport/block_transport.h b/tpu_sync/transport/block_transport.h index 72d3a214c..c4fdd9568 100644 --- a/tpu_sync/transport/block_transport.h +++ b/tpu_sync/transport/block_transport.h @@ -133,6 +133,13 @@ class BlockTransport final { expected_layer_chunks); } + void SetTestOnlyRateLimiters( + std::shared_ptr egress, + std::shared_ptr ingress) { + raw_transport_.SetTestOnlyRateLimiters(std::move(egress), + std::move(ingress)); + } + private: lib::Request BuildBlockRequest( uint8_t socket_opcode, uint8_t* laddr, size_t len, uint32_t count_or_size, diff --git a/tpu_sync/transport/lib/BUILD b/tpu_sync/transport/lib/BUILD index 975060134..ffb01577b 100644 --- a/tpu_sync/transport/lib/BUILD +++ b/tpu_sync/transport/lib/BUILD @@ -162,6 +162,29 @@ cc_test( ], ) +cc_library( + name = "test_only_rate_limiter", + srcs = ["test_only_rate_limiter.cc"], + hdrs = ["test_only_rate_limiter.h"], + visibility = ["//visibility:public"], + deps = [ + "@com_google_absl//absl/base:core_headers", + "@com_google_absl//absl/synchronization", + "@com_google_absl//absl/time", + ], +) + +cc_test( + name = "test_only_rate_limiter_test", + srcs = ["test_only_rate_limiter_test.cc"], + deps = [ + ":test_only_rate_limiter", + "@com_google_absl//absl/time", + "@com_google_googletest//:gtest", + "@com_google_googletest//:gtest_main", + ], +) + cc_library( name = "raw_buffer_transport", srcs = ["raw_buffer_transport.cc"], @@ -176,6 +199,7 @@ cc_library( ":chunk_serializer", ":histogram", ":raw_buffer_transport_delegate", + ":test_only_rate_limiter", ":transport_adapter", "//tpu_sync/common:detached_thread_group", "//tpu_sync/core:numa_thread_pool", diff --git a/tpu_sync/transport/lib/raw_buffer_transport.cc b/tpu_sync/transport/lib/raw_buffer_transport.cc index 673ec2043..adbfb2560 100644 --- a/tpu_sync/transport/lib/raw_buffer_transport.cc +++ b/tpu_sync/transport/lib/raw_buffer_transport.cc @@ -65,6 +65,7 @@ #include "tpu_sync/transport/lib/histogram.h" #include "tpu_sync/transport/lib/raw_buffer_transport_delegate.h" #include "tpu_sync/transport/lib/socket/tcp_psp_helper.h" +#include "tpu_sync/transport/lib/test_only_rate_limiter.h" #include "tpu_sync/transport/peregrine/src/api/socket_util.h" namespace tpu_raiden::transport::lib { @@ -267,6 +268,12 @@ absl::Status RawBufferTransport::ProcessPeerRequest(int client_fd) { uint8_t* const dest_ptr = base_host_ptr + dst_offset; ABSL_RETURN_IF_ERROR(ReadExact(client_fd, dest_ptr, size_bytes)); + TestOnlyRateLimiter* const limiter = + test_only_ingress_rate_limiter_raw_.load(std::memory_order_relaxed); + if (ABSL_PREDICT_FALSE(limiter != nullptr)) { + limiter->Consume(size_bytes); + } + const uint8_t ack = 1; ABSL_RETURN_IF_ERROR(WriteExact(client_fd, &ack, 1)); @@ -359,6 +366,12 @@ absl::Status RawBufferTransport::ProcessPeerRequest(int client_fd) { ABSL_RETURN_IF_ERROR(ReadVExact(client_fd, iovs)); } + TestOnlyRateLimiter* const limiter = + test_only_ingress_rate_limiter_raw_.load(std::memory_order_relaxed); + if (ABSL_PREDICT_FALSE(limiter != nullptr)) { + limiter->Consume(total_bytes); + } + const uint8_t ack = 1; ABSL_RETURN_IF_ERROR(WriteExact(client_fd, &ack, 1)); @@ -706,6 +719,12 @@ absl::Status RawBufferTransport::ProcessSocketBufferPush( return absl::InternalError("PushBuffer verification failed"); } + TestOnlyRateLimiter* const limiter = + test_only_egress_rate_limiter_raw_.load(std::memory_order_relaxed); + if (ABSL_PREDICT_FALSE(limiter != nullptr)) { + limiter->Consume(request.len); + } + ok_to_pool = true; return absl::OkStatus(); } @@ -953,10 +972,27 @@ absl::Status RawBufferTransport::ProcessSocketBufferBatchPush( "ProcessSocketBufferBatchPush verification failed"); } + TestOnlyRateLimiter* const limiter = + test_only_egress_rate_limiter_raw_.load(std::memory_order_relaxed); + if (ABSL_PREDICT_FALSE(limiter != nullptr)) { + limiter->Consume(total_bytes); + } + ok_to_pool = true; return absl::OkStatus(); } +void RawBufferTransport::SetTestOnlyRateLimiters( + std::shared_ptr egress, + std::shared_ptr ingress) { + test_only_egress_rate_limiter_ = std::move(egress); + test_only_ingress_rate_limiter_ = std::move(ingress); + test_only_egress_rate_limiter_raw_.store(test_only_egress_rate_limiter_.get(), + std::memory_order_relaxed); + test_only_ingress_rate_limiter_raw_.store( + test_only_ingress_rate_limiter_.get(), std::memory_order_relaxed); +} + void RawBufferTransport::ForgetPushProgress(uint64_t uuid) { absl::MutexLock lock(raw_progress_mu_); raw_progress_.erase(uuid); diff --git a/tpu_sync/transport/lib/raw_buffer_transport.h b/tpu_sync/transport/lib/raw_buffer_transport.h index 9f7fb0ecc..015a552b0 100644 --- a/tpu_sync/transport/lib/raw_buffer_transport.h +++ b/tpu_sync/transport/lib/raw_buffer_transport.h @@ -43,6 +43,7 @@ #include "tpu_sync/transport/lib/conn/pool.h" #include "tpu_sync/transport/lib/raw_buffer_transport_delegate.h" #include "tpu_sync/transport/lib/socket/tcp_psp_helper.h" +#include "tpu_sync/transport/lib/test_only_rate_limiter.h" #include "tpu_sync/transport/lib/transport_adapter.h" namespace tpu_raiden::transport::lib { @@ -150,6 +151,9 @@ class RawBufferTransport final { absl::StatusOr RegisterPspPeer(uint32_t client_spi, absl::string_view client_key); + void SetTestOnlyRateLimiters(std::shared_ptr egress, + std::shared_ptr ingress); + private: // Pulls a buffer request from `peer` over a borrowed TCP connection. absl::Status ProcessSocketBufferPull(absl::string_view peer, @@ -205,6 +209,16 @@ class RawBufferTransport final { // To protect multiple gRPC threads can call RegisterPspPeer concurrently absl::Mutex psp_mu_; + // Note: Google3's libc++ does not implement the C++20 + // std::atomic> specialization, and atomic raw pointers + // allow lock-free relaxed loads on the socket hot path. + std::shared_ptr test_only_egress_rate_limiter_ = nullptr; + std::shared_ptr test_only_ingress_rate_limiter_ = + nullptr; + std::atomic test_only_egress_rate_limiter_raw_{nullptr}; + std::atomic test_only_ingress_rate_limiter_raw_{ + nullptr}; + std::thread listener_thread_; }; diff --git a/tpu_sync/transport/lib/test_only_rate_limiter.cc b/tpu_sync/transport/lib/test_only_rate_limiter.cc new file mode 100644 index 000000000..9b436d46b --- /dev/null +++ b/tpu_sync/transport/lib/test_only_rate_limiter.cc @@ -0,0 +1,48 @@ +// Copyright 2026 Google LLC. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "tpu_sync/transport/lib/test_only_rate_limiter.h" + +#include + +#include "absl/synchronization/mutex.h" +#include "absl/time/clock.h" +#include "absl/time/time.h" + +namespace tpu_raiden::transport::lib { + +absl::Time TestOnlyRateLimiter::Update(size_t bytes) { + const absl::Duration required_time = + absl::Seconds(static_cast(bytes) * sec_per_byte_); + absl::MutexLock lock(mu_); + const absl::Time now = absl::Now(); + if (next_available_time_ < now) { + next_available_time_ = now; + } + next_available_time_ += required_time; + return next_available_time_; +} + +void TestOnlyRateLimiter::Consume(size_t bytes) { + if (bandwidth_bytes_per_sec_ == 0 || bytes == 0) { + return; + } + const absl::Time sleep_until = Update(bytes); + const absl::Time now = absl::Now(); + if (sleep_until > now) { + absl::SleepFor(sleep_until - now); + } +} + +} // namespace tpu_raiden::transport::lib diff --git a/tpu_sync/transport/lib/test_only_rate_limiter.h b/tpu_sync/transport/lib/test_only_rate_limiter.h new file mode 100644 index 000000000..a3117f32a --- /dev/null +++ b/tpu_sync/transport/lib/test_only_rate_limiter.h @@ -0,0 +1,50 @@ +// Copyright 2026 Google LLC. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#ifndef THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_LIB_TEST_ONLY_RATE_LIMITER_H_ +#define THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_LIB_TEST_ONLY_RATE_LIMITER_H_ + +#include +#include + +#include "absl/base/thread_annotations.h" +#include "absl/synchronization/mutex.h" +#include "absl/time/time.h" + +namespace tpu_raiden::transport::lib { + +class TestOnlyRateLimiter { + public: + explicit TestOnlyRateLimiter(uint64_t bandwidth_bytes_per_sec) + : bandwidth_bytes_per_sec_(bandwidth_bytes_per_sec), + sec_per_byte_(bandwidth_bytes_per_sec > 0 + ? 1.0 / static_cast(bandwidth_bytes_per_sec) + : 0.0) {} + + uint64_t bandwidth_bytes_per_sec() const { return bandwidth_bytes_per_sec_; } + + void Consume(size_t bytes); + + private: + absl::Time Update(size_t bytes); + + const uint64_t bandwidth_bytes_per_sec_; + const double sec_per_byte_; + absl::Mutex mu_; + absl::Time next_available_time_ ABSL_GUARDED_BY(mu_) = absl::InfinitePast(); +}; + +} // namespace tpu_raiden::transport::lib + +#endif // THIRD_PARTY_TPU_RAIDEN_TPU_SYNC_TRANSPORT_LIB_TEST_ONLY_RATE_LIMITER_H_ diff --git a/tpu_sync/transport/lib/test_only_rate_limiter_test.cc b/tpu_sync/transport/lib/test_only_rate_limiter_test.cc new file mode 100644 index 000000000..49563fa2a --- /dev/null +++ b/tpu_sync/transport/lib/test_only_rate_limiter_test.cc @@ -0,0 +1,67 @@ +// Copyright 2026 Google LLC. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "tpu_sync/transport/lib/test_only_rate_limiter.h" + +#include // NOLINT + +#include +#include "absl/time/clock.h" +#include "absl/time/time.h" + +namespace tpu_raiden::transport::lib { +namespace { + +TEST(TestOnlyRateLimiterTest, ZeroBandwidthDoesNotSleep) { + TestOnlyRateLimiter limiter(0); + absl::Time start = absl::Now(); + limiter.Consume(1024 * 1024); + absl::Duration elapsed = absl::Now() - start; + EXPECT_LT(elapsed, absl::Milliseconds(10)); +} + +TEST(TestOnlyRateLimiterTest, ZeroBytesDoesNotSleep) { + TestOnlyRateLimiter limiter(1000); + absl::Time start = absl::Now(); + limiter.Consume(0); + absl::Duration elapsed = absl::Now() - start; + EXPECT_LT(elapsed, absl::Milliseconds(10)); +} + +TEST(TestOnlyRateLimiterTest, SequentialPacingAccuracy) { + // 100,000 bytes per second = 100 bytes / millisecond + TestOnlyRateLimiter limiter(100000); + absl::Time start = absl::Now(); + limiter.Consume(5000); // 50ms + absl::Duration elapsed = absl::Now() - start; + EXPECT_GE(elapsed, absl::Milliseconds(40)); + EXPECT_LE(elapsed, absl::Milliseconds(150)); +} + +TEST(TestOnlyRateLimiterTest, ConcurrentTimelineAdvancement) { + // 1,000,000 bytes per second. Two threads consume 25,000 bytes each (50,000 + // bytes total = 50ms). + TestOnlyRateLimiter limiter(1000000); + absl::Time start = absl::Now(); + std::thread t1([&] { limiter.Consume(25000); }); + std::thread t2([&] { limiter.Consume(25000); }); + t1.join(); + t2.join(); + absl::Duration elapsed = absl::Now() - start; + EXPECT_GE(elapsed, absl::Milliseconds(40)); + EXPECT_LE(elapsed, absl::Milliseconds(200)); +} + +} // namespace +} // namespace tpu_raiden::transport::lib diff --git a/tpu_sync/weight_sync/BUILD b/tpu_sync/weight_sync/BUILD index 3c42e9b31..957e33601 100644 --- a/tpu_sync/weight_sync/BUILD +++ b/tpu_sync/weight_sync/BUILD @@ -66,6 +66,7 @@ cc_library( "//tpu_sync/core:xla_raw_transfer_headers", "//tpu_sync/rpc:raiden_service_cc_proto", "//tpu_sync/transport:buffer_push_task", + "//tpu_sync/transport/lib:test_only_rate_limiter", "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/container:flat_hash_set", @@ -165,3 +166,20 @@ py_test( "@com_google_absl_py//absl/testing:absltest", ], ) + +py_test( + name = "weight_sync_fanout_perf_test", + size = "large", + srcs = ["weight_sync_fanout_perf_test.py"], + tags = ["requires-mem:28g"], + deps = [ + "//tpu_sync/api:common", + "//tpu_sync/api/jax:weight_synchronizer_jax_py", + "//tpu_sync/rpc:raiden_controller", + "//tpu_sync/rpc:raiden_service_py_pb2", + "@com_google_absl_py//absl/flags", + "@com_google_absl_py//absl/testing:absltest", + "@com_google_absl_py//absl/testing:parameterized", + "@pypi//numpy", + ], +) diff --git a/tpu_sync/weight_sync/weight_sync_fanout_perf_test.py b/tpu_sync/weight_sync/weight_sync_fanout_perf_test.py new file mode 100644 index 000000000..5f0f6c999 --- /dev/null +++ b/tpu_sync/weight_sync/weight_sync_fanout_perf_test.py @@ -0,0 +1,633 @@ +# Copyright 2026 Google LLC. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Performance unit test for weight syncing fan-out and resharding. + +Benchmarks 1-to-4 Flat Direct Push (broadcast_k=64) versus Tree Broadcast +(broadcast_k=2) side-by-side using host DRAM loopback networking on scaled +Qwen-35B model specs, verifying micro-block fragmentation realism and parity. +""" + +import asyncio +import time +from typing import Any, Dict, List, Optional, Tuple + +from absl import flags +from absl.testing import absltest +from absl.testing import parameterized +import numpy as np + +from tpu_sync.api.common import RaidenId +from tpu_sync.api.jax import weight_synchronizer +from tpu_sync.rpc import raiden_controller +from tpu_sync.rpc import raiden_service_pb2 + +_NUM_LAYERS = flags.DEFINE_integer( + "num_layers", 40, "Number of layers to simulate." +) +_NUM_ROUTED_EXPERTS = flags.DEFINE_integer( + "num_routed_experts", 8, "Number of routed experts." +) +_NUM_DESTINATIONS = flags.DEFINE_integer( + "num_destinations", 4, "Number of destination replicas." +) +_TEST_ONLY_SIMULATED_NIC_GBPS = flags.DEFINE_float( + "test_only_simulated_nic_gbps", + 0.0, + "Simulated NIC line rate in Gbps (0.0 = unlimited). When 0.0, " + "test_rate_limited_flat_vs_tree_comparison benchmarks at 10.0 Gbps.", +) + + +def _make_scaled_qwen_specs( + num_layers: Optional[int] = None, + num_routed_experts: Optional[int] = None, + role: str = "source", +) -> List[Tuple[Tuple[int, ...], List[str], str, int]]: + """Generates parameter specs for scaled Qwen 3.5 35B.""" + if num_layers is None: + num_layers = _NUM_LAYERS.value + if num_routed_experts is None: + num_routed_experts = _NUM_ROUTED_EXPERTS.value + + dim = 2048 + routed_mlp_dim = 512 + shared_mlp_dim = 512 + attn_dim = 2048 + linear_ba_dim = 64 + linear_conv_dim = 256 + + is_dest = role == "destination" + specs = [] + + for l in range(num_layers): + if l % 2 == 0: + # Even layer: Full Attention (GQA) + MoE Block + specs.append(( + (dim, 2, 256), + ["", "", ""] if is_dest else ["", "", ""], + f"decoder.layers.{l}.attention.attention.key.kernel", + l, + )) + specs.append(( + (dim, attn_dim), + ["", "tp_out"] if is_dest else ["", ""], + f"decoder.layers.{l}.attention.attention.out.kernel", + l, + )) + specs.append(( + (dim, attn_dim), + ["", "tp_out"] if is_dest else ["", ""], + f"decoder.layers.{l}.attention.attention.query.kernel", + l, + )) + specs.append(( + (dim, 2, 256), + ["", "", ""] if is_dest else ["", "", ""], + f"decoder.layers.{l}.attention.attention.value.kernel", + l, + )) + else: + # Odd layer: GDN Linear Attention + MoE Block + specs.append(( + (linear_ba_dim, linear_conv_dim), + ["", ""] if is_dest else ["", ""], + f"decoder.layers.{l}.attention.linear_attn.b_kernel", + l, + )) + specs.append(( + (dim, linear_ba_dim), + ["", ""] if is_dest else ["", ""], + f"decoder.layers.{l}.attention.linear_attn.ba_kernel", + l, + )) + specs.append(( + (4, 1, linear_conv_dim), + ["", "", ""] if is_dest else ["", "", ""], + f"decoder.layers.{l}.attention.linear_attn.conv1d.kernel", + l, + )) + specs.append(( + (dim, attn_dim), + ["", "tp_out"] if is_dest else ["", ""], + f"decoder.layers.{l}.attention.linear_attn.g_kernel", + l, + )) + specs.append(( + (attn_dim, dim), + ["", "tp_out"] if is_dest else ["", ""], + f"decoder.layers.{l}.attention.linear_attn.out_kernel", + l, + )) + specs.append(( + (dim, attn_dim), + ["", "tp_out"] if is_dest else ["", ""], + f"decoder.layers.{l}.attention.linear_attn.qkvz_kernel", + l, + )) + + # Shared MoE block across both even and odd layers + specs.append(( + (num_routed_experts, routed_mlp_dim, dim), + ["", "", "tp_wo"] if is_dest else ["", "", ""], + f"decoder.layers.{l}.mlp.experts.down_proj.kernel", + l, + )) + specs.append(( + (num_routed_experts, dim, routed_mlp_dim), + ["", "", "tp"] if is_dest else ["", "", ""], + f"decoder.layers.{l}.mlp.experts.gate_proj.kernel", + l, + )) + specs.append(( + (num_routed_experts, dim, routed_mlp_dim), + ["", "", "tp"] if is_dest else ["", "", ""], + f"decoder.layers.{l}.mlp.experts.up_proj.kernel", + l, + )) + specs.append(( + (shared_mlp_dim, dim), + ["", "tp_wo"] if is_dest else ["", ""], + f"decoder.layers.{l}.mlp.shared_expert.down_proj.kernel", + l, + )) + specs.append(( + (dim, shared_mlp_dim), + ["", "tp"] if is_dest else ["", ""], + f"decoder.layers.{l}.mlp.shared_expert.gate_proj.kernel", + l, + )) + specs.append(( + (dim, shared_mlp_dim), + ["", "tp"] if is_dest else ["", ""], + f"decoder.layers.{l}.mlp.shared_expert.up_proj.kernel", + l, + )) + + return specs + + +def build_variable_protos( + specs: List[Tuple[Tuple[int, ...], List[str], str, int]], + mesh_shape_dict: Dict[str, int], + item_size: int = 2, +) -> List[raiden_service_pb2.VariableMetadataProto]: + """Constructs VariableMetadataProtos with explicit sharding specs and layouts.""" + protos = [] + for idx, (global_shape, spec_axes, name, _) in enumerate(specs): + sharding_shape = [mesh_shape_dict.get(axis, 1) for axis in spec_axes] + layout = list(range(len(global_shape) - 1, -1, -1)) + protos.append( + raiden_service_pb2.VariableMetadataProto( + name=name, + shape=list(global_shape), + mesh_shape=sharding_shape, + layout=layout, + item_size=item_size, + layer_idx=idx, + sharding_spec=spec_axes, + global_shard_indices=[0], + ) + ) + return protos + + +class WeightSyncFanoutPerfTest(parameterized.TestCase): + """Performance and micro-block fragmentation tests for weight sync fan-out.""" + + def setUp(self): + super().setUp() + self.num_destinations = _NUM_DESTINATIONS.value + self.num_layers_flag = _NUM_LAYERS.value + self.num_routed_experts_flag = _NUM_ROUTED_EXPERTS.value + + # 1. Model specs & variable protos + self.src_mesh_dict = { + "tp": 1, + "tp_wo": 1, + "tp_out": 1, + } + self.dst_mesh_dict = { + "tp": 2, + "tp_wo": 32, + "tp_out": 8, + } + + self.src_specs = _make_scaled_qwen_specs( + num_layers=self.num_layers_flag, + num_routed_experts=self.num_routed_experts_flag, + role="source", + ) + self.dst_specs = _make_scaled_qwen_specs( + num_layers=self.num_layers_flag, + num_routed_experts=self.num_routed_experts_flag, + role="destination", + ) + + self.num_layers = len(self.src_specs) + + self.src_var_protos = build_variable_protos( + self.src_specs, self.src_mesh_dict + ) + self.dst_var_protos = build_variable_protos( + self.dst_specs, self.dst_mesh_dict + ) + + # Calculate model volume per destination replica + self.total_model_bytes = sum( + int(np.prod(proto.shape) // np.prod(proto.mesh_shape)) * proto.item_size + for proto in self.dst_var_protos + ) + + # Calculate buffer sizes needed per variable + self.src_slice_byte_sizes = [ + int(np.prod(proto.shape)) * proto.item_size + for proto in self.src_var_protos + ] + self.dst_slice_byte_sizes = [ + int(np.prod(proto.shape) // np.prod(proto.mesh_shape)) * proto.item_size + for proto in self.dst_var_protos + ] + + # Calculate max transferred offset per variable + self.layer_max_bytes: Dict[int, int] = {} + for proto in self.dst_var_protos: + var_bytes = ( + int(np.prod(proto.shape) // np.prod(proto.mesh_shape)) + * proto.item_size + ) + self.layer_max_bytes[proto.layer_idx] = var_bytes + + # 2. Start centralized Controller Server on loopback + self.controller_network_client = ( + raiden_controller.WeightSyncWorkerRpcClient(name_resolver=None) + ) + self.addCleanup(self.controller_network_client.close) + + self.controller = raiden_controller.RaidenController( + port=0, + worker_rpc_client=self.controller_network_client, + ) + self.controller_server = raiden_controller.RaidenControllerServer( + self.controller + ) + self.controller_server.start() + self.addCleanup(self.controller_server.stop) + self.controller_port = self.controller_server.port + + self.ctrl_client = raiden_controller.RaidenControllerClientFacade( + f"127.0.0.1:{self.controller_port}", + name_resolver=None, + ) + if hasattr(self.ctrl_client, "_control_pipe_client") and hasattr( + self.ctrl_client._control_pipe_client, "close" + ): + self.addCleanup(self.ctrl_client._control_pipe_client.close) + + # 3. Instantiate 1 source worker + 4 destination workers + self.ws_src = ( + weight_synchronizer.WeightSynchronizer.test_only_create_cpu_instance( + num_layers=self.num_layers, + num_shards=1, + slice_byte_size=self.src_slice_byte_sizes, + local_port=0, + listener_port=0, + bind_ip="127.0.0.1", + ) + ) + self.addCleanup(self.ws_src.shutdown) + + self.ws_dsts: List[weight_synchronizer.WeightSynchronizer] = [] + for _ in range(self.num_destinations): + ws_dst = ( + weight_synchronizer.WeightSynchronizer.test_only_create_cpu_instance( + num_layers=self.num_layers, + num_shards=1, + slice_byte_size=self.dst_slice_byte_sizes, + local_port=0, + listener_port=0, + bind_ip="127.0.0.1", + ) + ) + self.addCleanup(ws_dst.shutdown) + self.ws_dsts.append(ws_dst) + + # 4. Register work units + self.src_unit = RaidenId("trainer", "0", "weights") + self.dst_units = [ + RaidenId("sampler", str(i), "weights") + for i in range(self.num_destinations) + ] + + mesh_axes = ["tp", "tp_wo", "tp_out"] + mesh_shape = [1] * len(mesh_axes) + + self.ctrl_client.register_work_unit( + self.src_unit, + [f"127.0.0.1:{self.ws_src.local_port}"], + f"127.0.0.1:{self.ws_src.listener_port}", + mesh_shape=mesh_shape, + variables=self.src_var_protos, + mesh_axes=mesh_axes, + ) + + for i, ws_dst in enumerate(self.ws_dsts): + self.ctrl_client.register_work_unit( + self.dst_units[i], + [f"127.0.0.1:{ws_dst.local_port}"], + f"127.0.0.1:{ws_dst.listener_port}", + mesh_shape=mesh_shape, + variables=self.dst_var_protos, + mesh_axes=mesh_axes, + ) + + def _get_schedule_and_task_counts(self) -> Tuple[int, int]: + if hasattr(self, "_cached_task_counts"): + return self._cached_task_counts + old_k = self.controller.broadcast_k + self.controller.broadcast_k = 64 + self.controller._plan_cache.clear() + try: + loop = asyncio.new_event_loop() + try: + sched = loop.run_until_complete( + self.controller._compute_transfer_schedule( + src_units=[self.src_unit], + dst_units=self.dst_units, + skip_tiling={l: False for l in range(self.num_layers)}, + ) + ) + finally: + loop.close() + finally: + self.controller.broadcast_k = old_k + self.controller._plan_cache.clear() + + self.assertIsNotNone(sched) + self.assertIsNotNone(sched.direct_schedules) + + total_tasks = 0 + le_512_tasks = 0 + + for src_unit, sched_by_shard in sched.direct_schedules.items(): + for shard_idx, entries in sched_by_shard.items(): + for entry in entries: + size = entry[4] + src_stride = entry[7] + dst_stride = entry[8] + count = entry[9] + is_contiguous = (count == 1) or ( + src_stride == size and dst_stride == size + ) + num_tasks = 1 if is_contiguous else count + total_tasks += num_tasks + if size <= 512: + le_512_tasks += num_tasks + + self._cached_task_counts = (total_tasks, le_512_tasks) + return total_tasks, le_512_tasks + + def test_scaled_spec_chunk_distribution(self): + """Verifies that >95% of scheduled copy tasks have size_bytes <= 512 bytes.""" + total_tasks, le_512_tasks = self._get_schedule_and_task_counts() + self.assertGreater(total_tasks, 0, "No copy tasks were scheduled") + pct_le_512 = (le_512_tasks / total_tasks) * 100.0 + print( + f"\nChunk Distribution: {le_512_tasks}/{total_tasks} tasks" + f" ({pct_le_512:.2f}%) have size <= 512B" + ) + self.assertGreater( + pct_le_512, + 95.0, + f"Expected >95% tasks with size <= 512 bytes, got {pct_le_512:.2f}%", + ) + + def _run_flat_direct_push_perf(self) -> Tuple[float, float]: + """Executes 1-to-4 Flat Direct Push (broadcast_k=64) with byte parity check.""" + # 1. Fill source buffers with 0xAB + for l in range(self.num_layers): + buf = self.ws_src.get_host_buffer(layer_idx=l, shard_idx=0) + buf[:] = 0xAB + + # 2. Clear destination buffers to 0x00 + for ws_dst in self.ws_dsts: + for l in range(self.num_layers): + buf = ws_dst.get_host_buffer(layer_idx=l, shard_idx=0) + buf[:] = 0x00 + + # 3. Start flat direct push transfer + uuid = 1001 + self.controller.broadcast_k = 64 + self.controller._plan_cache.clear() + t0 = time.perf_counter() + future = self.controller.start_transfer( + src_units=[self.src_unit], + dst_units=self.dst_units, + dst_mem_type=raiden_controller.RaidenMemoryType.DRAM, + use_block_chunks=True, + is_sender=True, + uuid=uuid, + req_id="flat_perf", + skip_d2h=True, + skip_tiling={l: False for l in range(self.num_layers)}, + ) + loop = asyncio.new_event_loop() + try: + loop.run_until_complete(future.wait()) + finally: + loop.close() + + for ws_dst in self.ws_dsts: + ws_dst.wait_for_transfer_completion(uuid=uuid) + + elapsed = time.perf_counter() - t0 + + # 4. Verify byte parity across all 4 destinations + for ws_dst in self.ws_dsts: + for l in range(self.num_layers): + dst_buf = ws_dst.get_host_buffer(layer_idx=l, shard_idx=0) + valid_bytes = self.layer_max_bytes[l] + self.assertTrue( + np.all(dst_buf[:valid_bytes] == 0xAB), + f"Flat push byte parity mismatch in layer {l}", + ) + + throughput_gb_s = ( + (self.total_model_bytes * self.num_destinations) / 1e9 + ) / max(elapsed, 1e-9) + print( + f"\n[Flat Push (k=64)] Elapsed: {elapsed:.3f}s, Throughput:" + f" {throughput_gb_s:.2f} GB/s, Parity: PASS" + ) + return elapsed, throughput_gb_s + + def test_flat_direct_push_perf(self): + """Benchmarks 1-to-4 Flat Direct Push (broadcast_k=64) with byte parity check.""" + self._run_flat_direct_push_perf() + + def _run_tree_broadcast_perf(self) -> Tuple[float, float]: + """Executes 1-to-4 Tree Broadcast (broadcast_k=2) awaiting controller future.""" + # 1. Fill source buffers with 0xCD + for l in range(self.num_layers): + buf = self.ws_src.get_host_buffer(layer_idx=l, shard_idx=0) + buf[:] = 0xCD + + # 2. Re-zero destination buffers to 0x00 + for ws_dst in self.ws_dsts: + for l in range(self.num_layers): + buf = ws_dst.get_host_buffer(layer_idx=l, shard_idx=0) + buf[:] = 0x00 + + # 3. Start tree broadcast transfer (broadcast_k=2) + uuid = 1002 + self.controller.broadcast_k = 2 + self.controller._plan_cache.clear() + t0 = time.perf_counter() + future = self.controller.start_transfer( + src_units=[self.src_unit], + dst_units=self.dst_units, + dst_mem_type=raiden_controller.RaidenMemoryType.DRAM, + use_block_chunks=True, + is_sender=True, + uuid=uuid, + req_id="tree_perf", + skip_d2h=True, + skip_tiling={l: False for l in range(self.num_layers)}, + ) + loop = asyncio.new_event_loop() + try: + # Awaits controller tree broadcast future until all hops complete + loop.run_until_complete(future.wait()) + finally: + loop.close() + + elapsed = time.perf_counter() - t0 + + # 4. Verify byte parity across all 4 destinations + for ws_dst in self.ws_dsts: + for l in range(self.num_layers): + dst_buf = ws_dst.get_host_buffer(layer_idx=l, shard_idx=0) + valid_bytes = self.layer_max_bytes[l] + self.assertTrue( + np.all(dst_buf[:valid_bytes] == 0xCD), + f"Tree broadcast byte parity mismatch in layer {l}", + ) + + throughput_gb_s = ( + (self.total_model_bytes * self.num_destinations) / 1e9 + ) / max(elapsed, 1e-9) + print( + f"\n[Tree Broadcast (k=2)] Elapsed: {elapsed:.3f}s, Throughput:" + f" {throughput_gb_s:.2f} GB/s, Parity: PASS" + ) + return elapsed, throughput_gb_s + + def test_tree_broadcast_perf(self): + """Benchmarks 1-to-4 Tree Broadcast (broadcast_k=2) awaiting controller future.""" + self._run_tree_broadcast_perf() + + def test_sxs_flat_vs_tree_performance_comparison(self): + """Executes flat vs tree modes side-by-side and prints comparative summary table.""" + t_flat, bw_flat = self._run_flat_direct_push_perf() + t_tree, bw_tree = self._run_tree_broadcast_perf() + + total_tasks, _ = self._get_schedule_and_task_counts() + if total_tasks >= 1_000_000: + task_str = f"~{total_tasks / 1e6:.1f}M" + elif total_tasks >= 1000: + task_str = f"~{total_tasks // 1000}k" + else: + task_str = str(total_tasks) + model_mb = self.total_model_bytes / 1e6 + + print("\n" + "=" * 70) + print( + f"Weight Sync Fan-out Benchmark Results (Model: ~{model_mb:.1f} MB," + f" N={self.num_destinations})" + ) + print("=" * 70) + print( + f"{'Mode':<18} {'Time (s)':<12} {'Throughput (GB/s)':<19} {'Tasks':<10}" + f" {'Parity':<6}" + ) + print("-" * 70) + print( + f"{'Flat (k=64)':<18} {t_flat:.2f} s {bw_flat:.2f} GB/s " + f" {task_str:<10} PASS" + ) + print( + f"{'Tree (k=2)':<18} {t_tree:.2f} s {bw_tree:.2f} GB/s " + f" {task_str:<10} PASS" + ) + print("=" * 70 + "\n") + + def test_rate_limited_flat_vs_tree_comparison(self): + """Benchmarks Flat Direct Push vs Tree Broadcast under simulated NIC bandwidth cap.""" + rate_gbps = ( + _TEST_ONLY_SIMULATED_NIC_GBPS.value + if _TEST_ONLY_SIMULATED_NIC_GBPS.value > 0.0 + else 10.0 + ) + self.ws_src.test_only_set_bandwidth_limit( + test_only_simulated_egress_gbps=rate_gbps, + test_only_simulated_ingress_gbps=0.0, + ) + for ws_dst in self.ws_dsts: + ws_dst.test_only_set_bandwidth_limit( + test_only_simulated_egress_gbps=rate_gbps, + test_only_simulated_ingress_gbps=0.0, + ) + + t_flat, bw_flat = self._run_flat_direct_push_perf() + t_tree, bw_tree = self._run_tree_broadcast_perf() + + print("\n" + "=" * 70) + print( + f"Rate-Limited Weight Sync Fan-out Benchmark ({rate_gbps:.1f} Gbps NIC," + f" N={self.num_destinations})" + ) + print("=" * 70) + print( + f"{'Mode':<18} {'Time (s)':<12} {'Throughput (GB/s)':<19} {'Parity':<6}" + ) + print("-" * 70) + print( + f"{'Flat (k=64)':<18} {t_flat:.2f} s {bw_flat:.2f} GB/s " + " PASS" + ) + print( + f"{'Tree (k=2)':<18} {t_tree:.2f} s {bw_tree:.2f} GB/s " + " PASS" + ) + print("=" * 70 + "\n") + + if self.num_destinations >= 16: + self.assertLess( + t_tree, + t_flat, + f"Tree Broadcast ({t_tree:.3f}s) must outperform Flat Direct Push" + f" ({t_flat:.3f}s) under constrained NIC line rate at" + f" N={self.num_destinations}.", + ) + else: + # At small fan-out (N < 16, e.g. N=4), Tree broadcast only saves 1 serialized + # copy (from 4 to 3), which is outweighed by multi-hop store-and-forward + # Python scheduling latency. Verify that both modes executed cleanly and + # satisfied byte-exact parity. + self.assertGreater(t_flat, 0.0) + self.assertGreater(t_tree, 0.0) + + +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 a58c1bb93..86495fb97 100644 --- a/tpu_sync/weight_sync/weight_synchronizer_base.cc +++ b/tpu_sync/weight_sync/weight_synchronizer_base.cc @@ -56,6 +56,7 @@ #include "tpu_sync/core/raw_transfer_core.h" #include "tpu_sync/rpc/raiden_service.pb.h" #include "tpu_sync/transport/buffer_push_task.h" +#include "tpu_sync/transport/lib/test_only_rate_limiter.h" #include "tpu_sync/weight_sync/tiling_utils.h" #include "tpu_sync/weight_sync/weight_synchronizer_listener.h" @@ -1429,5 +1430,12 @@ size_t WeightSynchronizerBase::GetHostSize(size_t layer_idx, return layers_[layer_idx].shards[local_idx].host_size; } +void WeightSynchronizerBase::SetTestOnlyRateLimiters( + std::shared_ptr egress, + std::shared_ptr ingress) { + RaidenManagerBase::SetTestOnlyRateLimiters(std::move(egress), + std::move(ingress)); +} + } // namespace weight_sync } // namespace tpu_raiden diff --git a/tpu_sync/weight_sync/weight_synchronizer_base.h b/tpu_sync/weight_sync/weight_synchronizer_base.h index b0c9da016..7e242ac8f 100644 --- a/tpu_sync/weight_sync/weight_synchronizer_base.h +++ b/tpu_sync/weight_sync/weight_synchronizer_base.h @@ -35,6 +35,7 @@ #include "tpu_sync/core/raiden_manager_base.h" #include "tpu_sync/core/raiden_transfer_endpoint.h" #include "tpu_sync/core/raw_transfer_core.h" +#include "tpu_sync/transport/lib/test_only_rate_limiter.h" namespace tpu_sync { namespace rpc { @@ -269,6 +270,10 @@ class WeightSynchronizerBase : public tpu_raiden::RaidenManagerBase { metrics_ = m; } + void SetTestOnlyRateLimiters( + std::shared_ptr egress, + std::shared_ptr ingress); + void SetPipelineGroupSize(std::optional group_size) { pipeline_group_size_override_ = group_size; }