diff --git a/tpu_raiden/core/controller/raiden_controller.cc b/tpu_raiden/core/controller/raiden_controller.cc index 671169a2..f532d9b4 100644 --- a/tpu_raiden/core/controller/raiden_controller.cc +++ b/tpu_raiden/core/controller/raiden_controller.cc @@ -377,6 +377,7 @@ absl::Status RaidenController::DeallocateBuffers( absl::StatusOr RaidenController::BuildTransferBuffersRequest( absl::Span src_buffers, absl::Span dst_buffers, + absl::Span staging_host_buffers, absl::Span copy_sizes) { if (src_buffers.empty() || src_buffers.size() != dst_buffers.size()) { return absl::InvalidArgumentError( @@ -416,6 +417,15 @@ RaidenController::BuildTransferBuffersRequest( for (int64_t size : copy_sizes) { transfer->add_copy_sizes(size); } + for (const auto& buf : staging_host_buffers) { + if (buf.index() < 0) { + return absl::InvalidArgumentError(absl::StrCat( + "Staging host buffer has invalid negative index: ", buf.index())); + } + auto* added_buf = transfer->add_staging_host_buffers(); + *added_buf = buf.ToProto(); + added_buf->set_index(buf.index()); + } return request; } @@ -423,9 +433,10 @@ RaidenController::BuildTransferBuffersRequest( tsl::Future<> RaidenController::TransferBuffers( absl::string_view worker_id, absl::Span src_buffers, absl::Span dst_buffers, + absl::Span staging_host_buffers, absl::Span copy_sizes) { - auto request_or = - BuildTransferBuffersRequest(src_buffers, dst_buffers, copy_sizes); + auto request_or = BuildTransferBuffersRequest( + src_buffers, dst_buffers, staging_host_buffers, copy_sizes); if (!request_or.ok()) { return tsl::Future<>(request_or.status()); } @@ -447,6 +458,7 @@ tsl::Future<> RaidenController::TransferBuffers( tsl::Future<> RaidenController::TransferBuffers( absl::Span src_buffers, absl::Span dst_buffers, + absl::Span staging_host_buffers, absl::Span copy_sizes) { if (src_buffers.empty() || src_buffers.size() != dst_buffers.size()) { return tsl::Future<>(absl::InvalidArgumentError( @@ -523,8 +535,10 @@ tsl::Future<> RaidenController::TransferBuffers( if (worker_src.empty()) continue; - auto req_or = - BuildTransferBuffersRequest(worker_src, worker_dst, worker_copy_sizes); + // Every worker owns a shard of every block, so the (host) staging offsets + // are identical across workers. + auto req_or = BuildTransferBuffersRequest( + worker_src, worker_dst, staging_host_buffers, worker_copy_sizes); if (!req_or.ok()) { return tsl::Future<>(req_or.status()); } diff --git a/tpu_raiden/core/controller/raiden_controller.h b/tpu_raiden/core/controller/raiden_controller.h index ab5a0d39..a011d927 100644 --- a/tpu_raiden/core/controller/raiden_controller.h +++ b/tpu_raiden/core/controller/raiden_controller.h @@ -118,16 +118,23 @@ class RaidenController { // is performed. absl::Status AllocateTargetBlockIds(absl::Span block_ids); - // Targeted worker transfer - tsl::Future<> TransferBuffers(absl::string_view worker_id, - absl::Span src_buffers, - absl::Span dst_buffers, - absl::Span copy_sizes = {}); + // Targeted worker transfer. + // staging_host_buffers: host DRAM staging (bridge) block offsets, required + // for 2-stage remote transfers (remote H2D read/write, remote D2H write); + // always the middle hop of the data flow. Unused for local transfers. + tsl::Future<> TransferBuffers( + absl::string_view worker_id, absl::Span src_buffers, + absl::Span dst_buffers, + absl::Span staging_host_buffers = {}, + absl::Span copy_sizes = {}); - // Broadcast transfer to all registered workers - tsl::Future<> TransferBuffers(absl::Span src_buffers, - absl::Span dst_buffers, - absl::Span copy_sizes = {}); + // Broadcast transfer to all registered workers (staging_host_buffers as + // above). + tsl::Future<> TransferBuffers( + absl::Span src_buffers, + absl::Span dst_buffers, + absl::Span staging_host_buffers = {}, + absl::Span copy_sizes = {}); // Initiates remote read from source controller. block_hashes (parallel to the // block ids) let the source verify/pin the blocks in its LRU before transfer. @@ -169,7 +176,8 @@ class RaidenController { absl::StatusOr BuildTransferBuffersRequest( absl::Span src_buffers, absl::Span dst_buffers, - absl::Span copy_sizes); + absl::Span staging_host_buffers = {}, + absl::Span copy_sizes = {}); void Init(absl::Span worker_addresses, absl::string_view raiden_orchestrator_address, diff --git a/tpu_raiden/core/controller/raiden_controller_test.cc b/tpu_raiden/core/controller/raiden_controller_test.cc index dce779ca..508ed181 100644 --- a/tpu_raiden/core/controller/raiden_controller_test.cc +++ b/tpu_raiden/core/controller/raiden_controller_test.cc @@ -326,7 +326,7 @@ TEST_F(RaidenControllerTest, TransferBuffersValidationMismatchedCopySizes) { auto status = controller .TransferBuffers({src_buf1, src_buf2}, {dst_buf1, dst_buf2}, - copy_sizes) + /*staging_host_buffers=*/{}, copy_sizes) .Await(); EXPECT_FALSE(status.ok()); EXPECT_EQ(status.code(), absl::StatusCode::kInvalidArgument); @@ -431,7 +431,8 @@ TEST_F(RaidenControllerTest, TransferBuffersD2HSuccess) { auto status = controller .TransferBuffers("worker_0", {src_buf1, src_buf2}, - {dst_buf1, dst_buf2}, copy_sizes) + {dst_buf1, dst_buf2}, + /*staging_host_buffers=*/{}, copy_sizes) .Await(); ASSERT_TRUE(status.ok()); EXPECT_EQ(mock_mgr.d2h_calls, 1); diff --git a/tpu_raiden/core/controller/test_util.h b/tpu_raiden/core/controller/test_util.h index 9a52cb8e..9be061e4 100644 --- a/tpu_raiden/core/controller/test_util.h +++ b/tpu_raiden/core/controller/test_util.h @@ -53,6 +53,7 @@ struct MockTransferManager { std::string last_peer; std::vector last_src_offsets; std::vector last_dst_offsets; + std::vector last_staging_offsets; std::vector last_copy_sizes; absl::StatusOr D2h( @@ -67,13 +68,15 @@ struct MockTransferManager { } absl::StatusOr D2hWrite( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_device_offsets, + const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, const std::vector& copy_sizes) { d2h_write_calls++; last_peer = std::string(peer); - last_src_offsets = src_offsets; - last_dst_offsets = dst_offsets; + last_src_offsets = src_device_offsets; + last_staging_offsets = src_host_offsets; + last_dst_offsets = dst_host_offsets; last_copy_sizes = copy_sizes; return raiden::PjRtCopyFuture(); } @@ -102,25 +105,29 @@ struct MockTransferManager { } absl::StatusOr H2dWrite( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, + const std::vector& dst_device_offsets, const std::vector& copy_sizes) { h2d_write_calls++; last_peer = std::string(peer); - last_src_offsets = src_offsets; - last_dst_offsets = dst_offsets; + last_src_offsets = src_host_offsets; + last_staging_offsets = dst_host_offsets; + last_dst_offsets = dst_device_offsets; last_copy_sizes = copy_sizes; return raiden::PjRtCopyFuture(); } absl::StatusOr H2dRead( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, + const std::vector& dst_device_offsets, const std::vector& copy_sizes) { h2d_read_calls++; last_peer = std::string(peer); - last_src_offsets = src_offsets; - last_dst_offsets = dst_offsets; + last_src_offsets = src_host_offsets; + last_staging_offsets = dst_host_offsets; + last_dst_offsets = dst_device_offsets; last_copy_sizes = copy_sizes; return raiden::PjRtCopyFuture(); } diff --git a/tpu_raiden/core/controller/worker_service_impl.cc b/tpu_raiden/core/controller/worker_service_impl.cc index 1c9a6995..a3fa698d 100644 --- a/tpu_raiden/core/controller/worker_service_impl.cc +++ b/tpu_raiden/core/controller/worker_service_impl.cc @@ -220,6 +220,15 @@ grpc::Status WorkerServiceImpl::TransferBuffers( copy_sizes.assign(src_offsets.size(), 1); } + // Host staging (bridge) offsets for 2-stage remote transfers -- the middle + // hop of the data flow. Required for remote H2D read/write and remote D2H + // write; the transfer manager rejects those transfers when it is missing. + std::vector staging_host_offsets; + staging_host_offsets.reserve(transfer.staging_host_buffers_size()); + for (const auto& buf : transfer.staging_host_buffers()) { + staging_host_offsets.push_back(buf.index()); + } + std::vector dst_remote_descriptors; if (transfer.dst_buffers_size() > 0 && transfer.dst_buffers(0).remote_descriptors_size() > 0) { @@ -256,8 +265,9 @@ grpc::Status WorkerServiceImpl::TransferBuffers( } else if (transfer.dst_buffers_size() > 0 && !transfer.dst_buffers(0).remote_address().empty()) { std::string dst_peer = transfer.dst_buffers(0).remote_address(); - future_or = transfer_manager_.D2hWrite(dst_peer, src_offsets, dst_offsets, - copy_sizes); + // Flow order: local device src -> local host staging -> remote host dst. + future_or = transfer_manager_.D2hWrite( + dst_peer, src_offsets, staging_host_offsets, dst_offsets, copy_sizes); } else { future_or = transfer_manager_.D2h(src_offsets, dst_offsets, copy_sizes); } @@ -265,13 +275,15 @@ grpc::Status WorkerServiceImpl::TransferBuffers( if (transfer.src_buffers_size() > 0 && !transfer.src_buffers(0).remote_address().empty()) { std::string src_peer = transfer.src_buffers(0).remote_address(); - future_or = transfer_manager_.H2dRead(src_peer, src_offsets, dst_offsets, - copy_sizes); + // Flow order: remote host src -> local host staging -> local device dst. + future_or = transfer_manager_.H2dRead( + src_peer, src_offsets, staging_host_offsets, dst_offsets, copy_sizes); } else if (transfer.dst_buffers_size() > 0 && !transfer.dst_buffers(0).remote_address().empty()) { std::string dst_peer = transfer.dst_buffers(0).remote_address(); - future_or = transfer_manager_.H2dWrite(dst_peer, src_offsets, dst_offsets, - copy_sizes); + // Flow order: local host src -> remote host staging -> remote device dst. + future_or = transfer_manager_.H2dWrite( + dst_peer, src_offsets, staging_host_offsets, dst_offsets, copy_sizes); } else { future_or = transfer_manager_.H2d(src_offsets, dst_offsets, copy_sizes); } diff --git a/tpu_raiden/core/controller/worker_service_test.cc b/tpu_raiden/core/controller/worker_service_test.cc index ac42c626..98c29d39 100644 --- a/tpu_raiden/core/controller/worker_service_test.cc +++ b/tpu_raiden/core/controller/worker_service_test.cc @@ -132,7 +132,7 @@ TEST_F(WorkerServiceTest, TransferBuffersH2hSuccess) { transfer->add_dst_buffers()->set_remote_address("localhost:8080"); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 0); EXPECT_EQ(mock_mgr.h2d_calls, 0); EXPECT_EQ(mock_mgr.h2h_calls, 1); @@ -263,7 +263,7 @@ TEST_F(WorkerServiceTest, TransferBuffersD2HSuccess) { transfer->add_copy_sizes(2); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 1); EXPECT_EQ(mock_mgr.h2d_calls, 0); EXPECT_THAT(mock_mgr.last_src_offsets, ElementsAre(10, 30)); @@ -283,13 +283,17 @@ TEST_F(WorkerServiceTest, TransferBuffersRemoteD2hWithPeerSuccess) { transfer->add_dst_offsets(200); transfer->add_dst_buffers()->set_remote_address("remote_host:1234"); + auto* staging_buf = transfer->add_staging_host_buffers(); + staging_buf->set_index(300); + auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 0); EXPECT_EQ(mock_mgr.d2h_write_calls, 1); EXPECT_EQ(mock_mgr.h2d_calls, 0); EXPECT_EQ(mock_mgr.last_peer, "remote_host:1234"); EXPECT_THAT(mock_mgr.last_src_offsets, ElementsAre(100)); + EXPECT_THAT(mock_mgr.last_staging_offsets, ElementsAre(300)); EXPECT_THAT(mock_mgr.last_dst_offsets, ElementsAre(200)); EXPECT_THAT(mock_mgr.last_copy_sizes, ElementsAre(1)); } @@ -311,7 +315,7 @@ TEST_F(WorkerServiceTest, dst_buf->set_remote_address("remote_host:5678"); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 0); EXPECT_EQ(mock_mgr.d2h_write_calls, 1); EXPECT_EQ(mock_mgr.h2d_calls, 0); @@ -338,7 +342,7 @@ TEST_F(WorkerServiceTest, dst_buf->set_index(200); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 0); EXPECT_EQ(mock_mgr.d2h_write_calls, 0); EXPECT_EQ(mock_mgr.d2h_read_calls, 1); @@ -361,7 +365,7 @@ TEST_F(WorkerServiceTest, TransferBuffersLocalD2hFallbackSuccess) { transfer->add_dst_offsets(200); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 1); EXPECT_EQ(mock_mgr.d2h_write_calls, 0); EXPECT_THAT(mock_mgr.last_src_offsets, ElementsAre(100)); @@ -381,7 +385,7 @@ TEST_F(WorkerServiceTest, TransferBuffersH2DSuccess) { transfer->add_dst_offsets(200); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 0); EXPECT_EQ(mock_mgr.h2d_calls, 1); EXPECT_THAT(mock_mgr.last_src_offsets, ElementsAre(100)); @@ -400,14 +404,16 @@ TEST_F(WorkerServiceTest, TransferBuffersRemoteH2dWithPeerSuccess) { transfer->add_src_offsets(100); transfer->add_dst_offsets(200); transfer->add_dst_buffers()->set_remote_address("remote_host:1234"); + transfer->add_staging_host_buffers()->set_index(300); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 0); EXPECT_EQ(mock_mgr.h2d_calls, 0); EXPECT_EQ(mock_mgr.h2d_write_calls, 1); EXPECT_EQ(mock_mgr.last_peer, "remote_host:1234"); EXPECT_THAT(mock_mgr.last_src_offsets, ElementsAre(100)); + EXPECT_THAT(mock_mgr.last_staging_offsets, ElementsAre(300)); EXPECT_THAT(mock_mgr.last_dst_offsets, ElementsAre(200)); EXPECT_THAT(mock_mgr.last_copy_sizes, ElementsAre(1)); } @@ -428,13 +434,17 @@ TEST_F(WorkerServiceTest, dst_buf->set_index(200); dst_buf->set_remote_address("remote_host:5678"); + auto* staging_buf = transfer->add_staging_host_buffers(); + staging_buf->set_index(300); + auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 0); EXPECT_EQ(mock_mgr.h2d_calls, 0); EXPECT_EQ(mock_mgr.h2d_write_calls, 1); EXPECT_EQ(mock_mgr.last_peer, "remote_host:5678"); EXPECT_THAT(mock_mgr.last_src_offsets, ElementsAre(100)); + EXPECT_THAT(mock_mgr.last_staging_offsets, ElementsAre(300)); EXPECT_THAT(mock_mgr.last_dst_offsets, ElementsAre(200)); EXPECT_THAT(mock_mgr.last_copy_sizes, ElementsAre(1)); } @@ -455,14 +465,18 @@ TEST_F(WorkerServiceTest, auto* dst_buf = transfer->add_dst_buffers(); dst_buf->set_index(200); + auto* staging_buf = transfer->add_staging_host_buffers(); + staging_buf->set_index(300); + auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 0); EXPECT_EQ(mock_mgr.h2d_calls, 0); EXPECT_EQ(mock_mgr.h2d_write_calls, 0); EXPECT_EQ(mock_mgr.h2d_read_calls, 1); EXPECT_EQ(mock_mgr.last_peer, "remote_host:5678"); EXPECT_THAT(mock_mgr.last_src_offsets, ElementsAre(100)); + EXPECT_THAT(mock_mgr.last_staging_offsets, ElementsAre(300)); EXPECT_THAT(mock_mgr.last_dst_offsets, ElementsAre(200)); EXPECT_THAT(mock_mgr.last_copy_sizes, ElementsAre(1)); } @@ -479,7 +493,7 @@ TEST_F(WorkerServiceTest, TransferBuffersLocalH2dFallbackSuccess) { transfer->add_dst_offsets(200); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 0); EXPECT_EQ(mock_mgr.h2d_calls, 1); EXPECT_EQ(mock_mgr.h2d_write_calls, 0); @@ -502,7 +516,7 @@ TEST_F(WorkerServiceTest, TransferBuffersWithBufferProtosSuccess) { dst_buf->set_index(20); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.d2h_calls, 1); EXPECT_THAT(mock_mgr.last_src_offsets, ElementsAre(10)); EXPECT_THAT(mock_mgr.last_dst_offsets, ElementsAre(20)); @@ -524,7 +538,7 @@ TEST_F(WorkerServiceTest, dst_buf->set_remote_address("localhost:8080"); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.h2h_calls, 1); EXPECT_EQ(mock_mgr.h2h_read_calls, 0); EXPECT_EQ(mock_mgr.h2h_write_calls, 1); @@ -548,7 +562,7 @@ TEST_F(WorkerServiceTest, TransferBuffersH2hReadRemoteSrcSuccess) { dst_buf->set_memory_type(rpc::MEMORY_TYPE_DRAM); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); - ASSERT_TRUE(status.ok()); + ASSERT_TRUE(status.ok()) << status.message(); EXPECT_EQ(mock_mgr.h2h_calls, 1); EXPECT_EQ(mock_mgr.h2h_read_calls, 1); EXPECT_EQ(mock_mgr.h2h_write_calls, 0); diff --git a/tpu_raiden/core/kv_cache_manager_with_transfer.cc b/tpu_raiden/core/kv_cache_manager_with_transfer.cc index 679d24ae..62de7d73 100644 --- a/tpu_raiden/core/kv_cache_manager_with_transfer.cc +++ b/tpu_raiden/core/kv_cache_manager_with_transfer.cc @@ -2284,7 +2284,7 @@ absl::Status KVCacheManagerWithTransfer::OnBlocksReceived( if (!found) { // Forward to base class for direct pull operations - return RaidenManagerBase::OnBlocksReceived(block_ids, uuid); + return KVCacheManagerBase::OnBlocksReceived(block_ids, uuid); } { @@ -2400,7 +2400,9 @@ absl::Status KVCacheManagerWithTransfer::OnLayerReceived(size_t layer_idx, absl::MutexLock lock(mu_); auto it = active_recv_entries_.find(uuid); if (it == active_recv_entries_.end()) { - return absl::OkStatus(); + // Not a disagg recv entry: hand to the base (in-band op-7 H2D plans + // live in KVCacheManagerBase). + return KVCacheManagerBase::OnLayerReceived(layer_idx, uuid); } auto& entry = it->second; h2d_copy = entry.h2d_copy; diff --git a/tpu_raiden/core/kv_manager_holder.h b/tpu_raiden/core/kv_manager_holder.h index ac75dc48..c37779be 100644 --- a/tpu_raiden/core/kv_manager_holder.h +++ b/tpu_raiden/core/kv_manager_holder.h @@ -21,6 +21,7 @@ #include #include #include +#include #include #include "absl/status/status.h" @@ -43,6 +44,7 @@ struct has_h2d_write().H2dWrite( std::declval(), std::declval&>(), std::declval&>(), + std::declval&>(), std::declval&>()))>> : std::true_type {}; @@ -57,6 +59,7 @@ struct has_h2d_read().H2dRead( std::declval(), std::declval&>(), std::declval&>(), + std::declval&>(), std::declval&>()))>> : std::true_type {}; @@ -71,6 +74,7 @@ struct has_d2h_write().D2hWrite( std::declval(), std::declval&>(), std::declval&>(), + std::declval&>(), std::declval&>()))>> : std::true_type {}; @@ -148,16 +152,19 @@ class KVManagerHolder { const std::vector& src_offsets, const std::vector& dst_offsets) = 0; virtual absl::StatusOr H2dWrite( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, + const std::vector& dst_device_offsets, const std::vector& copy_sizes) = 0; virtual absl::StatusOr H2dRead( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, + const std::vector& dst_device_offsets, const std::vector& copy_sizes) = 0; virtual absl::StatusOr D2hWrite( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_device_offsets, + const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, const std::vector& copy_sizes) = 0; virtual absl::StatusOr D2hRead( absl::string_view peer, const std::vector& src_offsets, @@ -233,33 +240,39 @@ class KVManagerHolder { } } absl::StatusOr H2dWrite( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, + const std::vector& dst_device_offsets, const std::vector& copy_sizes) override { if constexpr (internal::has_h2d_write_v) { - return impl_->H2dWrite(peer, src_offsets, dst_offsets, copy_sizes); + return impl_->H2dWrite(peer, src_host_offsets, dst_host_offsets, + dst_device_offsets, copy_sizes); } else { return absl::UnimplementedError( "H2dWrite is not implemented by the underlying transfer manager."); } } absl::StatusOr H2dRead( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, + const std::vector& dst_device_offsets, const std::vector& copy_sizes) override { if constexpr (internal::has_h2d_read_v) { - return impl_->H2dRead(peer, src_offsets, dst_offsets, copy_sizes); + return impl_->H2dRead(peer, src_host_offsets, dst_host_offsets, + dst_device_offsets, copy_sizes); } else { return absl::UnimplementedError( "H2dRead is not implemented by the underlying transfer manager."); } } absl::StatusOr D2hWrite( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_device_offsets, + const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, const std::vector& copy_sizes) override { if constexpr (internal::has_d2h_write_v) { - return impl_->D2hWrite(peer, src_offsets, dst_offsets, copy_sizes); + return impl_->D2hWrite(peer, src_device_offsets, src_host_offsets, + dst_host_offsets, copy_sizes); } else { return absl::UnimplementedError( "D2hWrite is not implemented by the underlying transfer manager."); @@ -362,33 +375,39 @@ class KVManagerHolder { } absl::StatusOr H2dWrite( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, + const std::vector& dst_device_offsets, const std::vector& copy_sizes) const { if (!self_) { return absl::InternalError("KVManagerHolder is null"); } - return self_->H2dWrite(peer, src_offsets, dst_offsets, copy_sizes); + return self_->H2dWrite(peer, src_host_offsets, dst_host_offsets, + dst_device_offsets, copy_sizes); } absl::StatusOr H2dRead( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, + const std::vector& dst_device_offsets, const std::vector& copy_sizes) const { if (!self_) { return absl::InternalError("KVManagerHolder is null"); } - return self_->H2dRead(peer, src_offsets, dst_offsets, copy_sizes); + return self_->H2dRead(peer, src_host_offsets, dst_host_offsets, + dst_device_offsets, copy_sizes); } absl::StatusOr D2hWrite( - absl::string_view peer, const std::vector& src_offsets, - const std::vector& dst_offsets, + absl::string_view peer, const std::vector& src_device_offsets, + const std::vector& src_host_offsets, + const std::vector& dst_host_offsets, const std::vector& copy_sizes) const { if (!self_) { return absl::InternalError("KVManagerHolder is null"); } - return self_->D2hWrite(peer, src_offsets, dst_offsets, copy_sizes); + return self_->D2hWrite(peer, src_device_offsets, src_host_offsets, + dst_host_offsets, copy_sizes); } absl::StatusOr D2hRead( diff --git a/tpu_raiden/core/raiden_manager_base.cc b/tpu_raiden/core/raiden_manager_base.cc index 5481df5e..8380689e 100644 --- a/tpu_raiden/core/raiden_manager_base.cc +++ b/tpu_raiden/core/raiden_manager_base.cc @@ -245,7 +245,8 @@ void RaidenManagerBase::SetExternalHostPointers( absl::StatusOr> RaidenManagerBase::H2hWriteDirect( const std::vector& peers, const std::vector& src_block_ids, - const std::vector& dst_block_ids, uint64_t uuid, int layer_idx) { + const std::vector& dst_block_ids, uint64_t uuid, int layer_idx, + const std::vector& dst_device_block_ids) { InitTransportServer(); absl::MutexLock lock(server_init_mu_); if (!server_) { @@ -253,14 +254,15 @@ absl::StatusOr> RaidenManagerBase::H2hWriteDirect( } return server_->SyncPush(peers, src_block_ids, dst_block_ids, parallelism_, tpu_raiden::transport::MajorOrder::kLayerMajor, uuid, - layer_idx); + layer_idx, dst_device_block_ids); } void RaidenManagerBase::H2hWriteDirectAsync( const std::vector& peers, const std::vector& src_block_ids, const std::vector& dst_block_ids, uint64_t uuid, int layer_idx, - std::function>)> on_complete) { + std::function>)> on_complete, + const std::vector& dst_device_block_ids) { InitTransportServer(); absl::MutexLock lock(server_init_mu_); if (!server_) { @@ -270,7 +272,7 @@ void RaidenManagerBase::H2hWriteDirectAsync( } server_->AsyncPush(peers, src_block_ids, dst_block_ids, parallelism_, tpu_raiden::transport::MajorOrder::kLayerMajor, uuid, - layer_idx, std::move(on_complete)); + layer_idx, std::move(on_complete), dst_device_block_ids); } absl::StatusOr> RaidenManagerBase::H2hReadDirect( diff --git a/tpu_raiden/core/raiden_manager_base.h b/tpu_raiden/core/raiden_manager_base.h index a5d5b311..5d4d316d 100644 --- a/tpu_raiden/core/raiden_manager_base.h +++ b/tpu_raiden/core/raiden_manager_base.h @@ -44,18 +44,57 @@ class RaidenManagerBase : public tpu_raiden::transport::BlockTransportDelegate { ~RaidenManagerBase() override; - // Direct C++ H2H network write (Push) + /** + * @brief Direct C++ H2H network write (Push). + * + * Synchronously pushes host-memory blocks to a remote endpoint. + * If `dst_device_block_ids` is non-empty, the request escalates to an in-band + * pipelined H2D transfer (op=7) that will additionally instruct the remote + * peer to automatically execute an H2D memory transfer. + * + * @param peers The remote endpoints to push to. + * @param src_block_ids The source block identifiers on the local host. + * @param dst_block_ids The destination staging block identifiers on the + * remote host. + * @param uuid Uniquely identifies the transfer, used for acknowledging the + * correct session. + * @param layer_idx Overriding layer index offset. + * @param dst_device_block_ids Optional device HBM identifiers. When provided, + * signals the remote to initiate HBM DMA immediately following successful + * receipt. + * @return A status or vector of successfully pushed blocks. + */ absl::StatusOr> H2hWriteDirect( const std::vector& peers, const std::vector& src_block_ids, const std::vector& dst_block_ids = {}, uint64_t uuid = 0, - int layer_idx = -1); - - void H2hWriteDirectAsync( - const std::vector& peers, - const std::vector& src_block_ids, - const std::vector& dst_block_ids, uint64_t uuid, int layer_idx, - std::function>)> on_complete); + int layer_idx = -1, const std::vector& dst_device_block_ids = {}); + + /** + * @brief Asynchronous C++ H2H network write (Push). + * + * Asynchronously pushes host-memory blocks to a remote endpoint. + * Shares semantics with H2hWriteDirect, escalating to op=7 if + * `dst_device_block_ids` is present. + * + * @param peers The remote endpoints to push to. + * @param src_block_ids The source block identifiers on the local host. + * @param dst_block_ids The destination staging block identifiers on the + * remote host. + * @param uuid Uniquely identifies the transfer. + * @param layer_idx Overriding layer index offset. + * @param on_complete Callback fired upon completion or failure of the push + * operation. + * @param dst_device_block_ids Optional device HBM identifiers. Escalate to + * pipelined H2D transfer if present. + */ + void H2hWriteDirectAsync(const std::vector& peers, + const std::vector& src_block_ids, + const std::vector& dst_block_ids = {}, + uint64_t uuid = 0, int layer_idx = -1, + std::function>)> + on_complete = nullptr, + const std::vector& dst_device_block_ids = {}); // Direct C++ H2H network read (Pull) absl::StatusOr> H2hReadDirect( diff --git a/tpu_raiden/core/raw_transfer_core.h b/tpu_raiden/core/raw_transfer_core.h index 3e6fa053..90108455 100644 --- a/tpu_raiden/core/raw_transfer_core.h +++ b/tpu_raiden/core/raw_transfer_core.h @@ -509,6 +509,15 @@ struct PjRtCopyFuture { tsl::AsyncValue* av = future.async_value(); if (av->IsError()) { status = av->GetError(); + } else { + // A stateless xla::Future<> fulfilled via Promise::Set(error_status) + // stores the status as the future's VALUE (the async value itself is + // concrete, not in error state) -- extract it or the error is + // silently dropped. + absl::Status value_status = future.Await(); + if (!value_status.ok()) { + status = value_status; + } } } return status; diff --git a/tpu_raiden/frameworks/jax/kv_cache_manager_wrapper_test.cc b/tpu_raiden/frameworks/jax/kv_cache_manager_wrapper_test.cc index 91bc7341..b9bb8323 100644 --- a/tpu_raiden/frameworks/jax/kv_cache_manager_wrapper_test.cc +++ b/tpu_raiden/frameworks/jax/kv_cache_manager_wrapper_test.cc @@ -137,7 +137,8 @@ class MockSubManager : public KVCacheManagerWithTransfer { absl::StatusOr, raiden::PjRtCopyFuture>> H2hWrite( std::string peer, const std::vector& src_block_ids, const std::vector& dst_block_ids = {}, uint64_t uuid = 0, - int layer_idx = -1) override { + int layer_idx = -1, + const std::vector& dst_device_block_ids = {}) override { h2h_write_calls++; last_h2h_write_peer = std::move(peer); last_h2h_write_src_blocks = src_block_ids; @@ -547,10 +548,12 @@ TEST(KVCacheManagerWrapperTest, RaidenControllerTransferBuffersIntegration) { Buffer dst_d2h_1(20, {}, std::nullopt, rpc::MEMORY_TYPE_DRAM); Buffer dst_d2h_2(40, {}, std::nullopt, rpc::MEMORY_TYPE_DRAM); - auto status_d2h = controller - .TransferBuffers("worker_0", {src_d2h_1, src_d2h_2}, - {dst_d2h_1, dst_d2h_2}, copy_sizes) - .Await(); + auto status_d2h = + controller + .TransferBuffers("worker_0", {src_d2h_1, src_d2h_2}, + {dst_d2h_1, dst_d2h_2}, /*staging_host_buffers=*/{}, + copy_sizes) + .Await(); ASSERT_TRUE(status_d2h.ok()); EXPECT_EQ(ptr0->d2h_calls, 1); EXPECT_EQ(ptr0->h2d_calls, 0); @@ -563,10 +566,12 @@ TEST(KVCacheManagerWrapperTest, RaidenControllerTransferBuffersIntegration) { Buffer dst_h2d_1(20, {}, std::nullopt, rpc::MEMORY_TYPE_HBM); Buffer dst_h2d_2(40, {}, std::nullopt, rpc::MEMORY_TYPE_HBM); - auto status_h2d = controller - .TransferBuffers("worker_0", {src_h2d_1, src_h2d_2}, - {dst_h2d_1, dst_h2d_2}, copy_sizes) - .Await(); + auto status_h2d = + controller + .TransferBuffers("worker_0", {src_h2d_1, src_h2d_2}, + {dst_h2d_1, dst_h2d_2}, /*staging_host_buffers=*/{}, + copy_sizes) + .Await(); ASSERT_TRUE(status_h2d.ok()); EXPECT_EQ(ptr0->d2h_calls, 1); EXPECT_EQ(ptr0->h2d_calls, 1); diff --git a/tpu_raiden/kv_cache/kv_cache_manager_base.cc b/tpu_raiden/kv_cache/kv_cache_manager_base.cc index 40949285..5e974395 100644 --- a/tpu_raiden/kv_cache/kv_cache_manager_base.cc +++ b/tpu_raiden/kv_cache/kv_cache_manager_base.cc @@ -92,6 +92,21 @@ absl::Status ValidateOffsetsAndSizes(const std::vector& src_offsets, return absl::OkStatus(); } +// Converts int64 host-block offsets to validated int block ids. +absl::StatusOr> ToHostBlockIds( + const std::vector& offsets) { + std::vector ids; + ids.reserve(offsets.size()); + for (int64_t offset : offsets) { + if (offset < 0 || offset > std::numeric_limits::max()) { + return absl::InvalidArgumentError( + absl::StrCat("Invalid host block ID: ", offset)); + } + ids.push_back(static_cast(offset)); + } + return ids; +} + // Coalesce runs of adjacent copies into one, so a run of N consecutive // 1-block copies becomes one N-block copy. void CoalesceMajorDimCopies(const std::vector& src_offsets, @@ -690,137 +705,194 @@ absl::StatusOr KVCacheManagerBase::D2h( } absl::StatusOr KVCacheManagerBase::H2dWrite( - absl::string_view peer, const std::vector& src_offsets_major_dim, - const std::vector& dst_offsets_major_dim, + absl::string_view peer, + const std::vector& src_host_offsets_major_dim, + const std::vector& dst_host_offsets_major_dim, + const std::vector& dst_device_offsets_major_dim, const std::vector& copy_sizes_major_dim) { - const bool present = !src_offsets_major_dim.empty() || - !dst_offsets_major_dim.empty() || - !copy_sizes_major_dim.empty(); - if (present && - (src_offsets_major_dim.size() != dst_offsets_major_dim.size() || - src_offsets_major_dim.size() != copy_sizes_major_dim.size())) { + size_t num_chunks = src_host_offsets_major_dim.size(); + if (num_chunks == 0) { + return raiden::PjRtCopyFuture(std::vector{}); + } + if (dst_host_offsets_major_dim.size() != num_chunks || + dst_device_offsets_major_dim.size() != num_chunks || + copy_sizes_major_dim.size() != num_chunks) { return absl::InvalidArgumentError( - "src_offsets, dst_offsets, and sizes must have the same length"); + "src_host, dst_host (staging), dst_device offsets and copy_sizes " + "must have the same length"); } - std::vector src_block_ids; - src_block_ids.reserve(src_offsets_major_dim.size()); - for (int64_t offset : src_offsets_major_dim) { - if (offset < 0 || offset > std::numeric_limits::max()) { - return absl::InvalidArgumentError( - absl::StrCat("Invalid host block ID: ", offset)); - } - src_block_ids.push_back(static_cast(offset)); - } + ASSIGN_OR_RETURN(std::vector src_block_ids, + ToHostBlockIds(src_host_offsets_major_dim)); + ASSIGN_OR_RETURN(std::vector staging_block_ids, + ToHostBlockIds(dst_host_offsets_major_dim)); + TF_RETURN_IF_ERROR(ValidateOffsetsAndSizes(src_host_offsets_major_dim, + dst_device_offsets_major_dim, + copy_sizes_major_dim)); - std::vector dst_block_ids; - dst_block_ids.reserve(dst_offsets_major_dim.size()); - for (int64_t offset : dst_offsets_major_dim) { + std::vector device_block_ids; + device_block_ids.reserve(num_chunks); + for (int64_t offset : dst_device_offsets_major_dim) { if (offset < 0 || offset > std::numeric_limits::max()) { return absl::InvalidArgumentError( - absl::StrCat("Invalid host block ID: ", offset)); - } - dst_block_ids.push_back(static_cast(offset)); - } - - TF_RETURN_IF_ERROR(ValidateOffsetsAndSizes( - src_offsets_major_dim, dst_offsets_major_dim, copy_sizes_major_dim)); - - size_t num_chunks = src_offsets_major_dim.size(); - if (num_chunks == 0) { - return raiden::PjRtCopyFuture(std::vector{}); - } - - if (num_chunks == 1 || !push_pool_) { - ASSIGN_OR_RETURN(auto h2h_res, - H2hWrite(std::string(peer), src_block_ids, dst_block_ids)); + absl::StrCat("Invalid device block ID: ", offset)); + } + device_block_ids.push_back(static_cast(offset)); + } + + // Full in-band pipeline (transport op 7): ONE push carries the H2D plan + // alongside the payload. Data lands in the peer's EXPLICIT host staging + // blocks (dst_host_offsets_major_dim), and the RECEIVER copies each + // completed layer into dst_device_offsets_major_dim (per-layer H2D, + // pipelined with the network). The returned future resolves only after the + // remote H2D completed (the receiver delays the transport ack until then). + // A single push (with internal stream parallelism) keeps the receiver's + // stream accounting (header.reserved) exact. + const uint64_t uuid = + (1ULL << 63) | (static_cast(node_id() & 0x3FFF) << 48) | + ((static_cast(reinterpret_cast(this) >> 4) & 0xFFFF) + << 32) | + (wire_h2d_uuid_counter_.fetch_add(1) + 1); + + if (!push_pool_) { + ASSIGN_OR_RETURN( + auto h2h_res, + H2hWrite(std::string(peer), src_block_ids, staging_block_ids, uuid, + /*layer_idx=*/-1, device_block_ids)); return h2h_res.second; } auto [promise, aggregate_future] = xla::MakePromise(); auto state = std::make_shared( - num_chunks, std::move(promise), raiden::BufferHolders{}); + 1, std::move(promise), raiden::BufferHolders{}); std::shared_ptr pool = push_pool_; std::string peer_str(peer); - for (size_t i = 0; i < num_chunks; ++i) { - int src_block_id = src_block_ids[i]; - int dst_block_id = dst_block_ids.empty() ? src_block_id : dst_block_ids[i]; - pool->Schedule([this, state, peer_str, src_block_id, dst_block_id]() { - if (state->HasFailed()) { - state->MarkChunkComplete(); - return; - } - absl::Status status = - H2hWriteDirect(peer_str, {src_block_id}, {dst_block_id}).status(); - if (!status.ok()) { - state->SetError(status); - } - state->MarkChunkComplete(); - }); - } + pool->Schedule([this, state, peer_str, src_block_ids, staging_block_ids, + device_block_ids, uuid]() { + absl::Status status = + H2hWriteDirect({peer_str}, src_block_ids, staging_block_ids, uuid, + /*layer_idx=*/-1, device_block_ids) + .status(); + VLOG(1) << "H2dWrite push (uuid " << uuid << ") status: " << status; + if (!status.ok()) { + state->SetError(status); + } + state->MarkChunkComplete(); + }); return raiden::PjRtCopyFuture(std::move(aggregate_future), state->combined_holds, state); } -absl::StatusOr KVCacheManagerBase::H2dRead( - absl::string_view peer, const std::vector& src_offsets_major_dim, - const std::vector& dst_offsets_major_dim, - const std::vector& copy_sizes_major_dim) { - const bool present = !src_offsets_major_dim.empty() || - !dst_offsets_major_dim.empty() || - !copy_sizes_major_dim.empty(); - if (present && !dst_offsets_major_dim.empty() && - src_offsets_major_dim.size() != dst_offsets_major_dim.size()) { +absl::Status KVCacheManagerBase::ArmRecvH2dFromWire( + uint64_t uuid, absl::Span staging_block_ids, + absl::Span device_block_ids, int expected_streams) { + if (staging_block_ids.size() != device_block_ids.size()) { return absl::InvalidArgumentError( - "src_offsets and dst_offsets must have the same length"); + "in-band H2D plan: staging/device id count mismatch"); } - if (present && !copy_sizes_major_dim.empty() && - src_offsets_major_dim.size() != copy_sizes_major_dim.size()) { + if (expected_streams <= 0) { return absl::InvalidArgumentError( - "src_offsets and copy_sizes must have the same length"); + "in-band H2D plan: expected_streams must be positive"); } - - std::vector src_block_ids; - src_block_ids.reserve(src_offsets_major_dim.size()); - for (int64_t offset : src_offsets_major_dim) { - if (offset < 0 || offset > std::numeric_limits::max()) { - return absl::InvalidArgumentError( - absl::StrCat("Invalid host block ID: ", offset)); - } - src_block_ids.push_back(static_cast(offset)); + absl::MutexLock lock(wire_h2d_mu_); + WireH2dPlan& plan = wire_h2d_plans_[uuid]; + plan.expected_streams = expected_streams; + // Streams of one push carry disjoint subsets; merging is a plain append. + for (size_t i = 0; i < staging_block_ids.size(); ++i) { + plan.staging_offsets.push_back(staging_block_ids[i]); + plan.device_offsets.push_back(device_block_ids[i]); } + return absl::OkStatus(); +} - for (int64_t offset : dst_offsets_major_dim) { - if (offset < 0 || offset > std::numeric_limits::max()) { - return absl::InvalidArgumentError( - absl::StrCat("Invalid host block ID: ", offset)); +absl::Status KVCacheManagerBase::OnLayerReceived(size_t layer_idx, + uint64_t uuid) { + { + absl::MutexLock lock(wire_h2d_mu_); + if (wire_h2d_plans_.contains(uuid)) { + // Wire-plan pushes use layer_idx=-1 (whole blob), so the transport's + // per-layer completion collapses to a single "network complete" event; + // the H2D fires in OnBlocksReceived finalization instead. + return absl::OkStatus(); } } + return RaidenManagerBase::OnLayerReceived(layer_idx, uuid); +} - TF_RETURN_IF_ERROR(ValidateOffsetsAndSizes( - src_offsets_major_dim, dst_offsets_major_dim, copy_sizes_major_dim)); +absl::Status KVCacheManagerBase::OnBlocksReceived( + const std::vector& block_ids, uint64_t uuid) { + std::vector staging_offsets; + std::vector device_offsets; + { + absl::MutexLock lock(wire_h2d_mu_); + auto it = wire_h2d_plans_.find(uuid); + if (it == wire_h2d_plans_.end()) { + return RaidenManagerBase::OnBlocksReceived(block_ids, uuid); + } + WireH2dPlan& plan = it->second; + ++plan.completed_streams; + if (plan.completed_streams < plan.expected_streams) { + return absl::OkStatus(); + } + // Last stream: every stream delivered all its layers and armed its pairs, + // so the merged plan is complete and all staging bytes have landed. + staging_offsets = std::move(plan.staging_offsets); + device_offsets = std::move(plan.device_offsets); + wire_h2d_plans_.erase(it); + } + // Fire the device copy for ALL layers and await it BEFORE the transport + // ack, so the sender's future means "resident in device memory" and any + // H2D failure propagates back to the sender. + auto future_or = H2d(staging_offsets, device_offsets, + std::vector(staging_offsets.size(), 1)); + if (!future_or.ok()) { + LOG(ERROR) << "In-band H2D plan failed for uuid " << uuid << ": " + << future_or.status(); + return future_or.status(); + } + return future_or->Await(); +} - size_t num_chunks = src_offsets_major_dim.size(); +absl::StatusOr KVCacheManagerBase::H2dRead( + absl::string_view peer, + const std::vector& src_host_offsets_major_dim, + const std::vector& dst_host_offsets_major_dim, + const std::vector& dst_device_offsets_major_dim, + const std::vector& copy_sizes_major_dim) { + size_t num_chunks = src_host_offsets_major_dim.size(); if (num_chunks == 0) { return raiden::PjRtCopyFuture(std::vector{}); } + if (dst_host_offsets_major_dim.size() != num_chunks || + dst_device_offsets_major_dim.size() != num_chunks || + copy_sizes_major_dim.size() != num_chunks) { + return absl::InvalidArgumentError( + "src_host, dst_host (staging), dst_device offsets and copy_sizes " + "must have the same length"); + } + ASSIGN_OR_RETURN(std::vector src_block_ids, + ToHostBlockIds(src_host_offsets_major_dim)); + ASSIGN_OR_RETURN(std::vector staging_block_ids, + ToHostBlockIds(dst_host_offsets_major_dim)); + TF_RETURN_IF_ERROR(ValidateOffsetsAndSizes(dst_host_offsets_major_dim, + dst_device_offsets_major_dim, + copy_sizes_major_dim)); + + // Pull remote host blocks into the EXPLICIT local host staging blocks + // (dst_host_offsets_major_dim) -- never into an aliased copy of the remote + // src id -- then H2D the staging blocks into the local device destination. if (num_chunks == 1 || !pull_pool_) { - ASSIGN_OR_RETURN(auto h2h_fut, H2hReadExplicit(std::string(peer), - src_block_ids, src_block_ids, - /*explicit_dst_ptrs=*/{})); + ASSIGN_OR_RETURN( + auto h2h_fut, + H2hReadExplicit(std::string(peer), src_block_ids, staging_block_ids, + /*explicit_dst_ptrs=*/{})); RETURN_IF_ERROR(h2h_fut.Await()); - const std::vector& h2d_dst_offsets = dst_offsets_major_dim.empty() - ? src_offsets_major_dim - : dst_offsets_major_dim; - std::vector h2d_sizes = copy_sizes_major_dim.empty() - ? std::vector(num_chunks, 1) - : copy_sizes_major_dim; - - return H2d(src_offsets_major_dim, h2d_dst_offsets, h2d_sizes); + return H2d(dst_host_offsets_major_dim, dst_device_offsets_major_dim, + copy_sizes_major_dim); } auto [promise, aggregate_future] = xla::MakePromise(); @@ -838,20 +910,21 @@ absl::StatusOr KVCacheManagerBase::H2dRead( std::string peer_str(peer); for (size_t i = 0; i < num_chunks; ++i) { int src_block_id = src_block_ids[i]; - int64_t dst_offset = dst_offsets_major_dim.empty() - ? src_offsets_major_dim[i] - : dst_offsets_major_dim[i]; - int64_t size = copy_sizes_major_dim.empty() ? 1 : copy_sizes_major_dim[i]; + int staging_block_id = staging_block_ids[i]; + int64_t staging_offset = dst_host_offsets_major_dim[i]; + int64_t dst_device_offset = dst_device_offsets_major_dim[i]; + int64_t size = copy_sizes_major_dim[i]; - pull_pool_->Schedule([this, state, peer_str, src_block_id, dst_offset, - size]() { + pull_pool_->Schedule([this, state, peer_str, src_block_id, staging_block_id, + staging_offset, dst_device_offset, size]() { if (state->HasFailed()) { state->MarkChunkComplete(); return; } - auto h2h_fut_or = H2hReadExplicit( - peer_str, {src_block_id}, {src_block_id}, /*explicit_dst_ptrs=*/{}); + auto h2h_fut_or = + H2hReadExplicit(peer_str, {src_block_id}, {staging_block_id}, + /*explicit_dst_ptrs=*/{}); if (!h2h_fut_or.ok()) { state->SetError(h2h_fut_or.status()); state->MarkChunkComplete(); @@ -870,7 +943,7 @@ absl::StatusOr KVCacheManagerBase::H2dRead( return; } - auto h2d_fut_or = H2d({src_block_id}, {dst_offset}, {size}); + auto h2d_fut_or = H2d({staging_offset}, {dst_device_offset}, {size}); if (!h2d_fut_or.ok()) { state->SetError(h2d_fut_or.status()); state->MarkChunkComplete(); @@ -891,44 +964,42 @@ absl::StatusOr KVCacheManagerBase::H2dRead( } absl::StatusOr KVCacheManagerBase::D2hWrite( - absl::string_view peer, const std::vector& src_offsets_major_dim, - const std::vector& dst_offsets_major_dim, + absl::string_view peer, + const std::vector& src_device_offsets_major_dim, + const std::vector& src_host_offsets_major_dim, + const std::vector& dst_host_offsets_major_dim, const std::vector& copy_sizes_major_dim) { - const bool present = !src_offsets_major_dim.empty() || - !dst_offsets_major_dim.empty() || - !copy_sizes_major_dim.empty(); - if (present && - (src_offsets_major_dim.size() != dst_offsets_major_dim.size() || - src_offsets_major_dim.size() != copy_sizes_major_dim.size())) { - return absl::InvalidArgumentError( - "src_offsets, dst_offsets, and sizes must have the same length"); - } - - std::vector host_block_ids; - host_block_ids.reserve(dst_offsets_major_dim.size()); - for (int64_t offset : dst_offsets_major_dim) { - if (offset < 0 || offset > std::numeric_limits::max()) { - return absl::InvalidArgumentError( - absl::StrCat("Invalid host block ID: ", offset)); - } - host_block_ids.push_back(static_cast(offset)); - } - - TF_RETURN_IF_ERROR(ValidateOffsetsAndSizes( - src_offsets_major_dim, dst_offsets_major_dim, copy_sizes_major_dim)); - size_t num_chunks = src_offsets_major_dim.size(); + size_t num_chunks = src_device_offsets_major_dim.size(); if (num_chunks == 0) { return raiden::PjRtCopyFuture(std::vector{}); } + if (src_host_offsets_major_dim.size() != num_chunks || + dst_host_offsets_major_dim.size() != num_chunks || + copy_sizes_major_dim.size() != num_chunks) { + return absl::InvalidArgumentError( + "src_device, src_host (staging), dst_host offsets and copy_sizes " + "must have the same length"); + } + + ASSIGN_OR_RETURN(std::vector staging_block_ids, + ToHostBlockIds(src_host_offsets_major_dim)); + ASSIGN_OR_RETURN(std::vector dst_block_ids, + ToHostBlockIds(dst_host_offsets_major_dim)); + TF_RETURN_IF_ERROR(ValidateOffsetsAndSizes(src_device_offsets_major_dim, + src_host_offsets_major_dim, + copy_sizes_major_dim)); + // Stage local device blocks into the EXPLICIT local host staging blocks + // (src_host_offsets_major_dim) -- never into an aliased copy of the remote + // dst id -- then push the staging blocks to the peer's host destination. if (num_chunks == 1 || !push_pool_) { ASSIGN_OR_RETURN(auto d2h_future, - D2h(src_offsets_major_dim, dst_offsets_major_dim, - copy_sizes_major_dim)); + D2h(src_device_offsets_major_dim, + src_host_offsets_major_dim, copy_sizes_major_dim)); RETURN_IF_ERROR(d2h_future.Await()); - ASSIGN_OR_RETURN(auto h2h_res, H2hWrite(std::string(peer), host_block_ids, - host_block_ids)); + ASSIGN_OR_RETURN(auto h2h_res, H2hWrite(std::string(peer), + staging_block_ids, dst_block_ids)); return h2h_res.second; } @@ -937,22 +1008,24 @@ absl::StatusOr KVCacheManagerBase::D2hWrite( struct ChunkD2h { raiden::PjRtCopyFuture d2h_fut; - int host_block_id; + int staging_block_id; + int dst_block_id; }; std::vector chunks; chunks.reserve(num_chunks); for (size_t i = 0; i < num_chunks; ++i) { ASSIGN_OR_RETURN(auto chunk_futures, - DispatchD2hChunks({src_offsets_major_dim[i]}, - {dst_offsets_major_dim[i]}, + DispatchD2hChunks({src_device_offsets_major_dim[i]}, + {src_host_offsets_major_dim[i]}, {copy_sizes_major_dim[i]})); raiden::PjRtCopyFuture d2h_fut = raiden::JoinPjRtCopyFutures(absl::MakeSpan(chunk_futures)); for (const auto& h : d2h_fut.holds) { all_holds.push_back(h); } - chunks.push_back({std::move(d2h_fut), host_block_ids[i]}); + chunks.push_back( + {std::move(d2h_fut), staging_block_ids[i], dst_block_ids[i]}); } auto state = std::make_shared( @@ -961,21 +1034,23 @@ absl::StatusOr KVCacheManagerBase::D2hWrite( std::shared_ptr pool = push_pool_; std::string peer_str(peer); for (size_t i = 0; i < num_chunks; ++i) { - int host_block_id = chunks[i].host_block_id; - chunks[i].d2h_fut.OnReady([this, pool, state, peer_str, - host_block_id](auto status_or) { + int staging_block_id = chunks[i].staging_block_id; + int dst_block_id = chunks[i].dst_block_id; + chunks[i].d2h_fut.OnReady([this, pool, state, peer_str, staging_block_id, + dst_block_id](auto status_or) { if (!status_or.ok()) { state->SetError(status_or.status()); state->MarkChunkComplete(); return; } - pool->Schedule([this, state, peer_str, host_block_id]() { + pool->Schedule([this, state, peer_str, staging_block_id, dst_block_id]() { if (state->HasFailed()) { state->MarkChunkComplete(); return; } absl::Status status = - H2hWriteDirect(peer_str, {host_block_id}, {host_block_id}).status(); + H2hWriteDirect(peer_str, {staging_block_id}, {dst_block_id}) + .status(); if (!status.ok()) { state->SetError(status); } @@ -1043,10 +1118,11 @@ absl::StatusOr, raiden::PjRtCopyFuture>> KVCacheManagerBase::H2hWrite(const std::vector& peers, const std::vector& src_block_ids, const std::vector& dst_block_ids, - uint64_t uuid, int layer_idx) { - ASSIGN_OR_RETURN( - std::vector allocated_ids, - H2hWriteDirect(peers, src_block_ids, dst_block_ids, uuid, layer_idx)); + uint64_t uuid, int layer_idx, + const std::vector& dst_device_block_ids) { + ASSIGN_OR_RETURN(std::vector allocated_ids, + H2hWriteDirect(peers, src_block_ids, dst_block_ids, uuid, + layer_idx, dst_device_block_ids)); return std::make_pair( allocated_ids, raiden::PjRtCopyFuture(std::vector{})); @@ -1056,9 +1132,10 @@ absl::StatusOr, raiden::PjRtCopyFuture>> KVCacheManagerBase::H2hWrite(std::string peer, const std::vector& src_block_ids, const std::vector& dst_block_ids, - uint64_t uuid, int layer_idx) { + uint64_t uuid, int layer_idx, + const std::vector& dst_device_block_ids) { return H2hWrite(std::vector{std::move(peer)}, src_block_ids, - dst_block_ids, uuid, layer_idx); + dst_block_ids, uuid, layer_idx, dst_device_block_ids); } absl::StatusOr, raiden::PjRtCopyFuture>> diff --git a/tpu_raiden/kv_cache/kv_cache_manager_base.h b/tpu_raiden/kv_cache/kv_cache_manager_base.h index 058fba63..11da14e4 100644 --- a/tpu_raiden/kv_cache/kv_cache_manager_base.h +++ b/tpu_raiden/kv_cache/kv_cache_manager_base.h @@ -15,6 +15,7 @@ #ifndef THIRD_PARTY_TPU_RAIDEN_KV_CACHE_KV_CACHE_MANAGER_BASE_H_ #define THIRD_PARTY_TPU_RAIDEN_KV_CACHE_KV_CACHE_MANAGER_BASE_H_ +#include #include #include #include @@ -118,17 +119,30 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase { std::optional layer_idx = std::nullopt, std::optional shard_idx = std::nullopt); + // Pushes local Host DRAM blocks into a remote peer's device memory (TPU + // HBM). Parameters are flow-ordered: local host source -> REMOTE host + // staging (bridge, explicit; never aliased to another id) -> remote device + // destination. The full pipeline is self-contained in this one call: the + // push carries the H2D plan in-band (transport op 7), the receiver fires + // its local H2D once the pushed payload is complete, and the returned + // future resolves only once the data is resident in the peer's device + // blocks (receiver H2D failures fail this future). virtual absl::StatusOr H2dWrite( absl::string_view peer, - const std::vector& src_offsets_major_dim = {}, - const std::vector& dst_offsets_major_dim = {}, - const std::vector& copy_sizes_major_dim = {}); + const std::vector& src_host_offsets_major_dim, + const std::vector& dst_host_offsets_major_dim, + const std::vector& dst_device_offsets_major_dim, + const std::vector& copy_sizes_major_dim); + // Pulls remote Host DRAM blocks into local TPU HBM. Parameters are + // flow-ordered: remote host source -> LOCAL host staging (bridge, explicit; + // caller-owned, never aliased to the remote id) -> local HBM destination. virtual absl::StatusOr H2dRead( absl::string_view peer, - const std::vector& src_offsets_major_dim = {}, - const std::vector& dst_offsets_major_dim = {}, - const std::vector& copy_sizes_major_dim = {}); + const std::vector& src_host_offsets_major_dim, + const std::vector& dst_host_offsets_major_dim, + const std::vector& dst_device_offsets_major_dim, + const std::vector& copy_sizes_major_dim); // Async on-chip D2H offloads E2E virtual absl::StatusOr D2h( @@ -139,11 +153,15 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase { std::optional layer_idx = std::nullopt, std::optional shard_idx = std::nullopt); + // Pushes local TPU HBM blocks to a remote peer's Host DRAM. Parameters are + // flow-ordered: local HBM source -> LOCAL host staging (bridge, explicit; + // caller-owned, never aliased to the remote id) -> remote host destination. virtual absl::StatusOr D2hWrite( absl::string_view peer, - const std::vector& src_offsets_major_dim = {}, - const std::vector& dst_offsets_major_dim = {}, - const std::vector& copy_sizes_major_dim = {}); + const std::vector& src_device_offsets_major_dim, + const std::vector& src_host_offsets_major_dim, + const std::vector& dst_host_offsets_major_dim, + const std::vector& copy_sizes_major_dim); virtual absl::StatusOr D2hRead( absl::string_view peer, @@ -158,17 +176,22 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase { D2hAutoAllocate(const std::vector& src_offsets_major_dim = {}, const std::vector& copy_sizes_major_dim = {}); - // Symmetrical H2H writes E2E + // Symmetrical H2H writes E2E. + // dst_device_block_ids (optional): when non-empty, the push carries an + // in-band H2D plan (transport op 7) and the receiver copies each landed + // host block into the paired device block (see H2dWrite). virtual absl::StatusOr, raiden::PjRtCopyFuture>> H2hWrite(const std::vector& peers, const std::vector& src_block_ids, const std::vector& dst_block_ids = {}, uint64_t uuid = 0, - int layer_idx = -1); + int layer_idx = -1, + const std::vector& dst_device_block_ids = {}); virtual absl::StatusOr, raiden::PjRtCopyFuture>> H2hWrite(std::string peer, const std::vector& src_block_ids, const std::vector& dst_block_ids = {}, uint64_t uuid = 0, - int layer_idx = -1); + int layer_idx = -1, + const std::vector& dst_device_block_ids = {}); virtual absl::StatusOr, raiden::PjRtCopyFuture>> H2hRead(const std::vector& peers, @@ -307,6 +330,22 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase { bool use_block_chunks(uint64_t uuid) const override; + // In-band H2D plan hooks (transport op 7). ArmRecvH2dFromWire merges each + // push stream's staging->device pairs into the per-uuid plan; + // OnBlocksReceived finalizes on the LAST stream: fires the local H2D + // (staging -> device, all layers) and awaits it before the transport ack, + // so the sender's future means "resident in device memory" and receiver + // H2D failures propagate back to the sender. (Wire-plan pushes are + // whole-blob (layer_idx=-1), so OnLayerReceived is a single + // network-complete event, not per-layer -- it is a no-op for wire uuids.) + absl::Status ArmRecvH2dFromWire(uint64_t uuid, + absl::Span staging_block_ids, + absl::Span device_block_ids, + int expected_streams) override; + absl::Status OnLayerReceived(size_t layer_idx, uint64_t uuid) override; + absl::Status OnBlocksReceived(const std::vector& block_ids, + uint64_t uuid) override; + absl::StatusOr> GetPoolPushProgressSpec(size_t pool_idx, uint64_t uuid) const override; @@ -460,6 +499,24 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase { }; absl::flat_hash_map active_plans_ ABSL_GUARDED_BY(plans_mu_); + + // Receiver-side state for in-band H2D plans (transport op 7), keyed by the + // push uuid. Streams arm incrementally (disjoint subsets); the last + // stream's OnBlocksReceived awaits the per-layer H2D futures and erases the + // plan. + struct WireH2dPlan { + std::vector staging_offsets; + std::vector device_offsets; + int expected_streams = 0; + int completed_streams = 0; + }; + absl::Mutex wire_h2d_mu_; + absl::flat_hash_map wire_h2d_plans_ + ABSL_GUARDED_BY(wire_h2d_mu_); + // Sender-side uuid generator for op-7 pushes. Top bit marks wire-plan + // uuids; node id + instance salt keep concurrent senders from colliding in + // the receiver's shared uuid keyspace. + std::atomic wire_h2d_uuid_counter_{0}; }; } // namespace kv_cache diff --git a/tpu_raiden/kv_cache/kv_cache_manager_test.cc b/tpu_raiden/kv_cache/kv_cache_manager_test.cc index c9325c59..0a7b5bfe 100644 --- a/tpu_raiden/kv_cache/kv_cache_manager_test.cc +++ b/tpu_raiden/kv_cache/kv_cache_manager_test.cc @@ -17,13 +17,16 @@ #include #include #include +#include #include +#include #include #include #include #include "absl/status/status.h" #include "absl/strings/str_cat.h" +#include "absl/synchronization/mutex.h" #include "tpu_raiden/kv_cache/kv_cache_manager_base.h" #include "tpu_raiden/rpc/raiden_service.pb.h" #include "tpu_raiden/transport/block_transport.h" @@ -835,23 +838,39 @@ class TestH2dKVCacheManager : public TestKVCacheManager { std::optional slot_idx = std::nullopt, std::optional layer_idx = std::nullopt, std::optional shard_idx = std::nullopt) override { + absl::MutexLock lock(h2d_mu_); h2d_called_ = true; + ++h2d_call_count_; last_h2d_src_offsets_ = src_offsets_major_dim; last_h2d_dst_offsets_ = dst_offsets_major_dim; last_h2d_copy_sizes_ = copy_sizes_major_dim; + h2d_layer_calls_.push_back(layer_idx); + // Order-independent record of every staging->device pair, per call. + for (size_t i = 0; + i < src_offsets_major_dim.size() && i < dst_offsets_major_dim.size(); + ++i) { + h2d_pairs_.insert({src_offsets_major_dim[i], dst_offsets_major_dim[i]}); + } return raiden::PjRtCopyFuture(std::vector{}); } + void set_parallelism_for_test(int p) { parallelism_ = p; } + + mutable absl::Mutex h2d_mu_; bool h2d_called_ = false; + int h2d_call_count_ = 0; std::vector last_h2d_src_offsets_; std::vector last_h2d_dst_offsets_; std::vector last_h2d_copy_sizes_; + std::vector> h2d_layer_calls_; + std::set> h2d_pairs_; }; TEST(KVCacheManagerTest, D2hWriteFailsWithCpuOnlyManager) { KVCacheManagerBase manager(/*num_layers=*/1, /*num_shards=*/1, /*slice_byte_size=*/128); - auto res = manager.D2hWrite("127.0.0.1:8080", {0}, {0}, {1}); + auto res = manager.D2hWrite("127.0.0.1:8080", /*src_device=*/{0}, + /*src_host=*/{0}, /*dst_host=*/{0}, {1}); EXPECT_FALSE(res.ok()); EXPECT_EQ(res.status().code(), absl::StatusCode::kFailedPrecondition); EXPECT_THAT(res.status().message(), @@ -861,7 +880,8 @@ TEST(KVCacheManagerTest, D2hWriteFailsWithCpuOnlyManager) { TEST(KVCacheManagerTest, D2hWriteFailsWithInvalidHostBlockId) { TestD2hKVCacheManager manager(/*num_layers=*/1, /*num_shards=*/1, /*slice_byte_size=*/128, /*host_blocks=*/2); - auto res = manager.D2hWrite("127.0.0.1:8080", {0}, {-1}, {1}); + auto res = manager.D2hWrite("127.0.0.1:8080", /*src_device=*/{0}, + /*src_host=*/{-1}, /*dst_host=*/{0}, {1}); EXPECT_FALSE(res.ok()); EXPECT_EQ(res.status().code(), absl::StatusCode::kInvalidArgument); EXPECT_THAT(res.status().message(), @@ -879,16 +899,19 @@ TEST(KVCacheManagerTest, D2hWriteSuccessWithMockD2h) { std::string receiver_peer = absl::StrCat(receiver.local_ip(), ":", *receiver_port); - std::vector src_offsets = {0}; - std::vector dst_offsets = {1}; + std::vector src_device_offsets = {0}; + std::vector src_host_offsets = {0}; // local staging (bridge) + std::vector dst_host_offsets = {1}; // remote destination std::vector copy_sizes = {1}; - auto res = - sender.D2hWrite(receiver_peer, src_offsets, dst_offsets, copy_sizes); + auto res = sender.D2hWrite(receiver_peer, src_device_offsets, + src_host_offsets, dst_host_offsets, copy_sizes); ASSERT_TRUE(res.ok()) << res.status().ToString(); EXPECT_TRUE(sender.d2h_called_); - EXPECT_EQ(sender.last_src_offsets_, src_offsets); - EXPECT_EQ(sender.last_dst_offsets_, dst_offsets); + EXPECT_EQ(sender.last_src_offsets_, src_device_offsets); + // The D2H stage lands in the EXPLICIT local staging blocks, not in a local + // alias of the remote destination id. + EXPECT_EQ(sender.last_dst_offsets_, src_host_offsets); EXPECT_EQ(sender.last_copy_sizes_, copy_sizes); } @@ -913,12 +936,13 @@ TEST(KVCacheManagerTest, D2hWritePipelinedSuccess) { std::memset(sender_buf + 128, 0xCD, 128); std::memset(receiver_buf, 0, 256); - std::vector src_offsets = {0, 1}; - std::vector dst_offsets = {0, 1}; + std::vector src_device_offsets = {0, 1}; + std::vector src_host_offsets = {0, 1}; // local staging (bridge) + std::vector dst_host_offsets = {0, 1}; // remote destination std::vector copy_sizes = {1, 1}; - auto res = - sender.D2hWrite(receiver_peer, src_offsets, dst_offsets, copy_sizes); + auto res = sender.D2hWrite(receiver_peer, src_device_offsets, + src_host_offsets, dst_host_offsets, copy_sizes); ASSERT_TRUE(res.ok()) << res.status().ToString(); EXPECT_TRUE(res->Await().ok()); @@ -953,12 +977,15 @@ TEST(KVCacheManagerTest, H2dReadSuccess) { std::memset(receiver_buf, 0, 256); // Test empty src_offsets returns OK empty future - auto empty_res = receiver.H2dRead(sender_peer, {}); + auto empty_res = receiver.H2dRead(sender_peer, {}, {}, {}, {}); ASSERT_TRUE(empty_res.ok()) << empty_res.status().ToString(); EXPECT_TRUE(empty_res->Await().ok()); - // Test H2dRead reading sender block 0 into receiver block 0 - auto res = receiver.H2dRead(sender_peer, /*src_offsets_major_dim=*/{0}); + // Test H2dRead reading sender block 0 via local staging block 0 into + // receiver device block 0. + auto res = receiver.H2dRead(sender_peer, /*src_host=*/{0}, + /*dst_host(staging)=*/{0}, /*dst_device=*/{0}, + /*copy_sizes=*/{1}); ASSERT_TRUE(res.ok()) << res.status().ToString(); EXPECT_TRUE(res->Await().ok()); @@ -986,12 +1013,13 @@ TEST(KVCacheManagerTest, H2dReadPipelinedSuccess) { std::memset(sender_buf + 128, 0x66, 128); std::memset(receiver_buf, 0, 256); - std::vector src_offsets = {0, 1}; - std::vector dst_offsets = {0, 1}; + std::vector src_host_offsets = {0, 1}; + std::vector dst_host_offsets = {0, 1}; // local staging (bridge) + std::vector dst_device_offsets = {0, 1}; std::vector copy_sizes = {1, 1}; - auto res = - receiver.H2dRead(sender_peer, src_offsets, dst_offsets, copy_sizes); + auto res = receiver.H2dRead(sender_peer, src_host_offsets, dst_host_offsets, + dst_device_offsets, copy_sizes); ASSERT_TRUE(res.ok()) << res.status().ToString(); EXPECT_TRUE(res->Await().ok()); @@ -1015,13 +1043,14 @@ TEST(KVCacheManagerTest, H2dReadCallsH2dForTpuHbmDestination) { ASSERT_NE(sender_buf, nullptr); std::memset(sender_buf, 0x77, 128); - auto res = receiver.H2dRead(sender_peer, /*src_offsets_major_dim=*/{0}, - /*dst_offsets_major_dim=*/{1}, - /*copy_sizes_major_dim=*/{1}); + auto res = receiver.H2dRead(sender_peer, /*src_host=*/{0}, + /*dst_host(staging)=*/{0}, /*dst_device=*/{1}, + /*copy_sizes=*/{1}); ASSERT_TRUE(res.ok()) << res.status().ToString(); EXPECT_TRUE(res->Await().ok()); - // H2dRead MUST trigger Stage 2 H2d DMA into TPU HBM destination offset {1}. + // H2dRead MUST trigger Stage 2 H2d DMA from the explicit staging block {0} + // into TPU HBM destination offset {1}. EXPECT_TRUE(receiver.h2d_called_); EXPECT_EQ(receiver.last_h2d_src_offsets_, std::vector{0}); EXPECT_EQ(receiver.last_h2d_dst_offsets_, std::vector{1}); @@ -1031,8 +1060,8 @@ TEST(KVCacheManagerTest, H2dReadCallsH2dForTpuHbmDestination) { TEST(KVCacheManagerTest, H2dWriteSuccess) { TestKVCacheManager sender(/*num_layers=*/1, /*num_shards=*/1, /*slice_byte_size=*/128, /*host_blocks=*/2); - TestKVCacheManager receiver(/*num_layers=*/1, /*num_shards=*/1, - /*slice_byte_size=*/128, /*host_blocks=*/2); + TestH2dKVCacheManager receiver(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, /*host_blocks=*/2); const std::optional receiver_port = receiver.local_port(); ASSERT_TRUE(receiver_port.has_value()); @@ -1048,26 +1077,84 @@ TEST(KVCacheManagerTest, H2dWriteSuccess) { std::memset(sender_buf, 0xCD, 128); std::memset(receiver_buf, 0, 256); - std::vector src_offsets = {0}; - std::vector dst_offsets = {1}; + std::vector src_host_offsets = {0}; + std::vector dst_host_offsets = {1}; // remote staging (bridge) + std::vector dst_device_offsets = {5}; std::vector copy_sizes = {1}; - auto res = - sender.H2dWrite(receiver_peer, src_offsets, dst_offsets, copy_sizes); + auto res = sender.H2dWrite(receiver_peer, src_host_offsets, dst_host_offsets, + dst_device_offsets, copy_sizes); ASSERT_TRUE(res.ok()) << res.status().ToString(); EXPECT_TRUE(res->Await().ok()); + // Network stage: payload landed byte-exact in the explicit remote staging + // block {1}; block {0} untouched. EXPECT_TRUE(std::all_of(receiver_buf, receiver_buf + 128, [](uint8_t v) { return v == 0; })); EXPECT_TRUE(std::all_of(receiver_buf + 128, receiver_buf + 256, [](uint8_t v) { return v == 0xCD; })); + + // Remote H2D stage: fired automatically from the in-band plan, from the + // staging block into the device destination (all layers). This is the full + // pipeline triggered by the single H2dWrite() call. + absl::MutexLock lock(receiver.h2d_mu_); + EXPECT_TRUE(receiver.h2d_called_); + EXPECT_EQ(receiver.h2d_call_count_, 1); + EXPECT_EQ(receiver.last_h2d_src_offsets_, dst_host_offsets); + EXPECT_EQ(receiver.last_h2d_dst_offsets_, dst_device_offsets); + EXPECT_EQ(receiver.last_h2d_copy_sizes_, std::vector{1}); + ASSERT_EQ(receiver.h2d_layer_calls_.size(), 1u); + EXPECT_EQ(receiver.h2d_layer_calls_[0], std::nullopt); +} + +// A receiver whose device copy fails must FAIL the sender's whole H2dWrite +// future (no silent staging-only success): the remote H2D error propagates +// back through the transport ack path. +class TestFailingH2dKVCacheManager : public TestKVCacheManager { + public: + using TestKVCacheManager::TestKVCacheManager; + + absl::StatusOr H2d( + const std::vector& src_offsets_major_dim = {}, + const std::vector& dst_offsets_major_dim = {}, + const std::vector& copy_sizes_major_dim = {}, + std::optional slot_idx = std::nullopt, + std::optional layer_idx = std::nullopt, + std::optional shard_idx = std::nullopt) override { + return absl::FailedPreconditionError( + "receiver has no device backing for H2D"); + } +}; + +TEST(KVCacheManagerTest, H2dWriteFailsWhenReceiverH2dFails) { + TestKVCacheManager sender(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, /*host_blocks=*/2); + TestFailingH2dKVCacheManager receiver(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, + /*host_blocks=*/2); + + const std::optional receiver_port = receiver.local_port(); + ASSERT_TRUE(receiver_port.has_value()); + std::string receiver_peer = + absl::StrCat(receiver.local_ip(), ":", *receiver_port); + + uint8_t* sender_buf = sender.GetHostPointer(/*layer_idx=*/0, /*shard_idx=*/0); + ASSERT_NE(sender_buf, nullptr); + std::memset(sender_buf, 0xCD, 128); + + auto res = sender.H2dWrite(receiver_peer, /*src_host=*/{0}, + /*dst_host(staging)=*/{1}, /*dst_device=*/{0}, + /*copy_sizes=*/{1}); + if (res.ok()) { + EXPECT_FALSE(res->Await().ok()); + } } TEST(KVCacheManagerTest, H2dWritePipelinedSuccess) { TestKVCacheManager sender(/*num_layers=*/1, /*num_shards=*/1, /*slice_byte_size=*/128, /*host_blocks=*/2); - TestKVCacheManager receiver(/*num_layers=*/1, /*num_shards=*/1, - /*slice_byte_size=*/128, /*host_blocks=*/2); + TestH2dKVCacheManager receiver(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, /*host_blocks=*/2); const std::optional receiver_port = receiver.local_port(); ASSERT_TRUE(receiver_port.has_value()); @@ -1084,19 +1171,207 @@ TEST(KVCacheManagerTest, H2dWritePipelinedSuccess) { std::memset(sender_buf + 128, 0x44, 128); std::memset(receiver_buf, 0, 256); - std::vector src_offsets = {0, 1}; - std::vector dst_offsets = {0, 1}; + std::vector src_host_offsets = {0, 1}; + std::vector dst_host_offsets = {0, 1}; // remote staging (bridge) + std::vector dst_device_offsets = {6, 9}; std::vector copy_sizes = {1, 1}; - auto res = - sender.H2dWrite(receiver_peer, src_offsets, dst_offsets, copy_sizes); + auto res = sender.H2dWrite(receiver_peer, src_host_offsets, dst_host_offsets, + dst_device_offsets, copy_sizes); ASSERT_TRUE(res.ok()) << res.status().ToString(); EXPECT_TRUE(res->Await().ok()); + // Network stage byte-exact into the explicit staging blocks. EXPECT_TRUE(std::all_of(receiver_buf, receiver_buf + 128, [](uint8_t v) { return v == 0x33; })); EXPECT_TRUE(std::all_of(receiver_buf + 128, receiver_buf + 256, [](uint8_t v) { return v == 0x44; })); + + // Remote H2D stage: single layer -> one H2d call covering BOTH + // staging->device pairs (order-independent). + absl::MutexLock lock(receiver.h2d_mu_); + EXPECT_TRUE(receiver.h2d_called_); + EXPECT_EQ(receiver.h2d_call_count_, 1); + EXPECT_EQ(receiver.h2d_pairs_, + (std::set>{{0, 6}, {1, 9}})); +} + +// Full-pipeline data-correctness test: multiple layers, multiple blocks, +// multi-stream push (parallelism 2). Every (layer, block) region carries a +// distinct byte pattern; verifies (a) every staging block is byte-exact after +// the network stage, (b) the receiver fired exactly one H2D per layer, each +// covering every staging->device pair with the correct pairing. +TEST(KVCacheManagerTest, H2dWriteMultiLayerMultiStreamByteExact) { + constexpr size_t kLayers = 2; + constexpr int kBlocks = 3; + constexpr size_t kBlockBytes = 128; + TestH2dKVCacheManager sender(kLayers, /*num_shards=*/1, kBlockBytes, + /*host_blocks=*/kBlocks); + TestH2dKVCacheManager receiver(kLayers, /*num_shards=*/1, kBlockBytes, + /*host_blocks=*/kBlocks); + sender.set_parallelism_for_test(2); + + const std::optional receiver_port = receiver.local_port(); + ASSERT_TRUE(receiver_port.has_value()); + std::string receiver_peer = + absl::StrCat(receiver.local_ip(), ":", *receiver_port); + + // Distinct pattern per (layer, block): 0xA0 + l*16 + k. + for (size_t l = 0; l < kLayers; ++l) { + uint8_t* sbuf = sender.GetHostPointer(l, /*shard_idx=*/0); + uint8_t* rbuf = receiver.GetHostPointer(l, /*shard_idx=*/0); + ASSERT_NE(sbuf, nullptr); + ASSERT_NE(rbuf, nullptr); + std::memset(rbuf, 0, kBlocks * kBlockBytes); + for (int k = 0; k < kBlocks; ++k) { + std::memset(sbuf + k * kBlockBytes, static_cast(0xA0 + l * 16 + k), + kBlockBytes); + } + } + + std::vector src_host_offsets = {0, 1, 2}; + std::vector dst_host_offsets = {2, 0, 1}; // permuted staging + std::vector dst_device_offsets = {10, 11, 12}; + std::vector copy_sizes = {1, 1, 1}; + + auto res = sender.H2dWrite(receiver_peer, src_host_offsets, dst_host_offsets, + dst_device_offsets, copy_sizes); + ASSERT_TRUE(res.ok()) << res.status().ToString(); + ASSERT_TRUE(res->Await().ok()); + + // (a) Byte-exact landing: sender block k (pattern 0xA0+l*16+k) must sit in + // receiver staging block dst_host_offsets[k], for every layer. + for (size_t l = 0; l < kLayers; ++l) { + uint8_t* rbuf = receiver.GetHostPointer(l, /*shard_idx=*/0); + for (int k = 0; k < kBlocks; ++k) { + const uint8_t want = static_cast(0xA0 + l * 16 + k); + uint8_t* region = rbuf + dst_host_offsets[k] * kBlockBytes; + EXPECT_TRUE(std::all_of(region, region + kBlockBytes, + [want](uint8_t v) { return v == want; })) + << "layer " << l << " sender block " << k << " staging block " + << dst_host_offsets[k]; + } + } + + // (b) Exactly ONE H2D at network-complete, covering all layers (no layer + // restriction) and every staging->device pair with correct pairing. + absl::MutexLock lock(receiver.h2d_mu_); + EXPECT_EQ(receiver.h2d_call_count_, 1); + ASSERT_EQ(receiver.h2d_layer_calls_.size(), 1u); + EXPECT_EQ(receiver.h2d_layer_calls_[0], std::nullopt); + EXPECT_EQ(receiver.h2d_pairs_, + (std::set>{{2, 10}, {0, 11}, {1, 12}})); + EXPECT_EQ(receiver.last_h2d_src_offsets_.size(), 3u); +} + +// Anti-clobber regression: the local staging block for H2dRead is the +// EXPLICIT dst_host block. A local block whose id happens to equal the remote +// src id must NOT be touched (the pre-fix code aliased the remote src id as +// the local staging id and destroyed that block's contents). +TEST(KVCacheManagerTest, H2dReadExplicitStagingDoesNotClobberAliasedBlock) { + TestKVCacheManager sender(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, /*host_blocks=*/2); + TestKVCacheManager receiver(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, /*host_blocks=*/2); + + const std::optional sender_port = sender.local_port(); + ASSERT_TRUE(sender_port.has_value()); + std::string sender_peer = absl::StrCat(sender.local_ip(), ":", *sender_port); + + uint8_t* sender_buf = sender.GetHostPointer(/*layer_idx=*/0, /*shard_idx=*/0); + uint8_t* receiver_buf = + receiver.GetHostPointer(/*layer_idx=*/0, /*shard_idx=*/0); + ASSERT_NE(sender_buf, nullptr); + ASSERT_NE(receiver_buf, nullptr); + + std::memset(sender_buf, 0xEF, 128); + // Sentinel in receiver's local block 0 -- same id as the REMOTE src block. + // The pre-fix code staged into local block 0 and destroyed this. + std::memset(receiver_buf, 0x99, 128); + std::memset(receiver_buf + 128, 0, 128); + + auto res = receiver.H2dRead(sender_peer, /*src_host=*/{0}, + /*dst_host(staging)=*/{1}, /*dst_device=*/{0}, + /*copy_sizes=*/{1}); + ASSERT_TRUE(res.ok()) << res.status().ToString(); + EXPECT_TRUE(res->Await().ok()); + + // Data staged into the explicit staging block 1. + EXPECT_TRUE(std::all_of(receiver_buf + 128, receiver_buf + 256, + [](uint8_t v) { return v == 0xEF; })); + // The would-be-aliased local block 0 is untouched. + EXPECT_TRUE(std::all_of(receiver_buf, receiver_buf + 128, + [](uint8_t v) { return v == 0x99; })); +} + +// Anti-clobber regression: D2hWrite stages through the EXPLICIT src_host +// block and pushes THAT block to the peer. A local block whose id happens to +// equal the remote dst id must NOT be used (the pre-fix code staged into and +// pushed from the local alias of the remote dst id). +TEST(KVCacheManagerTest, D2hWriteExplicitStagingIsPushedNotAliasedBlock) { + TestD2hKVCacheManager sender(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, /*host_blocks=*/2); + TestKVCacheManager receiver(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, /*host_blocks=*/2); + + const std::optional receiver_port = receiver.local_port(); + ASSERT_TRUE(receiver_port.has_value()); + std::string receiver_peer = + absl::StrCat(receiver.local_ip(), ":", *receiver_port); + + uint8_t* sender_buf = sender.GetHostPointer(/*layer_idx=*/0, /*shard_idx=*/0); + uint8_t* receiver_buf = + receiver.GetHostPointer(/*layer_idx=*/0, /*shard_idx=*/0); + ASSERT_NE(sender_buf, nullptr); + ASSERT_NE(receiver_buf, nullptr); + + // Staging block 0 holds the payload (the mocked D2h stage is a no-op, so + // the pre-seeded content is what gets pushed). Local block 1 -- same id as + // the REMOTE dst block -- holds a sentinel the pre-fix code would have + // staged into and pushed. + std::memset(sender_buf, 0xAB, 128); + std::memset(sender_buf + 128, 0x99, 128); + std::memset(receiver_buf, 0, 256); + + auto res = sender.D2hWrite(receiver_peer, /*src_device=*/{0}, + /*src_host(staging)=*/{0}, /*dst_host=*/{1}, + /*copy_sizes=*/{1}); + ASSERT_TRUE(res.ok()) << res.status().ToString(); + EXPECT_TRUE(res->Await().ok()); + + // The peer received the STAGING block's payload, not the sentinel from the + // sender's local block 1 (the would-be alias of the remote dst id). + EXPECT_TRUE(std::all_of(receiver_buf + 128, receiver_buf + 256, + [](uint8_t v) { return v == 0xAB; })); + // The sender's local block 1 is untouched. + EXPECT_TRUE(std::all_of(sender_buf + 128, sender_buf + 256, + [](uint8_t v) { return v == 0x99; })); +} + +// The 2-stage remote APIs require the explicit staging list (same length as +// the other offset lists) -- no silent alias fallback. +TEST(KVCacheManagerTest, RemoteTwoStageApisRequireExplicitStaging) { + TestKVCacheManager manager(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, /*host_blocks=*/2); + + auto h2d_read = manager.H2dRead("localhost:1", /*src_host=*/{0}, + /*dst_host(staging)=*/{}, /*dst_device=*/{0}, + /*copy_sizes=*/{1}); + EXPECT_EQ(h2d_read.status().code(), absl::StatusCode::kInvalidArgument); + EXPECT_THAT(h2d_read.status().message(), testing::HasSubstr("same length")); + + auto d2h_write = manager.D2hWrite("localhost:1", /*src_device=*/{0}, + /*src_host(staging)=*/{}, /*dst_host=*/{0}, + /*copy_sizes=*/{1}); + EXPECT_EQ(d2h_write.status().code(), absl::StatusCode::kInvalidArgument); + EXPECT_THAT(d2h_write.status().message(), testing::HasSubstr("same length")); + + auto h2d_write = + manager.H2dWrite("localhost:1", /*src_host=*/{0}, + /*dst_host(staging)=*/{}, /*dst_device=*/{0}, + /*copy_sizes=*/{1}); + EXPECT_EQ(h2d_write.status().code(), absl::StatusCode::kInvalidArgument); + EXPECT_THAT(h2d_write.status().message(), testing::HasSubstr("same length")); } } // namespace diff --git a/tpu_raiden/kv_cache/kv_cache_store.cc b/tpu_raiden/kv_cache/kv_cache_store.cc index 71c806cb..96ac4187 100644 --- a/tpu_raiden/kv_cache/kv_cache_store.cc +++ b/tpu_raiden/kv_cache/kv_cache_store.cc @@ -459,8 +459,8 @@ absl::Status KVCacheStore::Save(const std::vector& block_hashes) { } // Trigger transfer - tsl::Future<> future = - raiden_controller_->TransferBuffers(src_buffers, dst_buffers); + tsl::Future<> future = raiden_controller_->TransferBuffers( + src_buffers, dst_buffers, /*staging_host_buffers=*/{}, /*copy_sizes=*/{}); { absl::MutexLock lock(mutex_); @@ -533,8 +533,8 @@ absl::Status KVCacheStore::Load(const std::vector& block_hashes, } // Trigger transfer - tsl::Future<> future = - raiden_controller_->TransferBuffers(src_buffers, dst_buffers); + tsl::Future<> future = raiden_controller_->TransferBuffers( + src_buffers, dst_buffers, /*staging_host_buffers=*/{}, /*copy_sizes=*/{}); { absl::MutexLock lock(mutex_); diff --git a/tpu_raiden/proto/worker_service.proto b/tpu_raiden/proto/worker_service.proto index 746ac6dc..499de1c1 100644 --- a/tpu_raiden/proto/worker_service.proto +++ b/tpu_raiden/proto/worker_service.proto @@ -143,6 +143,13 @@ message TransferBufferSpec { repeated BufferProto src_buffers = 7; // Destination buffers for transfer. repeated BufferProto dst_buffers = 8; + // Host DRAM staging (bridge) blocks for 2-stage remote transfers + // (remote H2D read/write and remote D2H write). Always the MIDDLE hop of the + // data flow: local host staging for H2dRead/D2hWrite, remote host staging + // for H2dWrite. Must be explicit and caller-owned -- never aliased to a + // peer's block id. Required (same length as src/dst offsets) for those + // transfer types; unused otherwise. + repeated BufferProto staging_host_buffers = 9; } // Request to transfer data across memory spaces on a transfer worker. diff --git a/tpu_raiden/transport/block_transport.cc b/tpu_raiden/transport/block_transport.cc index cbf4b149..1bb7e264 100644 --- a/tpu_raiden/transport/block_transport.cc +++ b/tpu_raiden/transport/block_transport.cc @@ -239,7 +239,7 @@ absl::Status BlockTransport::HandleCustomRequest(int client_fd, << ", uuid=" << header.uuid << ", numa=" << block_delegate_->node_id(); - if (header.op == 1 || header.op == 6) { + if (header.op == 1 || header.op == 6 || header.op == 7) { return HandleIncomingPush(client_fd, header); } else if (header.op == 2) { return HandleIncomingPull(client_fd, header); @@ -308,6 +308,17 @@ absl::Status BlockTransport::HandleIncomingPush(int client_fd, src_block_ids.resize(header.count_or_size, 0); RETURN_IF_ERROR(ReadExact(client_fd, src_block_ids.data(), header.count_or_size * sizeof(int))); + if (header.op == 7) { + // Op 7 ("push + H2D plan"): this stream's device-target subset travels + // in-band. Arm (merge) the receiver-side H2D plan BEFORE any payload + // lands so the per-layer H2D can fire as soon as a layer completes. + std::vector device_block_ids(header.count_or_size, 0); + RETURN_IF_ERROR(ReadExact(client_fd, device_block_ids.data(), + header.count_or_size * sizeof(int))); + RETURN_IF_ERROR(block_delegate_->ArmRecvH2dFromWire( + header.uuid, allocated_ids, device_block_ids, + /*expected_streams=*/header.reserved)); + } uint8_t ack = 1; RETURN_IF_ERROR(WriteExact(client_fd, &ack, 1)); } @@ -588,14 +599,18 @@ absl::StatusOr> BlockTransport::SyncPush( const std::vector& peers, const std::vector& src_block_ids, const std::vector& dst_block_ids, int parallelism, - MajorOrder major_order, uint64_t uuid, int layer_idx) { + MajorOrder major_order, uint64_t uuid, int layer_idx, + const std::vector& dst_device_block_ids) { auto promise = std::make_shared>>>(); auto future = promise->get_future(); - AsyncPush(peers, src_block_ids, dst_block_ids, parallelism, major_order, uuid, - layer_idx, [promise](absl::StatusOr> res) { - promise->set_value(std::move(res)); - }); + AsyncPush( + peers, src_block_ids, dst_block_ids, parallelism, major_order, uuid, + layer_idx, + [promise](absl::StatusOr> res) { + promise->set_value(std::move(res)); + }, + dst_device_block_ids); return future.get(); } @@ -604,7 +619,8 @@ void BlockTransport::AsyncPush( const std::vector& src_block_ids, const std::vector& dst_block_ids, int parallelism, MajorOrder major_order, uint64_t uuid, int layer_idx, - std::function>)> on_complete) { + std::function>)> on_complete, + const std::vector& dst_device_block_ids) { size_t num_blocks = src_block_ids.size(); if (num_blocks == 0) { on_complete(absl::InvalidArgumentError("Block list cannot be empty")); @@ -614,6 +630,13 @@ void BlockTransport::AsyncPush( on_complete(absl::InvalidArgumentError("Peer list cannot be empty")); return; } + if (!dst_device_block_ids.empty() && + (dst_device_block_ids.size() != num_blocks || + dst_block_ids.size() != num_blocks)) { + on_complete(absl::InvalidArgumentError( + "dst_device_block_ids, if provided, must match src/dst block lists")); + return; + } int P = parallelism; if (P <= 0) { @@ -624,6 +647,8 @@ void BlockTransport::AsyncPush( auto shared_src_block_ids = std::make_shared>(src_block_ids); auto shared_dst_block_ids = std::make_shared>(dst_block_ids); + auto shared_dst_device_block_ids = + std::make_shared>(dst_device_block_ids); auto allocated_ids = std::make_shared>(num_blocks, 0); auto statuses = std::make_shared>(P, absl::OkStatus()); @@ -641,13 +666,14 @@ void BlockTransport::AsyncPush( std::string remote_peer = peers[i % peers.size()]; auto task_run = [this, i, remote_peer, local_ip, block_offset, block_count, - shared_src_block_ids, shared_dst_block_ids, allocated_ids, - statuses, remaining_workers, major_order, uuid, layer_idx, - P, on_complete]() { + shared_src_block_ids, shared_dst_block_ids, + shared_dst_device_block_ids, allocated_ids, statuses, + remaining_workers, major_order, uuid, layer_idx, P, + on_complete]() { H2hWriteWorker(i, remote_peer, local_ip, block_offset, block_count, *shared_src_block_ids, *shared_dst_block_ids, - *allocated_ids, *statuses, major_order, uuid, layer_idx, - P); + *allocated_ids, *statuses, major_order, uuid, layer_idx, P, + *shared_dst_device_block_ids); if (remaining_workers->fetch_sub(1) == 1) { absl::Status final_status = absl::OkStatus(); @@ -776,15 +802,14 @@ absl::StatusOr> BlockTransport::SyncPull( return allocated_ids; } -void BlockTransport::H2hWriteWorker(int stream_idx, absl::string_view peer, - absl::string_view local_ip, - size_t block_offset, size_t block_count, - const std::vector& src_block_ids, - const std::vector& dst_block_ids, - std::vector& allocated_ids, - std::vector& statuses, - MajorOrder major_order, uint64_t uuid, - int layer_idx, int parallelism) { +void BlockTransport::H2hWriteWorker( + int stream_idx, absl::string_view peer, absl::string_view local_ip, + size_t block_offset, size_t block_count, + const std::vector& src_block_ids, + const std::vector& dst_block_ids, std::vector& allocated_ids, + std::vector& statuses, MajorOrder major_order, uint64_t uuid, + int layer_idx, int parallelism, + const std::vector& dst_device_block_ids) { auto status_or_fd = BorrowConnection(peer, local_ip); if (!status_or_fd.ok()) { statuses[stream_idx] = status_or_fd.status(); @@ -797,7 +822,10 @@ void BlockTransport::H2hWriteWorker(int stream_idx, absl::string_view peer, [&] { ReturnConnection(ok_to_pool, fd, peer, local_ip); }); PacketHeader header = {}; - header.op = dst_block_ids.empty() ? 1 : 6; + header.op = dst_block_ids.empty() ? 1 + : dst_device_block_ids.empty() + ? 6 + : 7; // 7 = push + in-band H2D plan header.flags = static_cast(major_order); header.buffer_id = 0; header.remote_id = static_cast(block_delegate_->node_id()); @@ -817,7 +845,7 @@ void BlockTransport::H2hWriteWorker(int stream_idx, absl::string_view peer, return; } - if (header.op == 6) { + if (header.op == 6 || header.op == 7) { ABSL_DCHECK_LE(block_offset + block_count, dst_block_ids.size()); s = WriteExact(fd, &dst_block_ids[block_offset], block_count * sizeof(int)); if (!s.ok()) { @@ -829,6 +857,15 @@ void BlockTransport::H2hWriteWorker(int stream_idx, absl::string_view peer, statuses[stream_idx] = s; return; } + if (header.op == 7) { + ABSL_DCHECK_LE(block_offset + block_count, dst_device_block_ids.size()); + s = WriteExact(fd, &dst_device_block_ids[block_offset], + block_count * sizeof(int)); + if (!s.ok()) { + statuses[stream_idx] = s; + return; + } + } uint8_t ack = 0; s = ReadExact(fd, &ack, 1); if (!s.ok() || ack != 1) { diff --git a/tpu_raiden/transport/block_transport.h b/tpu_raiden/transport/block_transport.h index dd04939f..2dadb739 100644 --- a/tpu_raiden/transport/block_transport.h +++ b/tpu_raiden/transport/block_transport.h @@ -170,6 +170,27 @@ class BlockTransportDelegate : public lib::RawBufferTransportDelegate { return absl::OkStatus(); } + /** + * @brief Hook to arm receiver-side H2D operations directly from the payload + * of an in-band H2dWrite request (op=7). + * + * @param uuid Transfer identifier connecting stream sets. + * @param staging_block_ids The array of host-memory block identifiers to + * populate. + * @param device_block_ids The remote target's device-memory block + * identifiers. + * @param expected_streams Tracks total concurrent streams processing this + * uuid block group. + * @return Status or UnimplementedError on host-only delegates lacking + * hardware. + */ + virtual absl::Status ArmRecvH2dFromWire( + uint64_t uuid, absl::Span staging_block_ids, + absl::Span device_block_ids, int expected_streams) { + return absl::UnimplementedError( + "Receiver does not support in-band H2D plans (op 7)"); + } + virtual absl::Status OnBlocksReceived(const std::vector& block_ids, uint64_t uuid = 0) { return OnDataReceived(); @@ -218,21 +239,47 @@ class BlockTransport : public lib::RawBufferTransport { int parallelism = 1); ~BlockTransport() override; - // Asynchronous Scatter-Gather Push + /** + * @brief Asynchronous Scatter-Gather Push + * + * @param peers Target endpoints. + * @param src_block_ids Logical source block IDs locally. + * @param dst_block_ids Logical target staging block IDs remotely. + * @param parallelism Socket concurrency limit. + * @param major_order Transposition format instructions for remote. + * @param uuid Transfer session UUID. + * @param layer_idx Layer index override. + * @param on_complete Resolution callback. + * @param dst_device_block_ids (Optional) Device block IDs. When non-empty + * escalates protocol to op=7 pipelined H2D execution. + */ void AsyncPush( const std::vector& peers, const std::vector& src_block_ids, const std::vector& dst_block_ids, int parallelism, MajorOrder major_order, uint64_t uuid, int layer_idx, - std::function>)> on_complete); - - // Synchronous Scatter-Gather Push (op = 1 / op = 6) + std::function>)> on_complete, + const std::vector& dst_device_block_ids = {}); + + /** + * @brief Synchronous Scatter-Gather Push + * + * @param peers Target endpoints. + * @param src_block_ids Logical source block IDs locally. + * @param dst_block_ids Logical target staging block IDs remotely. + * @param parallelism Socket concurrency limit. + * @param major_order Transposition format instructions for remote. + * @param uuid Transfer session UUID. + * @param layer_idx Layer index override. + * @param dst_device_block_ids (Optional) Device block IDs. When non-empty + * escalates protocol to op=7 pipelined H2D execution. + */ absl::StatusOr> SyncPush( const std::vector& peers, const std::vector& src_block_ids, const std::vector& dst_block_ids = {}, int parallelism = 1, MajorOrder major_order = MajorOrder::kLayerMajor, uint64_t uuid = 0, - int layer_idx = -1); + int layer_idx = -1, const std::vector& dst_device_block_ids = {}); // Synchronous Scatter-Gather Pull (op = 2) // When explicit_dst_ptrs is supplied it contains one base pointer per @@ -278,7 +325,8 @@ class BlockTransport : public lib::RawBufferTransport { std::vector& allocated_ids, std::vector& statuses, MajorOrder major_order, uint64_t uuid = 0, - int layer_idx = -1, int parallelism = 1); + int layer_idx = -1, int parallelism = 1, + const std::vector& dst_device_block_ids = {}); void H2hReadWorker(int stream_idx, absl::string_view peer, absl::string_view local_ip, size_t local_block_offset,