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
2 changes: 2 additions & 0 deletions tpu_raiden/kv_cache/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -280,6 +280,7 @@ cc_library(
"@com_google_absl//absl/status:statusor",
"@com_google_absl//absl/strings",
"@com_google_absl//absl/types:span",
"@xla//xla/tsl/concurrency:future",
],
)

Expand Down Expand Up @@ -354,6 +355,7 @@ cc_test(
"@com_google_absl//absl/status:statusor",
"@com_google_googletest//:gtest",
"@com_google_googletest//:gtest_main",
"@xla//xla/tsl/concurrency:future",
],
)

Expand Down
224 changes: 134 additions & 90 deletions tpu_raiden/kv_cache/host_offload_backend.cc
Original file line number Diff line number Diff line change
Expand Up @@ -946,119 +946,163 @@ tsl::Future<> HostOffloadBackend::Load(
}

controller::RaidenController* ctrl = nullptr;
bool is_remote = false;
{
absl::MutexLock lock(mutex_);
ctrl = raiden_controller_;
is_remote = !remote_id.empty() && remote_id != raiden_id_;
}

if (ctrl == nullptr) {
return tsl::Future<>(
absl::FailedPreconditionError("RaidenController is null"));
}

auto client_or = GetKVCacheStoreClient(remote_id);
if (!client_or.ok()) {
return tsl::Future<>(client_or.status());
}
std::shared_ptr<KVCacheStoreClient> client = std::move(client_or.value());
if (is_remote) {
auto client_or = GetKVCacheStoreClient(remote_id);
if (!client_or.ok()) {
return tsl::Future<>(client_or.status());
}
std::shared_ptr<KVCacheStoreClient> client = std::move(client_or.value());

auto host_blocks_or = ctrl->AllocateBlockIds(block_hashes.size());
if (!host_blocks_or.ok()) {
return tsl::Future<>(host_blocks_or.status());
}
std::vector<int32_t> dst_host_block_ids(host_blocks_or.value().begin(),
host_blocks_or.value().end());
auto host_blocks_or = ctrl->AllocateBlockIds(block_hashes.size());
if (!host_blocks_or.ok()) {
return tsl::Future<>(host_blocks_or.status());
}
std::vector<int32_t> dst_host_block_ids(host_blocks_or.value().begin(),
host_blocks_or.value().end());

auto [load_promise, load_future] = tsl::MakePromise<>();

tpu_raiden::rpc::RaidenIdProto client_raiden_id = ctrl->unit();
std::vector<::tpu_raiden::proto::RaidenWorkerEndpointsProto>
client_worker_endpoints = BuildLocalWorkerEndpoints(ctrl);
tsl::Future<proto::FetchResponse> fetch_future =
client->Fetch(block_hashes, device_block_ids, dst_host_block_ids,
client_raiden_id, client_worker_endpoints);

fetch_future.OnReady(
[this, remote_id, dst_host_block_ids,
dev_ids_vec = std::vector<int32_t>(device_block_ids.begin(),
device_block_ids.end()),
load_promise = std::move(load_promise)](
const absl::StatusOr<proto::FetchResponse>& response_or) mutable {
controller::RaidenController* ctrl_cb = nullptr;
{
absl::MutexLock lock(mutex_);
ctrl_cb = raiden_controller_;
}
if (!response_or.ok()) {
if (ctrl_cb) {
(void)ctrl_cb->DeallocateBlockIds(dst_host_block_ids);
}
// The peer may have restarted on a new port; drop the cached client
// so the next attempt re-resolves instead of redialling a dead one.
InvalidateStoreClient(remote_id);
load_promise.Set(response_or.status());
return;
}

auto [load_promise, load_future] = tsl::MakePromise<>();
const auto& response = response_or.value();
if (!response.failed_block_hashes().empty()) {
if (ctrl_cb) {
(void)ctrl_cb->DeallocateBlockIds(dst_host_block_ids);
}
std::string err_msg = response.error_message().empty()
? "Fetch RPC returned failed blocks"
: response.error_message();
load_promise.Set(absl::InternalError(err_msg));
return;
}

tpu_raiden::rpc::RaidenIdProto client_raiden_id = ctrl->unit();
std::vector<::tpu_raiden::proto::RaidenWorkerEndpointsProto>
client_worker_endpoints = BuildLocalWorkerEndpoints(ctrl);
tsl::Future<proto::FetchResponse> fetch_future = client->Fetch(
block_hashes, device_block_ids, dst_host_block_ids, client_raiden_id,
client_worker_endpoints);
if (dev_ids_vec.empty()) {
if (ctrl_cb) {
(void)ctrl_cb->DeallocateBlockIds(dst_host_block_ids);
}
load_promise.Set(absl::OkStatus());
return;
}

fetch_future.OnReady(
[this, remote_id, dst_host_block_ids,
dev_ids_vec = std::vector<int32_t>(device_block_ids.begin(),
device_block_ids.end()),
load_promise = std::move(load_promise)](
const absl::StatusOr<proto::FetchResponse>& response_or) mutable {
controller::RaidenController* ctrl_cb = nullptr;
{
absl::MutexLock lock(mutex_);
ctrl_cb = raiden_controller_;
}
if (!response_or.ok()) {
if (ctrl_cb) {
(void)ctrl_cb->DeallocateBlockIds(dst_host_block_ids);
std::vector<Buffer> src_buffers;
src_buffers.reserve(dst_host_block_ids.size());
for (int id : dst_host_block_ids) {
src_buffers.emplace_back(id, std::vector<BufferShard>{},
std::nullopt, rpc::MEMORY_TYPE_DRAM);
}
// The peer may have restarted on a new port; drop the cached client
// so the next attempt re-resolves instead of redialling a dead one.
InvalidateStoreClient(remote_id);
load_promise.Set(response_or.status());
return;
}

const auto& response = response_or.value();
if (!response.failed_block_hashes().empty()) {
if (ctrl_cb) {
(void)ctrl_cb->DeallocateBlockIds(dst_host_block_ids);
std::vector<Buffer> dst_buffers;
dst_buffers.reserve(dev_ids_vec.size());
for (int id : dev_ids_vec) {
dst_buffers.emplace_back(id, std::vector<BufferShard>{},
std::nullopt, rpc::MEMORY_TYPE_HBM);
}
std::string err_msg = response.error_message().empty()
? "Fetch RPC returned failed blocks"
: response.error_message();
load_promise.Set(absl::InternalError(err_msg));
return;
}

if (dev_ids_vec.empty()) {
if (ctrl_cb) {
(void)ctrl_cb->DeallocateBlockIds(dst_host_block_ids);
if (!ctrl_cb) {
load_promise.Set(
absl::FailedPreconditionError("RaidenController is null"));
return;
}
load_promise.Set(absl::OkStatus());
return;
}

std::vector<Buffer> src_buffers;
src_buffers.reserve(dst_host_block_ids.size());
for (int id : dst_host_block_ids) {
src_buffers.emplace_back(id, std::vector<BufferShard>{}, std::nullopt,
rpc::MEMORY_TYPE_DRAM);
}
tsl::Future<> h2d_future =
ctrl_cb->TransferBuffers(src_buffers, dst_buffers);

h2d_future.OnReady([this, dst_host_block_ids,
load_promise = std::move(load_promise)](
absl::Status status) mutable {
controller::RaidenController* ctrl_h2d = nullptr;
{
absl::MutexLock lock(mutex_);
ctrl_h2d = raiden_controller_;
}
if (ctrl_h2d) {
(void)ctrl_h2d->DeallocateBlockIds(dst_host_block_ids);
}
load_promise.Set(status);
});
});

std::vector<Buffer> dst_buffers;
dst_buffers.reserve(dev_ids_vec.size());
for (int id : dev_ids_vec) {
dst_buffers.emplace_back(id, std::vector<BufferShard>{}, std::nullopt,
rpc::MEMORY_TYPE_HBM);
}
return load_future;
}

if (!ctrl_cb) {
load_promise.Set(
absl::FailedPreconditionError("RaidenController is null"));
return;
}
// --- Local Host DRAM Branch ---
std::vector<int64_t> src_host_block_ids;
src_host_block_ids.reserve(block_hashes.size());
{
absl::MutexLock lock(mutex_);
for (const auto& hash : block_hashes) {
const RaidenBlockID* entry = lru_cache_.Peek(hash);
if (entry == nullptr) {
return tsl::Future<>(absl::NotFoundError(
absl::StrCat("Block hash not found in host backend: ", hash)));
}
if (entry->status != BlockStatus::HOST &&
entry->status != BlockStatus::HOST_AND_HBM) {
return tsl::Future<>(absl::FailedPreconditionError(
absl::StrCat("Block is not on host: ", hash)));
}
if (entry->host_block_id == -1) {
return tsl::Future<>(absl::FailedPreconditionError(
absl::StrCat("Block host_block_id is -1: ", hash)));
}
src_host_block_ids.push_back(entry->host_block_id);
}
}

std::vector<Buffer> 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<BufferShard>{}, std::nullopt,
rpc::MEMORY_TYPE_DRAM);
}

std::vector<Buffer> dst_buffers;
dst_buffers.reserve(device_block_ids.size());
for (int id : device_block_ids) {
dst_buffers.emplace_back(id, std::vector<BufferShard>{}, std::nullopt,
rpc::MEMORY_TYPE_HBM);
}

tsl::Future<> h2d_future =
ctrl_cb->TransferBuffers(src_buffers, dst_buffers);

h2d_future.OnReady(
[this, dst_host_block_ids, load_promise = std::move(load_promise)](
absl::Status status) mutable {
controller::RaidenController* ctrl_h2d = nullptr;
{
absl::MutexLock lock(mutex_);
ctrl_h2d = raiden_controller_;
}
if (ctrl_h2d) {
(void)ctrl_h2d->DeallocateBlockIds(dst_host_block_ids);
}
load_promise.Set(status);
});
});

return load_future;
return ctrl->TransferBuffers(src_buffers, dst_buffers);
}

