From 869b9d288da6b5939778138966c1eff1f95c7549 Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Tue, 29 Sep 2026 11:56:46 -0700 Subject: [PATCH] Support multi-pass profiling in OpenXLA FunctionalHloRunner. Add SimpleMultiPassTracer backend tracer and extract profiler classes (XSpaceProfilerInterface, HLORunnerProfiler, and multi-pass runners) into hlo_runner_profiler.{h,cc} with AllReduceOr collective synchronization. PiperOrigin-RevId: 990445598 --- tsl/profiler/lib/BUILD | 56 ++- .../lib/multi_pass_profiler_controller.cc | 203 ++++++++ .../lib/multi_pass_profiler_controller.h | 77 +++ tsl/profiler/lib/profiler_factory.cc | 39 ++ tsl/profiler/lib/profiler_factory.h | 14 + tsl/profiler/lib/profiler_factory_test.cc | 469 ++++++++++++++++++ tsl/profiler/lib/profiler_interface.h | 31 ++ tsl/profiler/lib/profiler_passes.cc | 227 +++++++++ tsl/profiler/lib/profiler_passes.h | 109 ++++ tsl/profiler/protobuf/profiler_options.proto | 5 +- 10 files changed, 1228 insertions(+), 2 deletions(-) create mode 100644 tsl/profiler/lib/multi_pass_profiler_controller.cc create mode 100644 tsl/profiler/lib/multi_pass_profiler_controller.h create mode 100644 tsl/profiler/lib/profiler_passes.cc create mode 100644 tsl/profiler/lib/profiler_passes.h diff --git a/tsl/profiler/lib/BUILD b/tsl/profiler/lib/BUILD index 5a98a3704..238caa763 100644 --- a/tsl/profiler/lib/BUILD +++ b/tsl/profiler/lib/BUILD @@ -109,10 +109,12 @@ cc_library( "//learning/brain/tfrc/executor/api:__pkg__", ]), deps = [ + ":multi_pass_profiler_controller", ":profiler_controller", ":profiler_interface", "//tsl/profiler/protobuf:profiler_options_proto_cc", "@com_google_absl//absl/base:core_headers", + "@com_google_absl//absl/log", "@com_google_absl//absl/synchronization", ], alwayslink = True, @@ -125,10 +127,12 @@ tsl_cc_test( ":profiler_factory", ":profiler_factory_impl", ":profiler_interface", + ":profiler_lock", + ":profiler_passes", "//tsl/profiler/protobuf:profiler_options_proto_cc", "//tsl/profiler/protobuf:xplane_proto_cc", - "@com_google_absl//absl/memory", "@com_google_absl//absl/status", + "@com_google_absl//absl/strings:string_view", "@com_google_googletest//:gtest_main", "@xla//xla/tsl/platform:macros", "@xla//xla/tsl/platform:test", @@ -150,6 +154,7 @@ cc_library( "//tsl/profiler/protobuf:xplane_proto_cc", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", + "@com_google_absl//absl/strings", ], ) @@ -200,6 +205,52 @@ tsl_cc_test( ], ) +cc_library( + name = "multi_pass_profiler_controller", + srcs = ["multi_pass_profiler_controller.cc"], + hdrs = ["multi_pass_profiler_controller.h"], + visibility = internal_visibility([ + "@xla//xla/tsl/profiler:internal", + ]), + deps = [ + ":profiler_interface", + "//tsl/profiler/protobuf:xplane_proto_cc", + "@com_google_absl//absl/status", + "@com_google_absl//absl/strings", + "@xla//xla/tsl/platform:logging", + ], +) + +cc_library( + name = "profiler_passes", + srcs = ["profiler_passes.cc"], + hdrs = ["profiler_passes.h"], + visibility = internal_visibility([ + "@xla//xla/tsl/profiler:internal", + "@xla//xla/tsl/profiler:friends", + ]), + deps = [ + ":profiler_factory", + ":profiler_interface", + ":profiler_lock", + "//tsl/platform:platform_port", + "//tsl/profiler/protobuf:profiler_options_proto_cc", + "//tsl/profiler/protobuf:xplane_proto_cc", + "@com_google_absl//absl/base:core_headers", + "@com_google_absl//absl/memory", + "@com_google_absl//absl/status", + "@com_google_absl//absl/strings", + "@com_google_absl//absl/synchronization", + "@xla//xla/tsl/platform:env", + "@xla//xla/tsl/platform:logging", + "@xla//xla/tsl/profiler/convert:post_process_single_host_xplane", + "@xla//xla/tsl/profiler/utils:time_utils", + "@xla//xla/tsl/profiler/utils:xplane_builder", + "@xla//xla/tsl/profiler/utils:xplane_schema", + "@xla//xla/tsl/profiler/utils:xplane_utils", + ], +) + cc_library( name = "profiler_session", hdrs = ["profiler_session.h"], @@ -399,10 +450,13 @@ cc_library( srcs = ["profiler_collection.cc"], hdrs = ["profiler_collection.h"], visibility = internal_visibility([ + "@xla//xla/backends/profiler:__pkg__", "@xla//xla/backends/profiler/plugin:__pkg__", "//learning/brain/tfrc/executor/api:__pkg__", "@xla//xla/backends/profiler/cpu:__pkg__", "@xla//xla/backends/profiler/subprocess:__pkg__", + "@xla//xla/backends/profiler/tpu:__pkg__", + "@xla//xla/tsl/profiler:internal", ]), deps = [ ":profiler_interface", diff --git a/tsl/profiler/lib/multi_pass_profiler_controller.cc b/tsl/profiler/lib/multi_pass_profiler_controller.cc new file mode 100644 index 000000000..803a7d0a5 --- /dev/null +++ b/tsl/profiler/lib/multi_pass_profiler_controller.cc @@ -0,0 +1,203 @@ +/* Copyright 2026 The TensorFlow Authors. All Rights Reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +#include "tsl/profiler/lib/multi_pass_profiler_controller.h" + +#include +#include + +#include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "xla/tsl/platform/logging.h" +#include "tsl/profiler/lib/profiler_interface.h" +#include "tsl/profiler/protobuf/xplane.pb.h" + +namespace tsl { +namespace profiler { + +MultiPassProfilerController::MultiPassProfilerController( + std::unique_ptr profiler) + : profiler_(std::move(profiler)) {} + +MultiPassProfilerController::~MultiPassProfilerController() { + // Ensure an active pass is stopped. + if (state_ == MultiPassProfilerState::kPassStarted) { + StopPass().IgnoreError(); + } + if (state_ == MultiPassProfilerState::kStarted || + state_ == MultiPassProfilerState::kPassStopped) { + profiler_->Stop().IgnoreError(); + state_ = MultiPassProfilerState::kStopped; + } +} + +bool MultiPassProfilerController::NeedMorePasses() { + if (state_ == MultiPassProfilerState::kPassStopped || + state_ == MultiPassProfilerState::kStarted || + state_ == MultiPassProfilerState::kInit) { + if (!status_.ok()) return false; + return profiler_->NeedMorePasses(); + } + return false; +} + +absl::Status MultiPassProfilerController::StartPass() { + absl::Status status; + if (state_ == MultiPassProfilerState::kStarted || + state_ == MultiPassProfilerState::kPassStopped) { + state_ = MultiPassProfilerState::kPassStarted; + if (status_.ok()) { + status = status_ = profiler_->StartPass(); + } else { + status = absl::AbortedError("Previous call returned an error."); + } + } else { + status = absl::AbortedError("StartPass called in the wrong order"); + } + if (!status.ok()) LOG(ERROR) << status; + return status; +} + +absl::Status MultiPassProfilerController::PushRange(absl::string_view name) { + absl::Status status; + if (state_ == MultiPassProfilerState::kPassStarted) { + if (status_.ok()) { + status = status_ = profiler_->PushRange(name); + if (status_.ok()) { + ++active_range_count_; + } + } else { + status = absl::AbortedError("Previous call returned an error."); + } + } else { + status = absl::AbortedError("PushRange called in the wrong order"); + } + if (!status.ok()) LOG(ERROR) << status; + return status; +} + +absl::Status MultiPassProfilerController::PopRange() { + absl::Status status; + if (state_ == MultiPassProfilerState::kPassStarted) { + if (active_range_count_ <= 0) { + if (status_.ok()) { + status = status_ = + absl::InternalError("PopRange called with no active ranges."); + } else { + status = absl::AbortedError("Previous call returned an error."); + } + } else { + --active_range_count_; + if (status_.ok()) { + status = status_ = profiler_->PopRange(); + } else { + profiler_->PopRange().IgnoreError(); + status = absl::AbortedError("Previous call returned an error."); + } + } + } else { + status = absl::AbortedError("PopRange called in the wrong order"); + } + if (!status.ok()) LOG(ERROR) << status; + return status; +} + +absl::Status MultiPassProfilerController::PopAllRanges() { + absl::Status status; + while (active_range_count_ > 0) { + status.Update(PopRange()); + } + return status; +} + +absl::Status MultiPassProfilerController::StopPass() { + absl::Status status; + if (state_ == MultiPassProfilerState::kPassStarted) { + status = PopAllRanges(); + state_ = MultiPassProfilerState::kPassStopped; + if (status_.ok()) { + status = status_ = profiler_->StopPass(); + } else { + profiler_->StopPass().IgnoreError(); + if (status.ok()) { + status = absl::AbortedError("Previous call returned an error."); + } + } + } else { + status = absl::AbortedError("StopPass called in the wrong order"); + } + if (!status.ok()) LOG(ERROR) << status; + return status; +} + +absl::Status MultiPassProfilerController::CollectData( + tensorflow::profiler::XSpace* space) { + absl::Status status; + if (state_ == MultiPassProfilerState::kInit) { + status = absl::OkStatus(); + state_ = MultiPassProfilerState::kCollectData; + } else if (state_ == MultiPassProfilerState::kStopped) { + state_ = MultiPassProfilerState::kCollectData; + if (status_.ok()) { + status = status_ = profiler_->CollectData(space); + } else { + status = absl::AbortedError("Previous call returned an error."); + } + } else { + status = absl::AbortedError("CollectData called in the wrong order."); + } + if (!status.ok()) LOG(ERROR) << status; + return status; +} + +absl::Status MultiPassProfilerController::Start() { + absl::Status status; + if (state_ == MultiPassProfilerState::kInit) { + state_ = MultiPassProfilerState::kStarted; + if (status_.ok()) { + status = status_ = profiler_->Start(); + } else { + status = absl::AbortedError("Previous call returned an error."); + } + } else { + status = absl::AbortedError("Start called in the wrong order"); + } + if (!status.ok()) LOG(ERROR) << status; + return status; +} + +absl::Status MultiPassProfilerController::Stop() { + absl::Status status; + if (state_ == MultiPassProfilerState::kPassStarted) { + status = StopPass(); + } + if (state_ == MultiPassProfilerState::kPassStopped || + state_ == MultiPassProfilerState::kStarted) { + state_ = MultiPassProfilerState::kStopped; + if (status_.ok()) { + status = status_ = profiler_->Stop(); + } else if (status.ok()) { + status = absl::AbortedError("Previous call returned an error."); + } + } else if (status.ok()) { + status = absl::AbortedError("Stop called in the wrong order"); + } + if (!status.ok()) LOG(ERROR) << status; + return status; +} + +} // namespace profiler + +} // namespace tsl diff --git a/tsl/profiler/lib/multi_pass_profiler_controller.h b/tsl/profiler/lib/multi_pass_profiler_controller.h new file mode 100644 index 000000000..9cd8a55c0 --- /dev/null +++ b/tsl/profiler/lib/multi_pass_profiler_controller.h @@ -0,0 +1,77 @@ +/* Copyright 2026 The TensorFlow Authors. All Rights Reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +#ifndef TENSORFLOW_TSL_PROFILER_LIB_MULTI_PASS_PROFILER_CONTROLLER_H_ +#define TENSORFLOW_TSL_PROFILER_LIB_MULTI_PASS_PROFILER_CONTROLLER_H_ + +#include + +#include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "tsl/profiler/lib/profiler_interface.h" +#include "tsl/profiler/protobuf/xplane.pb.h" + +namespace tsl { +namespace profiler { + +// Decorator for multi-pass profiler plugins. +// +// Tracks that calls to the underlying multi-pass profiler interface functions +// are made in the expected order: Start, StartPass, [PushRange/PopRange], +// StopPass, ..., Stop, and CollectData. Making calls in an invalid order causes +// them to be aborted. +// +// If a call to the decorated profiler fails, subsequent calls will be aborted +// and no further calls will be forwarded to the underlying profiler. +class MultiPassProfilerController : public MultiPassProfilerInterface { + public: + explicit MultiPassProfilerController( + std::unique_ptr profiler); + ~MultiPassProfilerController() override; + + // MultiPassProfilerInterface methods + bool NeedMorePasses() override; + absl::Status StartPass() override; + absl::Status PushRange(absl::string_view name) override; + absl::Status PopRange() override; + absl::Status StopPass() override; + + // ProfilerInterface methods + absl::Status Start() override; + absl::Status Stop() override; + absl::Status CollectData(tensorflow::profiler::XSpace* space) override; + + private: + enum class MultiPassProfilerState { + kInit = 0, + kStarted = 1, + kPassStarted = 2, + kPassStopped = 3, + kStopped = 4, + kCollectData = 5, + }; + + absl::Status PopAllRanges(); + + MultiPassProfilerState state_ = MultiPassProfilerState::kInit; + std::unique_ptr profiler_; + absl::Status status_; // result of calls to profiler_ + int active_range_count_ = 0; +}; + +} // namespace profiler +} // namespace tsl + +#endif // TENSORFLOW_TSL_PROFILER_LIB_MULTI_PASS_PROFILER_CONTROLLER_H_ diff --git a/tsl/profiler/lib/profiler_factory.cc b/tsl/profiler/lib/profiler_factory.cc index 09fc8f0fd..c80623073 100644 --- a/tsl/profiler/lib/profiler_factory.cc +++ b/tsl/profiler/lib/profiler_factory.cc @@ -19,7 +19,9 @@ limitations under the License. #include #include "absl/base/const_init.h" +#include "absl/log/log.h" #include "absl/synchronization/mutex.h" +#include "tsl/profiler/lib/multi_pass_profiler_controller.h" #include "tsl/profiler/lib/profiler_controller.h" #include "tsl/profiler/lib/profiler_interface.h" #include "tsl/profiler/protobuf/profiler_options.pb.h" @@ -29,12 +31,18 @@ namespace profiler { namespace { absl::Mutex mu(absl::kConstInit); +absl::Mutex mu_multipass(absl::kConstInit); std::vector* GetFactories() { static auto factories = new std::vector(); return factories; } +std::vector* GetMultiPassFactories() { + static auto factories = new std::vector(); + return factories; +} + } // namespace void RegisterProfilerFactory(ProfilerFactory factory) { @@ -42,6 +50,11 @@ void RegisterProfilerFactory(ProfilerFactory factory) { GetFactories()->push_back(std::move(factory)); } +void RegisterMultiPassProfilerFactory(MultiPassProfilerFactory factory) { + absl::MutexLock lock(mu_multipass); + GetMultiPassFactories()->push_back(std::move(factory)); +} + std::vector> CreateProfilers( const tensorflow::ProfileOptions& options) { std::vector> result; @@ -56,10 +69,36 @@ std::vector> CreateProfilers( return result; } +std::vector> +CreateMultiPassProfilers(const tensorflow::ProfileOptions& options) { + std::vector> multipass_profilers; + { + absl::MutexLock lock(mu_multipass); + VLOG(3) << "Creating MultiPassProfilers() with " + << GetMultiPassFactories()->size() << " factories."; + for (const auto& factory : *GetMultiPassFactories()) { + auto profiler = factory(options); + if (profiler == nullptr) { + continue; + } + multipass_profilers.emplace_back( + std::make_unique(std::move(profiler))); + } + VLOG(3) << "Created " << multipass_profilers.size() + << " multipass profilers."; + } + return multipass_profilers; +} + void ClearRegisteredProfilersForTest() { absl::MutexLock lock(mu); GetFactories()->clear(); } +void ClearRegisteredMultiPassProfilersForTest() { + absl::MutexLock lock(mu_multipass); + GetMultiPassFactories()->clear(); +} + } // namespace profiler } // namespace tsl diff --git a/tsl/profiler/lib/profiler_factory.h b/tsl/profiler/lib/profiler_factory.h index 8266f3216..082fa822d 100644 --- a/tsl/profiler/lib/profiler_factory.h +++ b/tsl/profiler/lib/profiler_factory.h @@ -30,16 +30,30 @@ namespace profiler { using ProfilerFactory = std::function( const tensorflow::ProfileOptions&)>; +using MultiPassProfilerFactory = + std::function( + const tensorflow::ProfileOptions&)>; + // Registers a profiler factory. Should be invoked at most once per factory. void RegisterProfilerFactory(ProfilerFactory factory); +// Registers a multi-pass profiler factory. Should be invoked at most once per +// factory. +void RegisterMultiPassProfilerFactory(MultiPassProfilerFactory factory); + // Invokes all registered profiler factories with the given options, and // returns the instantiated (non-null) profiler interfaces. std::vector> CreateProfilers( const tensorflow::ProfileOptions& options); +// Invokes all registered multi-pass profiler factories with the given options, +// and returns the instantiated (non-null) multi-pass profiler interfaces. +std::vector> +CreateMultiPassProfilers(const tensorflow::ProfileOptions& options); + // For testing only. void ClearRegisteredProfilersForTest(); +void ClearRegisteredMultiPassProfilersForTest(); } // namespace profiler } // namespace tsl diff --git a/tsl/profiler/lib/profiler_factory_test.cc b/tsl/profiler/lib/profiler_factory_test.cc index 2c504201f..57128e009 100644 --- a/tsl/profiler/lib/profiler_factory_test.cc +++ b/tsl/profiler/lib/profiler_factory_test.cc @@ -18,9 +18,12 @@ limitations under the License. #include #include "absl/status/status.h" +#include "absl/strings/string_view.h" #include "xla/tsl/platform/macros.h" #include "xla/tsl/platform/test.h" #include "tsl/profiler/lib/profiler_interface.h" +#include "tsl/profiler/lib/profiler_lock.h" +#include "tsl/profiler/lib/profiler_passes.h" #include "tsl/profiler/protobuf/profiler_options.pb.h" #include "tsl/profiler/protobuf/xplane.pb.h" @@ -97,6 +100,472 @@ TEST(ProfilerFactoryTest, FactoryClassCapturedByLambda) { EXPECT_EQ(profilers.size(), 1); } +class TestMultiPassProfiler : public MultiPassProfilerInterface { + public: + bool NeedMorePasses() override { return false; } + absl::Status StartPass() override { return absl::OkStatus(); } + absl::Status PushRange(absl::string_view name) override { + return absl::OkStatus(); + } + absl::Status PopRange() override { return absl::OkStatus(); } + absl::Status StopPass() override { return absl::OkStatus(); } + + absl::Status Start() override { return absl::OkStatus(); } + absl::Status Stop() override { return absl::OkStatus(); } + absl::Status CollectData(tensorflow::profiler::XSpace*) override { + return absl::OkStatus(); + } +}; + +std::unique_ptr TestMultiPassFactoryFunction( + const tensorflow::ProfileOptions& options) { + return std::make_unique(); +} + +TEST(ProfilerFactoryTest, MultiPassFactoryFunctionPointer) { + ClearRegisteredMultiPassProfilersForTest(); + RegisterMultiPassProfilerFactory(&TestMultiPassFactoryFunction); + auto profilers = CreateMultiPassProfilers(tensorflow::ProfileOptions()); + EXPECT_EQ(profilers.size(), 1); +} + +TEST(ProfilerFactoryTest, MultiPassFactoryLambda) { + ClearRegisteredMultiPassProfilersForTest(); + RegisterMultiPassProfilerFactory( + [](const tensorflow::ProfileOptions& options) { + return std::make_unique(); + }); + auto profilers = CreateMultiPassProfilers(tensorflow::ProfileOptions()); + EXPECT_EQ(profilers.size(), 1); +} + +std::unique_ptr NullMultiPassFactoryFunction( + const tensorflow::ProfileOptions& options) { + return nullptr; +} + +TEST(ProfilerFactoryTest, MultiPassFactoryReturnsNull) { + ClearRegisteredMultiPassProfilersForTest(); + RegisterMultiPassProfilerFactory(&NullMultiPassFactoryFunction); + auto profilers = CreateMultiPassProfilers(tensorflow::ProfileOptions()); + EXPECT_TRUE(profilers.empty()); +} + +class TrackingMultiPassProfiler : public MultiPassProfilerInterface { + public: + struct Tracker { + int start_called = 0; + int stop_called = 0; + int start_pass_called = 0; + int stop_pass_called = 0; + int push_range_called = 0; + int pop_range_called = 0; + int collect_data_called = 0; + bool fail_push_range = false; + bool fail_pop_range = false; + }; + + explicit TrackingMultiPassProfiler(Tracker* tracker = nullptr) + : tracker_(tracker) {} + + bool NeedMorePasses() override { return false; } + absl::Status StartPass() override { + start_pass_called_++; + if (tracker_) tracker_->start_pass_called++; + return absl::OkStatus(); + } + absl::Status PushRange(absl::string_view name) override { + push_range_called_++; + if (tracker_) tracker_->push_range_called++; + if (fail_push_range_ || (tracker_ && tracker_->fail_push_range)) { + return absl::InternalError("PushRange failed"); + } + return absl::OkStatus(); + } + absl::Status PopRange() override { + pop_range_called_++; + if (tracker_) tracker_->pop_range_called++; + if (fail_pop_range_ || (tracker_ && tracker_->fail_pop_range)) { + return absl::InternalError("PopRange failed"); + } + return absl::OkStatus(); + } + absl::Status StopPass() override { + stop_pass_called_++; + if (tracker_) tracker_->stop_pass_called++; + return absl::OkStatus(); + } + + absl::Status Start() override { + start_called_++; + if (tracker_) tracker_->start_called++; + return absl::OkStatus(); + } + absl::Status Stop() override { + stop_called_++; + if (tracker_) tracker_->stop_called++; + return absl::OkStatus(); + } + absl::Status CollectData(tensorflow::profiler::XSpace*) override { + collect_data_called_++; + if (tracker_) tracker_->collect_data_called++; + return absl::OkStatus(); + } + + Tracker* tracker_ = nullptr; + int start_called_ = 0; + int stop_called_ = 0; + int start_pass_called_ = 0; + int stop_pass_called_ = 0; + int push_range_called_ = 0; + int pop_range_called_ = 0; + int collect_data_called_ = 0; + bool fail_push_range_ = false; + bool fail_pop_range_ = false; +}; + +TEST(ProfilerFactoryTest, MultiPassControllerStopWhilePassActiveStopsPass) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler* raw_profiler = nullptr; + RegisterMultiPassProfilerFactory( + [&raw_profiler](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + raw_profiler = profiler.get(); + return profiler; + }); + auto profilers = CreateMultiPassProfilers(tensorflow::ProfileOptions()); + ASSERT_EQ(profilers.size(), 1); + ASSERT_TRUE(profilers[0]->Start().ok()); + ASSERT_TRUE(profilers[0]->StartPass().ok()); + EXPECT_EQ(raw_profiler->start_pass_called_, 1); + EXPECT_EQ(raw_profiler->stop_pass_called_, 0); + + // Calling Stop() while pass is started should automatically stop the pass. + EXPECT_TRUE(profilers[0]->Stop().ok()); + EXPECT_EQ(raw_profiler->stop_pass_called_, 1); + EXPECT_EQ(raw_profiler->stop_called_, 1); +} + +TEST(ProfilerFactoryTest, + MultiPassControllerErrorInPassTransitionsStateAndStopsPass) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler* raw_profiler = nullptr; + RegisterMultiPassProfilerFactory( + [&raw_profiler](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + raw_profiler = profiler.get(); + return profiler; + }); + auto profilers = CreateMultiPassProfilers(tensorflow::ProfileOptions()); + ASSERT_EQ(profilers.size(), 1); + ASSERT_TRUE(profilers[0]->Start().ok()); + ASSERT_TRUE(profilers[0]->StartPass().ok()); + raw_profiler->fail_push_range_ = true; + EXPECT_FALSE(profilers[0]->PushRange("fail").ok()); + + // StopPass() should still invoke StopPass() on underlying profiler and + // transition state. + EXPECT_FALSE(profilers[0]->StopPass().ok()); + EXPECT_EQ(raw_profiler->stop_pass_called_, 1); + + // Stop() should now handle stopping the profiler session cleanly. + EXPECT_FALSE(profilers[0]->Stop().ok()); // Latched error is returned + // No call forwarded since error latched. + EXPECT_EQ(raw_profiler->stop_called_, 0); +} + +TEST(ProfilerFactoryTest, + MultiPassControllerDestructorCleansUpActivePassOnError) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler* raw_profiler = nullptr; + RegisterMultiPassProfilerFactory( + [&raw_profiler](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + raw_profiler = profiler.get(); + return profiler; + }); + { + auto profilers = CreateMultiPassProfilers(tensorflow::ProfileOptions()); + ASSERT_EQ(profilers.size(), 1); + ASSERT_TRUE(profilers[0]->Start().ok()); + ASSERT_TRUE(profilers[0]->StartPass().ok()); + raw_profiler->fail_push_range_ = true; + EXPECT_FALSE(profilers[0]->PushRange("fail").ok()); + // Destructor runs when exiting this scope while state_ == kPassStarted with + // error. + } + EXPECT_EQ(raw_profiler->stop_pass_called_, 1); + EXPECT_EQ(raw_profiler->stop_called_, 1); +} + +#if !defined(IS_MOBILE_PLATFORM) +TEST(ProfilerPassesTest, CollectDataReleasesLockOnError) { + ClearRegisteredMultiPassProfilersForTest(); + RegisterMultiPassProfilerFactory( + [](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + profiler->fail_push_range_ = true; + return profiler; + }); + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + ASSERT_TRUE(passes->StartPass().ok()); + // Fail PushRange so status_ latches an error. + ASSERT_FALSE(passes->PushRange("fail").ok()); + ASSERT_FALSE(passes->StopPass().ok()); + + // CollectData should fail, but must release the profiler lock! + tensorflow::profiler::XSpace space; + EXPECT_FALSE(passes->CollectData(&space).ok()); + + // While passes is still in scope, another session must be able to acquire the + // lock. + EXPECT_FALSE(ProfilerLock::HasActiveSession()); + auto passes2 = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + EXPECT_TRUE(passes2->Status().ok()); +} + +TEST(ProfilerPassesTest, PopRangeWithoutActiveRangeFails) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler* raw_profiler = nullptr; + RegisterMultiPassProfilerFactory( + [&raw_profiler](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + raw_profiler = profiler.get(); + return profiler; + }); + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + ASSERT_TRUE(passes->StartPass().ok()); + EXPECT_EQ(raw_profiler->pop_range_called_, 0); + + // Calling PopRange when no range has been pushed should fail. + absl::Status status = passes->PopRange(); + EXPECT_FALSE(status.ok()); + EXPECT_EQ(status.code(), absl::StatusCode::kInternal); + // Underlying profiler PopRange should NOT have been called. + EXPECT_EQ(raw_profiler->pop_range_called_, 0); +} + +TEST(ProfilerPassesTest, StopPassPopsAllActiveRanges) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler* raw_profiler = nullptr; + RegisterMultiPassProfilerFactory( + [&raw_profiler](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + raw_profiler = profiler.get(); + return profiler; + }); + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + ASSERT_TRUE(passes->StartPass().ok()); + ASSERT_TRUE(passes->PushRange("range1").ok()); + ASSERT_TRUE(passes->PushRange("range2").ok()); + ASSERT_TRUE(passes->PushRange("range3").ok()); + EXPECT_EQ(raw_profiler->push_range_called_, 3); + EXPECT_EQ(raw_profiler->pop_range_called_, 0); + + // Calling StopPass should automatically pop all 3 active ranges. + EXPECT_TRUE(passes->StopPass().ok()); + EXPECT_EQ(raw_profiler->pop_range_called_, 3); +} + +TEST(ProfilerPassesTest, PushAndPopRangesTrackCount) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler* raw_profiler = nullptr; + RegisterMultiPassProfilerFactory( + [&raw_profiler](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + raw_profiler = profiler.get(); + return profiler; + }); + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + ASSERT_TRUE(passes->StartPass().ok()); + ASSERT_TRUE(passes->PushRange("range1").ok()); + ASSERT_TRUE(passes->PushRange("range2").ok()); + EXPECT_EQ(raw_profiler->push_range_called_, 2); + EXPECT_EQ(raw_profiler->pop_range_called_, 0); + + // Pop one range explicitly. + ASSERT_TRUE(passes->PopRange().ok()); + EXPECT_EQ(raw_profiler->pop_range_called_, 1); + + // Push another range. + ASSERT_TRUE(passes->PushRange("range3").ok()); + EXPECT_EQ(raw_profiler->push_range_called_, 3); + + // StopPass should pop the 2 remaining active ranges. + EXPECT_TRUE(passes->StopPass().ok()); + EXPECT_EQ(raw_profiler->pop_range_called_, 3); +} + +TEST(ProfilerPassesTest, PopRangeFailsWhenMorePopsThanPushes) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler* raw_profiler = nullptr; + RegisterMultiPassProfilerFactory( + [&raw_profiler](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + raw_profiler = profiler.get(); + return profiler; + }); + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + ASSERT_TRUE(passes->StartPass().ok()); + ASSERT_TRUE(passes->PushRange("range1").ok()); + ASSERT_TRUE(passes->PushRange("range2").ok()); + EXPECT_EQ(raw_profiler->push_range_called_, 2); + + // Pop both active ranges explicitly. + ASSERT_TRUE(passes->PopRange().ok()); + ASSERT_TRUE(passes->PopRange().ok()); + EXPECT_EQ(raw_profiler->pop_range_called_, 2); + + // Third PopRange should fail since active_range_count_ is 0. + absl::Status status = passes->PopRange(); + EXPECT_FALSE(status.ok()); + EXPECT_EQ(status.code(), absl::StatusCode::kInternal); + // Underlying pop_range_called_ should not have increased. + EXPECT_EQ(raw_profiler->pop_range_called_, 2); +} + +TEST(ProfilerPassesTest, DestructorPopsRemainingActiveRanges) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler::Tracker tracker; + RegisterMultiPassProfilerFactory( + [&tracker](const tensorflow::ProfileOptions& options) { + return std::make_unique(&tracker); + }); + { + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + ASSERT_TRUE(passes->StartPass().ok()); + ASSERT_TRUE(passes->PushRange("range1").ok()); + ASSERT_TRUE(passes->PushRange("range2").ok()); + EXPECT_EQ(tracker.push_range_called, 2); + EXPECT_EQ(tracker.pop_range_called, 0); + } + // Destructor should have popped the 2 active ranges. + EXPECT_EQ(tracker.pop_range_called, 2); +} + +TEST(ProfilerPassesTest, CollectDataPopsRemainingActiveRanges) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler::Tracker tracker; + RegisterMultiPassProfilerFactory( + [&tracker](const tensorflow::ProfileOptions& options) { + return std::make_unique(&tracker); + }); + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + ASSERT_TRUE(passes->StartPass().ok()); + ASSERT_TRUE(passes->PushRange("range1").ok()); + EXPECT_EQ(tracker.push_range_called, 1); + EXPECT_EQ(tracker.pop_range_called, 0); + EXPECT_EQ(tracker.stop_pass_called, 0); + tensorflow::profiler::XSpace space; + EXPECT_TRUE(passes->CollectData(&space).ok()); + EXPECT_EQ(tracker.pop_range_called, 1); + EXPECT_EQ(tracker.stop_pass_called, 1); + EXPECT_EQ(tracker.collect_data_called, 1); +} + +TEST(ProfilerPassesTest, CollectDataAfterStopPassDoesNotCallStopPassAgain) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler::Tracker tracker; + RegisterMultiPassProfilerFactory( + [&tracker](const tensorflow::ProfileOptions& options) { + return std::make_unique(&tracker); + }); + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + ASSERT_TRUE(passes->StartPass().ok()); + ASSERT_TRUE(passes->PushRange("range1").ok()); + ASSERT_TRUE(passes->StopPass().ok()); + EXPECT_EQ(tracker.pop_range_called, 1); + EXPECT_EQ(tracker.stop_pass_called, 1); + + tensorflow::profiler::XSpace space; + EXPECT_TRUE(passes->CollectData(&space).ok()); + EXPECT_EQ(tracker.pop_range_called, 1); + EXPECT_EQ(tracker.stop_pass_called, 1); + EXPECT_EQ(tracker.collect_data_called, 1); +} + +TEST(ProfilerPassesTest, MultiPassTracksAndPopsRangesPerPass) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler* raw_profiler = nullptr; + RegisterMultiPassProfilerFactory( + [&raw_profiler](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + raw_profiler = profiler.get(); + return profiler; + }); + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + + // Pass 1: Push 2 ranges, explicitly pop 1, StopPass should pop the remaining + // 1. + ASSERT_TRUE(passes->StartPass().ok()); + ASSERT_TRUE(passes->PushRange("p1_r1").ok()); + ASSERT_TRUE(passes->PushRange("p1_r2").ok()); + ASSERT_TRUE(passes->PopRange().ok()); + EXPECT_EQ(raw_profiler->push_range_called_, 2); + EXPECT_EQ(raw_profiler->pop_range_called_, 1); + EXPECT_TRUE(passes->StopPass().ok()); + EXPECT_EQ(raw_profiler->pop_range_called_, 2); + + // Pass 2: Push 3 ranges, StopPass should pop all 3. + ASSERT_TRUE(passes->StartPass().ok()); + ASSERT_TRUE(passes->PushRange("p2_r1").ok()); + ASSERT_TRUE(passes->PushRange("p2_r2").ok()); + ASSERT_TRUE(passes->PushRange("p2_r3").ok()); + EXPECT_EQ(raw_profiler->push_range_called_, 5); + EXPECT_EQ(raw_profiler->pop_range_called_, 2); + EXPECT_TRUE(passes->StopPass().ok()); + EXPECT_EQ(raw_profiler->pop_range_called_, 5); +} + +TEST(ProfilerPassesTest, PushRangeFailureDoesNotIncrementActiveRangeCount) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler* raw_profiler = nullptr; + RegisterMultiPassProfilerFactory( + [&raw_profiler](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + raw_profiler = profiler.get(); + return profiler; + }); + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + ASSERT_TRUE(passes->StartPass().ok()); + + // Make PushRange fail. + raw_profiler->fail_push_range_ = true; + EXPECT_FALSE(passes->PushRange("fail").ok()); + EXPECT_EQ(raw_profiler->push_range_called_, 1); + + // StopPass should not call PopRange since active_range_count_ was not + // incremented. + EXPECT_FALSE(passes->StopPass().ok()); + EXPECT_EQ(raw_profiler->pop_range_called_, 0); +} + +TEST(ProfilerPassesTest, StopPassPropagatesPopAllRangesError) { + ClearRegisteredMultiPassProfilersForTest(); + TrackingMultiPassProfiler* raw_profiler = nullptr; + RegisterMultiPassProfilerFactory( + [&raw_profiler](const tensorflow::ProfileOptions& options) { + auto profiler = std::make_unique(); + raw_profiler = profiler.get(); + return profiler; + }); + auto passes = ProfilerPasses::Create(ProfilerPasses::DefaultOptions()); + ASSERT_TRUE(passes->StartPass().ok()); + ASSERT_TRUE(passes->PushRange("range1").ok()); + ASSERT_TRUE(passes->PushRange("range2").ok()); + EXPECT_EQ(raw_profiler->push_range_called_, 2); + + // Make PopRange fail during StopPass -> PopAllRanges. + raw_profiler->fail_pop_range_ = true; + absl::Status status = passes->StopPass(); + EXPECT_FALSE(status.ok()); + EXPECT_EQ(status.code(), absl::StatusCode::kInternal); + // Both active ranges should still have been popped and StopPass called. + EXPECT_EQ(raw_profiler->pop_range_called_, 2); + EXPECT_EQ(raw_profiler->stop_pass_called_, 1); +} +#endif + } // namespace } // namespace profiler } // namespace tsl diff --git a/tsl/profiler/lib/profiler_interface.h b/tsl/profiler/lib/profiler_interface.h index 700166c8f..ed93bd28f 100644 --- a/tsl/profiler/lib/profiler_interface.h +++ b/tsl/profiler/lib/profiler_interface.h @@ -20,6 +20,7 @@ limitations under the License. #include "absl/status/status.h" #include "absl/status/statusor.h" +#include "absl/strings/string_view.h" #include "tsl/profiler/protobuf/xplane.pb.h" namespace tsl { @@ -63,6 +64,36 @@ class ProfilerInterface { } }; +// MultiPassProfilerInterface manages multi-pass profiling plugins. +// Implementations plan the passes (e.g. counter partitioning), configure the +// underlying hardware tracer for each pass, and aggregate results into XSpace. +// +// Unlike single-pass ProfilerInterface which is driven by Start()/Stop(), +// MultiPassProfilerInterface is driven by ProfilerPasses via NeedMorePasses(), +// StartPass(), and StopPass(), optionally with PushRange() / PopRange() to +// delimit profiling scopes within a pass. +class MultiPassProfilerInterface : public ProfilerInterface { + public: + // Returns true if there are more passes to profile. + virtual bool NeedMorePasses() = 0; + + // Starts a new profiling pass. + virtual absl::Status StartPass() = 0; + + // Pushes a named range to delimit profiling regions within the active pass. + virtual absl::Status PushRange(absl::string_view name) = 0; + + // Pops the innermost named range within the active pass. + virtual absl::Status PopRange() = 0; + + // Stops the current profiling pass. + virtual absl::Status StopPass() = 0; + + absl::Status Start() override { return absl::OkStatus(); } + + absl::Status Stop() override { return absl::OkStatus(); } +}; + } // namespace profiler } // namespace tsl diff --git a/tsl/profiler/lib/profiler_passes.cc b/tsl/profiler/lib/profiler_passes.cc new file mode 100644 index 000000000..a341a0e26 --- /dev/null +++ b/tsl/profiler/lib/profiler_passes.cc @@ -0,0 +1,227 @@ +/* Copyright 2026 The TensorFlow Authors. All Rights Reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +#include "tsl/profiler/lib/profiler_passes.h" + +#include +#include + +#include "absl/memory/memory.h" +#include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "absl/synchronization/mutex.h" +#include "xla/tsl/platform/logging.h" +#include "xla/tsl/profiler/utils/xplane_builder.h" +#include "xla/tsl/profiler/utils/xplane_schema.h" +#include "xla/tsl/profiler/utils/xplane_utils.h" +#include "tsl/profiler/lib/profiler_interface.h" +#include "tsl/profiler/protobuf/profiler_options.pb.h" +#include "tsl/profiler/protobuf/xplane.pb.h" + +#if !defined(IS_MOBILE_PLATFORM) +#include "xla/tsl/platform/env.h" +#include "xla/tsl/profiler/convert/post_process_single_host_xplane.h" +#include "xla/tsl/profiler/utils/time_utils.h" +#include "tsl/platform/host_info.h" +#include "tsl/profiler/lib/profiler_factory.h" +#include "tsl/profiler/lib/profiler_lock.h" +#endif + +namespace tsl { +namespace { + +using tensorflow::ProfileOptions; +using tensorflow::profiler::XSpace; +using ::tsl::profiler::XPlaneBuilder; + +void SetProfileOptionsIntoSpace(const ProfileOptions& options, XSpace* space) { + XPlaneBuilder xplane(profiler::FindOrAddMutablePlaneWithName( + space, tsl::profiler::kTaskEnvPlaneName)); + xplane.AddStatValue( + *xplane.GetOrCreateStatMetadata(tsl::profiler::GetTaskEnvStatTypeStr( + tsl::profiler::kEnvProfileOptions)), + options); +} + +} // namespace + +std::unique_ptr ProfilerPasses::Create( + const ProfileOptions& options) { + return absl::WrapUnique(new ProfilerPasses(options)); +} + +absl::Status ProfilerPasses::Status() { + absl::MutexLock l(&mutex_); + return status_; +} + +bool ProfilerPasses::NeedMorePasses() { + absl::MutexLock l(&mutex_); + if (!status_.ok()) return false; +#if !defined(IS_MOBILE_PLATFORM) + if (multi_pass_profiler_ == nullptr) { + return false; + } + return multi_pass_profiler_->NeedMorePasses(); +#else + return false; +#endif +} + +absl::Status ProfilerPasses::StartPass() { + absl::MutexLock l(&mutex_); +#if !defined(IS_MOBILE_PLATFORM) + if (!status_.ok()) { + return absl::InternalError("Previous operations failed"); + } + if (multi_pass_profiler_ == nullptr) { + return absl::OkStatus(); + } + + if (pass_count_++ == 0) { + first_pass_start_time_ns_ = profiler::GetCurrentTimeNanos(); + absl::Status status = multi_pass_profiler_->Start(); + if (options_.raise_error_on_start_failure()) { + status_ = status; + } else { + status.IgnoreError(); + } + if (!status_.ok()) { + return status_; + } + + status = multi_pass_profiler_->StartPass(); + if (options_.raise_error_on_start_failure()) { + status_ = status; + } else { + status.IgnoreError(); + } + return status_; + } + + // passes after first pass. + return status_ = multi_pass_profiler_->StartPass(); +#else + return status_; +#endif +} + +absl::Status ProfilerPasses::PushRange(absl::string_view name) { + absl::MutexLock l(&mutex_); + if (!status_.ok()) { + return status_; + } +#if !defined(IS_MOBILE_PLATFORM) + if (multi_pass_profiler_ != nullptr) { + status_ = multi_pass_profiler_->PushRange(name); + } +#endif + return status_; +} + +absl::Status ProfilerPasses::PopRange() { + absl::MutexLock l(&mutex_); + if (!status_.ok()) { + return status_; + } +#if !defined(IS_MOBILE_PLATFORM) + if (multi_pass_profiler_ != nullptr) { + status_ = multi_pass_profiler_->PopRange(); + } +#endif + return status_; +} + +absl::Status ProfilerPasses::StopPass() { + absl::MutexLock l(&mutex_); +#if !defined(IS_MOBILE_PLATFORM) + if (pass_count_ == 1) { + first_pass_stop_time_ns_ = profiler::GetCurrentTimeNanos(); + } + if (multi_pass_profiler_ != nullptr) { + status_.Update(multi_pass_profiler_->StopPass()); + } +#endif + return status_; +} + +#if !defined(IS_MOBILE_PLATFORM) +absl::Status ProfilerPasses::CollectDataInternal(XSpace* space) { + absl::MutexLock l(&mutex_); + VLOG(3) << "Profiler passes collecting data."; + if (multi_pass_profiler_ != nullptr) { + multi_pass_profiler_->Stop().IgnoreError(); + if (status_.ok()) { + multi_pass_profiler_->CollectData(space).IgnoreError(); + } + multi_pass_profiler_.reset(); + } + // Allow another session to start. + profiler_lock_.ReleaseIfActive(); + return status_; +} +#endif + +absl::Status ProfilerPasses::CollectData(XSpace* space) { +#if !defined(IS_MOBILE_PLATFORM) + space->add_hostnames(port::Hostname()); + if (absl::Status status = CollectDataInternal(space); !status.ok()) { + return status; + } + profiler::SetXSpacePidIfNotSet(*space, tsl::Env::Default()->GetProcessId()); + profiler::PostProcessSingleHostXSpace(space, first_pass_start_time_ns_, + first_pass_stop_time_ns_); +#endif + SetProfileOptionsIntoSpace(options_, space); + return absl::OkStatus(); +} + +ProfilerPasses::ProfilerPasses(const ProfileOptions& options) +#if defined(IS_MOBILE_PLATFORM) + : status_(absl::UnimplementedError( + "Profiler is unimplemented for mobile platforms.")) { +#else + : options_(options) { + auto profiler_lock = profiler::ProfilerLock::Acquire(); + if (!profiler_lock.ok()) { + status_ = profiler_lock.status(); + return; + } + profiler_lock_ = *std::move(profiler_lock); + + DCHECK(profiler_lock_.Active()); + VLOG(3) << "Profiler passes initializing. options.enable_multipass = " + << options_.enable_multipass(); + auto multi_pass_profilers = profiler::CreateMultiPassProfilers(options_); + if (!multi_pass_profilers.empty()) { + multi_pass_profiler_ = std::move(multi_pass_profilers[0]); + if (multi_pass_profilers.size() > 1) { + LOG(ERROR) + << "More than one multipass-planner found, only use one of them."; + } + VLOG(3) << "Found one multipass planner!"; + } else { + VLOG(3) << "No multipass planner found!"; + } +#endif +} + +ProfilerPasses::~ProfilerPasses() { +#if !defined(IS_MOBILE_PLATFORM) + VLOG(3) << "Profiler passes tear down."; +#endif +} + +} // namespace tsl diff --git a/tsl/profiler/lib/profiler_passes.h b/tsl/profiler/lib/profiler_passes.h new file mode 100644 index 000000000..f56046328 --- /dev/null +++ b/tsl/profiler/lib/profiler_passes.h @@ -0,0 +1,109 @@ +/* Copyright 2026 The TensorFlow Authors. All Rights Reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +#ifndef TENSORFLOW_TSL_PROFILER_LIB_PROFILER_PASSES_H_ +#define TENSORFLOW_TSL_PROFILER_LIB_PROFILER_PASSES_H_ + +#include +#include + +#include "absl/base/thread_annotations.h" +#include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "absl/synchronization/mutex.h" +#include "tsl/profiler/protobuf/profiler_options.pb.h" +#include "tsl/profiler/protobuf/xplane.pb.h" + +#if !defined(IS_MOBILE_PLATFORM) +#include "tsl/profiler/lib/profiler_interface.h" +#include "tsl/profiler/lib/profiler_lock.h" +#endif + +namespace tsl { + +// ProfilerPasses orchestrates multi-pass profiling for workloads that require +// programmatic replay across multiple iterations (e.g. collecting different +// sets of hardware performance counters or metrics across repeated passes). +// +// Like its counterpart `ProfilerSession` (see `profiler_session.h`), it starts +// profiling upon creation and holds a profiler lock so that at most one +// session profiles at a time. Unlike `ProfilerSession`, which performs a single +// continuous profiling session across its lifetime, `ProfilerPasses` supports +// stepping through passes via `StartPass()`, `StopPass()`, and +// `NeedMorePasses()`, and delimiting ranges with `PushRange()` / `PopRange()`. +// Profile data from all passes is aggregated into an XSpace when +// `CollectData()` is called. +class ProfilerPasses { + public: + // Creates a ProfilerPasses and starts profiling. + static std::unique_ptr Create( + const tensorflow::ProfileOptions& options); + + static tensorflow::ProfileOptions DefaultOptions() { + tensorflow::ProfileOptions options; + options.set_version(1); + options.set_device_tracer_level(1); + options.set_host_tracer_level(2); + options.set_device_type(tensorflow::ProfileOptions::UNSPECIFIED); + options.set_python_tracer_level(0); + options.set_enable_hlo_proto(true); + options.set_include_dataset_ops(true); + options.set_enable_multipass(true); + return options; + } + + ~ProfilerPasses(); + + absl::Status Status() ABSL_LOCKS_EXCLUDED(mutex_); + + bool NeedMorePasses() ABSL_LOCKS_EXCLUDED(mutex_); + + absl::Status StartPass() ABSL_LOCKS_EXCLUDED(mutex_); + + absl::Status PushRange(absl::string_view name) ABSL_LOCKS_EXCLUDED(mutex_); + + absl::Status PopRange() ABSL_LOCKS_EXCLUDED(mutex_); + + absl::Status StopPass() ABSL_LOCKS_EXCLUDED(mutex_); + + // Collects profile data into XSpace. + absl::Status CollectData(tensorflow::profiler::XSpace* space) + ABSL_LOCKS_EXCLUDED(mutex_); + + private: + explicit ProfilerPasses(const tensorflow::ProfileOptions& options); + + ProfilerPasses(const ProfilerPasses&) = delete; + ProfilerPasses& operator=(const ProfilerPasses&) = delete; + +#if !defined(IS_MOBILE_PLATFORM) + absl::Status CollectDataInternal(tensorflow::profiler::XSpace* space); + + profiler::ProfilerLock profiler_lock_ ABSL_GUARDED_BY(mutex_); + std::unique_ptr multi_pass_profiler_ + ABSL_GUARDED_BY(mutex_); +#endif + + absl::Mutex mutex_; + absl::Status status_ ABSL_GUARDED_BY(mutex_); + tensorflow::ProfileOptions options_; + int pass_count_ ABSL_GUARDED_BY(mutex_) = 0; + int64_t first_pass_start_time_ns_ = 0; + int64_t first_pass_stop_time_ns_ = 0; +}; + +} // namespace tsl + +#endif // TENSORFLOW_TSL_PROFILER_LIB_PROFILER_PASSES_H_ diff --git a/tsl/profiler/protobuf/profiler_options.proto b/tsl/profiler/protobuf/profiler_options.proto index ce1395ae9..4deee01f1 100644 --- a/tsl/profiler/protobuf/profiler_options.proto +++ b/tsl/profiler/protobuf/profiler_options.proto @@ -17,7 +17,7 @@ syntax = "proto3"; package tensorflow; -// Next ID: 16 +// Next ID: 17 message ProfileOptions { // Some default value of option are not proto3 default value. Use this version // to determine if we should use default option value instead of proto3 @@ -118,6 +118,9 @@ message ProfileOptions { // If set, this hostname will be used to name the profile file. string override_hostname = 15; + + // Programmatic replay multipass enabled or not. + bool enable_multipass = 16; } // Options for remote profiler session manager.