diff --git a/tpu_raiden/frameworks/jax/weight_synchronizer_ffi.cc b/tpu_raiden/frameworks/jax/weight_synchronizer_ffi.cc index 7399984e..5845647a 100644 --- a/tpu_raiden/frameworks/jax/weight_synchronizer_ffi.cc +++ b/tpu_raiden/frameworks/jax/weight_synchronizer_ffi.cc @@ -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 out) { @@ -200,6 +201,12 @@ xla::ffi::Error TriggerWeightSynchronizerInitAndD2hImpl( } int32_t shard_idx = *reinterpret_cast(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(logical_idx_buf.untyped_data()); if (shard_idx < 0 || shard_idx >= 32) { return xla::ffi::Error( xla::ffi::ErrorCode::kInvalidArgument, @@ -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(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(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(logical_idx); + VLOG(1) << "[D2H] device_ordinal=" << shard_idx + << " logical=" << logical_idx << " local_slot=" << local_slot; uint8_t* dst_host_ptr = const_cast( g_weight_synchronizers[shard_idx]->GetHostBufferPtr(0, local_slot)); const uint8_t* src_device_ptr = @@ -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 out) { (void)anchor; if (shard_idx_buf.untyped_data() == nullptr) { @@ -331,6 +346,12 @@ xla::ffi::Error TriggerH2DImpl(xla::ffi::AnyBuffer anchor, } int32_t shard_idx = *reinterpret_cast(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(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, @@ -338,8 +359,16 @@ xla::ffi::Error TriggerH2DImpl(xla::ffi::AnyBuffer anchor, } size_t size = g_weight_synchronizers[shard_idx]->slice_byte_size(); - size_t local_slot = static_cast(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(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(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(out->untyped_data()); @@ -382,6 +411,7 @@ XLA_FFI_DEFINE_HANDLER( xla::ffi::Ffi::Bind() .Arg() // anchor JAX input array (Arg 0) .Arg() // shard_idx JAX input array (Arg 1) + .Arg() // logical_shard_idx .Attr("slice_byte_size") .Attr("local_port") .Attr("parallelism") @@ -394,6 +424,7 @@ XLA_FFI_DEFINE_HANDLER( xla::ffi::Ffi::Bind() .Arg() // anchor .Arg() // shard_idx_buf + .Arg() // logical_shard_idx .Ret() // result (aliased to anchor) ); diff --git a/tpu_raiden/frameworks/jax/weight_synchronizer_ffi.py b/tpu_raiden/frameworks/jax/weight_synchronizer_ffi.py index 3e1effd2..e3b08a0f 100644 --- a/tpu_raiden/frameworks/jax/weight_synchronizer_ffi.py +++ b/tpu_raiden/frameworks/jax/weight_synchronizer_ffi.py @@ -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, @@ -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. @@ -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,) @@ -128,6 +139,7 @@ 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), @@ -135,6 +147,9 @@ def _local_init_and_d2h(anchor, s_idx): 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) @@ -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) @@ -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: @@ -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 @@ -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) diff --git a/tpu_raiden/frameworks/jax/weight_synchronizer_ffi_internal.h b/tpu_raiden/frameworks/jax/weight_synchronizer_ffi_internal.h index 50f0f498..bf0f6fa3 100644 --- a/tpu_raiden/frameworks/jax/weight_synchronizer_ffi_internal.h +++ b/tpu_raiden/frameworks/jax/weight_synchronizer_ffi_internal.h @@ -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 out); diff --git a/tpu_raiden/frameworks/jax/weight_synchronizer_ffi_test.cc b/tpu_raiden/frameworks/jax/weight_synchronizer_ffi_test.cc index 43406da7..84ee2076 100644 --- a/tpu_raiden/frameworks/jax/weight_synchronizer_ffi_test.cc +++ b/tpu_raiden/frameworks/jax/weight_synchronizer_ffi_test.cc @@ -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); } }