diff --git a/src/dolfinx_adjoint/edge.py b/src/dolfinx_adjoint/edge.py index a571c20..4c1e716 100644 --- a/src/dolfinx_adjoint/edge.py +++ b/src/dolfinx_adjoint/edge.py @@ -72,7 +72,7 @@ def calculate_adjoint(self): """ return self.input_value - def __call__(self, value: float or PETSc.Vec): + def __call__(self, value: float | PETSc.Vec): """ This method is used to perform the backpropagation of the edge. @@ -93,12 +93,17 @@ def __call__(self, value: float or PETSc.Vec): self.input_value = value grad_value = self.calculate_adjoint() - # Call next functions in the path - for function in self.next_functions: + # Extract all marked next functions + next_functions = [ + function + for function in self.next_functions + if getattr(function, "marked", True) + ] + for function in next_functions: function(grad_value) # Accumulate gradient if end of path - if self.next_functions == [] and type(self.predecessor) == Node: + if not next_functions and isinstance(self.predecessor, Node): self.predecessor.accumulate_grad(grad_value) def __str__(self): diff --git a/src/dolfinx_adjoint/graph.py b/src/dolfinx_adjoint/graph.py index 601ed21..21eae85 100644 --- a/src/dolfinx_adjoint/graph.py +++ b/src/dolfinx_adjoint/graph.py @@ -52,6 +52,8 @@ def add_node(self, node: Node): """ self.nodes.append(node) + if self._nx_graph is not None: + self._add_node_to_networkx(self._nx_graph, node) def add_edge(self, edge: Edge): """Add an edge to the graph @@ -61,6 +63,8 @@ def add_edge(self, edge: Edge): """ self.edges.append(edge) + if self._nx_graph is not None: + self._add_edge_to_networkx(self._nx_graph, edge) def get_node(self, id: int, version=None): """Get a node from the graph @@ -161,6 +165,29 @@ def __str__(self) -> str: """ return f"Graph object with {len(self.nodes)} nodes and {len(self.edges)} edges." + @staticmethod + def _edge_color(edge: Edge) -> str: + """Get the color of an edge, black if it is part of the marked path""" + return "black" if hasattr(edge, "marked") else "grey" + + @classmethod + def _add_edge_to_networkx(cls, nx_graph: DiGraph, edge: Edge) -> None: + """Add an edge of the graph to its networkx representation""" + tag = "" if edge.__class__.__name__ == "Edge" else edge.__class__.__name__ + nx_graph.add_edge( + id(edge.predecessor), + id(edge.successor), + tag=tag, + color=cls._edge_color(edge), + edge=edge, + ) + + @staticmethod + def _add_node_to_networkx(nx_graph: DiGraph, node: AbstractNode) -> None: + """Add a node of the graph to its networkx representation""" + color = "pink" if type(node) == AbstractNode else "lightblue" + nx_graph.add_node(id(node), name=node.name, node=node, color=color) + def to_networkx(self) -> DiGraph: """Convert the graph to a networkx graph @@ -170,38 +197,18 @@ def to_networkx(self) -> DiGraph: if self._nx_graph is not None: # Update dynamic edge colors based on marking for edge in self.edges: - u = id(edge.successor) - v = id(edge.predecessor) + u = id(edge.predecessor) + v = id(edge.successor) if self._nx_graph.has_edge(u, v): - self._nx_graph[u][v]["color"] = ( - "black" if hasattr(edge, "marked") else "grey" - ) + self._nx_graph[u][v]["color"] = self._edge_color(edge) return self._nx_graph nx_graph = nx.DiGraph() for node in self.nodes: - nx_graph.add_node(id(node), name=node.name, node=node) - if type(node) == AbstractNode: - nx_graph.nodes[id(node)]["color"] = "pink" - else: - nx_graph.nodes[id(node)]["color"] = "lightblue" - + self._add_node_to_networkx(nx_graph, node) for edge in self.edges: - if not edge.__class__.__name__ == "Edge": - tag = edge.__class__.__name__ - else: - tag = "" - if hasattr(edge, "marked"): - color = "black" - else: - color = "grey" - nx_graph.add_edge( - id(edge.successor), - id(edge.predecessor), - tag=tag, - color=color, - edge=edge, - ) + self._add_edge_to_networkx(nx_graph, edge) + self._nx_graph = nx_graph return nx_graph @@ -281,24 +288,62 @@ def backprop(self, function_id: int, variable_id=None): """ Perform backpropagation in the graph + The gradients of all previous calls are reset before the propagation is started. + + If a variable is given, it acts as the control: only the edges on the path from the + control to the function are marked and executed. The propagation therefore stops at + the control, whose gradient is stored and returned, even if the control is an + intermediate node of the graph. No gradients are stored in the nodes between the + function and the control. + + If no variable is given, all edges the function depends on are marked and the + propagation continues until it reaches the leaves of this dependency subgraph. + Gradients are then stored only in these leaves; intermediate nodes are passed + through without storing their gradients. The gradients can be retrieved afterwards + with :py:meth:`Node.get_grad` on the respective nodes. + Args: function_id (int): The id of the function to be differentiated - variable_id (int, optional): The id of the variable with respect to which the differentiation is performed. Defaults to None. - If None, the differentiation is performed with respect to all variables in the graph. + variable_id (int, optional): The id of the variable (control) with respect to which + the differentiation is performed. Defaults to None. If None, the propagation is + carried out down to the dependency leaves of the function. + + Returns: + float or PETSc.Vec: The gradient of the function with respect to the variable, + if a variable is given. Otherwise the gradients are only stored in the nodes. + + Note: + The gradients are reset and the marks are refreshed on every call. The result is + stored in the given variable, or, without a variable, in the dependency leaves. + Every edge that can be executed has to be registered with :py:meth:`add_edge`, + since an edge that has never been marked is executed, and the graph must not be + modified during the propagation. + + Note: + When the propagation includes collective operations, e.g. the adjoint equation of + a problem that is solved on the communicator of the mesh, the executed edges + depend on the marked path. All participating ranks therefore have to build the + same graph and to select the corresponding function and variable, so that the + operations are performed in a compatible order. """ + self.reset_grads() function_node = self.get_node(function_id) if variable_id is not None: variable_node = self.get_node(variable_id) self.get_path(id(variable_node), id(function_node)) else: - for edge in self.edges: - edge.marked = True - grad_func = function_node.get_gradFuncs()[0] - grad_func(1.0) + self.get_dependencies(id(function_node)) + + # The function can be the result of more than one operation, so all of its + # gradient functions on the path are seeded with the derivative of one + for grad_func in function_node.get_gradFuncs(): + if getattr(grad_func, "marked", True): + grad_func(1.0) + if variable_id is not None: - return self.get_node(variable_id).get_grad() + return variable_node.get_grad() def get_path(self, start_id: int, end_id: int): """ @@ -307,22 +352,51 @@ def get_path(self, start_id: int, end_id: int): Args: start_id (int): The id of the start node end_id (int): The id of the end node + + Raises: + ValueError: If the start or the end node is not part of the graph + """ + + nx_graph = self._get_networkx_graph() + if start_id not in nx_graph: + raise ValueError( + f"The start node with id {start_id} is not part of the graph." + ) + if end_id not in nx_graph: + raise ValueError(f"The end node with id {end_id} is not part of the graph.") + + descendants_of_start = nx.descendants(nx_graph, start_id) | {start_id} + ancestors_of_end = nx.ancestors(nx_graph, end_id) | {end_id} + for edge in self.edges: - edge.marked = False + edge.marked = ( + id(edge.predecessor) in descendants_of_start + and id(edge.successor) in ancestors_of_end + ) + + def get_dependencies(self, end_id: int): + """ + Get all operations the end node is the result of by marking the edges + + All other edges are unmarked in the same pass and a query that raises leaves the previous marking unchanged. + + Args: + end_id (int): The id of the end node + + Raises: + ValueError: If the end node is not part of the graph + + """ nx_graph = self._get_networkx_graph() + if end_id not in nx_graph: + raise ValueError(f"The end node with id {end_id} is not part of the graph.") - descendants_of_start = nx.descendants(nx_graph, start_id) - descendants_of_start.add(start_id) - ancestors_of_end = nx.ancestors(nx_graph, end_id) - ancestors_of_end.add(end_id) + upstream_of_end = nx.ancestors(nx_graph, end_id) | {end_id} for edge in self.edges: - succ_id = id(edge.successor) - pred_id = id(edge.predecessor) - if succ_id in descendants_of_start and pred_id in ancestors_of_end: - edge.marked = True + edge.marked = id(edge.successor) in upstream_of_end def reset_grads(self): """ diff --git a/tests/unittests/test_backprop.py b/tests/unittests/test_backprop.py new file mode 100644 index 0000000..04fbf96 --- /dev/null +++ b/tests/unittests/test_backprop.py @@ -0,0 +1,259 @@ +"""Unit tests for the backpropagation through the computational graph.""" + +import pytest + +import dolfinx_adjoint.graph as graph +from dolfinx_adjoint.edge import Edge +from dolfinx_adjoint.node import Node + + +class LinearEdge(Edge): + """Test double with a constant derivative and a flag indicating execution.""" + + def __init__(self, predecessor: Node, successor: Node, factor: float = 1.0): + super().__init__(predecessor, successor) + self.factor = factor + self.called = False + + def calculate_adjoint(self): + self.called = True + return self.input_value * self.factor + + +def _build(edges: list, node_types: dict = None): + """Build a graph from ``(name, predecessor, successor, factor)`` tuples. + + The edges are given in the order in which the forward operations are performed and + ``node_types`` overrides the node class of a node. Returns the graph and the objects, + nodes and edges, the last three as dictionaries of names. + """ + node_types = node_types or {} + + names = [] + for _, predecessor, successor, _ in edges: + for name in (predecessor, successor): + if name not in names: + names.append(name) + + _graph = graph.Graph() + objects = {} + nodes = {} + for name in names: + objects[name] = object() + nodes[name] = node_types.get(name, Node)(objects[name], name=name) + _graph.add_node(nodes[name]) + + built_edges = {} + for name, predecessor, successor, factor in edges: + edge = LinearEdge(nodes[predecessor], nodes[successor], factor=factor) + nodes[successor].append_gradFuncs(edge) + edge.set_next_functions(nodes[predecessor].get_gradFuncs()) + _graph.add_edge(edge) + built_edges[name] = edge + + return _graph, objects, nodes, built_edges + + +def _executed_edges(edges: dict): + """Get the names of the edges that have been evaluated.""" + return {name for name, edge in edges.items() if edge.called} + + +# Gradient values and storage + + +def test_backprop_repeats_the_chain_derivative_without_accumulating(): + """Each call returns the product of the derivatives, without stale gradients.""" + _graph, objects, _, _ = _build( + [ + ("e1", "variable", "mid", 2.0), + ("e2", "mid", "objective", 3.0), + ] + ) + + for _ in range(2): + gradient = _graph.backprop(id(objects["objective"]), id(objects["variable"])) + assert gradient == pytest.approx(6.0) + + +def test_backprop_sums_parallel_paths(): + """Gradients of parallel paths are accumulated in the variable.""" + _graph, objects, _, _ = _build( + [ + ("e_upper_in", "variable", "upper", 2.0), + ("e_upper_out", "upper", "form", 3.0), + ("e_lower_in", "variable", "lower", 5.0), + ("e_lower_out", "lower", "form", 7.0), + ("e_form", "form", "objective", 1.0), + ] + ) + + gradient = _graph.backprop(id(objects["objective"]), id(objects["variable"])) + + assert gradient == pytest.approx(2.0 * 3.0 + 5.0 * 7.0) + + +def test_backprop_stores_the_gradient_in_an_intermediate_control(): + """A selected intermediate node receives the gradient; its input does not.""" + _graph, objects, nodes, _ = _build( + [ + ("upstream", "upstream", "variable", 2.0), + ("e1", "variable", "objective", 3.0), + ] + ) + + gradient = _graph.backprop(id(objects["objective"]), id(objects["variable"])) + + assert gradient == pytest.approx(3.0) + assert nodes["variable"].get_grad() == pytest.approx(3.0) + assert nodes["upstream"].get_grad() is None + + +def test_backprop_accumulates_in_specialised_nodes(): + """Nodes of a derived node class receive their gradient as well.""" + + class DerivedNode(Node): + """Minimal test subclass for checking inherited gradient accumulation.""" + + _graph, objects, _, _ = _build( + [ + ("e1", "variable", "objective", 2.0), + ], + node_types={"variable": DerivedNode}, + ) + + gradient = _graph.backprop(id(objects["objective"]), id(objects["variable"])) + + assert gradient == pytest.approx(2.0) + + +def test_backprop_seeds_all_gradient_functions_of_the_function(): + """A function that is reached by several edges is seeded on all of them.""" + # The paths merge in the function itself, so the function has two gradient + # functions, in contrast to the paths merging in an intermediate node. + _graph, objects, _, _ = _build( + [ + ("e_upper_in", "variable", "upper", 2.0), + ("e_upper_out", "upper", "objective", 3.0), + ("e_lower_in", "variable", "lower", 5.0), + ("e_lower_out", "lower", "objective", 7.0), + ] + ) + + gradient = _graph.backprop(id(objects["objective"]), id(objects["variable"])) + + assert gradient == pytest.approx(2.0 * 3.0 + 5.0 * 7.0) + + +def test_backprop_sums_parallel_edges_between_the_same_nodes(): + """Two operations connecting the same pair of nodes both contribute.""" + _graph, objects, _, _ = _build( + [ + ("e_first", "variable", "objective", 2.0), + ("e_second", "variable", "objective", 3.0), + ] + ) + + gradient = _graph.backprop(id(objects["objective"]), id(objects["variable"])) + + assert gradient == pytest.approx(5.0) + + +def test_backprop_without_a_variable_stores_gradients_only_in_dependency_leaves(): + """Unrestricted propagation stores gradients in the objective's dependency leaves.""" + _graph, objects, nodes, _ = _build( + [ + ("e1", "variable", "form", 2.0), + ("other", "other_input", "form", 3.0), + ("e_form", "form", "objective", 1.0), + ("post", "objective", "post_processing", 5.0), + ("unrelated", "unrelated_input", "unrelated_output", 7.0), + ] + ) + + assert _graph.backprop(id(objects["objective"])) is None + + assert nodes["variable"].get_grad() == pytest.approx(2.0) + assert nodes["other_input"].get_grad() == pytest.approx(3.0) + assert nodes["form"].get_grad() is None + assert nodes["unrelated_input"].get_grad() is None + assert nodes["post_processing"].get_grad() is None + + +def test_backprop_clears_gradients_when_switching_controls(): + """Selecting another control clears the previously stored gradient.""" + _graph, objects, nodes, _ = _build( + [ + ("e1", "variable", "form", 2.0), + ("other", "other_input", "form", 3.0), + ("e_form", "form", "objective", 1.0), + ] + ) + + gradient = _graph.backprop(id(objects["objective"]), id(objects["variable"])) + assert gradient == pytest.approx(2.0) + assert nodes["variable"].get_grad() == pytest.approx(2.0) + + gradient = _graph.backprop(id(objects["objective"]), id(objects["other_input"])) + assert gradient == pytest.approx(3.0) + assert nodes["variable"].get_grad() is None + + +# Execution of marked edges + + +def test_backprop_skips_unmarked_objective_inputs(): + """Only the selected input's edge is executed when seeding the objective.""" + # Select the second input to catch propagation that always seeds the first. + _graph, objects, _, edges = _build( + [ + ("other", "other_input", "objective", 3.0), + ("e_variable", "variable", "objective", 2.0), + ] + ) + + _graph.backprop(id(objects["objective"]), id(objects["variable"])) + + assert _executed_edges(edges) == {"e_variable"} + + +def test_backprop_skips_unmarked_upstream_edges(): + """Execution stops at the selected control, even when it has an input edge.""" + _graph, objects, _, edges = _build( + [ + ("upstream", "upstream", "variable", 2.0), + ("e1", "variable", "objective", 3.0), + ] + ) + + _graph.backprop(id(objects["objective"]), id(objects["variable"])) + + assert _executed_edges(edges) == {"e1"} + + +def test_backprop_updates_execution_when_switching_controls_and_modes(): + """Each query executes its own paths without losing previously skipped branches.""" + _graph, objects, _, edges = _build( + [ + ("e1", "variable", "form", 2.0), + ("other", "other_input", "form", 3.0), + ("e_form", "form", "objective", 1.0), + ("post", "objective", "post_processing", 5.0), + ("unrelated", "unrelated_input", "unrelated_output", 7.0), + ] + ) + + # Switch controls directly, then expand to all dependencies and restrict again. + for control, expected in ( + ("variable", {"e1", "e_form"}), + ("other_input", {"other", "e_form"}), + (None, {"e1", "other", "e_form"}), + ("variable", {"e1", "e_form"}), + ): + for edge in edges.values(): + edge.called = False + + variable_id = None if control is None else id(objects[control]) + _graph.backprop(id(objects["objective"]), variable_id) + + assert _executed_edges(edges) == expected, f"control={control!r}" diff --git a/tests/unittests/test_graph_networkx.py b/tests/unittests/test_graph_networkx.py new file mode 100644 index 0000000..4b49deb --- /dev/null +++ b/tests/unittests/test_graph_networkx.py @@ -0,0 +1,67 @@ +"""Unit tests for the networkx representation of the computational graph.""" + +import dolfinx_adjoint.graph as graph +from dolfinx_adjoint.edge import Edge +from dolfinx_adjoint.node import AbstractNode + + +def _chain(): + """Build the graph ``predecessor -> successor`` with a single edge.""" + _graph = graph.Graph() + + predecessor = AbstractNode(object(), name="predecessor") + successor = AbstractNode(object(), name="successor") + for node in (predecessor, successor): + _graph.add_node(node) + + edge = Edge(predecessor, successor) + _graph.add_edge(edge) + + return _graph, predecessor, successor, edge + + +def test_to_networkx_preserves_the_data_flow(): + """The networkx edges point from the predecessor to the successor. + + The nodes and edges themselves are kept in the representation, so that they can be + accessed from it. + """ + _graph, predecessor, successor, edge = _chain() + + nx_graph = _graph.to_networkx() + + assert nx_graph.has_edge(id(predecessor), id(successor)) + assert not nx_graph.has_edge(id(successor), id(predecessor)) + + assert nx_graph.nodes[id(predecessor)]["node"] is predecessor + assert nx_graph.nodes[id(successor)]["node"] is successor + assert nx_graph[id(predecessor)][id(successor)]["edge"] is edge + + +def test_to_networkx_contains_nodes_added_after_the_first_call(): + """A node added after the representation was built is contained in it.""" + _graph, _, _, _ = _chain() + + _graph.to_networkx() + + added = AbstractNode(object(), name="added") + _graph.add_node(added) + + assert id(added) in _graph.to_networkx() + + +def test_to_networkx_contains_edges_added_after_the_first_call(): + """An edge added after the representation was built is contained in it. + + Only the edge is added afterwards, independently of the node it connects to. + """ + _graph, _, successor, _ = _chain() + + added = AbstractNode(object(), name="added") + _graph.add_node(added) + + _graph.to_networkx() + + _graph.add_edge(Edge(successor, added)) + + assert _graph.to_networkx().has_edge(id(successor), id(added)) diff --git a/tests/unittests/test_graph_path.py b/tests/unittests/test_graph_path.py index 553a698..f6636f5 100644 --- a/tests/unittests/test_graph_path.py +++ b/tests/unittests/test_graph_path.py @@ -1,60 +1,210 @@ +"""Unit tests for the path marking of the computational graph.""" + +import pytest + import dolfinx_adjoint.graph as graph from dolfinx_adjoint.edge import Edge from dolfinx_adjoint.node import AbstractNode -def test_get_path_marks_only_edges_on_path(): - g = graph.Graph() +def _build(edges: dict): + """Build a graph from ``{name: (predecessor, successor)}``, given in flow direction.""" + + names = [] + for predecessor, successor in edges.values(): + for name in (predecessor, successor): + if name not in names: + names.append(name) + + _graph = graph.Graph() + nodes = {} + for name in names: + nodes[name] = AbstractNode(object(), name=name) + _graph.add_node(nodes[name]) + + built_edges = {} + for key, (predecessor, successor) in edges.items(): + built_edges[key] = Edge(nodes[predecessor], nodes[successor]) + _graph.add_edge(built_edges[key]) + + return _graph, nodes, built_edges + + +def _assert_marked(edges: dict, expected: set): + """Assert that exactly the edges in ``expected`` are marked.""" + marked = {key for key, edge in edges.items() if getattr(edge, "marked", False)} + assert marked == expected + + +def test_get_path_does_not_mark_unrelated_branches(): + """Only the edges connecting the variable to the objective are marked.""" + # variable -> mid -> objective, with a second input to mid, a dead end + # branching off mid, an operation feeding the variable and an operation + # performed on the objective. + _graph, nodes, edges = _build( + { + "e1": ("variable", "mid"), + "e2": ("mid", "objective"), + "other": ("other_input", "mid"), + "dead_end": ("mid", "dead_end"), + "upstream": ("upstream", "variable"), + "post_processing": ("objective", "post_processing"), + } + ) - start = AbstractNode(object(), name="start") - mid = AbstractNode(object(), name="mid") - end = AbstractNode(object(), name="end") - other_a = AbstractNode(object(), name="other_a") - other_b = AbstractNode(object(), name="other_b") + _graph.get_path(id(nodes["variable"]), id(nodes["objective"])) - for node in [start, mid, end, other_a, other_b]: - g.add_node(node) + _assert_marked(edges, {"e1", "e2"}) - # Path in reverse graph: start -> mid -> end - # Edge direction in this codebase is predecessor -> successor, - # and the reverse graph uses successor -> predecessor. - e1 = Edge(mid, start) - e2 = Edge(end, mid) - e3 = Edge(other_b, other_a) - g.add_edge(e1) - g.add_edge(e2) - g.add_edge(e3) +def test_get_path_marks_all_parallel_paths(): + """Every path from the variable to the objective is marked.""" + # variable -> {upper, lower} -> objective and variable -> objective, so that + # paths of different lengths connect the variable and the objective + _graph, nodes, edges = _build( + { + "e_upper_in": ("variable", "upper"), + "e_upper_out": ("upper", "objective"), + "e_lower_in": ("variable", "lower"), + "e_lower_out": ("lower", "objective"), + "e_direct": ("variable", "objective"), + } + ) - g.get_path(id(start), id(end)) + _graph.get_path(id(nodes["variable"]), id(nodes["objective"])) - assert e1.marked is True - assert e2.marked is True - assert e3.marked is False + _assert_marked( + edges, + {"e_upper_in", "e_upper_out", "e_lower_in", "e_lower_out", "e_direct"}, + ) def test_get_path_with_no_connection_marks_none(): - g = graph.Graph() + """A variable that the objective does not depend on marks no edge. + + The disconnected query is performed before and after a successful query, so that + it also covers the complete clearing of the marks of the successful query. + """ + + _graph, nodes, edges = _build( + { + "e1": ("variable", "mid"), + "e2": ("mid", "objective"), + "e3": ("other_input", "other_output"), + } + ) + + _graph.get_path(id(nodes["other_input"]), id(nodes["objective"])) + _assert_marked(edges, set()) + + _graph.get_path(id(nodes["variable"]), id(nodes["objective"])) + _assert_marked(edges, {"e1", "e2"}) + + _graph.get_path(id(nodes["other_input"]), id(nodes["objective"])) + _assert_marked(edges, set()) + + +def test_get_path_resets_previous_marks(): + """A second call to get_path does not keep the marks of the first call.""" + _graph, nodes, edges = _build( + { + "e1": ("variable", "mid"), + "e2": ("mid", "objective"), + "e3": ("other_input", "mid"), + } + ) + + _graph.get_path(id(nodes["variable"]), id(nodes["objective"])) + _assert_marked(edges, {"e1", "e2"}) + + _graph.get_path(id(nodes["other_input"]), id(nodes["objective"])) + _assert_marked(edges, {"e3", "e2"}) + + +def test_get_path_in_reverse_marks_nothing(): + """The path is directed: querying it against the data flow marks no edge.""" + _graph, nodes, edges = _build( + { + "e1": ("variable", "mid"), + "e2": ("mid", "objective"), + } + ) + + # Variable and objective are exchanged, so no forward path exists + _graph.get_path(id(nodes["objective"]), id(nodes["variable"])) + + _assert_marked(edges, set()) + + +def test_get_path_with_isolated_control_marks_nothing(): + """A control that is registered in the graph but unused marks no edge.""" + _graph, nodes, edges = _build( + { + "e1": ("variable", "mid"), + "e2": ("mid", "objective"), + } + ) + + # A control that has been added to the graph but is not used in any operation + isolated = AbstractNode(object(), name="isolated") + _graph.add_node(isolated) + + _graph.get_path(id(isolated), id(nodes["objective"])) + + _assert_marked(edges, set()) + + +def test_get_path_uses_nodes_added_after_the_first_call(): + """The path marking sees the nodes and edges added after the first call.""" + _graph, nodes, edges = _build({"e1": ("predecessor", "successor")}) + + _graph.get_path(id(nodes["predecessor"]), id(nodes["successor"])) + + # Extend the graph after the networkx representation has been built + added = AbstractNode(object(), name="added") + _graph.add_node(added) + added_edge = Edge(nodes["successor"], added) + _graph.add_edge(added_edge) + + _graph.get_path(id(nodes["predecessor"]), id(added)) + + assert edges["e1"].marked is True + assert added_edge.marked is True + + +def test_get_path_with_an_unknown_node_keeps_the_previous_marks(): + """A query that raises does not change the marking of the previous query.""" + _graph, nodes, edges = _build( + { + "e1": ("variable", "mid"), + "e2": ("mid", "objective"), + } + ) + + _graph.get_path(id(nodes["variable"]), id(nodes["objective"])) + + unregistered = AbstractNode(object(), name="unregistered") + with pytest.raises(ValueError): + _graph.get_path(id(unregistered), id(nodes["objective"])) - start = AbstractNode(object(), name="start") - mid = AbstractNode(object(), name="mid") - end = AbstractNode(object(), name="end") - other_a = AbstractNode(object(), name="other_a") - other_b = AbstractNode(object(), name="other_b") + _assert_marked(edges, {"e1", "e2"}) - for node in [start, mid, end, other_a, other_b]: - g.add_node(node) - e1 = Edge(mid, start) - e2 = Edge(end, mid) - e3 = Edge(other_b, other_a) +def test_get_dependencies_marks_only_dependencies_and_clears_previous_marks(): + """Dependency marking excludes downstream edges and clears unrelated marks.""" + _graph, nodes, edges = _build( + { + "e1": ("variable", "form"), + "other": ("other_input", "form"), + "e_form": ("form", "objective"), + "post": ("objective", "post_processing"), + "unrelated": ("unrelated_input", "unrelated_output"), + } + ) - g.add_edge(e1) - g.add_edge(e2) - g.add_edge(e3) + _graph.get_path(id(nodes["unrelated_input"]), id(nodes["unrelated_output"])) + _assert_marked(edges, {"unrelated"}) - g.get_path(id(other_a), id(end)) + _graph.get_dependencies(id(nodes["objective"])) - assert e1.marked is False - assert e2.marked is False - assert e3.marked is False + _assert_marked(edges, {"e1", "other", "e_form"})