Skip to content
Open
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
13 changes: 11 additions & 2 deletions src/taskgraph/graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,21 +107,30 @@ 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
visited *after* any nodes it links to.

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):
Expand Down
4 changes: 2 additions & 2 deletions src/taskgraph/optimize/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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()

Expand Down
21 changes: 21 additions & 0 deletions test/test_graph_perf.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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

Expand Down Expand Up @@ -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
# ---------------------------------------------------------------------------
Expand All @@ -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

Expand Down
Loading