diff --git a/src/taskgraph/graph.py b/src/taskgraph/graph.py index 576325514..f77d71ef3 100644 --- a/src/taskgraph/graph.py +++ b/src/taskgraph/graph.py @@ -107,6 +107,15 @@ def _visit(self, reverse): f"Dependency loop detected involving the following nodes: {loopy_nodes}" ) + @functools.cache + def _visit_order(self, reverse): + """ + Return the order in which `_visit` yields nodes as a tuple. The graph + is immutable, so this is cached to avoid sorting it again every time + it is visited. + """ + return tuple(self._visit(reverse)) + def visit_postorder(self): """ Generate a sequence of nodes in postorder, such that every node is @@ -114,14 +123,14 @@ def visit_postorder(self): Raises an exception if the graph contains a cycle. """ - return self._visit(False) + return iter(self._visit_order(False)) def visit_preorder(self): """ Like visit_postorder, but in reverse: evrey node is visited *before* any nodes it links to. """ - return self._visit(True) + return iter(self._visit_order(True)) @functools.cache def links_and_reverse_links_dict(self): diff --git a/src/taskgraph/optimize/base.py b/src/taskgraph/optimize/base.py index ba07ca276..aeb6cf5ff 100644 --- a/src/taskgraph/optimize/base.py +++ b/src/taskgraph/optimize/base.py @@ -78,7 +78,7 @@ def optimize_task_graph( # Gather each relevant task's index indexes = set() - for label in target_task_graph.graph.visit_postorder(): + for label in target_task_graph.tasks: if label in do_not_optimize: continue _, strategy, arg = optimizations(label) @@ -157,7 +157,7 @@ def remove_tasks( opt_counts = defaultdict(int) opt_reasons = {} removed = set() - dependents_of = target_task_graph.graph.reverse_links_dict() + _, dependents_of = target_task_graph.graph.links_and_reverse_links_dict() tasks = target_task_graph.tasks prune_candidates = set() diff --git a/test/test_graph_perf.py b/test/test_graph_perf.py index 83d400f1b..6212033fc 100644 --- a/test/test_graph_perf.py +++ b/test/test_graph_perf.py @@ -124,6 +124,8 @@ def test_transitive_closure(geometry): @pytest.mark.parametrize("geometry", ["linear", "fan", "btree", "diamond"]) def test_visit_postorder(geometry): _, graph, _ = GEOMETRIES[geometry] + # Clear the functools.cache to measure actual computation each time + graph._visit_order.cache_clear() order = list(graph.visit_postorder()) assert len(order) == N @@ -132,6 +134,8 @@ def test_visit_postorder(geometry): @pytest.mark.parametrize("geometry", ["linear", "fan", "btree", "diamond"]) def test_visit_preorder(geometry): _, graph, _ = GEOMETRIES[geometry] + # Clear the functools.cache to measure actual computation each time + graph._visit_order.cache_clear() order = list(graph.visit_preorder()) assert len(order) == N @@ -161,11 +165,26 @@ def test_links_dict(geometry): @pytest.mark.parametrize("geometry", ["linear", "fan", "btree", "diamond"]) def test_for_each_task(geometry): _, _, tg = GEOMETRIES[geometry] + # Clear the functools.cache to measure actual computation each time + tg.graph._visit_order.cache_clear() visited = [] tg.for_each_task(lambda task, _tg: visited.append(task.label)) assert len(visited) == N +@pytest.mark.benchmark +@pytest.mark.parametrize("geometry", ["linear", "fan", "btree", "diamond"]) +def test_for_each_task_repeated(geometry): + """Visit the same graph several times, like the verifications do.""" + _, _, tg = GEOMETRIES[geometry] + # Clear the functools.cache to measure actual computation each time + tg.graph._visit_order.cache_clear() + visited = [] + for _ in range(10): + tg.for_each_task(lambda task, _tg: visited.append(task.label)) + assert len(visited) == 10 * N + + # --------------------------------------------------------------------------- # Benchmarks – TaskGraph JSON round-trip # --------------------------------------------------------------------------- @@ -175,6 +194,8 @@ def test_for_each_task(geometry): @pytest.mark.parametrize("geometry", ["linear", "fan", "btree"]) def test_taskgraph_to_json(geometry): _, _, tg = GEOMETRIES[geometry] + # Clear the functools.cache to measure actual computation each time + tg.graph._visit_order.cache_clear() data = tg.to_json() assert len(data) == N