diff --git a/asm/builtins.py b/asm/builtins.py index e8012d6..b212248 100644 --- a/asm/builtins.py +++ b/asm/builtins.py @@ -19,9 +19,8 @@ BUILTIN_MACROS = """\ ; Topology: counter -> compare -> body_fan (fanout) -> inc -> feedback ; compare L -> body_fan (pass as fanout) -> @ret_body + &inc ; compare R -> @ret_exit -; User wires: init -> counter:L, limit -> compare:R -; Call with: #loop_counted |> body=&process, exit=&done -#loop_counted |> { +; Call with: #loop_counted &init, &limit |> body=&process, exit=&done +#loop_counted init, limit |> { &counter <| add &compare <| brgt &counter |> &compare:L @@ -30,71 +29,83 @@ BUILTIN_MACROS = """\ &inc <| inc &body_fan |> &inc:L &inc |> &counter:R + ${init} |> &counter:L + ${limit} |> &compare:R &body_fan |> @ret_body &compare |> @ret_exit:R } ; --- Condition-tested loop --- ; Topology: &gate (gate) -> L=body, R=exit -; User wires: test_value -> gate:L -; Call with: #loop_while |> body=&process, exit=&done -#loop_while |> { +; Call with: #loop_while &test_src |> body=&process, exit=&done +#loop_while test |> { &gate <| gate + ${test} |> &gate:L &gate |> @ret_body &gate |> @ret_exit:R } ; --- Permit injection (per-arity variants) --- -; Each injects const 1 tokens. User wires outputs to their gate node. +; Each injects const 1 tokens. Outputs wired via @ret. +; Call with: #permit_inject_1 |> out=&gate_node #permit_inject_1 |> { &p0 <| const, 1 + &p0 |> @ret_out } +; Call with: #permit_inject_2 |> out0=&gate_a, out1=&gate_b #permit_inject_2 |> { &p0 <| const, 1 &p1 <| const, 1 + &p0 |> @ret_out0 + &p1 |> @ret_out1 } +; Call with: #permit_inject_3 |> out0=&g0, out1=&g1, out2=&g2 #permit_inject_3 |> { &p0 <| const, 1 &p1 <| const, 1 &p2 <| const, 1 - &merge <| merge - &p1 |> &merge:L - &p2 |> &merge:R + &p0 |> @ret_out0 + &p1 |> @ret_out1 + &p2 |> @ret_out2 } +; Call with: #permit_inject_4 |> out0=&g0, out1=&g1, out2=&g2, out3=&g3 #permit_inject_4 |> { &p0 <| const, 1 &p1 <| const, 1 &p2 <| const, 1 &p3 <| const, 1 - &merge_a <| merge - &p0 |> &merge_a:L - &p1 |> &merge_a:R - &merge_b <| merge - &p2 |> &merge_b:L - &p3 |> &merge_b:R + &p0 |> @ret_out0 + &p1 |> @ret_out1 + &p2 |> @ret_out2 + &p3 |> @ret_out3 } ; --- Binary reduction trees (parameterized opcode) --- -; Usage: #reduce_2 add, #reduce_3 sub, etc. +; Call with: #reduce_2 add |> &result_node #reduce_2 op |> { &r <| ${op} + &r |> @ret } +; Call with: #reduce_3 sub |> &result_node #reduce_3 op |> { &r0 <| ${op} &r1 <| ${op} &r0 |> &r1:L + &r1 |> @ret } +; Call with: #reduce_4 add |> &result_node #reduce_4 op |> { &r0 <| ${op} &r1 <| ${op} &r2 <| ${op} &r0 |> &r2:L &r1 |> &r2:R + &r2 |> @ret } """ diff --git a/asm/expand.py b/asm/expand.py index 15fca61..3d1f8f2 100644 --- a/asm/expand.py +++ b/asm/expand.py @@ -780,6 +780,23 @@ def _expand_call( rewritten_edges = [] all_outputs = list(call.output_dests) + # Build ordered list of positional outputs for bare @ret resolution. + # Each bare @ret consumes the next positional output in order, + # enabling variadic macros to wire each iteration to a separate dest: + # #macro *vals |> { $( &c <| const, ${vals}; &c |> @ret ),* } + # #macro 3, 4 |> &x, &y ← iteration 0 → &x, iteration 1 → &y + positional_outputs: list[str] = [] + for output in all_outputs: + if isinstance(output, dict): + if "name" in output and "ref" in output: + continue # Named output — not positional + name = output.get("name", None) + if name is not None: + positional_outputs.append(name) + else: + positional_outputs.append(str(output)) + positional_idx = 0 + for edge in expanded_edges: if not (isinstance(edge.dest, str) and edge.dest.startswith("@ret")): rewritten_edges.append(edge) @@ -799,19 +816,11 @@ def _expand_call( dest_name = ref.get("name", ref) if isinstance(ref, dict) else str(ref) break - # Bare @ret -> first positional output + # Bare @ret -> next positional output (advances counter) if dest_name is None and ret_dest == "@ret": - for output in all_outputs: - if isinstance(output, dict): - if "name" in output and "ref" in output: - # Named output — skip for bare @ret - continue - # Positional output (just a ref dict) - dest_name = output.get("name", None) - else: - dest_name = str(output) - if dest_name is not None: - break + if positional_idx < len(positional_outputs): + dest_name = positional_outputs[positional_idx] + positional_idx += 1 if dest_name is None: errors.append(AssemblyError( diff --git a/tests/test_builtins.py b/tests/test_builtins.py index d40902e..d0c46f0 100644 --- a/tests/test_builtins.py +++ b/tests/test_builtins.py @@ -107,7 +107,8 @@ class TestAC81_BuiltinAvailable: """Invoking a built-in macro through pipeline produces expanded nodes.""" source = """ @system pe=1, sm=0 - #permit_inject_1 + &sink <| pass + #permit_inject_1 |> out=&sink """ graph = run_pipeline(source) assert len(graph.errors) == 0 @@ -154,7 +155,8 @@ class TestAC82_UserMacroShadows: """Built-in macro is used when user doesn't define it.""" source = """ @system pe=1, sm=0 - #reduce_2 add + &sink <| pass + #reduce_2 add |> &sink """ graph = run_pipeline(source) assert graph is not None @@ -178,9 +180,11 @@ class TestAC83_LoopCountedTopology: """ source = """ @system pe=1, sm=0 + &init <| const, 0 + &limit <| const, 10 &body <| pass &exit <| pass - #loop_counted |> body=&body, exit=&exit + #loop_counted &init, &limit |> body=&body, exit=&exit """ graph = run_pipeline(source) assert len(graph.errors) == 0 @@ -206,9 +210,11 @@ class TestAC83_LoopCountedTopology: """ source = """ @system pe=1, sm=0 + &init <| const, 0 + &limit <| const, 10 &body <| pass &exit <| pass - #loop_counted |> body=&body, exit=&exit + #loop_counted &init, &limit |> body=&body, exit=&exit """ graph = run_pipeline(source) @@ -240,7 +246,8 @@ class TestAC84_EndToEnd: """#reduce_2 add invocation assembles through the full pipeline.""" source = """ @system pe=1, sm=0 - #reduce_2 add + &out <| pass + #reduce_2 add |> &out """ result = assemble(source) assert result is not None, "assemble() should succeed" @@ -250,22 +257,22 @@ class TestAC84_EndToEnd: """#reduce_2 add runs through emulator without error.""" source = """ @system pe=1, sm=0 - #reduce_2 add + &out <| pass + #reduce_2 add |> &out """ outputs = run_program_direct(source, until=500) assert 0 in outputs, "PE 0 should exist in outputs" def test_builtin_reduce_2_produces_output_when_wired(self): - """#reduce_2 add produces correct sum when inputs and output are wired.""" + """#reduce_2 add produces correct sum when inputs and output are wired via @ret.""" source = """ @system pe=1, sm=0 &a <| const, 3 &b <| const, 4 &out <| pass - #reduce_2 add + #reduce_2 add |> &out &a |> #reduce_2_0.&r:L &b |> #reduce_2_0.&r:R - #reduce_2_0.&r |> &out:L """ outputs = run_program_direct(source, until=500) all_values = [] @@ -277,7 +284,8 @@ class TestAC84_EndToEnd: """#permit_inject_1 assembles and expands to const node.""" source = """ @system pe=1, sm=0 - #permit_inject_1 + &sink <| pass + #permit_inject_1 |> out=&sink """ result = assemble(source) assert result is not None @@ -290,7 +298,8 @@ class TestAC84_EndToEnd: """#reduce_3 add invocation assembles through the full pipeline.""" source = """ @system pe=1, sm=0 - #reduce_3 add + &out <| pass + #reduce_3 add |> &out """ result = assemble(source) assert result is not None, "assemble() should succeed" @@ -300,7 +309,8 @@ class TestAC84_EndToEnd: """#reduce_3 add runs through emulator without error.""" source = """ @system pe=1, sm=0 - #reduce_3 add + &out <| pass + #reduce_3 add |> &out """ outputs = run_program_direct(source, until=500) assert 0 in outputs, "PE 0 should exist in outputs" @@ -389,8 +399,9 @@ class TestBuiltinComposition: &val <| const, 77 } + &sink <| pass #my_const - #permit_inject_1 + #permit_inject_1 |> out=&sink """ graph = run_pipeline(source) assert len(graph.errors) == 0 diff --git a/tests/test_macro_syntax.py b/tests/test_macro_syntax.py index 17c63b8..486e87c 100644 --- a/tests/test_macro_syntax.py +++ b/tests/test_macro_syntax.py @@ -196,15 +196,14 @@ class TestScopedReferences: def test_macro_scoped_ref_in_edge_source(self, parser): """Parse macro scoped reference as edge source. - Note: Macro resolution happens in Phase 2. Here we just verify - the scoped_ref syntax parses and creates the qualified name. + Lower creates the qualified name; expand resolves it to the + actual expanded node name (e.g., #macro_0.&label). """ graph = parse_and_lower(parser, """\ &dest <| pass #macro.&label |> &dest:L """) - # Edge source should contain the scoped_ref syntax assert len(graph.edges) > 0 edge = graph.edges[0] assert edge.source == "#macro.&label" @@ -250,26 +249,26 @@ class TestMacroRefGrammar: edge = graph.edges[0] assert edge.source == "$func.&inner" - def test_scoped_ref_with_node_ref(self, parser): - """Parse scoped_ref using node_ref as inner. + def test_node_ref_inside_function_stays_global(self, parser): + """@name refs inside function bodies are NOT scope-qualified. - Note: In practice, node_ref (@name) in scoped context is unusual. - This tests the grammar accepts it. + Only &label refs get qualified with function scope. @name refs + are global (with @ret/@ret_name as special exceptions handled + by the expand pass). """ - # This is a theoretical test - using @name in function scope - # might not be semantically valid, but grammar should accept it graph = parse_and_lower(parser, """\ $func |> { @inner <| pass + &local <| pass } - &dest <| pass - $func.@inner |> &dest:L """) - # Scoped ref with node_ref should parse - assert len(graph.edges) > 0 - edge = graph.edges[0] - assert edge.source == "$func.@inner" + func_region = next( + r for r in graph.regions if r.kind == RegionKind.FUNCTION + ) + node_names = set(func_region.body.nodes.keys()) + assert "@inner" in node_names, f"@inner should stay unqualified, got {node_names}" + assert "$func.&local" in node_names, f"&local should be qualified, got {node_names}" class TestMacroInContext: diff --git a/tests/test_variadic.py b/tests/test_variadic.py index fd36bfb..e2da4ef 100644 --- a/tests/test_variadic.py +++ b/tests/test_variadic.py @@ -7,19 +7,31 @@ Tests verify: - Empty variadic invocation: no error, nothing expanded - Single variadic invocation: one iteration - Variadic parameter not last: error at lower pass +- Full pipeline: variadic macro assembles and runs in emulator """ from pathlib import Path -from asm.expand import expand -from asm.lower import lower +import simpy +from lark import Lark + +from asm import assemble from asm.errors import ErrorCategory +from asm.expand import expand from asm.ir import ( - IRGraph, IRNode, IREdge, MacroDef, MacroParam, ParamRef, SourceLoc, - IRMacroCall, IRRepetitionBlock + IREdge, + IRGraph, + IRMacroCall, + IRNode, + IRRepetitionBlock, + MacroDef, + MacroParam, + ParamRef, + SourceLoc, ) +from asm.lower import lower from cm_inst import ArithOp -from lark import Lark +from emu import build_topology def _get_parser(): @@ -167,16 +179,21 @@ class TestVariadicIndexVariable: # After expansion, should have 3 nodes with names: # #maker_0_rep0.&node_0, #maker_0_rep1.&node_1, #maker_0_rep2.&node_2 node_names = list(expanded.nodes.keys()) - assert len(node_names) == 3, f"Expected 3 nodes, got {len(node_names)}: {node_names}" + assert len(node_names) == 3, ( + f"Expected 3 nodes, got {len(node_names)}: {node_names}" + ) # Verify that _idx was substituted correctly in node names # Each iteration should have node_0, node_1, node_2 respectively - assert any("&node_0" in name for name in node_names), \ + assert any("&node_0" in name for name in node_names), ( f"Expected node with &node_0 (iteration 0), got {node_names}" - assert any("&node_1" in name for name in node_names), \ + ) + assert any("&node_1" in name for name in node_names), ( f"Expected node with &node_1 (iteration 1), got {node_names}" - assert any("&node_2" in name for name in node_names), \ + ) + assert any("&node_2" in name for name in node_names), ( f"Expected node with &node_2 (iteration 2), got {node_names}" + ) class TestVariadicMixedParams: @@ -235,7 +252,9 @@ class TestVariadicEdgeCases: # No nodes should be created nodes = list(graph.nodes.keys()) - assert len(nodes) == 0, f"Expected 0 nodes for empty variadic, got {len(nodes)}: {nodes}" + assert len(nodes) == 0, ( + f"Expected 0 nodes for empty variadic, got {len(nodes)}: {nodes}" + ) def test_single_variadic_one_iteration(self): """Invoke with one variadic arg: one iteration.""" @@ -274,8 +293,9 @@ class TestVariadicGrammarValidation: graph = parse_and_lower(source) # Lower pass should catch this error - assert any(e.category == ErrorCategory.NAME for e in graph.errors), \ + assert any(e.category == ErrorCategory.NAME for e in graph.errors), ( f"Expected NAME error for variadic not last, got: {graph.errors}" + ) def test_multiple_variadic_is_error(self): """Multiple variadic parameters: parser/lower should reject.""" @@ -289,8 +309,9 @@ class TestVariadicGrammarValidation: graph = parse_and_lower(source) # Lower pass should catch this error - assert any(e.category == ErrorCategory.NAME for e in graph.errors), \ + assert any(e.category == ErrorCategory.NAME for e in graph.errors), ( f"Expected NAME error for multiple variadic, got: {graph.errors}" + ) class TestVariadicIntegration: @@ -342,8 +363,12 @@ class TestVariadicIntegration: # Second invocation should have #expand_1_rep0, #expand_1_rep1, #expand_1_rep2 expand_0 = [n for n in nodes if "#expand_0" in n] expand_1 = [n for n in nodes if "#expand_1" in n] - assert len(expand_0) == 2, f"Expected 2 nodes from first invocation, got {len(expand_0)}" - assert len(expand_1) == 3, f"Expected 3 nodes from second invocation, got {len(expand_1)}" + assert len(expand_0) == 2, ( + f"Expected 2 nodes from first invocation, got {len(expand_0)}" + ) + assert len(expand_1) == 3, ( + f"Expected 3 nodes from second invocation, got {len(expand_1)}" + ) def test_variadic_nested_with_other_macros(self): """Variadic macro combined with non-variadic macros.""" @@ -372,3 +397,120 @@ class TestVariadicIntegration: expand_nodes = [n for n in nodes if "#expand" in n] assert len(simple_nodes) == 1, f"Expected 1 #simple node" assert len(expand_nodes) == 2, f"Expected 2 #expand nodes" + + +class TestVariadicPositionalRet: + """Test positional @ret wiring in variadic repetition blocks.""" + + def test_bare_ret_maps_to_positional_outputs_by_iteration(self): + """@ret in iteration N maps to Nth positional output at call site.""" + source = """ + @system pe=1, sm=1 + + #fan *vals |> { + $( &v <| pass + &v |> @ret ),* + } + + &a <| pass + &b <| pass + &c <| pass + #fan &a, &b, &c |> &x, &y, &z + &x <| pass + &y <| pass + &z <| pass + """ + graph = parse_lower_expand(source) + assert not graph.errors, f"Unexpected errors: {graph.errors}" + + # Check edges from expanded nodes to positional outputs + ret_edges = [e for e in graph.edges if e.dest in ("&x", "&y", "&z")] + dests = [e.dest for e in ret_edges] + assert "&x" in dests, f"Expected &x in ret edge dests: {dests}" + assert "&y" in dests, f"Expected &y in ret edge dests: {dests}" + assert "&z" in dests, f"Expected &z in ret edge dests: {dests}" + + def test_positional_ret_with_two_outputs(self): + """Two variadic iterations map to two positional outputs.""" + source = """ + @system pe=1, sm=1 + + #pair *vals |> { + $( &v <| pass + &v |> @ret ),* + } + + &a <| pass + &b <| pass + &left <| pass + &right <| pass + #pair &a, &b |> &left, &right + """ + graph = parse_lower_expand(source) + assert not graph.errors, f"Unexpected errors: {graph.errors}" + + ret_edges = [e for e in graph.edges if e.dest in ("&left", "&right")] + assert len(ret_edges) == 2, ( + f"Expected 2 ret edges, got {len(ret_edges)}: {ret_edges}" + ) + + def test_fewer_outputs_than_iterations_errors(self): + """More @ret iterations than positional outputs produces errors.""" + source = """ + @system pe=1, sm=1 + + #too_many *vals |> { + $( &v <| pass + &v |> @ret ),* + } + + &a <| pass + &b <| pass + &c <| pass + &only_one <| pass + #too_many &a, &b, &c |> &only_one + """ + graph = parse_lower_expand(source) + macro_errors = [e for e in graph.errors if e.category == ErrorCategory.MACRO] + assert len(macro_errors) >= 1, ( + f"Expected error for unmatched @ret, got: {graph.errors}" + ) + + +class TestVariadicFullPipeline: + """Full pipeline test: variadic macro through assemble and emulator.""" + + def test_variadic_macro_assembles_and_runs(self): + """Variadic macro with positional @ret wiring through full pipeline. + + Each iteration's @ret maps to the next positional output at the + call site, so `#multi_const 3, 4 |> &sum:L, &sum:R` wires + iteration 0 → &sum:L and iteration 1 → &sum:R. + """ + source = """ + @system pe=1, sm=0 + + #multi_const *vals |> { + $( &c <| const, ${vals} + &c |> @ret ),* + } + + &sum <| add + &out <| pass + #multi_const 3, 4 |> &sum:L, &sum:R + &sum |> &out:L + """ + result = assemble(source) + assert result is not None + assert len(result.pe_configs) > 0 + + env = simpy.Environment() + sys = build_topology(env, result.pe_configs, result.sm_configs) + for seed in result.seed_tokens: + sys.inject(seed) + env.run(until=500) + + all_values = [] + for pe in sys.pes.values(): + all_values.extend(t.data for t in pe.output_log if hasattr(t, "data")) + assert 7 in all_values, f"Expected 3+4=7 in outputs, got {all_values}"