Skip to content
Open
Show file tree
Hide file tree
Changes from 4 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
8 changes: 7 additions & 1 deletion dace/transformation/passes/dead_state_elimination.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,8 @@
from dace import SDFG, InterstateEdge, SDFGState, symbolic, properties
from dace.properties import CodeBlock
from dace.sdfg.graph import Edge
from dace.sdfg.state import ConditionalBlock, ControlFlowBlock, ControlFlowRegion
from dace.sdfg.state import (BreakBlock, ConditionalBlock, ContinueBlock, ControlFlowBlock, ControlFlowRegion,
ReturnBlock)
from dace.sdfg.validation import InvalidSDFGInterstateEdgeError, InvalidSDFGNodeError
from dace.transformation import pass_pipeline as ppl, transformation

Expand Down Expand Up @@ -119,6 +120,11 @@ def find_dead_control_flow(
continue
visited.add(node)

# Control flow does not continue past a return, break, or continue block, so anything
# reachable only through one of them is dead.
if isinstance(node, (ReturnBlock, BreakBlock, ContinueBlock)):
continue

# First, check for unconditional edges
unconditional = None
for e in cfg.out_edges(node):
Expand Down
74 changes: 73 additions & 1 deletion tests/passes/dead_code_elimination_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
import pytest
import dace
from dace.properties import CodeBlock
from dace.sdfg.state import ConditionalBlock, ControlFlowRegion, LoopRegion
from dace.sdfg.state import (BreakBlock, ConditionalBlock, ContinueBlock, ControlFlowRegion, LoopRegion, ReturnBlock)
from dace.sdfg.validation import InvalidSDFGNodeError
from dace.transformation.pass_pipeline import Pipeline
from dace.transformation.passes.dead_state_elimination import DeadStateElimination
Expand Down Expand Up @@ -146,6 +146,74 @@ def test_dse_malformed_conditional_block():
DeadStateElimination().apply_pass(sdfg, {})


def test_dse_after_return():
sdfg = dace.SDFG('dse_after_return')
start = sdfg.add_state(is_start_block=True)
ret = ReturnBlock('ret')
sdfg.add_node(ret)
unreachable = sdfg.add_state()
sdfg.add_edge(start, ret, dace.InterstateEdge())
sdfg.add_edge(ret, unreachable, dace.InterstateEdge())

DeadStateElimination().apply_pass(sdfg, {})
assert set(sdfg.nodes()) == {start, ret}


def test_dse_after_break():
sdfg = dace.SDFG('dse_after_break')
loop = LoopRegion('loop', 'i < 10', 'i', 'i = 0', 'i = i + 1')
start = sdfg.add_state(is_start_block=True)
sdfg.add_node(loop)
end = sdfg.add_state()
sdfg.add_edge(start, loop, dace.InterstateEdge())
sdfg.add_edge(loop, end, dace.InterstateEdge())
body = loop.add_state(is_start_block=True)
brk = BreakBlock('brk')
loop.add_node(brk)
unreachable = loop.add_state()
loop.add_edge(body, brk, dace.InterstateEdge())
loop.add_edge(brk, unreachable, dace.InterstateEdge())

DeadStateElimination().apply_pass(sdfg, {})
assert set(loop.nodes()) == {body, brk}
assert set(sdfg.states()) == {start, body, end}


def test_dse_after_continue():
sdfg = dace.SDFG('dse_after_continue')
loop = LoopRegion('loop', 'i < 10', 'i', 'i = 0', 'i = i + 1')
start = sdfg.add_state(is_start_block=True)
sdfg.add_node(loop)
end = sdfg.add_state()
sdfg.add_edge(start, loop, dace.InterstateEdge())
sdfg.add_edge(loop, end, dace.InterstateEdge())
body = loop.add_state(is_start_block=True)
cont = ContinueBlock('cont')
loop.add_node(cont)
unreachable = loop.add_state()
loop.add_edge(body, cont, dace.InterstateEdge())
loop.add_edge(cont, unreachable, dace.InterstateEdge())

DeadStateElimination().apply_pass(sdfg, {})
assert set(loop.nodes()) == {body, cont}
assert set(sdfg.states()) == {start, body, end}


def test_dse_reachable_around_return():
sdfg = dace.SDFG('dse_reachable_around_return')
sdfg.add_symbol('a', dace.int32)
start = sdfg.add_state(is_start_block=True)
ret = ReturnBlock('ret')
sdfg.add_node(ret)
after = sdfg.add_state()
sdfg.add_edge(start, ret, dace.InterstateEdge('a > 0'))
sdfg.add_edge(start, after, dace.InterstateEdge('a <= 0'))
sdfg.add_edge(ret, after, dace.InterstateEdge())

DeadStateElimination().apply_pass(sdfg, {})
assert set(sdfg.nodes()) == {start, ret, after}


def test_dde_simple():

@dace.program
Expand Down Expand Up @@ -526,6 +594,10 @@ def loop_condition(A: dace.float64[20]):
test_dse_inside_loop()
test_dse_inside_loop_conditional()
test_dse_malformed_conditional_block()
test_dse_after_return()
test_dse_after_break()
test_dse_after_continue()
test_dse_reachable_around_return()
test_dde_simple()
test_dde_libnode()
test_dde_access_node_in_scope(False)
Expand Down
Loading