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
3 changes: 3 additions & 0 deletions tpu_sync/core/controller/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ cc_library(
"//tpu_sync/core:transfer_program_reshard",
"//tpu_sync/kv_cache/backends:backend",
"//tpu_sync/kv_cache/backends/storage:posix_backend",
"//tpu_sync/kv_cache/backends/storage:tds_backend",
"//tpu_sync/proto:transfer_program_cc_proto",
"//tpu_sync/proto:worker_service_cc_grpc",
"//tpu_sync/proto:worker_service_cc_proto",
Expand Down Expand Up @@ -217,6 +218,7 @@ cc_library(
"//tpu_sync/kv_cache:kv_cache_store_backend_factory",
"//tpu_sync/kv_cache/backends:backend",
"//tpu_sync/kv_cache/backends/storage:posix_backend",
"//tpu_sync/kv_cache/backends/storage:tds_backend",
"@com_github_grpc_grpc//:grpc++",
"@com_google_absl//absl/container:flat_hash_map",
"@com_google_absl//absl/log:check",
Expand Down Expand Up @@ -262,6 +264,7 @@ cc_test(
"//tpu_sync/kv_cache:kv_cache_store_backend_factory",
"//tpu_sync/kv_cache/backends:backend",
"//tpu_sync/kv_cache/backends/storage:posix_backend",
"//tpu_sync/kv_cache/backends/storage:tds_backend",
"//tpu_sync/proto:worker_service_cc_proto",
"//tpu_sync/rpc:raiden_service_cc_proto",
"@com_google_absl//absl/container:flat_hash_map",
Expand Down
68 changes: 67 additions & 1 deletion tpu_sync/core/controller/raiden_controller_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,6 @@
#include <functional>
#include <memory>
#include <optional>
#include <stdexcept>
#include <string>
#include <utility>
#include <vector>
Expand Down Expand Up @@ -1517,6 +1516,73 @@ TEST_F(RaidenControllerTest, TransferBackendBuffersDispatchesOffloadAndRecall) {
EXPECT_THAT(mock_mgr.last_backend_dst_block_ids, ElementsAre(4));
}

// -----------------------------------------------------------------------------
// Verifies that `RaidenController::TransferBackendBuffers` dispatches both
// OFFLOAD ([hbm] -> [dram: User host_buf] -> storage) and RECALL
// (storage -> [dram: User host_buf] -> [hbm]) to the worker when the secondary
// backend is `"tds"` (`TdsKVBackend`).
//
// Specifically checks that:
// 1. `MockTransferManager::RegisterKVBackends` (test_util.h) creates a
// `TdsKVBackend` for backend type `"tds"`.
// 2. For `TRANSFER_DIR_OFFLOAD`, the worker maps `"tds_block0"` through the
// backend's `mapper()->MapKey()` to a path ending in the hex encoding
// (`7464735f626c6f636b30.bin`) and calls `D2hWriteToBackend` once.
// 3. For `TRANSFER_DIR_RECALL`, the worker maps `"tds_recall0"` to
// `7464735f726563616c6c30.bin` and calls `H2dReadFromBackend` once.
//
// The mock only records the calls; no storage I/O is performed.
// -----------------------------------------------------------------------------
TEST_F(RaidenControllerTest,
TransferBackendBuffersDispatchesOffloadAndRecallWithTdsBackend) {
// Step 1: Register the `"tds"` secondary storage backend on the worker's
// transfer manager so `GetKVBackend("tds")` returns a `TdsKVBackend`.
kv_cache::BackendConfig tds_cfg;
tds_cfg.type = "tds";
tds_cfg.parallelism.tp_rank = 0;
tds_cfg.parallelism.tp_size = 1;
tds_cfg.SetProperty("tp_size", "1");

MockTransferManager mock_mgr;
mock_mgr.RegisterKVBackends({tds_cfg});
test_server_->service->SetTransferManager(KVManagerHolder(&mock_mgr));

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

::tpu_sync::proto::BackendTransferSpec backend_spec;
backend_spec.set_name("tds");

// Step 2: Dispatch TRANSFER_DIR_OFFLOAD (`[hbm]` block 2 -> `[dram: User
// host_buf]` staging block 3 -> `"tds"` storage backend) and verify the
// resolved hex storage path passed to `D2HAndWriteToBackend`.
auto offload_status = controller->TransferBackendBuffers(
::tpu_sync::proto::TRANSFER_DIR_OFFLOAD, {"tds_block0"},
/*hbm_block_ids=*/{2}, /*host_block_ids=*/{3}, {backend_spec});
ABSL_EXPECT_OK(offload_status.Await());
EXPECT_EQ(mock_mgr.d2h_write_to_backend_calls, 1);
ASSERT_EQ(mock_mgr.last_d2h_backend_keys.size(), 1);
EXPECT_EQ(mock_mgr.last_d2h_backend_keys[0].block_hash, "tds_block0");
EXPECT_THAT(mock_mgr.last_d2h_backend_keys[0].resolved_key,
HasSubstr("7464735f626c6f636b30.bin")); // hex("tds_block0")

// Step 3: Dispatch TRANSFER_DIR_RECALL (`"tds"` storage backend -> `[dram:
// User host_buf]` staging block 1 -> `[hbm]` block 4) and verify the
// resolved hex storage path passed to `ReadFromBackendAndH2D`.
auto recall_status = controller->TransferBackendBuffers(
::tpu_sync::proto::TRANSFER_DIR_RECALL, {"tds_recall0"},
/*hbm_block_ids=*/{4}, /*host_block_ids=*/{1}, {backend_spec});
ABSL_EXPECT_OK(recall_status.Await());
EXPECT_EQ(mock_mgr.h2d_read_from_backend_calls, 1);
ASSERT_EQ(mock_mgr.last_h2d_backend_keys.size(), 1);
EXPECT_EQ(mock_mgr.last_h2d_backend_keys[0].block_hash, "tds_recall0");
EXPECT_THAT(mock_mgr.last_h2d_backend_keys[0].resolved_key,
HasSubstr("7464735f726563616c6c30.bin")); // hex("tds_recall0")
}

TEST_F(RaidenControllerTest, TransferBuffersBackendSpecValidationRejections) {
MockTransferManager mock_mgr;
test_server_->service->SetTransferManager(KVManagerHolder(&mock_mgr));
Expand Down
24 changes: 17 additions & 7 deletions tpu_sync/core/controller/test_util.h
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@
#include "tpu_sync/core/raw_transfer_core.h"
#include "tpu_sync/kv_cache/backends/backend.h"
#include "tpu_sync/kv_cache/backends/storage/posix_backend.h"
#include "tpu_sync/kv_cache/backends/storage/tds_backend.h"
#include "tpu_sync/kv_cache/kv_cache_store_backend_factory.h"

namespace tpu_raiden {
Expand Down Expand Up @@ -214,20 +215,29 @@ struct MockTransferManager {
void RegisterKVBackends(
absl::Span<const kv_cache::BackendConfig> backend_configs) {
for (const auto& cfg : backend_configs) {
if (!absl::EqualsIgnoreCase(
cfg.type, kv_cache::backends::storage::kPosixBackendName)) {
const bool is_posix = absl::EqualsIgnoreCase(
cfg.type, kv_cache::backends::storage::kPosixBackendName);
const bool is_tds = absl::EqualsIgnoreCase(
cfg.type, kv_cache::backends::storage::kTdsBackendName);
if (!is_posix && !is_tds) {
continue;
}
if (cfg.parallelism.tp_rank < 0) continue;
const std::string canonical_name =
std::string(kv_cache::backends::storage::kPosixBackendName);
std::string(is_tds ? kv_cache::backends::storage::kTdsBackendName
: kv_cache::backends::storage::kPosixBackendName);
if (GetKVBackend(canonical_name) != nullptr) continue;
auto props = cfg.properties;
props["tp_rank"] = absl::StrCat(cfg.parallelism.tp_rank);
auto backend =
std::make_shared<kv_cache::backends::storage::PosixKVBackend>(
canonical_name, props);
backends[canonical_name] = std::move(backend);
if (is_tds) {
backends[canonical_name] =
std::make_shared<kv_cache::backends::storage::TdsKVBackend>(
canonical_name, props);
} else {
backends[canonical_name] =
std::make_shared<kv_cache::backends::storage::PosixKVBackend>(
canonical_name, props);
}
}
}

Expand Down
5 changes: 5 additions & 0 deletions tpu_sync/kv_cache/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -149,6 +149,7 @@ cc_library(
":kv_cache_store_backend_factory",
"//tpu_sync/common:raiden_id",
"//tpu_sync/kv_cache/backends/storage:posix_backend",
"//tpu_sync/kv_cache/backends/storage:tds_backend",
"@com_google_absl//absl/log",
"@com_google_absl//absl/status",
"@com_google_absl//absl/status:statusor",
Expand All @@ -170,6 +171,7 @@ cc_test(
":kv_cache_store_wrapper",
"//tpu_sync/common:raiden_id",
"//tpu_sync/kv_cache/backends/storage:posix_backend",
"//tpu_sync/kv_cache/backends/storage:tds_backend",
"//tpu_sync/kv_cache/global_registry:test_util",
"@com_google_absl//absl/status:status_matchers",
"@com_google_absl//absl/strings",
Expand Down Expand Up @@ -283,6 +285,7 @@ cc_library(
"//tpu_sync/core:xla_raw_transfer_headers",
"//tpu_sync/kv_cache/backends:backend",
"//tpu_sync/kv_cache/backends/storage:posix_backend",
"//tpu_sync/kv_cache/backends/storage:tds_backend",
"//tpu_sync/rpc:raiden_service_cc_proto",
"//tpu_sync/telemetry:metrics_api",
"//tpu_sync/telemetry:metrics_backend",
Expand Down Expand Up @@ -562,6 +565,7 @@ cc_library(
"//tpu_sync/core/controller:raiden_controller",
"//tpu_sync/kv_cache/backends:backend",
"//tpu_sync/kv_cache/backends/storage:posix_backend",
"//tpu_sync/kv_cache/backends/storage:tds_backend",
"//tpu_sync/kv_cache/global_registry:global_registry_client_cc",
"//tpu_sync/kv_cache/reshard:reshard_service",
"//tpu_sync/rpc:raiden_service_cc_proto",
Expand Down Expand Up @@ -611,6 +615,7 @@ cc_test(
"//tpu_sync/core/controller:test_util",
"//tpu_sync/kv_cache/backends:backend",
"//tpu_sync/kv_cache/backends/storage:posix_backend",
"//tpu_sync/kv_cache/backends/storage:tds_backend",
"//tpu_sync/kv_cache/global_registry:global_registry_cc_grpc",
"//tpu_sync/kv_cache/global_registry:global_registry_client_cc",
"//tpu_sync/kv_cache/global_registry:global_registry_server_lib",
Expand Down
65 changes: 65 additions & 0 deletions tpu_sync/kv_cache/backends/storage/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -80,3 +80,68 @@ cc_test(
"@xla//xla/tsl/platform:statusor",
],
)

cc_library(
name = "tds_backend",
srcs = ["tds_backend.cc"],
hdrs = ["tds_backend.h"],
copts = [
"-fno-strict-aliasing",
"-fexceptions",
],
features = [
"-use_header_modules",
"-layering_check",
],
visibility = ["//visibility:public"],
deps = [
":posix_backend",
"//tpu_sync/core:numa_thread_pool",
"//tpu_sync/core/controller:raiden_controller",
"//tpu_sync/kv_cache:kv_cache_store_backend",
"//tpu_sync/kv_cache:kv_cache_store_backend_factory",
"//tpu_sync/kv_cache/backends:backend",
"//tpu_sync/tpudirect_storage:tdsul",
"@com_google_absl//absl/base",
"@com_google_absl//absl/base:core_headers",
"@com_google_absl//absl/cleanup",
"@com_google_absl//absl/container:flat_hash_map",
"@com_google_absl//absl/status",
"@com_google_absl//absl/status:status_macros",
"@com_google_absl//absl/status:statusor",
"@com_google_absl//absl/strings",
"@com_google_absl//absl/strings:string_view",
"@com_google_absl//absl/synchronization",
"@com_google_absl//absl/types:span",
"@xla//xla/tsl/platform:logging",
],
alwayslink = 1,
)

cc_test(
name = "tds_backend_test",
srcs = ["tds_backend_test.cc"],
copts = [
"-fno-strict-aliasing",
"-fexceptions",
],
features = [
"-use_header_modules",
"-layering_check",
],
deps = [
":posix_backend",
":tds_backend",
"//tpu_sync/kv_cache:kv_cache_store_backend",
"//tpu_sync/kv_cache:kv_cache_store_backend_factory",
"//tpu_sync/kv_cache/backends:backend",
"@com_google_absl//absl/container:flat_hash_map",
"@com_google_absl//absl/status",
"@com_google_absl//absl/status:status_matchers",
"@com_google_absl//absl/status:statusor",
"@com_google_absl//absl/strings",
"@com_google_absl//absl/synchronization",
"@com_google_absl//absl/types:span",
"@com_google_googletest//:gtest_main",
],
)
Loading
Loading