diff --git a/dwave/optimization/include/dwave-optimization/graph.hpp b/dwave/optimization/include/dwave-optimization/graph.hpp index 22fbc3d6..7f0bff59 100644 --- a/dwave/optimization/include/dwave-optimization/graph.hpp +++ b/dwave/optimization/include/dwave-optimization/graph.hpp @@ -40,6 +40,8 @@ class Node; struct Decision {}; class Graph { + friend class State; + public: Graph() noexcept = default; ~Graph() noexcept = default; @@ -115,6 +117,25 @@ class Graph { std::span inputs() noexcept { return inputs_; } std::span inputs() const noexcept { return inputs_; } + /// Return the decision nodes that have been "mutated", meaning that they + /// have pending changes that must be propagated before committing or + /// reverting. Equivalently, the combined descendants of the returned nodes + /// are guaranteed to be a superset of the nodes that require having + /// `propagate()` and then `commit()` or `revert()` called on them. + /// + /// Notes: + /// - Using this method in conjunction with `descendants()` and + /// `propagate()`/`commit()`/`revert()` may be inefficient compared to + /// doing these manually in the case where you know only some descendants + /// of one or more of the decision nodes are relevant, e.g. only one of the + /// `DisjointListNode` successors of `DisjointListsNode` has pending + /// changes and the rest of the `DisjointListNode`s (and their descendants) + /// can be ignored + /// - This method will return the same nodes before and after calling + /// `propagate()`. Only after committing/reverting will the returned list + /// be empty again. + std::span mutated(State& state) const; + /// All of the nodes in the graph. std::span> nodes() const { return nodes_; } @@ -139,7 +160,7 @@ class Graph { void propagate(State& state) const; /// Call the propagate method on each node in changed. Note this does not call propagate on - /// the descendents of changed. + /// the descendants of changed. void propagate(State& state, std::span changed) const; void propagate(State& state, std::vector&& changed) const; @@ -349,7 +370,7 @@ class Node { StateData* data_ptr_(State& state) const { const ssize_t index = topological_index(); assert(index >= 0 and "must be topologically sorted"); - assert(state.size() > static_cast(index) and "unexpected state length"); + assert(state.size() > index and "unexpected state length"); assert(state[index] != nullptr and "uninitialized state"); return static_cast(state[index].get()); @@ -358,7 +379,7 @@ class Node { const StateData* data_ptr_(const State& state) const { const ssize_t index = topological_index(); assert(index >= 0 and "must be topologically sorted"); - assert(state.size() > static_cast(index) and "unexpected state length"); + assert(state.size() > index and "unexpected state length"); assert(state[index] != nullptr and "uninitialized state"); return static_cast(state[index].get()); @@ -368,7 +389,7 @@ class Node { void emplace_data_ptr_(State& state, Args&&... args) const { const ssize_t index = topological_index(); assert(index >= 0 and "must be topologically sorted"); - assert(state.size() > static_cast(index) and "unexpected state length"); + assert(state.size() > index and "unexpected state length"); assert(state[index] == nullptr and "already initialized state"); state[index] = std::make_unique(std::forward(args)...); diff --git a/dwave/optimization/include/dwave-optimization/state.hpp b/dwave/optimization/include/dwave-optimization/state.hpp index 884bf972..7ebd7218 100644 --- a/dwave/optimization/include/dwave-optimization/state.hpp +++ b/dwave/optimization/include/dwave-optimization/state.hpp @@ -14,9 +14,9 @@ #pragma once +#include #include #include -#include namespace dwave::optimization { @@ -34,6 +34,32 @@ struct NodeStateData { bool mark = false; }; -using State = typename std::vector>; +// Foward declaration for storing the mutated decision nodes on State +class DecisionNode; + +class State { + friend class Graph; + + public: + State(ssize_t size = 0) : node_data_(size) {} + + template + auto& operator[](index_type index) { + return node_data_[index]; + } + + template + auto& operator[](index_type index) const { + return node_data_[index]; + } + + void resize(ssize_t size) { node_data_.resize(size); } + + ssize_t size() const { return node_data_.size(); } + + private: + std::vector> node_data_; + std::vector mutated_nodes_; +}; } // namespace dwave::optimization diff --git a/dwave/optimization/src/graph.cpp b/dwave/optimization/src/graph.cpp index 733e94ae..819195ca 100644 --- a/dwave/optimization/src/graph.cpp +++ b/dwave/optimization/src/graph.cpp @@ -27,6 +27,7 @@ #endif #include "dwave-optimization/array.hpp" +#include "dwave-optimization/nodes/collections.hpp" #include "dwave-optimization/nodes/constants.hpp" #include "dwave-optimization/nodes/inputs.hpp" @@ -103,9 +104,9 @@ std::vector Graph::descendants(State& state, std::vector Graph::descendants(std::vector sources) const { - State state; + State state(num_nodes()); for (ssize_t i = 0, stop = num_nodes(); i < stop; ++i) { - state.emplace_back(std::make_unique()); + state[i] = std::make_unique(); } return descendants(state, sources); } @@ -156,6 +157,39 @@ void Graph::initialize_state(State& state) { static_cast(this)->initialize_state(state); } +std::span Graph::mutated(State& state) const { + // We will want to eventually replace this implementation with an approach where + // decision nodes "eagerly" add themselves to the list of mutated nodes after they + // are mutated. This will avoid the need to iterate over all decision nodes every + // time this method is called. + state.mutated_nodes_.clear(); + + for (const DecisionNode* dec_ptr : decisions()) { + if (auto* arr_ptr = dynamic_cast(dec_ptr); arr_ptr) { + if (not arr_ptr->diff(state).empty()) { + state.mutated_nodes_.push_back(dec_ptr); + } + } else if ( + dynamic_cast(dec_ptr) or + dynamic_cast(dec_ptr) + ) { + for (const Node* suc_ptr : dec_ptr->successors()) { + const ArrayNode* arr_ptr = dynamic_cast(suc_ptr); + assert(arr_ptr and "all successors should be array nodes"); + if (not arr_ptr->diff(state).empty()) { + state.mutated_nodes_.push_back(dec_ptr); + break; + } + } + } else { + assert(false and "unknown decision node type"); + unreachable(); + } + } + + return state.mutated_nodes_; +} + void Graph::propagate(State& state) const { std::ranges::for_each(nodes(), [&state](const auto& ptr) { ptr->propagate(state); }); } @@ -284,7 +318,7 @@ ssize_t Graph::remove_unused_nodes(bool ignore_listeners) { for (auto& uptr : nodes_ | std::views::reverse) { if (uptr->topological_index_ == keep) continue; // we marked these to keep - if (uptr->successors().size() > 0) continue; // this node is used by other nodes + if (uptr->successors().size() > 0) continue; // this node is used by other nodes // We have a node with no successors and that we haven't marked it as important. // So let's mark it to be dropped later. diff --git a/dwave/optimization/src/nodes/lambda.cpp b/dwave/optimization/src/nodes/lambda.cpp index bdb0bd4d..ab4781ac 100644 --- a/dwave/optimization/src/nodes/lambda.cpp +++ b/dwave/optimization/src/nodes/lambda.cpp @@ -257,8 +257,7 @@ void AccumulateZipNode::initialize_state(State& state) const { ssize_t start_size = this->size(state); ssize_t num_args = operands_.size(); std::vector values; - State reg; - reg = expression_ptr_->empty_state(); + State reg = expression_ptr_->empty_state(); std::vector iterators; for (const ArrayNode* array_ptr : operands_) { diff --git a/tests/cpp/test_graph.cpp b/tests/cpp/test_graph.cpp index cfad9aa2..fc8f1f96 100644 --- a/tests/cpp/test_graph.cpp +++ b/tests/cpp/test_graph.cpp @@ -233,7 +233,9 @@ TEST_CASE("Graph constructors, assignment operators, and swapping") { } } -TEST_CASE("Graph::commit(), Graph::descendants(), Graph::propagate(), and Graph::revert") { +TEST_CASE( + "Graph::commit(), Graph::descendants(), Graph::mutated(), Graph::propagate(), and Graph::revert" +) { auto graph = Graph(); auto* x_ptr = graph.emplace_node(); auto* y_ptr = graph.emplace_node(); @@ -246,17 +248,23 @@ TEST_CASE("Graph::commit(), Graph::descendants(), Graph::propagate(), and Graph: CHECK_THAT(descendants, Catch::Matchers::RangeEquals(std::vector{x_ptr, z_ptr})); } SECTION("Propagate all") { + CHECK_THAT(graph.mutated(state), Catch::Matchers::RangeEquals(std::vector{})); + CHECK(x_ptr->view(state).front() == 0); CHECK(y_ptr->view(state).front() == 0); CHECK(z_ptr->view(state).front() == 0); x_ptr->flip(state, 0); + CHECK_THAT(graph.mutated(state), Catch::Matchers::RangeEquals({x_ptr})); + y_ptr->flip(state, 0); CHECK(x_ptr->diff(state).size()); CHECK(y_ptr->diff(state).size()); CHECK(z_ptr->diff(state).empty()); // not yet propagated to + CHECK_THAT(graph.mutated(state), Catch::Matchers::RangeEquals({x_ptr, y_ptr})); + graph.propagate(state); CHECK(x_ptr->view(state).front() == 1); @@ -267,6 +275,8 @@ TEST_CASE("Graph::commit(), Graph::descendants(), Graph::propagate(), and Graph: CHECK(y_ptr->diff(state).size()); CHECK(z_ptr->diff(state).size()); // now has pending changes + CHECK_THAT(graph.mutated(state), Catch::Matchers::RangeEquals({x_ptr, y_ptr})); + SECTION("Commit all") { graph.commit(state); @@ -278,6 +288,9 @@ TEST_CASE("Graph::commit(), Graph::descendants(), Graph::propagate(), and Graph: CHECK(x_ptr->diff(state).empty()); CHECK(y_ptr->diff(state).empty()); CHECK(z_ptr->diff(state).empty()); + + // Committing should reset the mutated nodes + CHECK_THAT(graph.mutated(state), Catch::Matchers::RangeEquals(std::vector{})); } SECTION("Revert all") { @@ -291,6 +304,9 @@ TEST_CASE("Graph::commit(), Graph::descendants(), Graph::propagate(), and Graph: CHECK(x_ptr->diff(state).empty()); CHECK(y_ptr->diff(state).empty()); CHECK(z_ptr->diff(state).empty()); + + // Reverting should reset the mutated nodes + CHECK_THAT(graph.mutated(state), Catch::Matchers::RangeEquals(std::vector{})); } } }