From 3ef42eb3e394a40e0a2c847f31ad88e960923cde Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:07:18 -0500 Subject: [PATCH 01/16] Add ForLoop AST node type --- dagrt/codegen/dag_ast.py | 25 +++++++++++++++++++++++++ 1 file changed, 25 insertions(+) diff --git a/dagrt/codegen/dag_ast.py b/dagrt/codegen/dag_ast.py index 6107865..1249ef0 100644 --- a/dagrt/codegen/dag_ast.py +++ b/dagrt/codegen/dag_ast.py @@ -65,6 +65,29 @@ def __getinitargs__(self): mapper_method = "map_IfThenElse" +class ForLoop(Expression): + """ + Bounds are a half-open interval as in Python + + .. attribute: loop_var_name + .. attribute: lbound + .. attribute: ubound + .. attribute: body + """ + init_args_names = ("loop_var_name", "lbound", "ubound", "body") + + def __init__(self, loop_var_name, lbound, ubound, body): + self.loop_var_name = loop_var_name + self.lbound = lbound + self.ubound = ubound + self.body = body + + def __getinitargs__(self): + return self.loop_var_name, self.lbound, self.ubound, self.body + + mapper_method = "map_ForLoop" + + class Block(Expression): """ .. attribute: children @@ -120,6 +143,8 @@ def get_statements_in_ast(ast): children = (ast.then,) elif isinstance(ast, IfThenElse): children = (ast.then, ast.else_) + elif isinstance(ast, ForLoop): + children = (ast.body,) elif isinstance(ast, Block): children = ast.children else: From 67e7dffedfce8f07cfc1ac32a2da927de1903b21 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:07:55 -0500 Subject: [PATCH 02/16] Implement/move AST stringifier, identity mapper --- dagrt/codegen/dag_ast.py | 145 ++++++++++++++++++++++++++++++++------- 1 file changed, 122 insertions(+), 23 deletions(-) diff --git a/dagrt/codegen/dag_ast.py b/dagrt/codegen/dag_ast.py index 1249ef0..ebedb29 100644 --- a/dagrt/codegen/dag_ast.py +++ b/dagrt/codegen/dag_ast.py @@ -1,8 +1,4 @@ """Abstract syntax""" -from pymbolic.mapper import IdentityMapper -from pymbolic.primitives import Expression, LogicalNot -from dagrt.language import Nop - __copyright__ = "Copyright (C) 2015 Matt Wala" @@ -26,6 +22,13 @@ THE SOFTWARE. """ +from pymbolic.mapper import IdentityMapper, Collector +from pymbolic.mapper.stringifier import StringifyMapper +from pymbolic.primitives import Expression, LogicalNot +from dagrt.language import Nop + + +# {{{ ast node types class IfThen(Expression): """ @@ -129,6 +132,117 @@ def __getinitargs__(self): mapper_method = "map_StatementWrapper" +# }}} + + +# {{{ ast mappers + +class ASTCollector(Collector): + def map_IfThenElse(self, expr): + return self.combine([ + self.rec(expr.condition), + self.rec(expr.then), + self.rec(expr.else_), + ]) + + def map_IfThen(self, expr): + return self.combine([ + self.rec(expr.condition), + self.rec(expr.then), + ]) + + def map_ForLoop(self, expr): + return self.combine([ + self.rec(expr.lbound), + self.rec(expr.ubound), + self.rec(expr.body)]) + + def map_Block(self, expr): + return self.combine([ + self.rec(ch) + for ch in expr.children]) + + +class LoopVariableFinder(ASTCollector): + def map_constant(self, expr): + return set() + + def map_variable(self, expr): + return set() + + def map_ForLoop(self, expr): + return {expr.loop_var_name} | super().map_ForLoop(expr) + + def map_StatementWrapper(self, expr): + return set() + + +class ASTIdentityMapper(IdentityMapper): + def map_IfThenElse(self, expr): + return type(expr)(self.rec(expr.condition), self.rec(expr.then), + self.rec(expr.else_)) + + def map_IfThen(self, expr): + return type(expr)(self.rec(expr.condition), self.rec(expr.then)) + + def map_ForLoop(self, expr): + return type(expr)( + loop_var_name=expr.loop_var_name, + lbound=self.rec(expr.lbound), + ubound=self.rec(expr.ubound), + body=self.rec(expr.body)) + + def map_Block(self, expr): + return type(expr)(*[self.rec(child) for child in expr.children]) + + def map_NullASTNode(self, expr): + return type(expr)() + + def map_StatementWrapper(self, expr): + return type(expr)(expr.statement) + + +class ASTStringifier(StringifyMapper): + indent_str = " " + + def map_IfThenElse(self, expr, indent): + istr = self.indent_str*indent + return ( + istr + f"if {expr.condition}:\n" + + self.rec(expr.then, indent+1) + "\n" + + + istr + "else:\n" + + self.rec(expr.else_, indent+1) + ) + + def map_IfThen(self, expr, indent): + istr = self.indent_str*indent + return ( + istr + f"if {expr.condition}:\n" + + self.rec(expr.then, indent+1)) + + def map_ForLoop(self, expr, indent): + istr = self.indent_str*indent + return ( + istr + f"for {expr.loop_var_name} " + f"in [{expr.lbound}, {expr.ubound}):\n" + + self.rec(expr.body, indent+1)) + + def map_Block(self, expr, indent): + istr = self.indent_str*indent + return ( + istr + "{\n" + + "\n".join(self.rec(ch, indent+1) for ch in expr.children) + + "\n" + + istr + "}") + + def map_NullASTNode(self, expr, indent): + return "**NULL**" + + def map_StatementWrapper(self, expr, indent): + return self.indent_str*indent + str(expr.statement) + +# }}} + def get_statements_in_ast(ast): """ @@ -237,25 +351,6 @@ def apply_pass(ast, pass_): return reduce(apply_pass, passes, ast) -class ASTIdentityMapper(IdentityMapper): - - def map_IfThenElse(self, expr): - return type(expr)(self.rec(expr.condition), self.rec(expr.then), - self.rec(expr.else_)) - - def map_IfThen(self, expr): - return type(expr)(self.rec(expr.condition), self.rec(expr.then)) - - def map_Block(self, expr): - return type(expr)(*[self.rec(child) for child in expr.children]) - - def map_NullASTNode(self, expr): - return type(expr)() - - def map_StatementWrapper(self, expr): - return type(expr)(expr.statement) - - class ASTPreSimplifyMapper(ASTIdentityMapper): def map_IfThen(self, expr): @@ -381,3 +476,7 @@ def flat_Block(*nodes): if len(children) == 1: return children[0] return Block(*children) + +# }}} + +# vim: foldmethod=marker From 9113321ae0a1a0471c20edf41e3d9e5a52048019 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:09:21 -0500 Subject: [PATCH 03/16] run_fortran: add debug flag --- dagrt/utils.py | 31 ++++++++++++++++++++----------- 1 file changed, 20 insertions(+), 11 deletions(-) diff --git a/dagrt/utils.py b/dagrt/utils.py index 39a13d3..a455767 100644 --- a/dagrt/utils.py +++ b/dagrt/utils.py @@ -172,7 +172,11 @@ def __del__(self): # {{{ run_fortran -def run_fortran(sources, fortran_options=None, fortran_libraries=None): +class DebuggerExit(Exception): + pass + + +def run_fortran(sources, fortran_options=None, fortran_libraries=None, debug=False): if fortran_options is None: fortran_options = [] if fortran_libraries is None: @@ -202,16 +206,21 @@ def run_fortran(sources, fortran_options=None, fortran_libraries=None): + ["-l"+lib for lib in fortran_libraries], cwd=tmpdir) - p = Popen([join(tmpdir, "runtest")], stdout=PIPE, stderr=PIPE, - close_fds=True) - stdout_data, stderr_data = p.communicate() - - if stdout_data: - print("Fortran code said this on stdout: -----------------------------", - file=sys.stderr) - print(stdout_data.decode(), file=sys.stderr) - print("---------------------------------------------------------------", - file=sys.stderr) + if debug: + p = Popen(["gdb", "--args", join(tmpdir, "runtest")]) + p.wait() + raise DebuggerExit + else: + p = Popen([join(tmpdir, "runtest")], stdout=PIPE, stderr=PIPE, + close_fds=True) + stdout_data, stderr_data = p.communicate() + + if stdout_data: + print("Fortran code said this on stdout: -------------------------", + file=sys.stderr) + print(stdout_data.decode(), file=sys.stderr) + print("-----------------------------------------------------------", + file=sys.stderr) if stderr_data: raise RuntimeError( From c6eefcba6970bed649266d39f28529e406b86206 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:09:56 -0500 Subject: [PATCH 04/16] SymbolKindFinder: allow forced_kinds --- dagrt/data.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/dagrt/data.py b/dagrt/data.py index 342f384..298a10c 100644 --- a/dagrt/data.py +++ b/dagrt/data.py @@ -444,7 +444,7 @@ class SymbolKindFinder: def __init__(self, function_registry): self.function_registry = function_registry - def __call__(self, names, phases): + def __call__(self, names, phases, forced_kinds=None): """Infer the kinds of all the symbols in a program. :arg names: a list of phase names @@ -463,6 +463,10 @@ def __call__(self, names, phases): result = SymbolKindTable() + if forced_kinds is not None: + for phase_name, ident, kind in forced_kinds: + result.set(phase_name, ident, kind=kind) + def make_kim(phase_name, check): return KindInferenceMapper( result.global_table, From 6e2ca0002668bae9677ef514df0b08aef302ca36 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:10:28 -0500 Subject: [PATCH 05/16] StructuredCodeGenerator: support ForLoop --- dagrt/codegen/codegen_base.py | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/dagrt/codegen/codegen_base.py b/dagrt/codegen/codegen_base.py index 30f5fe1..d044b9b 100644 --- a/dagrt/codegen/codegen_base.py +++ b/dagrt/codegen/codegen_base.py @@ -22,7 +22,8 @@ THE SOFTWARE. """ -from dagrt.codegen.dag_ast import Block, IfThen, IfThenElse, StatementWrapper +from dagrt.codegen.dag_ast import \ + Block, IfThen, IfThenElse, StatementWrapper, ForLoop class StructuredCodeGenerator: @@ -55,6 +56,11 @@ def lower_node(self, node): self.lower_node(node.else_) self.emit_if_end() + elif isinstance(node, ForLoop): + self.emit_for_begin(node.loop_var_name, node.lbound, node.ubound) + self.lower_node(node.body) + self.emit_for_end(node.loop_var_name) + elif isinstance(node, Block): for child in node.children: self.lower_node(child) @@ -101,5 +107,11 @@ def emit_if_end(self): def emit_else_begin(self): raise NotImplementedError() + def emit_for_begin(self, loop_var_name, lbound, ubount): + raise NotImplementedError() + + def emit_for_end(self, loop_var_name): + raise NotImplementedError() + def emit_return(self): raise NotImplementedError() From 54d95385ce8ef30cd69cfd8e67622fa5be3b672e Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:13:04 -0500 Subject: [PATCH 06/16] Generate for loops via AST - do funcall isolation on AST - do self dep isolation on AST - do if isolation on AST - implement ForLoop node codegen in Fortran - add selfdep-in-loop test --- dagrt/codegen/dag_ast.py | 65 ++++---- dagrt/codegen/fortran.py | 77 +++++----- dagrt/codegen/transform.py | 283 +++++++++++++++++++---------------- test/test_codegen_fortran.py | 35 +++++ 4 files changed, 265 insertions(+), 195 deletions(-) diff --git a/dagrt/codegen/dag_ast.py b/dagrt/codegen/dag_ast.py index ebedb29..c91b222 100644 --- a/dagrt/codegen/dag_ast.py +++ b/dagrt/codegen/dag_ast.py @@ -268,6 +268,33 @@ def get_statements_in_ast(ast): yield from get_statements_in_ast(child) +def statement_to_ast(statement): + return StatementWrapper(statement) + + +def conditional_to_ast(statement): + if statement.condition is not True: + new_statement = statement.copy(condition=True) + return IfThenElse(statement.condition, + statement_to_ast(new_statement), + NullASTNode()) + else: + return statement_to_ast(statement) + + +def loop_to_ast_node(statement): + if statement.loops: + loop_var_name, lower, upper = statement.loops[0] + new_statement = statement.copy(loops=statement.loops[1:]) + return ForLoop( + loop_var_name=loop_var_name, + lbound=lower, + ubound=upper, + body=loop_to_ast_node(new_statement)) + else: + return conditional_to_ast(statement) + + def create_ast_from_phase(code, phase_name): """ Return an AST representation of the statements corresponding to the phase @@ -276,7 +303,7 @@ def create_ast_from_phase(code, phase_name): phase = code.phases[phase_name] - # Construct a topological order of the statements. + # {{{ Construct a topological order of the statements. stack = [] statement_map = {inst.id: inst for inst in phase.statements} visiting = set() @@ -299,41 +326,25 @@ def create_ast_from_phase(code, phase_name): stack.extend( sorted(statement_map[statement].depends_on)) - # Convert the topological order to an AST. - main_block = [] + # }}} - from pymbolic.primitives import LogicalAnd + # {{{ Convert the topological order to an AST. + + main_block = [] + for top_order_id in topological_order: + statement = statement_map[top_order_id] - for statement in map(statement_map.__getitem__, topological_order): if isinstance(statement, Nop): continue - # Statements become AST nodes. An unconditional statement is wrapped - # into an StatementWrapper, while conditional statements are wrapped - # using IfThens. - - if isinstance(statement.condition, LogicalAnd): - # LogicalAnd(c1, c2, ...) => IfThen(c1, IfThen(c2, ...)) - conditions = reversed(statement.condition.children) - inst = IfThenElse(next(conditions), - StatementWrapper(statement.copy(condition=True)), - NullASTNode()) - for next_cond in conditions: - inst = IfThenElse(next_cond, inst, NullASTNode()) - main_block.append(inst) - - elif statement.condition is not True: - main_block.append(IfThenElse(statement.condition, - StatementWrapper(statement.copy(condition=True)), - NullASTNode())) + main_block.append(loop_to_ast_node(statement)) - else: - main_block.append(StatementWrapper(statement)) + # }}} - ast = Block(*main_block) + return simplify_ast(Block(*main_block)) - return simplify_ast(ast) +# {{{ ast simplification def simplify_ast(ast): """Return an optimized copy of the AST `ast`.""" diff --git a/dagrt/codegen/fortran.py b/dagrt/codegen/fortran.py index 41c7490..8562a89 100644 --- a/dagrt/codegen/fortran.py +++ b/dagrt/codegen/fortran.py @@ -1048,16 +1048,6 @@ def __call__(self, dag): from dagrt.codegen.analysis import verify_code verify_code(dag) - from dagrt.codegen.transform import ( - eliminate_self_dependencies, - isolate_function_arguments, - isolate_function_calls, - expand_IfThenElse) - dag = eliminate_self_dependencies(dag) - dag = isolate_function_arguments(dag) - dag = isolate_function_calls(dag) - dag = expand_IfThenElse(dag) - # from dagrt.language import show_dependency_graph # show_dependency_graph(dag) @@ -1070,17 +1060,39 @@ def __call__(self, dag): NameASTPair = namedtuple("NameASTPair", "name, ast") # noqa fdescrs = [] + def process_ast(ast, print_ast=False): + from dagrt.codegen.transform import ( + eliminate_self_dependencies, + isolate_function_arguments, + isolate_function_calls, + expand_IfThenElse) + ast = eliminate_self_dependencies(ast) + ast = isolate_function_arguments(ast) + ast = isolate_function_calls(ast) + ast = expand_IfThenElse(ast) + + if print_ast: + from dagrt.codegen.dag_ast import ASTStringifier + print(ASTStringifier()(ast, 0)) + + return ast + for phase_name in sorted(dag.phases.keys()): ast = create_ast_from_phase(dag, phase_name) - fdescrs.append(NameASTPair(phase_name, ast)) + fdescrs.append(NameASTPair(phase_name, process_ast(ast))) # }}} - from dagrt.data import SymbolKindFinder + from dagrt.data import SymbolKindFinder, Integer + from dagrt.codegen.dag_ast import LoopVariableFinder self.sym_kind_table = SymbolKindFinder(self.function_registry)( - [fd.name for fd in fdescrs], - [get_statements_in_ast(fd.ast) for fd in fdescrs]) + [fd.name for fd in fdescrs], + [get_statements_in_ast(fd.ast) for fd in fdescrs], + forced_kinds=[ + (fd.name, loop_var, Integer()) + for fd in fdescrs + for loop_var in LoopVariableFinder()(fd.ast)]) from dagrt.codegen.analysis import ( collect_ode_component_names_from_dag, @@ -1427,7 +1439,7 @@ def finish_emit(self, dag): self.module_emitter.__exit__(None, None, None) - self.emit("! vim:foldmethod=marker") + self.emit("! vim:foldmethod=marker:filetype=fortran") # }}} @@ -2070,6 +2082,19 @@ def emit_if_end(self,): def emit_else_begin(self): self.emitter.emit_else() # pylint:disable=no-member + def emit_for_begin(self, loop_var_name, lbound, ubound): + em = FortranDoEmitter( + self.emitter, + self.name_manager[loop_var_name], + "int({}), int({})".format( + self.expr(lbound), + self.expr(ubound-1)), + code_generator=self) + em.__enter__() + + def emit_for_end(self, loop_var_name): + self.emitter.__exit__(None, None, None) + def emit_assign_expr(self, assignee_sym, assignee_subscript, expr): from dagrt.data import UserType, Array @@ -2120,32 +2145,12 @@ def lower_inst(self, inst): # {{{ emit_inst_Assign def emit_inst_Assign(self, inst): - start_em = self.emitter - - for iloop, (ident, start, stop) in enumerate(inst.loops): - # Only parallelize the innermost loop - pdp = None - if iloop + 1 == len(inst.loops): - pdp = self.parallel_do_preamble - - em = FortranDoEmitter( - self.emitter, - self.name_manager[ident], - "int({}), int({})".format(self.expr(start), self.expr(stop-1)), - code_generator=self, - parallel_do_preamble=pdp) - em.__enter__() + assert not inst.loops self.emit_assign_expr( inst.assignee, inst.assignee_subscript, inst.expression) - - for _ident, _start, _stop in inst.loops[::-1]: - self.emitter.__exit__(None, None, None) - self.emit_deinit_for_last_usage_of_vars(inst) - assert start_em is self.emitter - # }}} # {{{ emit_inst_AssignFunctionCall diff --git a/dagrt/codegen/transform.py b/dagrt/codegen/transform.py index 1417af5..2489c90 100644 --- a/dagrt/codegen/transform.py +++ b/dagrt/codegen/transform.py @@ -24,6 +24,10 @@ """ from pymbolic.mapper import IdentityMapper +from pytools import UniqueNameGenerator +from dagrt.codegen.dag_ast import ( + ASTIdentityMapper, get_statements_in_ast, + Block, StatementWrapper) __doc__ = """ .. autofunction:: eliminate_self_dependencies @@ -33,65 +37,101 @@ """ +def get_stmt_id_generator(statements): + return UniqueNameGenerator({stmt.id for stmt in statements}) + + +def get_var_name_generator(statements): + existing_variables = set() + for stmt in statements: + existing_variables.update(stmt.get_written_variables()) + existing_variables.update(stmt.get_read_variables()) + return UniqueNameGenerator(existing_variables) + + +# {{{ ast statement rewriter + +class ASTStatementRewriter(ASTIdentityMapper): + def __init__(self, stmt_id_gen, var_name_gen): + self.stmt_id_gen = stmt_id_gen + self.var_name_gen = var_name_gen + + def map_StatementWrapper(self, expr): + new_statements = [ + StatementWrapper(stmt) + for stmt in self.map_statement(expr.statement)] + + if len(new_statements) > 1: + return Block(*new_statements) + else: + return new_statements[0] + + +def apply_statement_rewriter(rewriter_cls, phase_ast): + statements = list(get_statements_in_ast(phase_ast)) + rewriter = rewriter_cls( + stmt_id_gen=get_stmt_id_generator(statements), + var_name_gen=get_var_name_generator(statements)) + + return rewriter(phase_ast) + +# }}} + + # {{{ eliminate self dependencies -def eliminate_self_dependencies(dag): - stmt_id_gen = dag.get_stmt_id_generator() - var_name_gen = dag.get_var_name_generator() +class SelfDependencyEliminator(ASTStatementRewriter): + def map_statement(self, stmt): + read_and_written = ( + stmt.get_read_variables() & stmt.get_written_variables()) + + if not read_and_written: + return [stmt] + + substs = [] + tmp_stmt_ids = [] - new_phases = {} - for phase_name, phase in dag.phases.items(): new_statements = [] - for stmt in sorted(phase.statements, key=lambda stmt: stmt.id): - read_and_written = ( - stmt.get_read_variables() & stmt.get_written_variables()) - - if not read_and_written: - new_statements.append(stmt) - continue - - substs = [] - tmp_stmt_ids = [] - - from dagrt.language import Assign - from pymbolic import var - for var_name in read_and_written: - tmp_var_name = var_name_gen( - "temp_" - + var_name.replace("<", "_").replace(">", "_")) - substs.append((var_name, var(tmp_var_name))) - - tmp_stmt_id = stmt_id_gen("temp") - tmp_stmt_ids.append(tmp_stmt_id) - - new_tmp_stmt = Assign( - tmp_var_name, (), var(var_name), - condition=stmt.condition, - id=tmp_stmt_id, - depends_on=stmt.depends_on) - new_statements.append(new_tmp_stmt) - - from pymbolic import substitute - new_stmt = (stmt - .map_expressions( - lambda expr: substitute(expr, dict(substs)), - include_lhs=False) - .copy( - # lhs will be rewritten, but we don't want that. - depends_on=stmt.depends_on | frozenset(tmp_stmt_ids))) - - new_statements.append(new_stmt) - - new_phases[phase_name] = phase.copy(statements=new_statements) - - return dag.copy(phases=new_phases) + from dagrt.language import Assign + from pymbolic import var + for var_name in read_and_written: + tmp_var_name = self.var_name_gen( + "temp_" + + var_name.replace("<", "_").replace(">", "_")) + substs.append((var_name, var(tmp_var_name))) + + tmp_stmt_id = self.stmt_id_gen("temp") + tmp_stmt_ids.append(tmp_stmt_id) + + new_tmp_stmt = Assign( + tmp_var_name, (), var(var_name), + condition=stmt.condition, + id=tmp_stmt_id, + depends_on=stmt.depends_on) + new_statements.append(new_tmp_stmt) + + from pymbolic import substitute + new_stmt = (stmt + .map_expressions( + lambda expr: substitute(expr, dict(substs)), + include_lhs=False) + .copy( + # lhs will be rewritten, but we don't want that. + depends_on=stmt.depends_on | frozenset(tmp_stmt_ids))) + new_statements.append(new_stmt) + + return new_statements + + +def eliminate_self_dependencies(phase_ast): + return apply_statement_rewriter(SelfDependencyEliminator, phase_ast) # }}} # {{{ isolate function arguments -class FunctionArgumentIsolator(IdentityMapper): +class ExprFunctionArgumentIsolator(IdentityMapper): def __init__(self, new_statements, stmt_id_gen, var_name_gen): super().__init__() @@ -142,40 +182,37 @@ def map_call_with_kwargs(self, expr, base_condition, base_deps, extra_deps): ) -def isolate_function_arguments(dag): - stmt_id_gen = dag.get_stmt_id_generator() - var_name_gen = dag.get_var_name_generator() - - new_phases = {} - for phase_name, phase in dag.phases.items(): +class StatementFunctionArgumentIsolator(ASTStatementRewriter): + def map_statement(self, stmt): new_statements = [] - fai = FunctionArgumentIsolator( + fai = ExprFunctionArgumentIsolator( new_statements=new_statements, - stmt_id_gen=stmt_id_gen, - var_name_gen=var_name_gen) + stmt_id_gen=self.stmt_id_gen, + var_name_gen=self.var_name_gen) + + base_deps = stmt.depends_on + new_deps = [] - for stmt in sorted(phase.statements, key=lambda stmt: stmt.id): - base_deps = stmt.depends_on - new_deps = [] + new_statements.append( + stmt + .map_expressions( + lambda expr: fai( + expr, stmt.condition, base_deps, new_deps)) + .copy(depends_on=stmt.depends_on | frozenset(new_deps))) - new_statements.append( - stmt - .map_expressions( - lambda expr: fai( - expr, stmt.condition, base_deps, new_deps)) - .copy(depends_on=stmt.depends_on | frozenset(new_deps))) + return new_statements - new_phases[phase_name] = phase.copy(statements=new_statements) - return dag.copy(phases=new_phases) +def isolate_function_arguments(phase_ast): + return apply_statement_rewriter(StatementFunctionArgumentIsolator, phase_ast) # }}} # {{{ isolate function calls -class FunctionCallIsolator(IdentityMapper): +class ExpressionFunctionCallIsolator(IdentityMapper): def __init__(self, new_statements, stmt_id_gen, var_name_gen): super().__init__() @@ -235,49 +272,35 @@ def map_call_with_kwargs(self, expr, base_condition, base_deps, extra_deps): .map_call_with_kwargs) -def isolate_function_calls_in_phase(phase, stmt_id_gen, var_name_gen): - new_statements = [] - - fci = FunctionCallIsolator( - new_statements=new_statements, - stmt_id_gen=stmt_id_gen, - var_name_gen=var_name_gen) - - for stmt in sorted(phase.statements, key=lambda stmt: stmt.id): - new_deps = [] - +class StatementFunctionCallIsolator(ASTStatementRewriter): + def map_statement(self, stmt): from dagrt.language import Assign - if isinstance(stmt, Assign): - new_statements.append( - stmt - .map_expressions( - lambda expr: fci( - expr, stmt.condition, stmt.depends_on, new_deps)) - .copy(depends_on=stmt.depends_on | frozenset(new_deps))) - from pymbolic.primitives import Call, CallWithKwargs - assert not isinstance(new_statements[-1].rhs, - (Call, CallWithKwargs)) - else: - new_statements.append(stmt) - - return phase.copy(statements=new_statements) + if not isinstance(stmt, Assign): + return stmt + new_deps = [] + new_statements = [] -def isolate_function_calls(dag): - """ - :func:`isolate_function_arguments` should be - called before this. - """ + fci = ExpressionFunctionCallIsolator( + new_statements=new_statements, + stmt_id_gen=self.stmt_id_gen, + var_name_gen=self.var_name_gen) + + new_statements.append( + stmt + .map_expressions( + lambda expr: fci( + expr, stmt.condition, stmt.depends_on, new_deps)) + .copy(depends_on=stmt.depends_on | frozenset(new_deps))) + from pymbolic.primitives import Call, CallWithKwargs + assert not isinstance(new_statements[-1].rhs, + (Call, CallWithKwargs)) - stmt_id_gen = dag.get_stmt_id_generator() - var_name_gen = dag.get_var_name_generator() + return new_statements - new_phases = {} - for phase_name, phase in dag.phases.items(): - new_phases[phase_name] = isolate_function_calls_in_phase( - phase, stmt_id_gen, var_name_gen) - return dag.copy(phases=new_phases) +def isolate_function_calls(phase_ast): + return apply_statement_rewriter(StatementFunctionCallIsolator, phase_ast) # }}} @@ -295,7 +318,7 @@ def flat_LogicalAnd(*children): # noqa # {{{ expand IfThenElse expressions -class IfThenElseExpander(IdentityMapper): +class ExprIfThenElseExpander(IdentityMapper): def __init__(self, new_statements, stmt_id_gen, var_name_gen): super().__init__() @@ -363,37 +386,33 @@ def map_if(self, expr, base_condition, base_deps, extra_deps): return var(tmp_result) -def expand_IfThenElse(dag): # noqa - """ - Turn IfThenElse expressions into values that are computed as a result of an - If statement. This is useful for targets that do not support ternary - operators. - """ - - stmt_id_gen = dag.get_stmt_id_generator() - var_name_gen = dag.get_var_name_generator() - - new_phases = {} - for phase_name, phase in dag.phases.items(): +class StatementIfThenElseExpander(ASTStatementRewriter): + def map_statement(self, stmt): new_statements = [] - expander = IfThenElseExpander( + expander = ExprIfThenElseExpander( new_statements=new_statements, - stmt_id_gen=stmt_id_gen, - var_name_gen=var_name_gen) + stmt_id_gen=self.stmt_id_gen, + var_name_gen=self.var_name_gen) - for stmt in phase.statements: - base_deps = stmt.depends_on - new_deps = [] + base_deps = stmt.depends_on + new_deps = [] - new_statements.append( - stmt.map_expressions( - lambda expr: expander(expr, stmt.condition, base_deps, new_deps)) - .copy(depends_on=stmt.depends_on | frozenset(new_deps))) + new_statements.append( + stmt.map_expressions( + lambda expr: expander(expr, stmt.condition, base_deps, new_deps)) + .copy(depends_on=stmt.depends_on | frozenset(new_deps))) - new_phases[phase_name] = phase.copy(statements=new_statements) + return new_statements - return dag.copy(phases=new_phases) + +def expand_IfThenElse(phase_ast): # noqa + """ + Turn IfThenElse expressions into values that are computed as a result of an + If statement. This is useful for targets that do not support ternary + operators. + """ + return apply_statement_rewriter(StatementIfThenElseExpander, phase_ast) # }}} diff --git a/test/test_codegen_fortran.py b/test/test_codegen_fortran.py index 0341b33..420efd3 100755 --- a/test/test_codegen_fortran.py +++ b/test/test_codegen_fortran.py @@ -115,6 +115,41 @@ def test_arrays_and_linalg(): fortran_libraries=["lapack", "blas"]) +def test_self_dep_in_loop(): + with CodeBuilder(name="primary") as cb: + cb("y", "y") + cb("y", "f(0, 2*i*f(0, y if i > 2 else 2*y))", + loops=(("i", 0, 5),)) + cb("y", "y") + + code = create_DAGCode_with_steady_phase(cb.statements) + + rhs_function = "f" + + from dagrt.function_registry import ( + base_function_registry, register_ode_rhs) + freg = register_ode_rhs(base_function_registry, "ytype", + identifier=rhs_function, + input_names=("y",)) + freg = freg.register_codegen(rhs_function, "fortran", + f.CallCode(""" + ${result} = -2*${y} + """)) + + codegen = f.CodeGenerator( + "selfdep", + function_registry=freg, + user_type_map={"ytype": f.ArrayType((100,), f.BuiltinType("real*8"))}, + timing_function="second") + + code_str = codegen(code) + run_fortran([ + ("selfdep.f90", code_str), + ("test_selfdep.f90", read_file("test_selfdep.f90")), + ], + fortran_libraries=["lapack", "blas"]) + + if __name__ == "__main__": if len(sys.argv) > 1: exec(sys.argv[1]) From 7590d18c0c908582b62d5dec479f38d2195e68f8 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:39:55 -0500 Subject: [PATCH 07/16] Fix previously-unused path in _FunctioNameCollector --- dagrt/codegen/analysis.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/dagrt/codegen/analysis.py b/dagrt/codegen/analysis.py index 153b2f2..a99ec9b 100644 --- a/dagrt/codegen/analysis.py +++ b/dagrt/codegen/analysis.py @@ -180,11 +180,11 @@ def map_variable(self, expr): return set() def map_call(self, expr): - return ({expr.function} + return ({expr.function.name} | super().map_call(expr)) def map_call_with_kwargs(self, expr): - return ({expr.function} + return ({expr.function.name} | super().map_call_with_kwargs(expr)) From 66fbce31dece380f936c7985d4829acf64b10be3 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:40:22 -0500 Subject: [PATCH 08/16] AST creation: Only Assign nodes have loops --- dagrt/codegen/dag_ast.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/dagrt/codegen/dag_ast.py b/dagrt/codegen/dag_ast.py index c91b222..bdd5776 100644 --- a/dagrt/codegen/dag_ast.py +++ b/dagrt/codegen/dag_ast.py @@ -25,7 +25,7 @@ from pymbolic.mapper import IdentityMapper, Collector from pymbolic.mapper.stringifier import StringifyMapper from pymbolic.primitives import Expression, LogicalNot -from dagrt.language import Nop +from dagrt.language import Nop, Assign # {{{ ast node types @@ -283,7 +283,7 @@ def conditional_to_ast(statement): def loop_to_ast_node(statement): - if statement.loops: + if isinstance(statement, Assign) and statement.loops: loop_var_name, lower, upper = statement.loops[0] new_statement = statement.copy(loops=statement.loops[1:]) return ForLoop( From 5f9319ccbbe1e24dbb509f85f9e133dd2ab25755 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:40:52 -0500 Subject: [PATCH 09/16] Fortran function name collection: functions may now be hiding in expressions --- dagrt/codegen/fortran.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/dagrt/codegen/fortran.py b/dagrt/codegen/fortran.py index 8562a89..99cbb5b 100644 --- a/dagrt/codegen/fortran.py +++ b/dagrt/codegen/fortran.py @@ -993,7 +993,7 @@ def __init__(self, module_name, def get_called_function_names(self, dag): from dagrt.codegen.analysis import collect_function_names_from_dag - result = collect_function_names_from_dag(dag, no_expressions=True) + result = collect_function_names_from_dag(dag) if self.call_before_state_update: result.add(self.call_before_state_update) From 1519e5ed2a1c25103f708b3ee659a6d15b215b2c Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:41:15 -0500 Subject: [PATCH 10/16] Python codegen: support for loop --- dagrt/codegen/python.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/dagrt/codegen/python.py b/dagrt/codegen/python.py index 3e65ca1..00e6f42 100644 --- a/dagrt/codegen/python.py +++ b/dagrt/codegen/python.py @@ -447,12 +447,20 @@ def emit_def_end(self): del self._emitter def emit_if_begin(self, expr): - self._emit("if {expr}:".format(expr=self._expr(expr))) + self._emit(f"if {self._expr(expr)}:") self._emitter.indent() def emit_if_end(self): self._emitter.dedent() + def emit_for_begin(self, loop_var_name, lbound, ubound): + self._emit(f"for {self._name_manager[loop_var_name]} in " + f"range({self._expr(lbound)}, {self._expr(ubound)}):") + self._emitter.indent() + + def emit_for_end(self, loop_var_name): + self._emitter.dedent() + def emit_else_begin(self): self._emitter.dedent() self._emit("else:") From c04aa0c45a86a9446903641d5b3c6b7723f134b9 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:41:43 -0500 Subject: [PATCH 11/16] Placate pylint: ASTStatementRewriter needs a map_statement method --- dagrt/codegen/transform.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/dagrt/codegen/transform.py b/dagrt/codegen/transform.py index 2489c90..a87f1e5 100644 --- a/dagrt/codegen/transform.py +++ b/dagrt/codegen/transform.py @@ -66,6 +66,9 @@ def map_StatementWrapper(self, expr): else: return new_statements[0] + def map_statement(self, statement): + raise NotImplementedError() + def apply_statement_rewriter(rewriter_cls, phase_ast): statements = list(get_statements_in_ast(phase_ast)) From c4069975e7a5892ca1bc59ac20ed581514a3d654 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:42:10 -0500 Subject: [PATCH 12/16] Drop test_IfThenElse_expansion --- test/test_codegen_python.py | 9 --------- 1 file changed, 9 deletions(-) diff --git a/test/test_codegen_python.py b/test/test_codegen_python.py index 155a38b..6bdc4b4 100755 --- a/test/test_codegen_python.py +++ b/test/test_codegen_python.py @@ -324,15 +324,6 @@ def test_IfThenElse(python_method_impl): assert result == expected_result -def test_IfThenElse_expansion(python_method_impl): - from utils import execute_and_return_single_result - code, expected_result = get_IfThenElse_test_code_and_expected_result() - from dagrt.codegen.transform import expand_IfThenElse - code = expand_IfThenElse(code) - result = execute_and_return_single_result(python_method_impl, code) - assert result == expected_result - - def test_arrays_and_looping(python_method_impl): with CodeBuilder(name="primary") as cb: cb("myarray", "`array`(20)") From 0b54d4c58fc3e36ffa5f0fa91df164df0d81c002 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 18:42:40 -0500 Subject: [PATCH 13/16] Fix early escape in StatementFunctionCallIsolator --- dagrt/codegen/transform.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/dagrt/codegen/transform.py b/dagrt/codegen/transform.py index a87f1e5..9a7ce80 100644 --- a/dagrt/codegen/transform.py +++ b/dagrt/codegen/transform.py @@ -279,7 +279,7 @@ class StatementFunctionCallIsolator(ASTStatementRewriter): def map_statement(self, stmt): from dagrt.language import Assign if not isinstance(stmt, Assign): - return stmt + return [stmt] new_deps = [] new_statements = [] From 068d4d301b339fbfdf70c2663a62794b4750632e Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 19:00:10 -0500 Subject: [PATCH 14/16] Code builder: Don't silently drop loops on the floor if generating AssignFunctionCall --- dagrt/language.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/dagrt/language.py b/dagrt/language.py index b61c72d..8ed7ee2 100644 --- a/dagrt/language.py +++ b/dagrt/language.py @@ -1048,7 +1048,7 @@ def parse_if_necessary(s): from pymbolic.primitives import Call, CallWithKwargs, Variable - if isinstance(expression, (Call, CallWithKwargs)): + if isinstance(expression, (Call, CallWithKwargs)) and not loops: assignee_names = [] for a in assignees: if not isinstance(a, Variable): From 6c2b7cb7cab2d1f89f68d682b122484956e0a82a Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 19:00:50 -0500 Subject: [PATCH 15/16] Enable direct printing of ASTs --- dagrt/codegen/dag_ast.py | 17 +++++++++++------ dagrt/codegen/fortran.py | 3 +-- 2 files changed, 12 insertions(+), 8 deletions(-) diff --git a/dagrt/codegen/dag_ast.py b/dagrt/codegen/dag_ast.py index bdd5776..5058785 100644 --- a/dagrt/codegen/dag_ast.py +++ b/dagrt/codegen/dag_ast.py @@ -30,7 +30,12 @@ # {{{ ast node types -class IfThen(Expression): +class ASTNode(Expression): # not really, but it lets us abuse pymbolic's machinery + def __str__(self): + return ASTStringifier()(self, 0) + + +class IfThen(ASTNode): """ .. attribute: condition .. attribute: then @@ -48,7 +53,7 @@ def __getinitargs__(self): mapper_method = "map_IfThen" -class IfThenElse(Expression): +class IfThenElse(ASTNode): """ .. attribute: condition .. attribute: then @@ -68,7 +73,7 @@ def __getinitargs__(self): mapper_method = "map_IfThenElse" -class ForLoop(Expression): +class ForLoop(ASTNode): """ Bounds are a half-open interval as in Python @@ -91,7 +96,7 @@ def __getinitargs__(self): mapper_method = "map_ForLoop" -class Block(Expression): +class Block(ASTNode): """ .. attribute: children """ @@ -107,7 +112,7 @@ def __getinitargs__(self): mapper_method = "map_Block" -class NullASTNode(Expression): +class NullASTNode(ASTNode): init_arg_names = () @@ -117,7 +122,7 @@ def __getinitargs__(self): mapper_method = "map_NullASTNode" -class StatementWrapper(Expression): +class StatementWrapper(ASTNode): """ .. attribute: statement """ diff --git a/dagrt/codegen/fortran.py b/dagrt/codegen/fortran.py index 99cbb5b..9361a59 100644 --- a/dagrt/codegen/fortran.py +++ b/dagrt/codegen/fortran.py @@ -1072,8 +1072,7 @@ def process_ast(ast, print_ast=False): ast = expand_IfThenElse(ast) if print_ast: - from dagrt.codegen.dag_ast import ASTStringifier - print(ASTStringifier()(ast, 0)) + print(ast) return ast From 5e8e6d10ce3e803f686ceeb70cecdb944ebdce96 Mon Sep 17 00:00:00 2001 From: Andreas Kloeckner Date: Thu, 27 May 2021 19:02:44 -0500 Subject: [PATCH 16/16] Add missing test_selfdep.f90 --- test/test_selfdep.f90 | 31 +++++++++++++++++++++++++++++++ 1 file changed, 31 insertions(+) create mode 100644 test/test_selfdep.f90 diff --git a/test/test_selfdep.f90 b/test/test_selfdep.f90 new file mode 100644 index 0000000..1b37013 --- /dev/null +++ b/test/test_selfdep.f90 @@ -0,0 +1,31 @@ +program test_selfdep + + use selfdep, only: dagrt_state_type, & + timestep_initialize => initialize, & + timestep_run => run, & + timestep_shutdown => shutdown + + implicit none + + type(dagrt_state_type), target :: dagrt_state + type(dagrt_state_type), pointer :: dagrt_state_ptr + + real*8, dimension(100) :: y0 + + integer i + + ! start code ---------------------------------------------------------------- + + dagrt_state_ptr => dagrt_state + + + do i = 1, 100 + y0 = i + end do + + call timestep_initialize(dagrt_state=dagrt_state_ptr, state_y=y0) + call timestep_run(dagrt_state=dagrt_state_ptr) + call timestep_shutdown(dagrt_state=dagrt_state_ptr) + +end program +