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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 18 additions & 4 deletions tpu_raiden/core/controller/raiden_controller.cc
Original file line number Diff line number Diff line change
Expand Up @@ -377,6 +377,7 @@ absl::Status RaidenController::DeallocateBuffers(
absl::StatusOr<proto::TransferBuffersRequest>
RaidenController::BuildTransferBuffersRequest(
absl::Span<const Buffer> src_buffers, absl::Span<const Buffer> dst_buffers,
absl::Span<const Buffer> staging_host_buffers,
absl::Span<const int64_t> copy_sizes) {
if (src_buffers.empty() || src_buffers.size() != dst_buffers.size()) {
return absl::InvalidArgumentError(
Expand Down Expand Up @@ -416,16 +417,26 @@ 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;
}

tsl::Future<> RaidenController::TransferBuffers(
absl::string_view worker_id, absl::Span<const Buffer> src_buffers,
absl::Span<const Buffer> dst_buffers,
absl::Span<const Buffer> staging_host_buffers,
absl::Span<const int64_t> 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());
}
Expand All @@ -447,6 +458,7 @@ tsl::Future<> RaidenController::TransferBuffers(

tsl::Future<> RaidenController::TransferBuffers(
absl::Span<const Buffer> src_buffers, absl::Span<const Buffer> dst_buffers,
absl::Span<const Buffer> staging_host_buffers,
absl::Span<const int64_t> copy_sizes) {
if (src_buffers.empty() || src_buffers.size() != dst_buffers.size()) {
return tsl::Future<>(absl::InvalidArgumentError(
Expand Down Expand Up @@ -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());
}
Expand Down
28 changes: 18 additions & 10 deletions tpu_raiden/core/controller/raiden_controller.h
Original file line number Diff line number Diff line change
Expand Up @@ -118,16 +118,23 @@ class RaidenController {
// is performed.
absl::Status AllocateTargetBlockIds(absl::Span<const int> block_ids);

// Targeted worker transfer
tsl::Future<> TransferBuffers(absl::string_view worker_id,
absl::Span<const Buffer> src_buffers,
absl::Span<const Buffer> dst_buffers,
absl::Span<const int64_t> 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<const Buffer> src_buffers,
absl::Span<const Buffer> dst_buffers,
absl::Span<const Buffer> staging_host_buffers = {},
absl::Span<const int64_t> copy_sizes = {});

// Broadcast transfer to all registered workers
tsl::Future<> TransferBuffers(absl::Span<const Buffer> src_buffers,
absl::Span<const Buffer> dst_buffers,
absl::Span<const int64_t> copy_sizes = {});
// Broadcast transfer to all registered workers (staging_host_buffers as
// above).
tsl::Future<> TransferBuffers(
absl::Span<const Buffer> src_buffers,
absl::Span<const Buffer> dst_buffers,
absl::Span<const Buffer> staging_host_buffers = {},
absl::Span<const int64_t> 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.
Expand Down Expand Up @@ -169,7 +176,8 @@ class RaidenController {
absl::StatusOr<proto::TransferBuffersRequest> BuildTransferBuffersRequest(
absl::Span<const Buffer> src_buffers,
absl::Span<const Buffer> dst_buffers,
absl::Span<const int64_t> copy_sizes);
absl::Span<const Buffer> staging_host_buffers = {},
absl::Span<const int64_t> copy_sizes = {});

void Init(absl::Span<const std::string> worker_addresses,
absl::string_view raiden_orchestrator_address,
Expand Down
5 changes: 3 additions & 2 deletions tpu_raiden/core/controller/raiden_controller_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down Expand Up @@ -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);
Expand Down
31 changes: 19 additions & 12 deletions tpu_raiden/core/controller/test_util.h
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@ struct MockTransferManager {
std::string last_peer;
std::vector<int64_t> last_src_offsets;
std::vector<int64_t> last_dst_offsets;
std::vector<int64_t> last_staging_offsets;
std::vector<int64_t> last_copy_sizes;

absl::StatusOr<raiden::PjRtCopyFuture> D2h(
Expand All @@ -67,13 +68,15 @@ struct MockTransferManager {
}

absl::StatusOr<raiden::PjRtCopyFuture> D2hWrite(
absl::string_view peer, const std::vector<int64_t>& src_offsets,
const std::vector<int64_t>& dst_offsets,
absl::string_view peer, const std::vector<int64_t>& src_device_offsets,
const std::vector<int64_t>& src_host_offsets,
const std::vector<int64_t>& dst_host_offsets,
const std::vector<int64_t>& 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();
}
Expand Down Expand Up @@ -102,25 +105,29 @@ struct MockTransferManager {
}

absl::StatusOr<raiden::PjRtCopyFuture> H2dWrite(
absl::string_view peer, const std::vector<int64_t>& src_offsets,
const std::vector<int64_t>& dst_offsets,
absl::string_view peer, const std::vector<int64_t>& src_host_offsets,
const std::vector<int64_t>& dst_host_offsets,
const std::vector<int64_t>& dst_device_offsets,
const std::vector<int64_t>& 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<raiden::PjRtCopyFuture> H2dRead(
absl::string_view peer, const std::vector<int64_t>& src_offsets,
const std::vector<int64_t>& dst_offsets,
absl::string_view peer, const std::vector<int64_t>& src_host_offsets,
const std::vector<int64_t>& dst_host_offsets,
const std::vector<int64_t>& dst_device_offsets,
const std::vector<int64_t>& 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();
}
Expand Down
24 changes: 18 additions & 6 deletions tpu_raiden/core/controller/worker_service_impl.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<int64_t> 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<RaidenTransferEndpoint> dst_remote_descriptors;
if (transfer.dst_buffers_size() > 0 &&
transfer.dst_buffers(0).remote_descriptors_size() > 0) {
Expand Down Expand Up @@ -256,22 +265,25 @@ 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);
}
} else if (is_h2d) {
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);
}
Expand Down
42 changes: 28 additions & 14 deletions tpu_raiden/core/controller/worker_service_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down Expand Up @@ -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));
Expand All @@ -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));
}
Expand All @@ -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);
Expand All @@ -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);
Expand All @@ -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));
Expand All @@ -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));
Expand All @@ -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));
}
Expand All @@ -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));
}
Expand All @@ -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));
}
Expand All @@ -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);
Expand All @@ -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));
Expand All @@ -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);
Expand All @@ -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);
Expand Down
Loading
Loading