Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 31 additions & 10 deletions backends/vulkan/runtime/graph/ComputeGraph.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@

#include <executorch/backends/vulkan/runtime/graph/ComputeGraph.h>

#include <algorithm>

#include <executorch/backends/vulkan/runtime/api/containers/StagingBuffer.h>

#include <executorch/backends/vulkan/runtime/graph/ops/impl/Staging.h>
Expand Down Expand Up @@ -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<size_t>(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) {
Expand All @@ -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<size_t>(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<int64_t>& sizes) {
if (config_.enable_memory_layout_override) {
Expand Down Expand Up @@ -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);
}
}

Expand Down Expand Up @@ -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<uint32_t>(std::ceil( \
std::max( \
Expand Down Expand Up @@ -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;
Expand All @@ -1281,7 +1303,7 @@ void ComputeGraph::resize_input(
const std::vector<int64_t>& 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(
Expand All @@ -1290,8 +1312,7 @@ void ComputeGraph::virtual_resize(
std::vector<int64_t> 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);
}
}

Expand Down
8 changes: 6 additions & 2 deletions backends/vulkan/runtime/graph/ComputeGraph.h
Original file line number Diff line number Diff line change
Expand Up @@ -204,8 +204,8 @@ class ComputeGraph final {
// List of command buffers deferred for submission
std::vector<vkapi::CommandBuffer> deferred_cmd_list_;

// Set to track which ValueRefs were updated during inference
std::unordered_set<ValueRef> updated_values_;
std::vector<uint32_t> 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).
Expand Down Expand Up @@ -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
//
Expand Down
96 changes: 96 additions & 0 deletions backends/vulkan/test/vulkan_compute_api_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<float> data_##name(utils::multiply_integers(sizes)); \
std::fill(data_##name.begin(), data_##name.end(), val); \
Expand Down
Loading