From 4f122ae12f65b3fd067bccae8fcd7b19d59506f2 Mon Sep 17 00:00:00 2001 From: Stephen Jia Date: Wed, 2 Sep 2026 18:35:10 -0700 Subject: [PATCH] [ET-VK][runtime] Replace resize update set with generation stamps Pull Request resolved: https://github.com/pytorch/executorch/pull/22496 Dense generation stamps replace allocation-heavy `std::unordered_set` update tracking while preserving recursive `ValueList` semantics. Authored with Codex. ghstack-source-id: 423933048 @exported-using-ghexport Differential Revision: [D118543865](https://our.internmc.facebook.com/intern/diff/D118543865/) --- .../vulkan/runtime/graph/ComputeGraph.cpp | 41 ++++++-- backends/vulkan/runtime/graph/ComputeGraph.h | 8 +- .../vulkan/test/vulkan_compute_api_test.cpp | 96 +++++++++++++++++++ 3 files changed, 133 insertions(+), 12 deletions(-) diff --git a/backends/vulkan/runtime/graph/ComputeGraph.cpp b/backends/vulkan/runtime/graph/ComputeGraph.cpp index f23d1f19c66..e337f1b9791 100644 --- a/backends/vulkan/runtime/graph/ComputeGraph.cpp +++ b/backends/vulkan/runtime/graph/ComputeGraph.cpp @@ -11,6 +11,8 @@ #include +#include + #include #include @@ -245,12 +247,12 @@ bool ComputeGraph::was_value_updated(const ValueRef idx) const noexcept { return false; } - // Check if this ValueRef itself was updated - if (updated_values_.find(idx) != updated_values_.end()) { + const size_t value_idx = static_cast(idx); + if (value_idx < value_update_generations_.size() && + value_update_generations_[value_idx] == current_update_generation_) { return true; } - // If this is a ValueList, check each ValueRef in the list if (val_is_value_list(idx)) { const auto& value_list = values_.at(idx).toConstValueList(); for (const auto& nested_idx : value_list) { @@ -263,6 +265,26 @@ bool ComputeGraph::was_value_updated(const ValueRef idx) const noexcept { return false; } +void ComputeGraph::mark_value_updated(const ValueRef idx) { + if (!is_valid_value_idx(idx)) { + return; + } + if (value_update_generations_.size() < values_.size()) { + value_update_generations_.resize(values_.size()); + } + value_update_generations_[static_cast(idx)] = + current_update_generation_; +} + +void ComputeGraph::advance_update_generation() noexcept { + current_update_generation_++; + if (current_update_generation_ == 0) { + std::fill( + value_update_generations_.begin(), value_update_generations_.end(), 0); + current_update_generation_ = 1; + } +} + utils::GPUMemoryLayout ComputeGraph::suggested_memory_layout( const std::vector& sizes) { if (config_.enable_memory_layout_override) { @@ -775,8 +797,7 @@ void ComputeGraph::set_symint(const ValueRef idx, const int32_t val) { int32_t cur_val = read_symint(idx); if (cur_val != val) { get_symint(idx)->set(val); - // Track that this ValueRef was updated - updated_values_.insert(idx); + mark_value_updated(idx); } } @@ -1047,6 +1068,8 @@ void ComputeGraph::maybe_cast_and_copy_from_staging( } void ComputeGraph::prepare() { + value_update_generations_.resize(values_.size()); + #define MERGE_FIELD(field) \ static_cast(std::ceil( \ std::max( \ @@ -1258,8 +1281,7 @@ void ComputeGraph::execute() { execute_count_++; - // Clear the set of updated values at the end of inference - updated_values_.clear(); + advance_update_generation(); // Reset the re-encoding flag at the end of inference requires_reencode_ = false; @@ -1281,7 +1303,7 @@ void ComputeGraph::resize_input( const std::vector& new_sizes) { IOValueRef io_val = inputs_.at(idx); virtual_resize(io_val.value, new_sizes); - updated_values_.insert(io_val.staging); + mark_value_updated(io_val.staging); } void ComputeGraph::virtual_resize( @@ -1290,8 +1312,7 @@ void ComputeGraph::virtual_resize( std::vector cur_sizes = sizes_of(idx); if (cur_sizes != new_sizes) { get_tensor(idx)->virtual_resize(new_sizes); - // Track that this ValueRef was updated - updated_values_.insert(idx); + mark_value_updated(idx); } } diff --git a/backends/vulkan/runtime/graph/ComputeGraph.h b/backends/vulkan/runtime/graph/ComputeGraph.h index 1de890efb38..22cd3d9e692 100644 --- a/backends/vulkan/runtime/graph/ComputeGraph.h +++ b/backends/vulkan/runtime/graph/ComputeGraph.h @@ -204,8 +204,8 @@ class ComputeGraph final { // List of command buffers deferred for submission std::vector deferred_cmd_list_; - // Set to track which ValueRefs were updated during inference - std::unordered_set updated_values_; + std::vector value_update_generations_; + uint32_t current_update_generation_ = 1; // Cache to prevent duplicate prepacking of the same weight tensor with the // same kernel. Key is (inputValueRef, kernel_name). @@ -1222,6 +1222,10 @@ class ComputeGraph final { void print_readable(); + private: + void mark_value_updated(const ValueRef idx); + void advance_update_generation() noexcept; + // // Friend classes // diff --git a/backends/vulkan/test/vulkan_compute_api_test.cpp b/backends/vulkan/test/vulkan_compute_api_test.cpp index 95776e42304..c64a330e767 100644 --- a/backends/vulkan/test/vulkan_compute_api_test.cpp +++ b/backends/vulkan/test/vulkan_compute_api_test.cpp @@ -2152,6 +2152,102 @@ TEST(VulkanComputeGraphTest, test_simple_graph_with_symint) { } } +TEST(VulkanComputeGraphTest, was_value_updated_tracks_tensor_changes) { + GraphConfig config; + ComputeGraph graph(config); + + const ValueRef tensor = graph.add_tensor({2, 4}, vkapi::kFloat); + + EXPECT_FALSE(graph.was_value_updated(kDummyValueRef)); + EXPECT_FALSE(graph.was_value_updated(tensor)); + + graph.virtual_resize(tensor, {2, 4}); + EXPECT_FALSE(graph.was_value_updated(tensor)); + + graph.virtual_resize(tensor, {1, 4}); + EXPECT_TRUE(graph.was_value_updated(tensor)); +} + +TEST(VulkanComputeGraphTest, was_value_updated_tracks_symint_changes) { + GraphConfig config; + ComputeGraph graph(config); + + const ValueRef symint = graph.add_symint(3); + + EXPECT_FALSE(graph.was_value_updated(symint)); + + graph.set_symint(symint, 3); + EXPECT_FALSE(graph.was_value_updated(symint)); + + graph.set_symint(symint, 5); + EXPECT_TRUE(graph.was_value_updated(symint)); +} + +TEST(VulkanComputeGraphTest, was_value_updated_checks_nested_value_lists) { + GraphConfig config; + ComputeGraph graph(config); + + const ValueRef unchanged = graph.add_symint(1); + const ValueRef changed = graph.add_symint(2); + const ValueRef inner_list = graph.add_value_list({unchanged, changed}); + const ValueRef outer_list = + graph.add_value_list({kDummyValueRef, inner_list}); + + EXPECT_FALSE(graph.was_value_updated(inner_list)); + EXPECT_FALSE(graph.was_value_updated(outer_list)); + + graph.set_symint(changed, 3); + + EXPECT_FALSE(graph.was_value_updated(unchanged)); + EXPECT_TRUE(graph.was_value_updated(changed)); + EXPECT_TRUE(graph.was_value_updated(inner_list)); + EXPECT_TRUE(graph.was_value_updated(outer_list)); +} + +TEST(VulkanComputeGraphTest, resize_input_marks_staging_value_updated) { + GraphConfig config; + ComputeGraph graph(config); + + const IOValueRef input = graph.add_input_tensor({2, 4}, vkapi::kFloat); + + EXPECT_FALSE(graph.was_value_updated(input.value)); + EXPECT_FALSE(graph.was_value_updated(input.staging)); + + graph.resize_input(0, {2, 4}); + + EXPECT_FALSE(graph.was_value_updated(input.value)); + EXPECT_TRUE(graph.was_value_updated(input.staging)); +} + +TEST(VulkanComputeGraphTest, execute_advances_value_update_generation) { + GraphConfig config; + ComputeGraph graph(config); + + const ValueRef symint = graph.add_symint(1); + const ValueRef values = graph.add_value_list({symint}); + + graph.prepare(); + graph.set_symint(symint, 2); + + EXPECT_TRUE(graph.was_value_updated(symint)); + EXPECT_TRUE(graph.was_value_updated(values)); + + graph.execute(); + + EXPECT_FALSE(graph.was_value_updated(symint)); + EXPECT_FALSE(graph.was_value_updated(values)); + + graph.set_symint(symint, 3); + + EXPECT_TRUE(graph.was_value_updated(symint)); + EXPECT_TRUE(graph.was_value_updated(values)); + + graph.execute(); + + EXPECT_FALSE(graph.was_value_updated(symint)); + EXPECT_FALSE(graph.was_value_updated(values)); +} + #define CREATE_WEIGHT_TENSOR(name, sizes, dtype, val) \ std::vector data_##name(utils::multiply_integers(sizes)); \ std::fill(data_##name.begin(), data_##name.end(), val); \