Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
61 changes: 60 additions & 1 deletion tpu_sync/api/jax/weight_synchronizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
49 changes: 49 additions & 0 deletions tpu_sync/api/jax/weight_synchronizer_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
1 change: 1 addition & 0 deletions tpu_sync/core/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
18 changes: 18 additions & 0 deletions tpu_sync/core/raiden_manager_base.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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 {

Expand Down Expand Up @@ -104,6 +105,18 @@ void RaidenManagerBase::StopTransportServer() {
}
}

void RaidenManagerBase::SetTestOnlyRateLimiters(
std::shared_ptr<transport::lib::TestOnlyRateLimiter> egress,
std::shared_ptr<transport::lib::TestOnlyRateLimiter> 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<HostNicAddress> RaidenManagerBase::GetHostNics() const {
return GetLocalHostNicAddresses();
}
Expand Down Expand Up @@ -176,6 +189,11 @@ RaidenManagerBase::InitTransportServer() {

server_ = std::make_unique<tpu_raiden::transport::BlockTransport>(
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();
}

Expand Down
9 changes: 9 additions & 0 deletions tpu_sync/core/raiden_manager_base.h
Original file line number Diff line number Diff line change
Expand Up @@ -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 {

Expand Down Expand Up @@ -96,6 +97,10 @@ class RaidenManagerBase : public tpu_raiden::transport::BlockTransportDelegate {

virtual void ForgetPushProgress(uint64_t uuid);

void SetTestOnlyRateLimiters(
std::shared_ptr<transport::lib::TestOnlyRateLimiter> egress,
std::shared_ptr<transport::lib::TestOnlyRateLimiter> ingress);

// Stops and joins the underlying raw transport server if active.
void StopTransportServer();

Expand Down Expand Up @@ -163,6 +168,10 @@ class RaidenManagerBase : public tpu_raiden::transport::BlockTransportDelegate {
mutable absl::Mutex server_init_mu_;
std::unique_ptr<tpu_raiden::transport::BlockTransport> server_
ABSL_GUARDED_BY(server_init_mu_);
std::shared_ptr<transport::lib::TestOnlyRateLimiter>
test_only_egress_rate_limiter_ ABSL_GUARDED_BY(server_init_mu_);
std::shared_ptr<transport::lib::TestOnlyRateLimiter>
test_only_ingress_rate_limiter_ ABSL_GUARDED_BY(server_init_mu_);

std::vector<LayerInfoBase> layers_;

Expand Down
2 changes: 2 additions & 0 deletions tpu_sync/frameworks/jax/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down
62 changes: 58 additions & 4 deletions tpu_sync/frameworks/jax/tpu_raiden_jax_module.cc
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
#include <nanobind/stl/string.h> // IWYU pragma: keep
#include <nanobind/stl/string_view.h> // IWYU pragma: keep
#include <nanobind/stl/tuple.h> // IWYU pragma: keep
#include <nanobind/stl/unique_ptr.h> // IWYU pragma: keep
#include <nanobind/stl/vector.h> // IWYU pragma: keep
#include "xla/pjrt/status_casters.h"
#include "tpu_sync/core/raiden_future.h"
Expand Down Expand Up @@ -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<int> local_port, int parallelism,
std::optional<int> listener_port,
std::optional<std::string> bind_ip, bool auto_h2d,
std::optional<std::vector<int64_t>> 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<size_t> slice_byte_sizes,
std::optional<int> local_port, int parallelism,
std::optional<int> listener_port,
std::optional<std::string> bind_ip, bool auto_h2d,
std::optional<std::vector<int64_t>> 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) {
Expand Down Expand Up @@ -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<uint8_t, nb::numpy, nb::c_contig>(
const_cast<uint8_t*>(ptr), 1, shape,
nb::handle() /* view only, no ownership copy */
);
const_cast<uint8_t*>(ptr), 1, shape, nb::find(&self));
},
nb::arg("layer_idx") = 0, nb::arg("shard_idx") = 0)
.def_prop_ro("local_port", &WeightSynchronizer::local_port)
Expand Down
Loading
Loading