From 79f908ce02770ee4b61c75dc26f08b9666098724 Mon Sep 17 00:00:00 2001 From: justinlu Date: Tue, 21 Jul 2026 13:32:54 -0700 Subject: [PATCH] Deprecate legacy raw-offset TransferBuffers overloads in RaidenController. PiperOrigin-RevId: 951669979 --- .../core/controller/raiden_controller.cc | 172 ++++++------------ .../core/controller/raiden_controller.h | 25 +-- .../core/controller/raiden_controller_test.cc | 50 +---- .../core/controller/worker_service_impl.cc | 21 +-- .../core/controller/worker_service_test.cc | 10 +- tpu_raiden/kv_cache/BUILD | 1 + tpu_raiden/kv_cache/kv_cache_store.cc | 42 +++-- tpu_raiden/proto/worker_service.proto | 3 +- 8 files changed, 108 insertions(+), 216 deletions(-) diff --git a/tpu_raiden/core/controller/raiden_controller.cc b/tpu_raiden/core/controller/raiden_controller.cc index 01ae7381..40692c15 100644 --- a/tpu_raiden/core/controller/raiden_controller.cc +++ b/tpu_raiden/core/controller/raiden_controller.cc @@ -151,8 +151,51 @@ void RaidenController::Init(absl::Span worker_addresses, absl::Span dst_offsets, absl::Span copy_sizes, absl::Span peers) { - return this->TransferBuffers(src_mem_type, dst_mem_type, src_offsets, - dst_offsets, copy_sizes, peers); + if (src_offsets.empty() || src_offsets.size() != dst_offsets.size()) { + return tsl::Future<>( + absl::InvalidArgumentError("Source and destination offsets must " + "have the same non-zero length")); + } + if (!copy_sizes.empty() && copy_sizes.size() != src_offsets.size()) { + return tsl::Future<>(absl::InvalidArgumentError( + "copy_sizes, if provided, must match the length of src_offsets")); + } + auto workers = worker_registry_->GetRegisteredWorkers(); + if (workers.empty()) { + return tsl::Future<>(absl::FailedPreconditionError( + "No registered workers available for TransferBuffers")); + } + std::sort(workers.begin(), workers.end(), + [](const core::controller::WorkerRegistration& a, + const core::controller::WorkerRegistration& b) { + return CompareWorkerIds(a.worker_id, b.worker_id); + }); + if (!peers.empty() && workers.size() != peers.size()) { + return tsl::Future<>(absl::InvalidArgumentError( + absl::StrCat("Peers count mismatch: workers has ", workers.size(), + ", peers has ", peers.size()))); + } + std::vector> worker_futures; + worker_futures.reserve(workers.size()); + for (size_t i = 0; i < workers.size(); ++i) { + std::optional peer = + peers.empty() ? std::nullopt : std::make_optional(peers[i]); + std::vector src_buffers; + src_buffers.reserve(src_offsets.size()); + for (int64_t offset : src_offsets) { + src_buffers.emplace_back(offset, std::vector{}, + std::nullopt, src_mem_type); + } + std::vector dst_buffers; + dst_buffers.reserve(dst_offsets.size()); + for (int64_t offset : dst_offsets) { + dst_buffers.emplace_back(offset, std::vector{}, peer, + dst_mem_type); + } + worker_futures.push_back(this->TransferBuffers( + workers[i].worker_id, src_buffers, dst_buffers, copy_sizes)); + } + return tsl::JoinFutures(absl::MakeSpan(worker_futures)); }); // 3. Register static workers @@ -172,8 +215,10 @@ void RaidenController::Init(absl::Span worker_addresses, orchestrator_client_ = std::make_unique( grpc::CreateChannel(std::string(raiden_orchestrator_address), grpc::InsecureChannelCredentials())); - std::string my_endpoint = absl::StrCat("localhost:", raiden_controller_port_); - absl::Status status = orchestrator_client_->RegisterController(unit_, my_endpoint); + std::string my_endpoint = + absl::StrCat("localhost:", raiden_controller_port_); + absl::Status status = + orchestrator_client_->RegisterController(unit_, my_endpoint); if (!status.ok()) { throw std::runtime_error(absl::StrCat( "Failed to register with orchestrator: ", status.message())); @@ -190,7 +235,8 @@ RaidenController::RaidenController( shard_size_bytes_(shard_size_bytes), num_total_blocks_(num_blocks), worker_registry_(std::make_shared()), - block_manager_(std::make_unique(num_blocks)), + block_manager_( + std::make_unique(num_blocks)), raiden_controller_port_(raiden_controller_port) { Init(/*worker_addresses=*/{}, raiden_orchestrator_address); } @@ -205,7 +251,8 @@ RaidenController::RaidenController( shard_size_bytes_(shard_size_bytes), num_total_blocks_(num_blocks), worker_registry_(std::make_shared()), - block_manager_(std::make_unique(num_blocks)), + block_manager_( + std::make_unique(num_blocks)), raiden_controller_port_(raiden_controller_port) { Init(worker_addresses, raiden_orchestrator_address); } @@ -372,22 +419,11 @@ RaidenController::BuildTransferBuffersRequest( rpc::MemoryType src_mem_type = src_buffers[0].memory_type(); rpc::MemoryType dst_mem_type = dst_buffers[0].memory_type(); - std::string peer; - if (dst_buffers[0].remote_address().has_value() && - !dst_buffers[0].remote_address()->empty()) { - peer = *dst_buffers[0].remote_address(); - } else if (src_buffers[0].remote_address().has_value() && - !src_buffers[0].remote_address()->empty()) { - peer = *src_buffers[0].remote_address(); - } proto::TransferBuffersRequest request; auto* transfer = request.mutable_transfer(); transfer->set_src_mem_type(src_mem_type); transfer->set_dst_mem_type(dst_mem_type); - if (!peer.empty()) { - transfer->set_peer(peer); - } for (const auto& buf : src_buffers) { if (buf.index() < 0) { @@ -414,41 +450,6 @@ RaidenController::BuildTransferBuffersRequest( return request; } -absl::StatusOr -RaidenController::BuildRawTransferBuffersRequest( - rpc::MemoryType src_mem_type, rpc::MemoryType dst_mem_type, - absl::Span src_offsets, - absl::Span dst_offsets, absl::Span copy_sizes, - absl::string_view peer) { - if (src_offsets.empty() || src_offsets.size() != dst_offsets.size()) { - return absl::InvalidArgumentError( - "Source and destination offsets must have the same non-zero length"); - } - if (!copy_sizes.empty() && copy_sizes.size() != src_offsets.size()) { - return absl::InvalidArgumentError( - "copy_sizes, if provided, must match the length of src_offsets"); - } - - proto::TransferBuffersRequest request; - auto* transfer = request.mutable_transfer(); - transfer->set_src_mem_type(src_mem_type); - transfer->set_dst_mem_type(dst_mem_type); - if (!peer.empty()) { - transfer->set_peer(std::string(peer)); - } - - for (int64_t offset : src_offsets) { - transfer->add_src_offsets(offset); - } - for (int64_t offset : dst_offsets) { - transfer->add_dst_offsets(offset); - } - for (int64_t size : copy_sizes) { - transfer->add_copy_sizes(size); - } - return request; -} - tsl::Future<> RaidenController::TransferBuffers( absl::string_view worker_id, absl::Span src_buffers, absl::Span dst_buffers, @@ -504,74 +505,7 @@ tsl::Future<> RaidenController::TransferBuffers( return tsl::JoinFutures(absl::MakeSpan(worker_futures)); } -tsl::Future<> RaidenController::TransferBuffers( - absl::string_view worker_id, rpc::MemoryType src_mem_type, - rpc::MemoryType dst_mem_type, absl::Span src_offsets, - absl::Span dst_offsets, absl::Span copy_sizes, - absl::string_view peer) { - auto request_or = BuildRawTransferBuffersRequest( - src_mem_type, dst_mem_type, src_offsets, dst_offsets, copy_sizes, peer); - if (!request_or.ok()) { - return tsl::Future<>(request_or.status()); - } - auto worker_or = worker_registry_->GetWorker(worker_id); - if (!worker_or.ok()) { - return tsl::Future<>(absl::FailedPreconditionError( - absl::StrCat("Worker ", worker_id, " is not registered. ", - "Did you wait for the worker to register?"))); - } - auto worker_client = worker_or->worker_service_client; - if (!worker_client) { - return tsl::Future<>(absl::FailedPreconditionError(absl::StrCat( - "WorkerServiceClient for ", worker_id, " is not initialized."))); - } - - return worker_client->TransferBuffers(*request_or); -} - -tsl::Future<> RaidenController::TransferBuffers( - rpc::MemoryType src_mem_type, rpc::MemoryType dst_mem_type, - absl::Span src_offsets, - absl::Span dst_offsets, absl::Span copy_sizes, - absl::Span peers) { - if (src_offsets.empty() || src_offsets.size() != dst_offsets.size()) { - return tsl::Future<>(absl::InvalidArgumentError( - "Source and destination offsets must have the same non-zero length")); - } - if (!copy_sizes.empty() && copy_sizes.size() != src_offsets.size()) { - return tsl::Future<>(absl::InvalidArgumentError( - "copy_sizes, if provided, must match the length of src_offsets")); - } - auto workers = worker_registry_->GetRegisteredWorkers(); - if (workers.empty()) { - return tsl::Future<>(absl::FailedPreconditionError( - "No registered workers available for TransferBuffers")); - } - - std::sort(workers.begin(), workers.end(), - [](const core::controller::WorkerRegistration& a, - const core::controller::WorkerRegistration& b) { - return CompareWorkerIds(a.worker_id, b.worker_id); - }); - - if (!peers.empty() && workers.size() != peers.size()) { - return tsl::Future<>(absl::InvalidArgumentError( - absl::StrCat("Peers count mismatch: workers has ", workers.size(), - ", peers has ", peers.size()))); - } - - std::vector> worker_futures; - worker_futures.reserve(workers.size()); - for (size_t i = 0; i < workers.size(); ++i) { - std::string peer = peers.empty() ? "" : peers[i]; - worker_futures.push_back(TransferBuffers(workers[i].worker_id, src_mem_type, - dst_mem_type, src_offsets, - dst_offsets, copy_sizes, peer)); - } - - return tsl::JoinFutures(absl::MakeSpan(worker_futures)); -} absl::StatusOr RaidenController::ResolvePeerController( const rpc::RaidenIdProto& peer_id) { diff --git a/tpu_raiden/core/controller/raiden_controller.h b/tpu_raiden/core/controller/raiden_controller.h index 47bc8210..ec85e9c5 100644 --- a/tpu_raiden/core/controller/raiden_controller.h +++ b/tpu_raiden/core/controller/raiden_controller.h @@ -127,22 +127,7 @@ class RaidenController { absl::Span dst_buffers, absl::Span copy_sizes = {}); - // Legacy targeted worker transfer using raw block offsets - tsl::Future<> TransferBuffers(absl::string_view worker_id, - rpc::MemoryType src_mem_type, - rpc::MemoryType dst_mem_type, - absl::Span src_offsets, - absl::Span dst_offsets, - absl::Span copy_sizes = {}, - absl::string_view peer = ""); - - // Legacy broadcast transfer using raw block offsets - tsl::Future<> TransferBuffers(rpc::MemoryType src_mem_type, - rpc::MemoryType dst_mem_type, - absl::Span src_offsets, - absl::Span dst_offsets, - absl::Span copy_sizes = {}, - absl::Span peers = {}); + // Initiates remote read from source controller. tsl::Future<> ReadRemote(const kv_cache::RaidenId& src_raiden_id, @@ -176,16 +161,12 @@ class RaidenController { absl::Span dst_buffers, absl::Span copy_sizes); - absl::StatusOr BuildRawTransferBuffersRequest( - rpc::MemoryType src_mem_type, rpc::MemoryType dst_mem_type, - absl::Span src_offsets, - absl::Span dst_offsets, - absl::Span copy_sizes, absl::string_view peer); void Init(absl::Span worker_addresses, absl::string_view raiden_orchestrator_address); - absl::Status InitializeWorkerBuffers(core::controller::WorkerRegistration& reg); + absl::Status InitializeWorkerBuffers( + core::controller::WorkerRegistration& reg); rpc::RaidenIdProto unit_; int num_shards_; diff --git a/tpu_raiden/core/controller/raiden_controller_test.cc b/tpu_raiden/core/controller/raiden_controller_test.cc index 12381f3e..a6294c49 100644 --- a/tpu_raiden/core/controller/raiden_controller_test.cc +++ b/tpu_raiden/core/controller/raiden_controller_test.cc @@ -681,7 +681,7 @@ TEST_F(RaidenControllerTest, TransferBuffersBufferProtoSuccess) { EXPECT_THAT(mock_mgr.last_dst_offsets, ElementsAre(0, 1)); } -TEST_F(RaidenControllerTest, LegacyTransferBuffersTargetedSuccess) { +TEST_F(RaidenControllerTest, TransferBuffersBroadcastSuccess) { MockTransferManager mock_mgr; test_server_->service->SetTransferManager(KVManagerHolder(&mock_mgr)); @@ -689,39 +689,10 @@ TEST_F(RaidenControllerTest, LegacyTransferBuffersTargetedSuccess) { /*shard_size_bytes=*/512); RegisterAndInitWorker(controller, "worker_0", test_server_->server_address); - std::vector src_offsets = {10, 30}; - std::vector dst_offsets = {20, 40}; - std::vector copy_sizes = {1, 2}; - - auto status = controller - .TransferBuffers("worker_0", rpc::MEMORY_TYPE_HBM, - rpc::MEMORY_TYPE_DRAM, src_offsets, - dst_offsets, copy_sizes, "") - .Await(); - ASSERT_TRUE(status.ok()); - EXPECT_EQ(mock_mgr.d2h_calls, 1); - EXPECT_EQ(mock_mgr.h2d_calls, 0); - EXPECT_THAT(mock_mgr.last_src_offsets, ElementsAre(10, 30)); - EXPECT_THAT(mock_mgr.last_dst_offsets, ElementsAre(20, 40)); - EXPECT_THAT(mock_mgr.last_copy_sizes, ElementsAre(1, 2)); -} - -TEST_F(RaidenControllerTest, LegacyTransferBuffersBroadcastSuccess) { - MockTransferManager mock_mgr; - test_server_->service->SetTransferManager(KVManagerHolder(&mock_mgr)); - - RaidenController controller(unit_, /*num_blocks=*/5, /*num_shards=*/1, - /*shard_size_bytes=*/512); - RegisterAndInitWorker(controller, "worker_0", test_server_->server_address); - - std::vector src_offsets = {100}; - std::vector dst_offsets = {200}; + Buffer src_buf(100, {}, std::nullopt, rpc::MEMORY_TYPE_DRAM); + Buffer dst_buf(200, {}, std::nullopt, rpc::MEMORY_TYPE_HBM); - auto status = - controller - .TransferBuffers(rpc::MEMORY_TYPE_DRAM, rpc::MEMORY_TYPE_HBM, - src_offsets, dst_offsets, {}, {}) - .Await(); + auto status = controller.TransferBuffers({src_buf}, {dst_buf}).Await(); ASSERT_TRUE(status.ok()); EXPECT_EQ(mock_mgr.d2h_calls, 0); EXPECT_EQ(mock_mgr.h2d_calls, 1); @@ -729,7 +700,7 @@ TEST_F(RaidenControllerTest, LegacyTransferBuffersBroadcastSuccess) { EXPECT_THAT(mock_mgr.last_dst_offsets, ElementsAre(200)); } -TEST_F(RaidenControllerTest, LegacyTransferBuffersH2HSuccess) { +TEST_F(RaidenControllerTest, TransferBuffersBroadcastH2HSuccess) { MockTransferManager mock_mgr; test_server_->service->SetTransferManager(KVManagerHolder(&mock_mgr)); @@ -737,15 +708,10 @@ TEST_F(RaidenControllerTest, LegacyTransferBuffersH2HSuccess) { /*shard_size_bytes=*/512); RegisterAndInitWorker(controller, "worker_0", test_server_->server_address); - std::vector src_offsets = {5}; - std::vector dst_offsets = {6}; + Buffer src_buf(5, {}, std::nullopt, rpc::MEMORY_TYPE_DRAM); + Buffer dst_buf(6, {}, "localhost:8080", rpc::MEMORY_TYPE_DRAM); - auto status = - controller - .TransferBuffers(rpc::MEMORY_TYPE_DRAM, rpc::MEMORY_TYPE_DRAM, - src_offsets, dst_offsets, {}, - std::vector{"localhost:8080"}) - .Await(); + auto status = controller.TransferBuffers({src_buf}, {dst_buf}).Await(); ASSERT_TRUE(status.ok()); EXPECT_EQ(mock_mgr.h2h_write_calls, 1); EXPECT_EQ(mock_mgr.last_peer, "localhost:8080"); diff --git a/tpu_raiden/core/controller/worker_service_impl.cc b/tpu_raiden/core/controller/worker_service_impl.cc index 364ba9a9..34d711ea 100644 --- a/tpu_raiden/core/controller/worker_service_impl.cc +++ b/tpu_raiden/core/controller/worker_service_impl.cc @@ -226,12 +226,9 @@ grpc::Status WorkerServiceImpl::TransferBuffers( std::string src_peer = transfer.src_buffers(0).remote_address(); future_or = transfer_manager_.D2hRead(src_peer, src_offsets, dst_offsets, copy_sizes); - } else if (!transfer.peer().empty() || - (transfer.dst_buffers_size() > 0 && - !transfer.dst_buffers(0).remote_address().empty())) { - std::string dst_peer = !transfer.peer().empty() - ? transfer.peer() - : transfer.dst_buffers(0).remote_address(); + } 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); } else { @@ -243,12 +240,9 @@ grpc::Status WorkerServiceImpl::TransferBuffers( std::string src_peer = transfer.src_buffers(0).remote_address(); future_or = transfer_manager_.H2dRead(src_peer, src_offsets, dst_offsets, copy_sizes); - } else if (!transfer.peer().empty() || - (transfer.dst_buffers_size() > 0 && - !transfer.dst_buffers(0).remote_address().empty())) { - std::string dst_peer = !transfer.peer().empty() - ? transfer.peer() - : transfer.dst_buffers(0).remote_address(); + } 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); } else { @@ -262,9 +256,6 @@ grpc::Status WorkerServiceImpl::TransferBuffers( !transfer.dst_buffers(0).remote_address().empty()) { future_or = transfer_manager_.H2hWrite( transfer.dst_buffers(0).remote_address(), src_offsets, dst_offsets); - } else if (!transfer.peer().empty()) { - future_or = - transfer_manager_.H2hWrite(transfer.peer(), src_offsets, dst_offsets); } else { response->set_success(false); response->set_message("Peer address must be provided for H2H transfers"); diff --git a/tpu_raiden/core/controller/worker_service_test.cc b/tpu_raiden/core/controller/worker_service_test.cc index f6c508c3..816a6827 100644 --- a/tpu_raiden/core/controller/worker_service_test.cc +++ b/tpu_raiden/core/controller/worker_service_test.cc @@ -129,7 +129,7 @@ TEST_F(WorkerServiceTest, TransferBuffersH2hSuccess) { transfer->set_dst_mem_type(rpc::MEMORY_TYPE_DRAM); transfer->add_src_offsets(10); transfer->add_dst_offsets(20); - transfer->set_peer("localhost:8080"); + transfer->add_dst_buffers()->set_remote_address("localhost:8080"); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); ASSERT_TRUE(status.ok()); @@ -170,7 +170,7 @@ TEST_F(WorkerServiceTest, TransferBuffersH2hInvalidCopySizeFails) { transfer->add_src_offsets(10); transfer->add_dst_offsets(20); transfer->add_copy_sizes(2); - transfer->set_peer("localhost:8080"); + transfer->add_dst_buffers()->set_remote_address("localhost:8080"); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); EXPECT_FALSE(status.ok()); @@ -188,7 +188,7 @@ TEST_F(WorkerServiceTest, TransferBuffersH2hOverflowFails) { transfer->set_dst_mem_type(rpc::MEMORY_TYPE_DRAM); transfer->add_src_offsets(2147483648L); transfer->add_dst_offsets(20); - transfer->set_peer("localhost:8080"); + transfer->add_dst_buffers()->set_remote_address("localhost:8080"); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); EXPECT_FALSE(status.ok()); @@ -280,7 +280,7 @@ TEST_F(WorkerServiceTest, TransferBuffersRemoteD2hWithPeerSuccess) { transfer->set_dst_mem_type(rpc::MEMORY_TYPE_DRAM); transfer->add_src_offsets(100); transfer->add_dst_offsets(200); - transfer->set_peer("remote_host:1234"); + transfer->add_dst_buffers()->set_remote_address("remote_host:1234"); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); ASSERT_TRUE(status.ok()); @@ -398,7 +398,7 @@ TEST_F(WorkerServiceTest, TransferBuffersRemoteH2dWithPeerSuccess) { transfer->set_dst_mem_type(rpc::MEMORY_TYPE_HBM); transfer->add_src_offsets(100); transfer->add_dst_offsets(200); - transfer->set_peer("remote_host:1234"); + transfer->add_dst_buffers()->set_remote_address("remote_host:1234"); auto status = test_server_->client->TransferBuffers(transfer_req).Await(); ASSERT_TRUE(status.ok()); diff --git a/tpu_raiden/kv_cache/BUILD b/tpu_raiden/kv_cache/BUILD index fa29dbcb..70763edd 100644 --- a/tpu_raiden/kv_cache/BUILD +++ b/tpu_raiden/kv_cache/BUILD @@ -251,6 +251,7 @@ cc_library( deps = [ ":lru_cache", ":raiden_id", + "//tpu_raiden/core:buffer", "//tpu_raiden/core:numa_thread_pool", "//tpu_raiden/core:raw_transfer_core", "//tpu_raiden/core/controller:raiden_controller", diff --git a/tpu_raiden/kv_cache/kv_cache_store.cc b/tpu_raiden/kv_cache/kv_cache_store.cc index 57d8e472..998bef89 100644 --- a/tpu_raiden/kv_cache/kv_cache_store.cc +++ b/tpu_raiden/kv_cache/kv_cache_store.cc @@ -21,7 +21,7 @@ #include #include #include -#include +#include // NOLINT(build/c++11) #include #include #include @@ -41,6 +41,7 @@ #include "grpcpp/create_channel.h" #include "grpcpp/security/credentials.h" #include "xla/tsl/concurrency/future.h" +#include "tpu_raiden/core/buffer.h" #include "tpu_raiden/core/controller/raiden_controller.h" #include "tpu_raiden/core/numa_thread_pool.h" #include "tpu_raiden/kv_cache/global_registry/global_registry_client.h" @@ -408,13 +409,23 @@ absl::Status KVCacheStore::Save(const std::vector& block_hashes) { return host_blocks_or.status(); } const auto& host_block_ids = host_blocks_or.value(); - std::vector host_block_ids_64(host_block_ids.begin(), - host_block_ids.end()); + + std::vector src_buffers; + src_buffers.reserve(src_device_block_ids.size()); + for (int64_t id : src_device_block_ids) { + src_buffers.emplace_back(id, std::vector{}, std::nullopt, + rpc::MEMORY_TYPE_HBM); + } + std::vector dst_buffers; + dst_buffers.reserve(host_block_ids.size()); + for (int id : host_block_ids) { + dst_buffers.emplace_back(id, std::vector{}, std::nullopt, + rpc::MEMORY_TYPE_DRAM); + } // Trigger transfer - tsl::Future<> future = raiden_controller_->TransferBuffers( - rpc::MEMORY_TYPE_HBM, rpc::MEMORY_TYPE_DRAM, src_device_block_ids, - host_block_ids_64); + tsl::Future<> future = + raiden_controller_->TransferBuffers(src_buffers, dst_buffers); { absl::MutexLock lock(mutex_); @@ -473,13 +484,22 @@ absl::Status KVCacheStore::Load(const std::vector& block_hashes, } } - std::vector dst_device_block_ids(device_block_ids.begin(), - device_block_ids.end()); + std::vector src_buffers; + src_buffers.reserve(src_host_block_ids.size()); + for (int64_t id : src_host_block_ids) { + src_buffers.emplace_back(id, std::vector{}, std::nullopt, + rpc::MEMORY_TYPE_DRAM); + } + std::vector dst_buffers; + dst_buffers.reserve(device_block_ids.size()); + for (int id : device_block_ids) { + dst_buffers.emplace_back(id, std::vector{}, std::nullopt, + rpc::MEMORY_TYPE_HBM); + } // Trigger transfer - tsl::Future<> future = raiden_controller_->TransferBuffers( - rpc::MEMORY_TYPE_DRAM, rpc::MEMORY_TYPE_HBM, src_host_block_ids, - dst_device_block_ids); + tsl::Future<> future = + raiden_controller_->TransferBuffers(src_buffers, dst_buffers); { absl::MutexLock lock(mutex_); diff --git a/tpu_raiden/proto/worker_service.proto b/tpu_raiden/proto/worker_service.proto index 42ee8bb7..d2ab265e 100644 --- a/tpu_raiden/proto/worker_service.proto +++ b/tpu_raiden/proto/worker_service.proto @@ -108,8 +108,7 @@ message TransferBufferSpec { tpu_raiden.rpc.MemoryType src_mem_type = 4; // Memory type of the destination buffer (e.g., DRAM for D2H, HBM for H2D). tpu_raiden.rpc.MemoryType dst_mem_type = 5; - // Optional peer address for H2H transfers. - string peer = 6; + reserved 6; // Source buffers for transfer. repeated BufferProto src_buffers = 7; // Destination buffers for transfer.