Skip to content
Open
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
45 changes: 38 additions & 7 deletions tpu_raiden/frameworks/jax/weight_synchronizer_ffi.cc
Original file line number Diff line number Diff line change
Expand Up @@ -190,6 +190,7 @@ xla::ffi::Error TriggerWeightSynchronizerInitImpl(
// FFI execution handler for WeightSynchronizer Init and D2H (Host CPU Executed)
xla::ffi::Error TriggerWeightSynchronizerInitAndD2hImpl(
xla::ffi::AnyBuffer x, xla::ffi::AnyBuffer shard_idx_buf,
xla::ffi::AnyBuffer logical_idx_buf,
int64_t slice_byte_size, int32_t local_port, int32_t parallelism,
int32_t num_layers, int32_t listener_port,
xla::ffi::Result<xla::ffi::AnyBuffer> out) {
Expand All @@ -200,6 +201,12 @@ xla::ffi::Error TriggerWeightSynchronizerInitAndD2hImpl(
}
int32_t shard_idx =
*reinterpret_cast<const int32_t*>(shard_idx_buf.untyped_data());
if (logical_idx_buf.untyped_data() == nullptr) {
return xla::ffi::Error(xla::ffi::ErrorCode::kInvalidArgument,
"logical_idx_buf null.");
}
int32_t logical_idx =
*reinterpret_cast<const int32_t*>(logical_idx_buf.untyped_data());
if (shard_idx < 0 || shard_idx >= 32) {
return xla::ffi::Error(
xla::ffi::ErrorCode::kInvalidArgument,
Expand Down Expand Up @@ -282,8 +289,16 @@ xla::ffi::Error TriggerWeightSynchronizerInitAndD2hImpl(

// --- D2H Part ---
size_t size = g_weight_synchronizers[shard_idx]->slice_byte_size();
size_t local_slot = static_cast<size_t>(shard_idx) %
g_weight_synchronizers[shard_idx]->num_shards();
size_t num_shards_v = g_weight_synchronizers[shard_idx]->num_shards();
if (logical_idx < 0 || static_cast<size_t>(logical_idx) >= num_shards_v) {
return xla::ffi::Error(xla::ffi::ErrorCode::kInvalidArgument,
absl::StrCat("logical_idx out of range: logical_idx=", logical_idx,
", num_shards=", num_shards_v,
", device_ordinal=", shard_idx));
}
size_t local_slot = static_cast<size_t>(logical_idx);
VLOG(1) << "[D2H] device_ordinal=" << shard_idx
<< " logical=" << logical_idx << " local_slot=" << local_slot;
uint8_t* dst_host_ptr = const_cast<uint8_t*>(
g_weight_synchronizers[shard_idx]->GetHostBufferPtr(0, local_slot));
const uint8_t* src_device_ptr =
Expand Down Expand Up @@ -318,11 +333,11 @@ xla::ffi::Error TriggerWeightSynchronizerInitAndD2hImpl(

// FFI custom call handler executing asynchronous Host-to-Device (H2D) memory
// transfers from local staging buffers (`GetHostBufferPtr`) directly onto
// device memory buffers. Uses multi-host modulo indexing (`shard_idx %
// num_shards()`) to ensure global shard indices map to the correct local
// staging slot.
// device memory buffers. Uses explicit logical index from JAX to ensure global
// shard indices map to the correct local staging slot.
xla::ffi::Error TriggerH2DImpl(xla::ffi::AnyBuffer anchor,
xla::ffi::AnyBuffer shard_idx_buf,
xla::ffi::AnyBuffer logical_idx_buf,
xla::ffi::Result<xla::ffi::AnyBuffer> out) {
(void)anchor;
if (shard_idx_buf.untyped_data() == nullptr) {
Expand All @@ -331,15 +346,29 @@ xla::ffi::Error TriggerH2DImpl(xla::ffi::AnyBuffer anchor,
}
int32_t shard_idx =
*reinterpret_cast<const int32_t*>(shard_idx_buf.untyped_data());
if (logical_idx_buf.untyped_data() == nullptr) {
return xla::ffi::Error(xla::ffi::ErrorCode::kInvalidArgument,
"logical_idx_buf null.");
}
int32_t logical_idx =
*reinterpret_cast<const int32_t*>(logical_idx_buf.untyped_data());
if (shard_idx < 0 || shard_idx >= 32 ||
g_weight_synchronizers[shard_idx] == nullptr) {
return xla::ffi::Error(xla::ffi::ErrorCode::kInternal,
"WS not initialized.");
}

size_t size = g_weight_synchronizers[shard_idx]->slice_byte_size();
size_t local_slot = static_cast<size_t>(shard_idx) %
g_weight_synchronizers[shard_idx]->num_shards();
size_t num_shards_v = g_weight_synchronizers[shard_idx]->num_shards();
if (logical_idx < 0 || static_cast<size_t>(logical_idx) >= num_shards_v) {
return xla::ffi::Error(xla::ffi::ErrorCode::kInvalidArgument,
absl::StrCat("logical_idx out of range: logical_idx=", logical_idx,
", num_shards=", num_shards_v,
", device_ordinal=", shard_idx));
}
size_t local_slot = static_cast<size_t>(logical_idx);
VLOG(1) << "[H2D] device_ordinal=" << shard_idx
<< " logical=" << logical_idx << " local_slot=" << local_slot;
const uint8_t* src_host_ptr =
g_weight_synchronizers[shard_idx]->GetHostBufferPtr(0, local_slot);
uint8_t* d_base = reinterpret_cast<uint8_t*>(out->untyped_data());
Expand Down Expand Up @@ -382,6 +411,7 @@ XLA_FFI_DEFINE_HANDLER(
xla::ffi::Ffi::Bind()
.Arg<xla::ffi::AnyBuffer>() // anchor JAX input array (Arg 0)
.Arg<xla::ffi::AnyBuffer>() // shard_idx JAX input array (Arg 1)
.Arg<xla::ffi::AnyBuffer>() // logical_shard_idx
.Attr<int64_t>("slice_byte_size")
.Attr<int32_t>("local_port")
.Attr<int32_t>("parallelism")
Expand All @@ -394,6 +424,7 @@ XLA_FFI_DEFINE_HANDLER(
xla::ffi::Ffi::Bind()
.Arg<xla::ffi::AnyBuffer>() // anchor
.Arg<xla::ffi::AnyBuffer>() // shard_idx_buf
.Arg<xla::ffi::AnyBuffer>() // logical_shard_idx
.Ret<xla::ffi::AnyBuffer>() // result (aliased to anchor)
);

Expand Down
34 changes: 26 additions & 8 deletions tpu_raiden/frameworks/jax/weight_synchronizer_ffi.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,16 @@
from tpu_raiden.frameworks.jax import _weight_synchronizer_ffi


def make_logical_shard_idx(mesh):
"""Canonical logical shard index per device (mesh row-major == controller
itertools.product order). Build once at setup, pass to the calls below to keep
it off the hot path."""
index_spec = jax.sharding.PartitionSpec(*mesh.axis_names)
return jax.device_put(
jnp.arange(mesh.size, dtype=jnp.int32).reshape(mesh.devices.shape),
jax.sharding.NamedSharding(mesh, index_spec))


def init_weight_synchronizer(
device_array,
shard_idx,
Expand Down Expand Up @@ -95,6 +105,7 @@ def init_weight_synchronizer_and_d2h(
parallelism: int = 1,
num_layers: int = 1,
listener_port: int = -1,
logical_shard_idx=None,
) -> jax.Array:
"""Registers and executes init_weight_synchronizer_and_d2h FFI custom call on each device rank.

Expand All @@ -117,7 +128,7 @@ def init_weight_synchronizer_and_d2h(
"""

@compute_on.compute_on("device_host")
def _local_init_and_d2h(anchor, s_idx):
def _local_init_and_d2h(anchor, s_idx, l_idx):
axis_names = mesh.axis_names
out_dim = 6 if listener_port >= 0 else 5
out_shape = tuple([1] * len(axis_names)) + (out_dim,)
Expand All @@ -128,13 +139,17 @@ def _local_init_and_d2h(anchor, s_idx):
)(
anchor,
s_idx,
l_idx,
slice_byte_size=slice_byte_size,
local_port=np.int32(local_port),
parallelism=np.int32(parallelism),
num_layers=np.int32(num_layers),
listener_port=np.int32(listener_port),
)

if logical_shard_idx is None:
logical_shard_idx = make_logical_shard_idx(mesh)

axis_names = mesh.axis_names
anchor_spec = device_array.sharding.spec
index_spec = jax.sharding.PartitionSpec(*axis_names)
Expand All @@ -143,9 +158,9 @@ def _local_init_and_d2h(anchor, s_idx):
return jax.shard_map(
_local_init_and_d2h,
mesh=mesh,
in_specs=(anchor_spec, index_spec),
in_specs=(anchor_spec, index_spec, index_spec),
out_specs=out_spec,
)(device_array, shard_idx)
)(device_array, shard_idx, logical_shard_idx)



Expand All @@ -170,7 +185,7 @@ def is_listener_active(shard_idx: int = 0) -> bool:
return _weight_synchronizer_ffi.is_listener_active(shard_idx)


def h2d(device_array, shard_idx, mesh) -> jax.Array:
def h2d(device_array, shard_idx, mesh, logical_shard_idx=None) -> jax.Array:
"""Executes asynchronous Host-to-Device (H2D) copy from local staging buffer directly onto device memory via FFI.

Args:
Expand All @@ -185,12 +200,15 @@ def h2d(device_array, shard_idx, mesh) -> jax.Array:
"""

@compute_on.compute_on("device_host")
def _local_h2d(anchor, s_idx):
def _local_h2d(anchor, s_idx, l_idx):
return jax.ffi.ffi_call(
"ws_h2d",
jax.ShapeDtypeStruct(anchor.shape, anchor.dtype),
has_side_effect=True,
)(anchor, s_idx)
)(anchor, s_idx, l_idx)

if logical_shard_idx is None:
logical_shard_idx = make_logical_shard_idx(mesh)

axis_names = mesh.axis_names
anchor_spec = device_array.sharding.spec
Expand All @@ -199,6 +217,6 @@ def _local_h2d(anchor, s_idx):
return jax.shard_map(
_local_h2d,
mesh=mesh,
in_specs=(anchor_spec, index_spec),
in_specs=(anchor_spec, index_spec, index_spec),
out_specs=anchor_spec,
)(device_array, shard_idx)
)(device_array, shard_idx, logical_shard_idx)
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@ xla::ffi::Error TriggerWeightSynchronizerInitImpl(

xla::ffi::Error TriggerWeightSynchronizerInitAndD2hImpl(
xla::ffi::AnyBuffer x, xla::ffi::AnyBuffer shard_idx_buf,
xla::ffi::AnyBuffer logical_idx_buf,
int64_t slice_byte_size, int32_t local_port, int32_t parallelism,
int32_t num_layers, int32_t listener_port,
xla::ffi::Result<xla::ffi::AnyBuffer> out);
Expand Down
2 changes: 1 addition & 1 deletion tpu_raiden/frameworks/jax/weight_synchronizer_ffi_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,7 @@ class WeightSynchronizerFfiParamTest : public WeightSynchronizerFfiTest,
num_layers, listener_port, out);
} else {
return TriggerWeightSynchronizerInitAndD2hImpl(
x, shard_idx_buf, slice_byte_size, local_port, parallelism,
x, shard_idx_buf, shard_idx_buf, slice_byte_size, local_port, parallelism,
num_layers, listener_port, out);
}
}
Expand Down
Loading