void HostOffloadBackend::SetMetadataEntry(absl::string_view hash,
Expand Down
2 changes: 1 addition & 1 deletion tpu_raiden/kv_cache/host_offload_backend.h
Original file line number Diff line number Diff line change
Expand Up @@ -131,7 +131,7 @@ class HostOffloadBackend : public KVCacheStoreBackend {

tsl::Future<> Load(const RaidenId& remote_id,
absl::Span<const std::string> block_hashes,
absl::Span<const int32_t> device_block_ids = {});
absl::Span<const int32_t> device_block_ids = {}) override;

// --- Remote write (WriteRemote); see KVCacheStoreBackend for why each of
// these exists rather than reusing Lookup/Insert/Delete.
Expand Down
98 changes: 98 additions & 0 deletions tpu_raiden/kv_cache/host_offload_backend_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -562,6 +562,104 @@ TEST(HostOffloadBackendTest, LoadSuccess) {
server->Shutdown();
}

TEST(HostOffloadBackendTest, LoadLocalSuccess) {
RaidenId node_id{"node_job", "0", "data", 0};
rpc::RaidenIdProto unit_proto;
unit_proto.set_job_name(node_id.job_name);
unit_proto.set_job_replica_id(node_id.job_replica_id);
unit_proto.set_data_name(node_id.data_name);
unit_proto.set_data_replica_idx(node_id.data_replica_idx);

controller::RaidenController controller(unit_proto, /*num_blocks=*/100,
/*num_shards=*/1,
/*shard_size_bytes=*/1024);
auto test_worker_server = controller::CreateTestWorkerServer();
auto transfer_mock =
std::make_unique<controller::ShardAwareMockTransferManager>();
test_worker_server->service->SetTransferManager(
KVManagerHolder(transfer_mock.get()));

core::controller::RaidenControllerClient controller_client(
controller.controller_address());
ASSERT_OK(controller_client.RegisterWorker(
"worker_0", test_worker_server->server_address,
{{test_worker_server->server_address, {}}}));

BackendConfig config;
config.type = "HostOffloadBackend";
config.capacity = 100;
config.raiden_id = node_id;

auto backend_or = HostOffloadBackend::Create(config, &controller);
ASSERT_OK(backend_or.status());
auto backend = std::dynamic_pointer_cast<HostOffloadBackend>(*backend_or);
ASSERT_NE(backend, nullptr);

backend->Insert({"local_hash_1"},
{RaidenBlockID(node_id, 10, BlockStatus::HOST)},
/*on_host=*/true);

auto load_future = backend->Load(RaidenId{}, {"local_hash_1"}, {5});
EXPECT_OK(load_future.Await());
}

TEST(HostOffloadBackendTest, LoadLocalMissingBlockError) {
RaidenId node_id{"node_job", "0", "data", 0};
rpc::RaidenIdProto unit_proto;
unit_proto.set_job_name(node_id.job_name);
unit_proto.set_job_replica_id(node_id.job_replica_id);
unit_proto.set_data_name(node_id.data_name);
unit_proto.set_data_replica_idx(node_id.data_replica_idx);

controller::RaidenController controller(unit_proto, /*num_blocks=*/100,
/*num_shards=*/1,
/*shard_size_bytes=*/1024);
BackendConfig config;
config.type = "HostOffloadBackend";
config.capacity = 100;
config.raiden_id = node_id;

auto backend_or = HostOffloadBackend::Create(config, &controller);
ASSERT_OK(backend_or.status());
auto backend = std::dynamic_pointer_cast<HostOffloadBackend>(*backend_or);
ASSERT_NE(backend, nullptr);

auto load_future = backend->Load(RaidenId{}, {"missing_hash"}, {5});
EXPECT_THAT(load_future.Await(),
absl_testing::StatusIs(absl::StatusCode::kNotFound));
}

TEST(HostOffloadBackendTest, LoadLocalNonHostBlockError) {
RaidenId node_id{"node_job", "0", "data", 0};
rpc::RaidenIdProto unit_proto;
unit_proto.set_job_name(node_id.job_name);
unit_proto.set_job_replica_id(node_id.job_replica_id);
unit_proto.set_data_name(node_id.data_name);
unit_proto.set_data_replica_idx(node_id.data_replica_idx);

controller::RaidenController controller(unit_proto, /*num_blocks=*/100,
/*num_shards=*/1,
/*shard_size_bytes=*/1024);
BackendConfig config;
config.type = "HostOffloadBackend";
config.capacity = 100;
config.raiden_id = node_id;

auto backend_or = HostOffloadBackend::Create(config, &controller);
ASSERT_OK(backend_or.status());
auto backend = std::dynamic_pointer_cast<HostOffloadBackend>(*backend_or);
ASSERT_NE(backend, nullptr);

RaidenId remote_id{"remote_job", "0", "data", 0};
backend->Insert({"remote_hash"},
{RaidenBlockID(remote_id, 10, BlockStatus::REMOTE)},
/*on_host=*/true);

auto load_future = backend->Load(RaidenId{}, {"remote_hash"}, {5});
EXPECT_THAT(load_future.Await(),
absl_testing::StatusIs(absl::StatusCode::kFailedPrecondition));
}

TEST(HostOffloadBackendTest, StoreServerOverride) {
RaidenId local_node_id{"override_job", "0", "cache", 0};
rpc::RaidenIdProto local_unit;
Expand Down
Loading
Loading