From 5cb44799c0cadd12bf6b6953312fe78dbb6db662 Mon Sep 17 00:00:00 2001 From: itarun Date: Mon, 28 Sep 2026 14:23:20 -0700 Subject: [PATCH] Declare and instrument PCIe H2D and D2H byte counter and per-device label telemetry metrics. PiperOrigin-RevId: 989843455 --- .../kv_cache_manager_with_transfer_test.cc | 67 +++-- tpu_sync/kv_cache/kv_cache_manager_base.cc | 186 ++++++++++++-- tpu_sync/kv_cache/kv_cache_manager_base.h | 9 + tpu_sync/kv_cache/kv_cache_manager_test.cc | 229 ++++++++++++++++-- tpu_sync/telemetry/metrics_api.cc | 10 +- tpu_sync/telemetry/metrics_api.h | 5 + tpu_sync/telemetry/metrics_api_test.cc | 88 +++++-- tpu_sync/telemetry/metrics_backend.h | 21 ++ .../python/telemetry_binding_test.py | 4 +- 9 files changed, 530 insertions(+), 89 deletions(-) diff --git a/tpu_sync/core/kv_cache_manager_with_transfer_test.cc b/tpu_sync/core/kv_cache_manager_with_transfer_test.cc index c1a1c37f6..1a72c892d 100644 --- a/tpu_sync/core/kv_cache_manager_with_transfer_test.cc +++ b/tpu_sync/core/kv_cache_manager_with_transfer_test.cc @@ -17,6 +17,7 @@ #include #include #include +#include #include #include #include @@ -34,6 +35,7 @@ #include "xla/tsl/platform/test.h" #include "tpu_sync/core/raw_transfer_core.h" #include "tpu_sync/core/tpu_pjrt_manager.h" +#include "tpu_sync/telemetry/metrics_api.h" #include "tpu_sync/telemetry/metrics_backend.h" #include "tpu_sync/telemetry/mock_metrics_backend.h" @@ -128,6 +130,46 @@ TEST(KVCacheManagerWithTransferTest, LocalOrchestratedTransfer) { ASSERT_THAT(buffer->GetReadyFuture().Await(), IsOk()); + struct ScopedLocalRank { + std::optional prev; + ScopedLocalRank() { + if (const char* p = std::getenv(telemetry::kLocalRankEnvVar)) { + prev = p; + } + unsetenv(telemetry::kLocalRankEnvVar); + } + ~ScopedLocalRank() { + if (prev.has_value()) { + setenv(telemetry::kLocalRankEnvVar, prev->c_str(), 1); + } else { + unsetenv(telemetry::kLocalRankEnvVar); + } + } + } scoped_rank; + + TF_ASSERT_OK_AND_ASSIGN(raiden::RaidenBufferHandle handle, + raiden::RaidenBufferHandle::Acquire(buffer.get())); + std::vector> layer_buffers = { + {std::move(handle)}}; + auto engine = std::make_unique( + layer_buffers, + /*local_port=*/std::nullopt, + /*host_blocks_to_allocate=*/std::nullopt, + /*unsafe_skip_buffer_lock=*/true, + /*parallelism=*/1, + /*host_allocator=*/nullptr, + /*node_id=*/0, + /*local_control_port=*/0, + /*max_blocks=*/2, + /*num_slots=*/2, + /*timeout_s=*/10.0); + const uint64_t expected_slice_bytes = engine->base()->slice_byte_size(); + + const auto pcie_labels_matcher = ElementsAre( + HasResolvedIpLabel(telemetry::metric_labels::kHostIp), + telemetry::MetricLabel{.key = telemetry::metric_labels::kLocalRank, + .value = "0"}); + auto mock_backend = std::make_unique(); telemetry::MockMetricsBackend* raw_mock = mock_backend.get(); EXPECT_CALL(*raw_mock, @@ -167,31 +209,22 @@ TEST(KVCacheManagerWithTransferTest, LocalOrchestratedTransfer) { ObserveHistogram(telemetry::metric_names::kD2hTransferTimeMs, IsEmpty(), Gt(0.0))) .Times(1); + EXPECT_CALL(*raw_mock, + IncrementCounter(telemetry::metric_names::kD2hBytesTotal, + pcie_labels_matcher, expected_slice_bytes)) + .Times(1); EXPECT_CALL(*raw_mock, ObserveHistogram(telemetry::metric_names::kH2dTransferTimeMs, IsEmpty(), Gt(0.0))) .Times(1); + EXPECT_CALL(*raw_mock, + IncrementCounter(telemetry::metric_names::kH2dBytesTotal, + pcie_labels_matcher, expected_slice_bytes)) + .Times(1); // Register mock backend telemetry::ScopedMetricsBackendReset scoped_metrics_reset( std::move(mock_backend)); - // Create KVCacheManagerWithTransfer - auto handle_or = raiden::RaidenBufferHandle::Acquire(buffer.get()); - std::vector> layer_buffers = { - {handle_or.value()}}; - auto engine = std::make_unique( - layer_buffers, - /*local_port=*/std::nullopt, - /*host_blocks_to_allocate=*/std::nullopt, - /*unsafe_skip_buffer_lock=*/true, - /*parallelism=*/1, - /*host_allocator=*/nullptr, - /*node_id=*/0, - /*local_control_port=*/0, - /*max_blocks=*/2, - /*num_slots=*/2, - /*timeout_s=*/10.0); - // Configure staging slots: 2 slots, max 2 blocks per slot ASSERT_THAT(engine->base()->ConfigureHostStagingSlots(2, 2), IsOk()); diff --git a/tpu_sync/kv_cache/kv_cache_manager_base.cc b/tpu_sync/kv_cache/kv_cache_manager_base.cc index a85eb8b0f..8ea2367eb 100644 --- a/tpu_sync/kv_cache/kv_cache_manager_base.cc +++ b/tpu_sync/kv_cache/kv_cache_manager_base.cc @@ -209,21 +209,95 @@ struct TransferPipelinedState { bool HasFailed() const { return has_failed.load(std::memory_order_acquire); } }; -// Joins the given PjRtCopyFutures and records the time taken to complete the -// join to the given metric name if telemetry is enabled. -raiden::PjRtCopyFuture JoinAndRecordTelemetry( - absl::Span futures, absl::Time start_time, - absl::string_view metric_name) { - auto joined_future = raiden::JoinPjRtCopyFutures(futures); - if (telemetry::RaidenMetricStore::GetGlobalMetricStore().HasBackends()) { - joined_future.OnReady([start_time, metric = std::string(metric_name)]( - const auto& result) { - if (result.ok()) { - telemetry::RaidenMetricStore::GetGlobalMetricStore().ObserveHistogram( - metric, {}, absl::ToDoubleMilliseconds(absl::Now() - start_time)); +// Computes the total bytes transferred per shard across the active layers. +uint64_t ComputeBytesPerShard( + const KVCacheManagerBase& manager, size_t max_physical_size, + absl::Span copy_sizes_major_dim, + std::optional target_layer_idx = std::nullopt) { + if (!telemetry::RaidenMetricStore::GetGlobalMetricStore().HasBackends()) { + return 0; + } + const auto& buffer_holds = manager.buffer_holds(); + uint64_t bytes_per_shard = 0; + if (copy_sizes_major_dim.empty()) { + for (size_t layer_idx = 0; layer_idx < manager.num_layers(); ++layer_idx) { + if (target_layer_idx.has_value() && *target_layer_idx != layer_idx) { + continue; } - }); + if (layer_idx < buffer_holds.size() && + buffer_holds[layer_idx].physical_size > 0) { + bytes_per_shard += buffer_holds[layer_idx].physical_size; + } else { + bytes_per_shard += max_physical_size; + } + } + return bytes_per_shard; } + + uint64_t total_blocks = 0; + for (int64_t blocks : copy_sizes_major_dim) { + if (blocks > 0) { + total_blocks += static_cast(blocks); + } + } + for (size_t layer_idx = 0; layer_idx < manager.num_layers(); ++layer_idx) { + if (target_layer_idx.has_value() && *target_layer_idx != layer_idx) { + continue; + } + const int64_t block_bytes = manager.LayerBlockByteSize(layer_idx); + if (block_bytes > 0) { + bytes_per_shard += total_blocks * static_cast(block_bytes); + } + } + return bytes_per_shard; +} + +// Joins the given PjRtCopyFutures and records transfer latency and per-device +// bytes transferred if telemetry is enabled. +raiden::PjRtCopyFuture JoinAndRecordTelemetry( + const KVCacheManagerBase& manager, + absl::Span futures, absl::Time start_time, + absl::string_view time_metric_name, absl::string_view bytes_metric_name, + uint64_t bytes_per_shard, + std::optional single_shard_idx = std::nullopt) { + raiden::PjRtCopyFuture joined_future = raiden::JoinPjRtCopyFutures(futures); + if (!telemetry::RaidenMetricStore::GetGlobalMetricStore().HasBackends() || + futures.empty()) { + return joined_future; + } + + const size_t num_shards = manager.num_shards(); + const size_t start_shard = single_shard_idx.value_or(0); + const size_t end_shard = std::min(num_shards, single_shard_idx.has_value() + ? (*single_shard_idx + 1) + : num_shards); + absl::Span active_local_ranks; + if (bytes_per_shard > 0 && end_shard > start_shard) { + const absl::Span shard_ranks = + manager.shard_local_ranks(); + const size_t clamped_start = std::min(start_shard, shard_ranks.size()); + const size_t clamped_end = std::min(end_shard, shard_ranks.size()); + active_local_ranks = + shard_ranks.subspan(clamped_start, clamped_end - clamped_start); + } + joined_future.OnReady( + [start_time, time_metric = time_metric_name, + bytes_metric = bytes_metric_name, bytes = bytes_per_shard, + host_ip = manager.local_ip(), local_ranks = active_local_ranks]( + const absl::StatusOr& result) { + if (result.ok()) { + auto& store = telemetry::RaidenMetricStore::GetGlobalMetricStore(); + store.ObserveHistogram( + time_metric, {}, + absl::ToDoubleMilliseconds(absl::Now() - start_time)); + for (absl::string_view local_rank : local_ranks) { + const telemetry::MetricLabel labels[] = { + {telemetry::metric_labels::kHostIp, host_ip}, + {telemetry::metric_labels::kLocalRank, local_rank}}; + store.IncrementCounter(bytes_metric, labels, bytes); + } + } + }); return joined_future; } @@ -399,6 +473,7 @@ KVCacheManagerBase::KVCacheManagerBase( dma_pool_ = std::make_unique(kPoolSize); push_pool_ = std::make_shared(kPoolSize); pull_pool_ = std::make_unique(kPoolSize); + InitShardLocalRanks(); InitBackgroundWorker(); UpdateAllocatedOccupancyMetric(); } @@ -500,6 +575,7 @@ KVCacheManagerBase::KVCacheManagerBase( push_pool_ = std::make_shared(kPoolSize); pull_pool_ = std::make_unique(kPoolSize); InitTransportServer(); + InitShardLocalRanks(); InitBackgroundWorker(); UpdateAllocatedOccupancyMetric(); } @@ -512,6 +588,27 @@ void KVCacheManagerBase::InitBackgroundWorker() { } } +void KVCacheManagerBase::InitShardLocalRanks() { + shard_local_ranks_.clear(); + shard_local_ranks_.reserve(num_shards_); + if (num_shards_ == 1) { + if (std::optional env_rank = + telemetry::ResolveEnvVar(telemetry::kLocalRankEnvVar)) { + shard_local_ranks_.push_back(*std::move(env_rank)); + return; + } + } + for (size_t sh = 0; sh < num_shards_; ++sh) { + const int hardware_id = + (!buffer_holds_.empty() && sh < buffer_holds_[0].holds.size() && + buffer_holds_[0].holds[sh].device != nullptr) + ? buffer_holds_[0].holds[sh].device->local_hardware_id().value() + : -1; + shard_local_ranks_.push_back(hardware_id >= 0 ? absl::StrCat(hardware_id) + : absl::StrCat(sh)); + } +} + KVCacheManagerBase::~KVCacheManagerBase() { StopTransportServer(); if (is_shared_memory_mapped()) { @@ -697,8 +794,13 @@ absl::StatusOr KVCacheManagerBase::H2dSyncDispatch( } VLOG(1) << "KVCacheManagerBase::H2d completed. Returning logical futures."; - return JoinAndRecordTelemetry(absl::MakeSpan(logical_futures), h2d_start, - telemetry::metric_names::kH2dTransferTimeMs); + return JoinAndRecordTelemetry( + *this, absl::MakeSpan(logical_futures), h2d_start, + telemetry::metric_names::kH2dTransferTimeMs, + telemetry::metric_names::kH2dBytesTotal, + ComputeBytesPerShard(*this, max_physical_size_, copy_sizes_major_dim, + target_layer_idx), + target_shard_idx); } absl::StatusOr> @@ -844,8 +946,13 @@ absl::StatusOr KVCacheManagerBase::D2hSyncDispatch( auto logical_futures, DispatchD2hChunks(src_offsets_major_dim, dst_offsets_major_dim, copy_sizes_major_dim, slot_idx, layer_idx, shard_idx)); - return JoinAndRecordTelemetry(absl::MakeSpan(logical_futures), d2h_start, - telemetry::metric_names::kD2hTransferTimeMs); + return JoinAndRecordTelemetry( + *this, absl::MakeSpan(logical_futures), d2h_start, + telemetry::metric_names::kD2hTransferTimeMs, + telemetry::metric_names::kD2hBytesTotal, + ComputeBytesPerShard(*this, max_physical_size_, copy_sizes_major_dim, + layer_idx), + shard_idx); } absl::StatusOr KVCacheManagerBase::H2dWrite( @@ -1098,8 +1205,11 @@ absl::StatusOr KVCacheManagerBase::D2hWrite( // asynchronous OnReady callback to record overall D2H transfer telemetry // once all chunk transfers complete. And it returns the // aggregated future, which we don't need in this case. - JoinAndRecordTelemetry(absl::MakeSpan(all_d2h_futures), d2h_start, - telemetry::metric_names::kD2hTransferTimeMs); + JoinAndRecordTelemetry( + *this, absl::MakeSpan(all_d2h_futures), d2h_start, + telemetry::metric_names::kD2hTransferTimeMs, + telemetry::metric_names::kD2hBytesTotal, + ComputeBytesPerShard(*this, max_physical_size_, copy_sizes_major_dim)); auto state = std::make_shared( num_chunks, std::move(promise), std::move(all_holds)); @@ -1453,9 +1563,16 @@ absl::StatusOr KVCacheManagerBase::H2dDirect( shard_futures_to_join.push_back(std::move(cf)); } } - return JoinAndRecordTelemetry(absl::MakeSpan(shard_futures_to_join), - h2d_start, - telemetry::metric_names::kH2dTransferTimeMs); + std::optional single_shard; + if (device_id >= 0) { + single_shard = static_cast(device_id); + } + return JoinAndRecordTelemetry( + *this, absl::MakeSpan(shard_futures_to_join), h2d_start, + telemetry::metric_names::kH2dTransferTimeMs, + telemetry::metric_names::kH2dBytesTotal, + ComputeBytesPerShard(*this, max_physical_size_, copy_sizes), + single_shard); } absl::StatusOr KVCacheManagerBase::D2hDirect( @@ -1468,8 +1585,16 @@ absl::StatusOr KVCacheManagerBase::D2hDirect( DispatchD2hChunks(src_offsets, dst_offsets, copy_sizes, /*slot_idx=*/std::nullopt, /*layer_idx=*/std::nullopt, /*shard_idx=*/std::nullopt, device_id)); - return JoinAndRecordTelemetry(absl::MakeSpan(futures), d2h_start, - telemetry::metric_names::kD2hTransferTimeMs); + std::optional single_shard; + if (device_id >= 0) { + single_shard = static_cast(device_id); + } + return JoinAndRecordTelemetry( + *this, absl::MakeSpan(futures), d2h_start, + telemetry::metric_names::kD2hTransferTimeMs, + telemetry::metric_names::kD2hBytesTotal, + ComputeBytesPerShard(*this, max_physical_size_, copy_sizes), + single_shard); } absl::Status KVCacheManagerBase::ConfigureHostStagingSlots( @@ -2227,6 +2352,12 @@ absl::StatusOr KVCacheManagerBase::CopyPoolBlocks( } std::vector shard_futures; + uint64_t bytes_per_shard = 0; + for (const HostDeviceExtent& extent : extents) { + if (extent.size > 0) { + bytes_per_shard += static_cast(extent.size); + } + } for (size_t sh = 0; sh < num_shards_; ++sh) { if (shard_idx.has_value() && sh != *shard_idx) { continue; @@ -2272,9 +2403,12 @@ absl::StatusOr KVCacheManagerBase::CopyPoolBlocks( } } return JoinAndRecordTelemetry( - absl::MakeSpan(shard_futures), copy_start, + *this, absl::MakeSpan(shard_futures), copy_start, device_to_host ? telemetry::metric_names::kD2hTransferTimeMs - : telemetry::metric_names::kH2dTransferTimeMs); + : telemetry::metric_names::kH2dTransferTimeMs, + device_to_host ? telemetry::metric_names::kD2hBytesTotal + : telemetry::metric_names::kH2dBytesTotal, + bytes_per_shard, shard_idx); } absl::StatusOr KVCacheManagerBase::D2hPoolBlocks( diff --git a/tpu_sync/kv_cache/kv_cache_manager_base.h b/tpu_sync/kv_cache/kv_cache_manager_base.h index 123f7817a..8f7733f8e 100644 --- a/tpu_sync/kv_cache/kv_cache_manager_base.h +++ b/tpu_sync/kv_cache/kv_cache_manager_base.h @@ -794,6 +794,11 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase { return buffer_holds_; } + // Returns the per-shard local_rank label values initialized at construction. + absl::Span shard_local_ranks() const { + return shard_local_ranks_; + } + bool has_device_buffers() const { return !buffer_holds_.empty(); } void AttachPlaceholderDeviceHoldForTest() { buffer_holds_.emplace_back(); } int parallelism() const { return parallelism_; } @@ -826,6 +831,7 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase { // Per-layer device buffer holds bundled with the layer's on-device size. // See the LayerDeviceInfo definition in the public section above. std::vector buffer_holds_; + std::vector shard_local_ranks_; // Pool table. Explicit after RegisterPools; otherwise lazily materialized // implicit pools (one per storage, tag "opaque"). pools_mu_ guards the lazy // build and replacement; hot transfer paths read the table without the lock @@ -1054,6 +1060,9 @@ class KVCacheManagerBase : public tpu_raiden::RaidenManagerBase { // enabled. void InitBackgroundWorker(); + // Initializes per-shard local_rank label strings at construction time. + void InitShardLocalRanks(); + mutable absl::Mutex backends_mu_; absl::flat_hash_map> backends_ ABSL_GUARDED_BY(backends_mu_); diff --git a/tpu_sync/kv_cache/kv_cache_manager_test.cc b/tpu_sync/kv_cache/kv_cache_manager_test.cc index 6f0916601..f3a4a7040 100644 --- a/tpu_sync/kv_cache/kv_cache_manager_test.cc +++ b/tpu_sync/kv_cache/kv_cache_manager_test.cc @@ -1269,25 +1269,77 @@ TEST(KVCacheManagerTest, BackgroundWorkerThreadDisabledByDefault) { EXPECT_EQ(manager.h2d_count_, 1); } +struct ScopedUnsetLocalRank { + std::optional saved_rank; + ScopedUnsetLocalRank() { + if (const char* prev_rank = std::getenv(telemetry::kLocalRankEnvVar)) { + saved_rank = prev_rank; + } + unsetenv(telemetry::kLocalRankEnvVar); + } + ~ScopedUnsetLocalRank() { + if (saved_rank.has_value()) { + setenv(telemetry::kLocalRankEnvVar, saved_rank->c_str(), 1); + } else { + unsetenv(telemetry::kLocalRankEnvVar); + } + } +}; + TEST(KVCacheManagerTest, TelemetryMetricsObservedWhenEnabled) { + ScopedUnsetLocalRank unset_local_rank; + auto mock_backend = std::make_unique(); auto* raw_backend = mock_backend.get(); - EXPECT_CALL( - *raw_backend, - ObserveHistogram(testing::Eq(telemetry::metric_names::kH2dTransferTimeMs), - testing::_, testing::Ge(0.0))) - .Times(testing::AtLeast(1)); - EXPECT_CALL( - *raw_backend, - ObserveHistogram(testing::Eq(telemetry::metric_names::kD2hTransferTimeMs), - testing::_, testing::Ge(0.0))) - .Times(testing::AtLeast(1)); + TestKVCacheManager manager(/*num_layers=*/2, /*num_shards=*/2, + /*slice_byte_size=*/128, /*host_blocks=*/4); + const std::string expected_ip = manager.local_ip(); + const telemetry::MetricLabel shard0_labels[] = { + {telemetry::metric_labels::kHostIp, expected_ip}, + {telemetry::metric_labels::kLocalRank, "0"}}; + const telemetry::MetricLabel shard1_labels[] = { + {telemetry::metric_labels::kHostIp, expected_ip}, + {telemetry::metric_labels::kLocalRank, "1"}}; + + EXPECT_CALL(*raw_backend, + ObserveHistogram(telemetry::metric_names::kH2dTransferTimeMs, + testing::IsEmpty(), testing::Ge(0.0))) + .Times(2); + EXPECT_CALL(*raw_backend, + IncrementCounter(telemetry::metric_names::kH2dBytesTotal, + testing::ElementsAreArray(shard0_labels), 256)) + .Times(1); + EXPECT_CALL(*raw_backend, + IncrementCounter(telemetry::metric_names::kH2dBytesTotal, + testing::ElementsAreArray(shard1_labels), 256)) + .Times(1); + // Multi-block partial transfer (sizes = {2, 1} -> 3 blocks * 2 layers * 128 B + // = 768 B per shard). + EXPECT_CALL(*raw_backend, + IncrementCounter(telemetry::metric_names::kH2dBytesTotal, + testing::ElementsAreArray(shard0_labels), 768)) + .Times(1); + EXPECT_CALL(*raw_backend, + IncrementCounter(telemetry::metric_names::kH2dBytesTotal, + testing::ElementsAreArray(shard1_labels), 768)) + .Times(1); + + EXPECT_CALL(*raw_backend, + ObserveHistogram(telemetry::metric_names::kD2hTransferTimeMs, + testing::IsEmpty(), testing::Ge(0.0))) + .Times(1); + EXPECT_CALL(*raw_backend, + IncrementCounter(telemetry::metric_names::kD2hBytesTotal, + testing::ElementsAreArray(shard0_labels), 256)) + .Times(1); + EXPECT_CALL(*raw_backend, + IncrementCounter(telemetry::metric_names::kD2hBytesTotal, + testing::ElementsAreArray(shard1_labels), 256)) + .Times(1); telemetry::ScopedMetricsBackendReset scoped_reset(std::move(mock_backend)); - TestKVCacheManager manager(/*num_layers=*/1, /*num_shards=*/1, - /*slice_byte_size=*/128, /*host_blocks=*/2); std::vector offsets = {0}; std::vector sizes = {1}; @@ -1295,11 +1347,132 @@ TEST(KVCacheManagerTest, TelemetryMetricsObservedWhenEnabled) { manager.H2d(offsets, offsets, sizes)); ABSL_EXPECT_OK(h2d_res.Await()); + TF_ASSERT_OK_AND_ASSIGN(raiden::PjRtCopyFuture h2d_multi_res, + manager.H2d({0, 2}, {0, 2}, {2, 1})); + ABSL_EXPECT_OK(h2d_multi_res.Await()); + TF_ASSERT_OK_AND_ASSIGN(raiden::PjRtCopyFuture d2h_res, manager.D2h(offsets, offsets, sizes)); ABSL_EXPECT_OK(d2h_res.Await()); } +TEST(KVCacheManagerTest, SingleShardTransferEmitsTelemetryOnlyForActiveShard) { + ScopedUnsetLocalRank unset_local_rank; + + auto mock_backend = std::make_unique(); + auto* raw_backend = mock_backend.get(); + + TestKVCacheManager manager(/*num_layers=*/2, /*num_shards=*/2, + /*slice_byte_size=*/128, /*host_blocks=*/2); + ABSL_ASSERT_OK(manager.ConfigureHostStagingSlots(/*num_slots=*/2, + /*max_major_per_slot=*/1)); + const std::string expected_ip = manager.local_ip(); + const telemetry::MetricLabel shard0_labels[] = { + {telemetry::metric_labels::kHostIp, expected_ip}, + {telemetry::metric_labels::kLocalRank, "0"}}; + const telemetry::MetricLabel shard1_labels[] = { + {telemetry::metric_labels::kHostIp, expected_ip}, + {telemetry::metric_labels::kLocalRank, "1"}}; + + EXPECT_CALL(*raw_backend, + ObserveHistogram(telemetry::metric_names::kH2dTransferTimeMs, + testing::IsEmpty(), testing::Ge(0.0))) + .Times(1); + EXPECT_CALL(*raw_backend, + IncrementCounter(telemetry::metric_names::kH2dBytesTotal, + testing::ElementsAreArray(shard1_labels), 256)) + .Times(1); + EXPECT_CALL(*raw_backend, + ObserveHistogram(telemetry::metric_names::kD2hTransferTimeMs, + testing::IsEmpty(), testing::Ge(0.0))) + .Times(1); + EXPECT_CALL(*raw_backend, + IncrementCounter(telemetry::metric_names::kD2hBytesTotal, + testing::ElementsAreArray(shard1_labels), 128)) + .Times(1); + + EXPECT_CALL( + *raw_backend, + IncrementCounter(telemetry::metric_names::kH2dBytesTotal, + testing::ElementsAreArray(shard0_labels), testing::_)) + .Times(0); + EXPECT_CALL( + *raw_backend, + IncrementCounter(telemetry::metric_names::kD2hBytesTotal, + testing::ElementsAreArray(shard0_labels), testing::_)) + .Times(0); + + telemetry::ScopedMetricsBackendReset scoped_reset(std::move(mock_backend)); + + TF_ASSERT_OK_AND_ASSIGN( + raiden::PjRtCopyFuture h2d_res, + manager.H2d({0}, {0}, {1}, /*slot_idx=*/0, /*layer_idx=*/std::nullopt, + /*shard_idx=*/1)); + ABSL_EXPECT_OK(h2d_res.Await()); + + TF_ASSERT_OK_AND_ASSIGN( + raiden::PjRtCopyFuture d2h_res, + manager.D2h({0}, {0}, {1}, /*slot_idx=*/0, /*layer_idx=*/0, + /*shard_idx=*/1)); + ABSL_EXPECT_OK(d2h_res.Await()); +} + +TEST(KVCacheManagerTest, SingleShardUsesLocalRankEnvWhileMultiShardIgnoresIt) { + ScopedUnsetLocalRank restore_local_rank; + setenv(telemetry::kLocalRankEnvVar, "3", 1); + + auto mock_backend = std::make_unique(); + auto* raw_backend = mock_backend.get(); + + TestKVCacheManager single_shard_mgr(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, + /*host_blocks=*/2); + TestKVCacheManager multi_shard_mgr(/*num_layers=*/1, /*num_shards=*/2, + /*slice_byte_size=*/128, + /*host_blocks=*/2); + const std::string single_ip = single_shard_mgr.local_ip(); + const std::string multi_ip = multi_shard_mgr.local_ip(); + const telemetry::MetricLabel single_shard_labels[] = { + {telemetry::metric_labels::kHostIp, single_ip}, + {telemetry::metric_labels::kLocalRank, "3"}}; + const telemetry::MetricLabel multi_shard0_labels[] = { + {telemetry::metric_labels::kHostIp, multi_ip}, + {telemetry::metric_labels::kLocalRank, "0"}}; + const telemetry::MetricLabel multi_shard1_labels[] = { + {telemetry::metric_labels::kHostIp, multi_ip}, + {telemetry::metric_labels::kLocalRank, "1"}}; + + EXPECT_CALL(*raw_backend, + ObserveHistogram(telemetry::metric_names::kH2dTransferTimeMs, + testing::IsEmpty(), testing::Ge(0.0))) + .Times(2); + EXPECT_CALL( + *raw_backend, + IncrementCounter(telemetry::metric_names::kH2dBytesTotal, + testing::ElementsAreArray(single_shard_labels), 128)) + .Times(1); + EXPECT_CALL( + *raw_backend, + IncrementCounter(telemetry::metric_names::kH2dBytesTotal, + testing::ElementsAreArray(multi_shard0_labels), 128)) + .Times(1); + EXPECT_CALL( + *raw_backend, + IncrementCounter(telemetry::metric_names::kH2dBytesTotal, + testing::ElementsAreArray(multi_shard1_labels), 128)) + .Times(1); + + telemetry::ScopedMetricsBackendReset scoped_reset(std::move(mock_backend)); + + TF_ASSERT_OK_AND_ASSIGN(raiden::PjRtCopyFuture single_res, + single_shard_mgr.H2d({0}, {0}, {1})); + ABSL_EXPECT_OK(single_res.Await()); + + TF_ASSERT_OK_AND_ASSIGN(raiden::PjRtCopyFuture multi_res, + multi_shard_mgr.H2d({0}, {0}, {1})); + ABSL_EXPECT_OK(multi_res.Await()); +} + TEST(KVCacheManagerTest, TelemetryMetricsSkippedWhenDisabled) { telemetry::ScopedMetricsBackendReset scoped_reset; telemetry::RaidenMetricStore::GetGlobalMetricStore().SetBackends({}); @@ -1321,9 +1494,25 @@ TEST(KVCacheManagerTest, TelemetryMetricsSkippedWhenDisabled) { } TEST(KVCacheManagerTest, D2hWritePipelinedTelemetryBatchObservation) { + ScopedUnsetLocalRank unset_local_rank; + + TestD2hKVCacheManager sender(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, /*host_blocks=*/2); + TestKVCacheManager receiver(/*num_layers=*/1, /*num_shards=*/1, + /*slice_byte_size=*/128, /*host_blocks=*/2); + const std::string expected_ip = sender.local_ip(); + const telemetry::MetricLabel sender_labels[] = { + {telemetry::metric_labels::kHostIp, expected_ip}, + {telemetry::metric_labels::kLocalRank, "0"}}; + auto mock_backend = std::make_unique(); auto* raw_backend = mock_backend.get(); + EXPECT_CALL( + *raw_backend, + IncrementCounter(testing::Ne(telemetry::metric_names::kD2hBytesTotal), + testing::_, testing::_)) + .Times(testing::AnyNumber()); EXPECT_CALL( *raw_backend, ObserveHistogram(testing::Ne(telemetry::metric_names::kD2hTransferTimeMs), @@ -1331,19 +1520,17 @@ TEST(KVCacheManagerTest, D2hWritePipelinedTelemetryBatchObservation) { .Times(2); // Exactly 1 observation for the entire batch of chunks, not 1 per chunk. - EXPECT_CALL( - *raw_backend, - ObserveHistogram(testing::Eq(telemetry::metric_names::kD2hTransferTimeMs), - testing::_, testing::Ge(0.0))) + EXPECT_CALL(*raw_backend, + ObserveHistogram(telemetry::metric_names::kD2hTransferTimeMs, + testing::IsEmpty(), testing::Ge(0.0))) + .Times(1); + EXPECT_CALL(*raw_backend, + IncrementCounter(telemetry::metric_names::kD2hBytesTotal, + testing::ElementsAreArray(sender_labels), 256)) .Times(1); telemetry::ScopedMetricsBackendReset scoped_reset(std::move(mock_backend)); - TestD2hKVCacheManager sender(/*num_layers=*/1, /*num_shards=*/1, - /*slice_byte_size=*/128, /*host_blocks=*/2); - TestKVCacheManager receiver(/*num_layers=*/1, /*num_shards=*/1, - /*slice_byte_size=*/128, /*host_blocks=*/2); - const std::optional receiver_port = receiver.local_port(); ASSERT_TRUE(receiver_port.has_value()); std::string receiver_peer = diff --git a/tpu_sync/telemetry/metrics_api.cc b/tpu_sync/telemetry/metrics_api.cc index dd694d0a5..8b05bd38c 100644 --- a/tpu_sync/telemetry/metrics_api.cc +++ b/tpu_sync/telemetry/metrics_api.cc @@ -42,9 +42,10 @@ namespace tpu_raiden::telemetry { -namespace { - std::optional ResolveEnvVar(const char* env_var) { + if (env_var == nullptr) { + return std::nullopt; + } const char* value = std::getenv(env_var); if (value != nullptr && *value != '\0') { absl::string_view trimmed = absl::StripAsciiWhitespace(value); @@ -55,6 +56,8 @@ std::optional ResolveEnvVar(const char* env_var) { return std::nullopt; } +namespace { + int ResolveExporterPort() { if (std::optional port_str = ResolveEnvVar(kPrometheusPortEnvVar)) { @@ -74,10 +77,9 @@ int ResolveExporterPort() { std::string ResolveEnvVar(const char* env_var, absl::string_view default_value) { - return ResolveEnvVar(env_var).value_or(std::string(default_value)); + return telemetry::ResolveEnvVar(env_var).value_or(std::string(default_value)); } - } // namespace RaidenMetricStore& RaidenMetricStore::GetGlobalMetricStore() { diff --git a/tpu_sync/telemetry/metrics_api.h b/tpu_sync/telemetry/metrics_api.h index 47ebb8d68..15b963a96 100644 --- a/tpu_sync/telemetry/metrics_api.h +++ b/tpu_sync/telemetry/metrics_api.h @@ -19,6 +19,7 @@ #include #include #include +#include #include #include @@ -46,6 +47,10 @@ inline constexpr absl::string_view kPrometheus = "prometheus"; inline constexpr absl::string_view kBuffered = "buffered"; // Backend names END. +// Reads and trims an environment variable, returning std::nullopt if unset or +// whitespace-only. +std::optional ResolveEnvVar(const char* env_var); + // Central Telemetry Facade for managing metrics across registered backends. // This class is thread-safe for all concurrent operations. class RaidenMetricStore { diff --git a/tpu_sync/telemetry/metrics_api_test.cc b/tpu_sync/telemetry/metrics_api_test.cc index c4f1bdb84..2b72390ca 100644 --- a/tpu_sync/telemetry/metrics_api_test.cc +++ b/tpu_sync/telemetry/metrics_api_test.cc @@ -22,6 +22,7 @@ #include // NOLINT(build/c++17) #include #include +#include #include #include // NOLINT(build/c++11) #include // NOLINT(build/c++11) @@ -160,6 +161,20 @@ TEST_F(MetricsApiTest, MetricMetadataConstants) { EXPECT_EQ(metric_metadata::kP2pTransferTimeMs.type, MetricType::kHistogram); EXPECT_THAT(metric_metadata::kP2pTransferTimeMs.label_names, IsEmpty()); + // H2dBytesTotal + EXPECT_EQ(metric_labels::kHostIp, "host_ip"); + EXPECT_EQ(metric_labels::kLocalRank, "local_rank"); + EXPECT_EQ(metric_names::kH2dBytesTotal, "h2d_bytes_total"); + EXPECT_EQ(metric_descriptions::kH2dBytesTotal, + "Cumulative bytes requested for Host DRAM to Device HBM transfers " + "that completed successfully."); + EXPECT_EQ(metric_metadata::kH2dBytesTotal.name, "h2d_bytes_total"); + EXPECT_EQ(metric_metadata::kH2dBytesTotal.description, + "Cumulative bytes requested for Host DRAM to Device HBM transfers " + "that completed successfully."); + EXPECT_EQ(metric_metadata::kH2dBytesTotal.type, MetricType::kCounter); + EXPECT_THAT(metric_metadata::kH2dBytesTotal.label_names, IsEmpty()); + // H2dTransferTimeMs EXPECT_EQ(metric_names::kH2dTransferTimeMs, "h2d_transfer_time_ms"); EXPECT_EQ(metric_descriptions::kH2dTransferTimeMs, @@ -170,6 +185,18 @@ TEST_F(MetricsApiTest, MetricMetadataConstants) { EXPECT_EQ(metric_metadata::kH2dTransferTimeMs.type, MetricType::kHistogram); EXPECT_THAT(metric_metadata::kH2dTransferTimeMs.label_names, IsEmpty()); + // D2hBytesTotal + EXPECT_EQ(metric_names::kD2hBytesTotal, "d2h_bytes_total"); + EXPECT_EQ(metric_descriptions::kD2hBytesTotal, + "Cumulative bytes requested for Device HBM to Host DRAM transfers " + "that completed successfully."); + EXPECT_EQ(metric_metadata::kD2hBytesTotal.name, "d2h_bytes_total"); + EXPECT_EQ(metric_metadata::kD2hBytesTotal.description, + "Cumulative bytes requested for Device HBM to Host DRAM transfers " + "that completed successfully."); + EXPECT_EQ(metric_metadata::kD2hBytesTotal.type, MetricType::kCounter); + EXPECT_THAT(metric_metadata::kD2hBytesTotal.label_names, IsEmpty()); + // D2hTransferTimeMs EXPECT_EQ(metric_names::kD2hTransferTimeMs, "d2h_transfer_time_ms"); EXPECT_EQ(metric_descriptions::kD2hTransferTimeMs, @@ -280,28 +307,33 @@ TEST_F(MetricsApiTest, MetricMetadataConstants) { EXPECT_EQ(metric_labels::kErrorCode, "error_code"); // All Metrics + // clang-format off EXPECT_THAT( metric_metadata::kAllMetrics, - ElementsAre(metric_metadata::kSentBytesTotal, - metric_metadata::kReceivedBytesTotal, - metric_metadata::kTransferFailuresTotal, - metric_metadata::kTransferDurationMs, - metric_metadata::kP2pTransferTimeMs, - metric_metadata::kH2dTransferTimeMs, - metric_metadata::kD2hTransferTimeMs, - metric_metadata::kBufferAllocatedBytes, - metric_metadata::kWeightSyncSentBytesTotal, - metric_metadata::kWeightSyncReceivedBytesTotal, - metric_metadata::kWeightSyncTransferFailuresTotal, - metric_metadata::kWeightSyncP2pTransferTimeMs, - metric_metadata::kWeightSyncD2hTransferTimeMs, - metric_metadata::kWeightSyncH2dTransferTimeMs, - metric_metadata::kWeightSyncPushDurationMs, - metric_metadata::kWeightSyncE2eBroadcastDurationMs, - metric_metadata::kWeightSyncBufferAllocatedBytes, - metric_metadata::kWeightSyncTilingTimeMs, - metric_metadata::kWeightSyncDetilingTimeMs, - metric_metadata::kWeightSyncScheduleGenerationTimeMs)); + ElementsAre( + metric_metadata::kSentBytesTotal, + metric_metadata::kReceivedBytesTotal, + metric_metadata::kTransferFailuresTotal, + metric_metadata::kTransferDurationMs, + metric_metadata::kP2pTransferTimeMs, + metric_metadata::kH2dBytesTotal, + metric_metadata::kH2dTransferTimeMs, + metric_metadata::kD2hBytesTotal, + metric_metadata::kD2hTransferTimeMs, + metric_metadata::kBufferAllocatedBytes, + metric_metadata::kWeightSyncSentBytesTotal, + metric_metadata::kWeightSyncReceivedBytesTotal, + metric_metadata::kWeightSyncTransferFailuresTotal, + metric_metadata::kWeightSyncP2pTransferTimeMs, + metric_metadata::kWeightSyncD2hTransferTimeMs, + metric_metadata::kWeightSyncH2dTransferTimeMs, + metric_metadata::kWeightSyncPushDurationMs, + metric_metadata::kWeightSyncE2eBroadcastDurationMs, + metric_metadata::kWeightSyncBufferAllocatedBytes, + metric_metadata::kWeightSyncTilingTimeMs, + metric_metadata::kWeightSyncDetilingTimeMs, + metric_metadata::kWeightSyncScheduleGenerationTimeMs)); + // clang-format on } TEST_F(MetricsApiTest, FastPathExitWhenNoBackends) { @@ -889,5 +921,21 @@ TEST_F(MetricsApiTest, InitializeWithLocalRankEnvironmentVariable) { EXPECT_TRUE(store_.HasBackends()); } +TEST_F(MetricsApiTest, ResolveEnvVarTrimsAndHandlesUnset) { + EXPECT_EQ(ResolveEnvVar(nullptr), std::nullopt); + { + ScopedEnvironmentVariable unset_env(kLocalRankEnvVar, nullptr); + EXPECT_EQ(ResolveEnvVar(kLocalRankEnvVar), std::nullopt); + } + { + ScopedEnvironmentVariable empty_env(kLocalRankEnvVar, " "); + EXPECT_EQ(ResolveEnvVar(kLocalRankEnvVar), std::nullopt); + } + { + ScopedEnvironmentVariable set_env(kLocalRankEnvVar, " 3 \t"); + EXPECT_EQ(ResolveEnvVar(kLocalRankEnvVar), "3"); + } +} + } // namespace } // namespace tpu_raiden::telemetry diff --git a/tpu_sync/telemetry/metrics_backend.h b/tpu_sync/telemetry/metrics_backend.h index c94bd2bd7..76d2beebf 100644 --- a/tpu_sync/telemetry/metrics_backend.h +++ b/tpu_sync/telemetry/metrics_backend.h @@ -88,7 +88,9 @@ inline constexpr absl::string_view kTransferFailuresTotal = inline constexpr absl::string_view kTransferDurationMs = "transfer_duration_ms"; inline constexpr absl::string_view kP2pTransferTimeMs = "p2p_transfer_time_ms"; +inline constexpr absl::string_view kH2dBytesTotal = "h2d_bytes_total"; inline constexpr absl::string_view kH2dTransferTimeMs = "h2d_transfer_time_ms"; +inline constexpr absl::string_view kD2hBytesTotal = "d2h_bytes_total"; inline constexpr absl::string_view kD2hTransferTimeMs = "d2h_transfer_time_ms"; inline constexpr absl::string_view kBufferAllocatedBytes = "buffer_allocated_bytes"; @@ -137,8 +139,14 @@ inline constexpr absl::string_view kBufferAllocatedBytes = "Current host DRAM buffer capacity allocated in bytes for KV cache staging " "across all layers and shards."; +inline constexpr absl::string_view kH2dBytesTotal = + "Cumulative bytes requested for Host DRAM to Device HBM transfers that " + "completed successfully."; inline constexpr absl::string_view kH2dTransferTimeMs = "Host-to-Device transfer latency in milliseconds."; +inline constexpr absl::string_view kD2hBytesTotal = + "Cumulative bytes requested for Device HBM to Host DRAM transfers that " + "completed successfully."; inline constexpr absl::string_view kD2hTransferTimeMs = "Device-to-Host transfer latency in milliseconds."; @@ -183,6 +191,7 @@ inline constexpr absl::string_view kLocalRank = "local_rank"; inline constexpr absl::string_view kSrcIp = "src_ip"; inline constexpr absl::string_view kDstIp = "dst_ip"; inline constexpr absl::string_view kUnknownIp = "unknown"; +inline constexpr absl::string_view kHostIp = "host_ip"; } // namespace metric_labels namespace metric_metadata { @@ -212,11 +221,21 @@ inline constexpr MetricMetadata kP2pTransferTimeMs{ .description = metric_descriptions::kP2pTransferTimeMs, .type = MetricType::kHistogram}; +inline constexpr MetricMetadata kH2dBytesTotal{ + .name = metric_names::kH2dBytesTotal, + .description = metric_descriptions::kH2dBytesTotal, + .type = MetricType::kCounter}; + inline constexpr MetricMetadata kH2dTransferTimeMs{ .name = metric_names::kH2dTransferTimeMs, .description = metric_descriptions::kH2dTransferTimeMs, .type = MetricType::kHistogram}; +inline constexpr MetricMetadata kD2hBytesTotal{ + .name = metric_names::kD2hBytesTotal, + .description = metric_descriptions::kD2hBytesTotal, + .type = MetricType::kCounter}; + inline constexpr MetricMetadata kD2hTransferTimeMs{ .name = metric_names::kD2hTransferTimeMs, .description = metric_descriptions::kD2hTransferTimeMs, @@ -293,7 +312,9 @@ inline constexpr MetricMetadata kAllMetrics[] = { kTransferFailuresTotal, kTransferDurationMs, kP2pTransferTimeMs, + kH2dBytesTotal, kH2dTransferTimeMs, + kD2hBytesTotal, kD2hTransferTimeMs, kBufferAllocatedBytes, kWeightSyncSentBytesTotal, diff --git a/tpu_sync/telemetry/python/telemetry_binding_test.py b/tpu_sync/telemetry/python/telemetry_binding_test.py index fea1fe140..21d21ca20 100644 --- a/tpu_sync/telemetry/python/telemetry_binding_test.py +++ b/tpu_sync/telemetry/python/telemetry_binding_test.py @@ -216,12 +216,14 @@ def test_metric_metadata_properties_repr_equality(self): self.assertEqual(metrics_by_name["transfer_failures_total"].label_names, []) self.assertEqual(metrics_by_name["transfer_duration_ms"].label_names, []) self.assertEqual(metrics_by_name["p2p_transfer_time_ms"].label_names, []) + self.assertEqual(metrics_by_name["h2d_bytes_total"].label_names, []) self.assertEqual(metrics_by_name["h2d_transfer_time_ms"].label_names, []) + self.assertEqual(metrics_by_name["d2h_bytes_total"].label_names, []) self.assertEqual(metrics_by_name["d2h_transfer_time_ms"].label_names, []) self.assertEqual(metrics_by_name["buffer_allocated_bytes"].label_names, []) # Verify weight sync metrics - self.assertEqual(len(metrics), 20) + self.assertEqual(len(metrics), 22) self.assertEqual( metrics_by_name["weight_sync_sent_bytes_total"].label_names, [] )