Skip to content
Merged
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
172 changes: 53 additions & 119 deletions tpu_raiden/core/controller/raiden_controller.cc
Original file line number Diff line number Diff line change
Expand Up @@ -151,8 +151,51 @@ void RaidenController::Init(absl::Span<const std::string> worker_addresses,
absl::Span<const int64_t> dst_offsets,
absl::Span<const int64_t> copy_sizes,
absl::Span<const std::string> 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<tsl::Future<>> worker_futures;
worker_futures.reserve(workers.size());
for (size_t i = 0; i < workers.size(); ++i) {
std::optional<std::string> peer =
peers.empty() ? std::nullopt : std::make_optional(peers[i]);
std::vector<Buffer> src_buffers;
src_buffers.reserve(src_offsets.size());
for (int64_t offset : src_offsets) {
src_buffers.emplace_back(offset, std::vector<BufferShard>{},
std::nullopt, src_mem_type);
}
std::vector<Buffer> dst_buffers;
dst_buffers.reserve(dst_offsets.size());
for (int64_t offset : dst_offsets) {
dst_buffers.emplace_back(offset, std::vector<BufferShard>{}, 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
Expand All @@ -172,8 +215,10 @@ void RaidenController::Init(absl::Span<const std::string> worker_addresses,
orchestrator_client_ = std::make_unique<OrchestratorServiceClient>(
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()));
Expand All @@ -190,7 +235,8 @@ RaidenController::RaidenController(
shard_size_bytes_(shard_size_bytes),
num_total_blocks_(num_blocks),
worker_registry_(std::make_shared<core::controller::WorkerRegistry>()),
block_manager_(std::make_unique<kv_cache::LogicalBlockManager>(num_blocks)),
block_manager_(
std::make_unique<kv_cache::LogicalBlockManager>(num_blocks)),
raiden_controller_port_(raiden_controller_port) {
Init(/*worker_addresses=*/{}, raiden_orchestrator_address);
}
Expand All @@ -205,7 +251,8 @@ RaidenController::RaidenController(
shard_size_bytes_(shard_size_bytes),
num_total_blocks_(num_blocks),
worker_registry_(std::make_shared<core::controller::WorkerRegistry>()),
block_manager_(std::make_unique<kv_cache::LogicalBlockManager>(num_blocks)),
block_manager_(
std::make_unique<kv_cache::LogicalBlockManager>(num_blocks)),
raiden_controller_port_(raiden_controller_port) {
Init(worker_addresses, raiden_orchestrator_address);
}
Expand Down Expand Up @@ -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) {
Expand All @@ -414,41 +450,6 @@ RaidenController::BuildTransferBuffersRequest(
return request;
}

absl::StatusOr<proto::TransferBuffersRequest>
RaidenController::BuildRawTransferBuffersRequest(
rpc::MemoryType src_mem_type, rpc::MemoryType dst_mem_type,
absl::Span<const int64_t> src_offsets,
absl::Span<const int64_t> dst_offsets, absl::Span<const int64_t> 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<const Buffer> src_buffers,
absl::Span<const Buffer> dst_buffers,
Expand Down Expand Up @@ -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<const int64_t> src_offsets,
absl::Span<const int64_t> dst_offsets, absl::Span<const int64_t> 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<const int64_t> src_offsets,
absl::Span<const int64_t> dst_offsets, absl::Span<const int64_t> copy_sizes,
absl::Span<const std::string> 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<tsl::Future<>> 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<std::string> RaidenController::ResolvePeerController(
const rpc::RaidenIdProto& peer_id) {
Expand Down
25 changes: 3 additions & 22 deletions tpu_raiden/core/controller/raiden_controller.h
Original file line number Diff line number Diff line change
Expand Up @@ -127,22 +127,7 @@ class RaidenController {
absl::Span<const Buffer> dst_buffers,
absl::Span<const int64_t> 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<const int64_t> src_offsets,
absl::Span<const int64_t> dst_offsets,
absl::Span<const int64_t> 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<const int64_t> src_offsets,
absl::Span<const int64_t> dst_offsets,
absl::Span<const int64_t> copy_sizes = {},
absl::Span<const std::string> peers = {});


// Initiates remote read from source controller.
tsl::Future<> ReadRemote(const kv_cache::RaidenId& src_raiden_id,
Expand Down Expand Up @@ -176,16 +161,12 @@ class RaidenController {
absl::Span<const Buffer> dst_buffers,
absl::Span<const int64_t> copy_sizes);

absl::StatusOr<proto::TransferBuffersRequest> BuildRawTransferBuffersRequest(
rpc::MemoryType src_mem_type, rpc::MemoryType dst_mem_type,
absl::Span<const int64_t> src_offsets,
absl::Span<const int64_t> dst_offsets,
absl::Span<const int64_t> copy_sizes, absl::string_view peer);

void Init(absl::Span<const std::string> 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_;
Expand Down
50 changes: 8 additions & 42 deletions tpu_raiden/core/controller/raiden_controller_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -681,71 +681,37 @@ 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));

RaidenController controller(unit_, /*num_blocks=*/5, /*num_shards=*/1,
/*shard_size_bytes=*/512);
RegisterAndInitWorker(controller, "worker_0", test_server_->server_address);

std::vector<int64_t> src_offsets = {10, 30};
std::vector<int64_t> dst_offsets = {20, 40};
std::vector<int64_t> 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<int64_t> src_offsets = {100};
std::vector<int64_t> 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);
EXPECT_THAT(mock_mgr.last_src_offsets, ElementsAre(100));
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));

RaidenController controller(unit_, /*num_blocks=*/5, /*num_shards=*/1,
/*shard_size_bytes=*/512);
RegisterAndInitWorker(controller, "worker_0", test_server_->server_address);

std::vector<int64_t> src_offsets = {5};
std::vector<int64_t> 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<std::string>{"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");
Expand Down
21 changes: 6 additions & 15 deletions tpu_raiden/core/controller/worker_service_impl.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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 {
Expand All @@ -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");
Expand Down
Loading
Loading