From 1f560fe64f2c60d2baa7b846e52e93793a68a928 Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Wed, 12 Aug 2026 21:46:36 -0700 Subject: [PATCH 01/68] Bounded stream and `stream.close` specification --- irspec/docs/spatial/routing.md | 73 ++++++++++++++++++++++++++- irspec/docs/spatial/spatial.md | 91 ++++++++++++++++++++++++++++++++-- 2 files changed, 158 insertions(+), 6 deletions(-) diff --git a/irspec/docs/spatial/routing.md b/irspec/docs/spatial/routing.md index 22f6130c..bea1f166 100644 --- a/irspec/docs/spatial/routing.md +++ b/irspec/docs/spatial/routing.md @@ -145,9 +145,10 @@ Next, we describe the condition under which the routing behavior is undefined: and write $(S_1, (i_1, j_1), S_2, (i_2, j_2)) \mapsto (S_3, (i_3, j_3), S_4, (i_4, j_4))$ if $S_2, (i_2, j_2) \longmapsto S_3, (i_3, j_3)$. -!!! danger "Error: Undefined Behavior" - If two paths $P_1$ and $P_1$ in the routing graph use the same channel, share a PE, and +!!! danger "Error: Concurrent Channel Use" + If two paths $P_1$ and $P_2$ in the routing graph use the same channel, share a PE, and their corresponding stream edges are not ordered by empties-before, then the behavior is undefined. + *This raises a compile error whenever the missing ordering can be established statically.* This is because the two messages may interfere with each other @@ -156,6 +157,29 @@ Recall that sending onto the same stream [must be synchronized using completions to avoid data races](../spatial#streaming-data-with-send). Hence, sending through the same stream multiple times in the same phases is ok as long as the sends (and receives) are correctly synchronized. +The constructive way to establish the ordering between two streams that share a channel is to +[close](../spatial#closing-streams-with-close) the earlier one. Closing a stream ends its *epoch*: +on every PE of the path, the channel is released and may be taken over by the next stream. + +!!! abstract "Definition: Channel Epoch" + An *epoch* of a channel $C$ at PE $(i, j)$ is a maximal interval during which a single stream + that uses $C$ occupies $(i, j)$. It begins at the first use of that stream and ends at its + `close` (which, for a [bounded](../spatial#streams) stream, is implicit after its `BOUND` + elements have been transferred, and, for any stream, is implicit at the end of its phase). + +!!! abstract "Lemma: Sufficient Condition for Channel Reuse" + Let $F_1$ and $F_2$ be two streams that use the same channel $C$, and let $(i, j)$ be a PE + shared by their paths. Let $S_c$ be the `close` of $F_1$ at $(i, j)$ and $S_u$ the first use of + $F_2$ at $(i, j)$. If $S_c, (i, j) \longmapsto S_u, (i, j)$ for every shared PE $(i, j)$, + then the stream edges of $F_1$ empty-before those of $F_2$ and the reuse of $C$ is well-defined. + +Note that a phase boundary satisfies the condition of the lemma at every PE, which is why streams +in different phases may share a channel without an explicit `close`. + +Within a single epoch, a stream may not be used in two different route configurations at the same +PE. In particular, a PE that both receives from and sends on the same channel must close the +channel in between, since the two uses require incompatible router configurations. + Keep in mind that PEs transition between phases asynchronously, that is, a PE may advance to the next phase before another PE has completed the current phase. We exploit here implicitly that routers back-pressure when @@ -167,3 +191,48 @@ to receive. where all streams are point-to-point paths. If multicasting is used, the correctness conditions must be adapted accordingly, especially when considering multiple phases. + + +## Lowering to Switches + +Channels are a scarce resource: each channel that is live at a PE occupies one of the hardware's +routing colors. Epochs are what makes it possible to reuse a channel, and hence a color, for +several streams. + +A stream induces, at each PE of its path, a *route configuration*: the set of directions the PE +receives from and the set of directions it transmits to (where `RAMP` denotes the PE's own compute +element). Consider a fixed channel $C$ and a fixed PE $(i, j)$. Ordering the streams that use $C$ +at $(i, j)$ by their epochs yields a sequence of route configurations +$R_0, R_1, \dotsc, R_{n-1}$, which is realized by the PE's *switch* for the color assigned to $C$: +$R_0$ is the initial configuration and the router *advances* to $R_{k+1}$ at the epoch boundary. + +!!! danger "Error: Too Many Route Configurations" + A router holds a bounded number of route configurations per color (four on both WSE-2 and + WSE-3). *If the streams sharing a channel require more configurations than that at a single PE, + a compile error is raised.* Assigning a different channel to some of the streams resolves it, + at the cost of an additional color. + +Consecutive configurations that are equal do not consume a position and do not require an advance. +This is a common case: two streams declared as `relative_stream(-2, 0)` in successive phases induce +the same configuration at every PE of their paths, so their shared channel needs no switching at +all. + +An advance is driven by the `close` that ends the epoch, and is emitted only at those PEs whose +next configuration differs. Two lowerings are available: + +- The sending PE marks the last transfer of a bounded stream so that its router advances once the + stream's `BOUND` elements have left the fabric. This adds no traffic to the channel. +- The sending PE emits a *switch-advance control message* on the channel. It follows the stream's + path using the configuration that is being retired, and advances the router of each PE it + traverses, after all data of the epoch. PEs on the path whose configuration does not change are + skipped. + +Because the control message travels the path of the retired configuration in order behind the data, +a receiving PE needs to emit nothing to advance its own router: the ordering required by the +[lemma above](#undefined-behavior) is provided by the fabric. A receiver's `close` therefore has no +runtime effect; it exists so that the lifetime of the stream — and hence the number of elements it +carries — is stated by every participant and can be checked. + +!!! note "Note: Number of Control Messages" + A single control message can carry advance commands for a bounded number of consecutive routers + (eight on both WSE-2 and WSE-3). *A path that would require more raises a compile error.* diff --git a/irspec/docs/spatial/spatial.md b/irspec/docs/spatial/spatial.md index db18b48e..591f4de2 100644 --- a/irspec/docs/spatial/spatial.md +++ b/irspec/docs/spatial/spatial.md @@ -22,11 +22,16 @@ A stream corresponds to an abstract way to communicate between PEs or the host d For any scalar type `T`, `stream` indicates the corresponding element type sent over the stream. -Streams do not send a predetermined number of elements, but the sender and receiver must agree on the number of elements sent and received. +A stream type may carry a second template parameter, its **bound**: `stream`. +If the bound is given, then exactly `BOUND` elements are transferred over the stream, after which +the stream [closes itself](#closing-streams-with-close). Such a stream is called *bounded*. +For kernel arguments, the bound also determines the size of the host-side transfer +(it is what enables, e.g., memcpy mode in CSL). + +A stream without a bound is *unbounded*: it does not send a predetermined number of elements, +but the sender and receiver must agree on the number of elements sent and received. This can be done explicitly (when the size is known from the parameters) or implicitly (by sending a completion signal with/after the last element). - -Kernel arguments that are streams may have a second template parameter `stream`. If the second parameter is given, then -exactly `K` elements are transferred over the stream. This is useful for enabling, e.g., memcpy mode in CSL. +An unbounded stream must be [closed explicitly](#closing-streams-with-close) before its channel can be reused. ### Arrays @@ -580,6 +585,9 @@ completion completion_name = async { // Statements } +// Close a stream (asynchronous) +completion completion_name = stream_name.close(); + // Await a completion await completion_name; ``` @@ -814,6 +822,75 @@ completion completion_name = foreach type k, type x in [0:K], receive(stream_nam } ``` +### Closing Streams with `close` + +Inside a `compute` block, the `close` statement ends the lifetime of a stream on the PE that +executes it. + +```rust +completion completion_name = stream_name.close(); +// Or, as a shorthand: +await stream_name.close(); +``` + +The completion semantics are the same as those of `send` and `receive`: the completion is triggered +when the close has been *issued* on this PE, not when every other participant has observed it. + +The interval between the first use of a stream on a PE and its `close` is called an *epoch* of the +stream. Only within an epoch may a stream be used. Once a stream has been closed, its +[`channel`](#routing-declarations) is free and may be reused by another stream; see +[Semantics of Routing Declarations](../routing#undefined-behavior). + +!!! danger "Error: Use After Close" + Using a stream after it has been closed on the same PE (with `send`, `receive`, `foreach`, or + another `close`) is an error. *This raises a compile error whenever the use is ordered after + the close in [local order](../async#local-order). A use that is not ordered with respect to a + `close` on another PE is undefined behavior.* + +!!! danger "Error: Unclosed Stream" + Closing a stream is *collective*: every PE that sends on or receives from a stream must close + it. *Failing to close a stream on one of its participants raises a compile error where + detectable.* + +A stream that carries a [bound](#streams) closes itself once `BOUND` elements have been transferred; +an explicit `close` on a bounded stream is redundant but legal. An unbounded stream is only closed +by an explicit `close`. + +!!! note "Note: Verification of Bounds" + Where the number of elements transferred over a bounded stream can be determined statically, it + must match the stream's bound, otherwise a compile error is raised. Where it cannot be + determined statically, no diagnostic is emitted. + +At the end of a [phase](#phases), every stream that is in scope is implicitly closed. This is +equivalent to injecting a `close` for each such stream on each participating PE immediately *after* +the phase's implicit `await` statements. The order matters: the implicit `await`s may be waiting on +operations that are still using those streams, so the streams are only closed once every such +operation has completed. + +??? example "Example: Reusing a channel within a phase" + Two streams that share a channel may be used one after the other in the same phase, as long as + the first is closed on every PE it passes through before the second is used. + ```rust + dataflow i32 i, i32 j in [0:4, 0] { + stream westwards = relative_stream(-1, 0) { + hops = [(-1, 0)], + channel = 0 + }; + stream eastwards = relative_stream(1, 0) { + hops = [(1, 0)], + channel = 0 + }; + } + compute i32 i, i32 j in [0:4, 0] { + await send(a, westwards) + await westwards.close() + await foreach i32 k, f32 x in [0:K], receive(eastwards) { + a[k] = x + } + await eastwards.close() + } + ``` + ### Processing arrays asynchronously with `map` Inside a `compute` block, the `map` statement is used to apply a computation to each element of an array. @@ -913,9 +990,15 @@ Within each phase, there can be at most one `compute` block defined per PE. If multiple `compute` blocks are defined per PE per phase, the behavior is undefined. After each `compute` block, there is a set of implicit `await` statements that wait for all completions to be triggered before starting the next `compute` block. +Every stream that is in scope and has not yet been closed is then implicitly +[closed](#closing-streams-with-close), after those `await`s, since the outstanding completions may +belong to operations on those very streams. Note that this does *not* imply that all PEs have executed the `compute` block. No `compute` block may be defined in the outermost scope. +Since streams do not outlive the phase in which they are used, a +[stream edge](../routing#the-routing-graph) must be entirely contained within a phase. + Phases run in the order they are defined in the code from each PE's point of view. That is, a PE goes through its phases in-order. A PEs may participate in some phases and not in others. From 866127144920a0a15e26478d196187bc8ea7f2a2 Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Wed, 12 Aug 2026 22:13:04 -0700 Subject: [PATCH 02/68] Language support for `stream.close()` --- spada/syntax/csl/statements.py | 4 + spada/syntax/spatial_ir/irnodes.py | 43 ++++++++- spada/syntax/spatial_ir/language.lark | 7 +- spada/syntax/spatial_ir/lark_to_ir.py | 11 +++ tests/spatial_ir/test_spatial_ir_parser.py | 100 +++++++++++++++++++++ 5 files changed, 162 insertions(+), 3 deletions(-) diff --git a/spada/syntax/csl/statements.py b/spada/syntax/csl/statements.py index f2eccc0e..8cb93140 100644 --- a/spada/syntax/csl/statements.py +++ b/spada/syntax/csl/statements.py @@ -49,6 +49,10 @@ def generate_csl_statement(statement: spir.Statement, elif isinstance(statement, (spir.AwaitCompletionStatement, spir.AwaitAllStatement)): # Skip (taken care of when tasks are defined) return "" + elif isinstance(statement, spir.CloseStatement): + # TODO(switching): Lower to a switch advance on the stream's channel. + raise NotImplementedError('Closing a stream is not yet supported by the CSL backend.\n' + f' In line {statement.lineinfo}') if op is None: return f'// TODO: Convert {statement} to CSL' diff --git a/spada/syntax/spatial_ir/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index 28d029ef..bc77eb64 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -112,18 +112,33 @@ def __hash__(self) -> int: class StreamType(SpatialNode, IRType): """ A stream type that sends elements of type T. + + The optional second template parameter is the stream's *bound*: exactly ``buffer_size`` elements + are transferred over the stream, after which the stream closes itself (see + ``CloseStatement``). A stream without a bound is unbounded and must be closed explicitly. + For kernel arguments, the bound also gives the size of the host-side transfer (which is what + enables memcpy mode in CSL), hence the field name. """ element_type: ScalarType buffer_size: Optional['Expression'] = None def validate(self) -> None: assert isinstance(self.element_type, ScalarType) + assert self.buffer_size is None or isinstance(self.buffer_size, Expression) def as_ir(self, indent: int = 0) -> str: if self.buffer_size is not None: return f'stream<{self.element_type.as_ir()}, {self.buffer_size.as_ir()}>' return f'stream<{self.element_type.as_ir()}>' + @property + def bound(self) -> Optional['Expression']: + """ + The number of elements transferred over this stream before it closes itself, or ``None`` if + the stream is unbounded. An alias of ``buffer_size`` that reads better in lifetime analyses. + """ + return self.buffer_size + @property def shape(self) -> list[Union[int, 'Expression']]: # This is here to consolidate data type analysis @@ -740,7 +755,7 @@ def validate(self) -> None: def as_ir(self, indent: int = 0) -> str: indent_str = ' ' * indent - return f'{indent_str}stream<{self.dtype.element_type.as_ir()}> {self.stream_name.as_ir()} = {self.stream.as_ir()}' + return f'{indent_str}{self.dtype.as_ir()} {self.stream_name.as_ir()} = {self.stream.as_ir()}' ### @@ -896,6 +911,32 @@ def as_ir(self, indent: int = 0) -> str: return f'receive({self.stream_name.as_ir()})' +# Close Statement +@dataclass +class CloseStatement(Statement): + """ + Close statement that ends the lifetime of a stream on the PE that executes it. + + Closing a stream releases its channel, which may then be reused by another stream (see the + routing specification). It is collective: every PE that sends on or receives from a stream must + close it. Bounded streams (``stream``) close themselves once ``BOUND`` elements have + been transferred, and every stream in scope is implicitly closed at the end of its phase. + """ + stream_name: Union[Identifier, ArraySlice] + completion_name: Optional[Completion] = None + + def validate(self) -> None: + assert isinstance(self.stream_name, (Identifier, ArraySlice)) + if self.completion_name: + assert isinstance(self.completion_name, Completion) + + def as_ir(self, indent: int = 0) -> str: + indent_str = ' ' * indent + if self.completion_name: + return f'{indent_str}{self.completion_name.as_ir()} = {self.stream_name.as_ir()}.close()' + return f'{indent_str}await {self.stream_name.as_ir()}.close()' + + # Foreach Loop (asynchronous) @dataclass class ForeachStatement(Statement): diff --git a/spada/syntax/spatial_ir/language.lark b/spada/syntax/spatial_ir/language.lark index 4b4a981b..5d41c17b 100644 --- a/spada/syntax/spatial_ir/language.lark +++ b/spada/syntax/spatial_ir/language.lark @@ -142,7 +142,10 @@ completion : "completion" identifier "=" ?prefix : completion | "await" // Free function call -function_call: prefix bare_id "(" [call_arguments] ")" +function_call: prefix bare_id "(" [call_arguments] ")" + +// Method call on a stream (e.g., `await s.close()`). +method_call : prefix (identifier | subscript) "." bare_id "(" [call_arguments] ")" await_completion : "await" identifier @@ -161,7 +164,7 @@ async_stmt : prefix "async" compute_body awaitall_stmt : "awaitall" // Top-level statements that are direct children of `compute` blocks -?base_stmt : (function_call | await_completion | map_stmt | foreach_stmt | async_stmt | for_stmt | awaitall_stmt) +?base_stmt : (function_call | method_call | await_completion | map_stmt | foreach_stmt | async_stmt | for_stmt | awaitall_stmt) assignment : (subscript | identifier) "=" value_expr diff --git a/spada/syntax/spatial_ir/lark_to_ir.py b/spada/syntax/spatial_ir/lark_to_ir.py index da877497..f1d62a51 100644 --- a/spada/syntax/spatial_ir/lark_to_ir.py +++ b/spada/syntax/spatial_ir/lark_to_ir.py @@ -159,6 +159,17 @@ def function_call(self, args, meta=None): return irnodes.ReceiveStatement(*arguments, completion_name=completion) raise SyntaxError(f'Unrecognized free function call to "{func}"') + # Method call on a stream object + def method_call(self, args, meta=None): + completion, stream, func, arguments = args + + if func == 'close': + if arguments: + raise SyntaxError(f'"close" takes no arguments, but {len(arguments)} were given in ' + f'"{stream.as_ir()}.close(...)"') + return irnodes.CloseStatement(stream, completion_name=completion) + raise SyntaxError(f'Unrecognized method call to "{func}" on stream "{stream.as_ir()}"') + subscript = irnodes.ArraySlice.from_lark subscript_expr = irnodes.ArraySlice.from_lark diff --git a/tests/spatial_ir/test_spatial_ir_parser.py b/tests/spatial_ir/test_spatial_ir_parser.py index 9376bb30..1d2ab45c 100644 --- a/tests/spatial_ir/test_spatial_ir_parser.py +++ b/tests/spatial_ir/test_spatial_ir_parser.py @@ -1,5 +1,6 @@ from spada.syntax.spatial_ir import irnodes as spast, parser import os +import pytest def test_spatial_roundtrip_laplacian(): @@ -153,6 +154,100 @@ def test_extern_stream(): assert out_decl.stream.routing.resolved_channel == 3 +def test_spatial_roundtrip_bounded_streams(): + """ + Tests that the bound of a ``stream`` survives a roundtrip on a dataflow declaration. + """ + file = os.path.join(os.path.dirname(__file__), 'samples', 'neighbor_exchange.sptl') + _rountrip_test(file) + + program = parser.parse_file(file) + df_block = next(stmt for stmt in program.body if isinstance(stmt, spast.DataflowBlock)) + assert 'stream eastwards' in df_block.statements[0].as_ir() + assert df_block.statements[0].dtype.bound is not None + + +def test_unbounded_stream_has_no_bound(): + """ + Tests that a stream declared without a second template parameter has no bound. + """ + code = """ + kernel @test(f32 coeff) { + dataflow u16 i, u16 j in [0:N, 0:N] { + stream eastwards = relative_stream(1, 0); + } + }""" + kernel = parser.parse_string(code) + df_block = next(stmt for stmt in kernel.body if isinstance(stmt, spast.DataflowBlock)) + assert df_block.statements[0].dtype.bound is None + assert df_block.statements[0].as_ir().strip() == 'stream eastwards = relative_stream(1, 0)' + + +def test_close_statement(): + """ + Tests parsing the ``close`` statement, both awaited and with a completion. + """ + code = """ + kernel @test(stream[N] readonly inp) { + place u16 i, u16 j in [0:N, 0:N] { + f32 a + } + dataflow u16 i, u16 j in [0:N, 0:N] { + stream eastwards = relative_stream(1, 0); + } + compute u16 i, u16 j in [0:N, 0:N] { + await send(a, eastwards) + await eastwards.close() + completion c = inp[i].close() + await c + } + }""" + kernel = parser.parse_string(code) + compute = next(stmt for stmt in kernel.body if isinstance(stmt, spast.ComputeBlock)) + _, awaited, with_completion, _ = compute.statements + + assert isinstance(awaited, spast.CloseStatement) + assert awaited.completion_name is None + assert awaited.stream_name.as_ir() == 'eastwards' + assert awaited.as_ir().strip() == 'await eastwards.close()' + + assert isinstance(with_completion, spast.CloseStatement) + assert with_completion.completion_name.name.as_ir() == 'c' + assert isinstance(with_completion.stream_name, spast.ArraySlice) + assert with_completion.as_ir().strip() == 'completion c = inp[i].close()' + + ir_1 = kernel.as_ir() + assert parser.parse_string(ir_1).as_ir() == ir_1 + + +def _method_call_kernel(call: str) -> str: + return f""" + kernel @test(f32 coeff) {{ + dataflow u16 i, u16 j in [0:N, 0:N] {{ + stream eastwards = relative_stream(1, 0); + }} + compute u16 i, u16 j in [0:N, 0:N] {{ + await {call} + }} + }}""" + + +def test_unknown_stream_method(): + """ + Tests that an unrecognized method on a stream raises a syntax error. + """ + with pytest.raises(Exception, match='open'): + parser.parse_string(_method_call_kernel('eastwards.open()')) + + +def test_close_rejects_arguments(): + """ + Tests that arguments to ``close`` are parsed, then rejected with a readable error. + """ + with pytest.raises(Exception, match='takes no arguments'): + parser.parse_string(_method_call_kernel('eastwards.close(3)')) + + if __name__ == '__main__': test_spatial_roundtrip_laplacian() test_spatial_visitor() @@ -163,3 +258,8 @@ def test_extern_stream(): test_spatial_roundtrip_backward() test_extern_field() test_extern_stream() + test_spatial_roundtrip_bounded_streams() + test_unbounded_stream_has_no_bound() + test_close_statement() + test_unknown_stream_method() + test_close_rejects_arguments() From 97ac1292ce75f0965c428acf4bc486f48b743691 Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Thu, 13 Aug 2026 00:11:35 -0700 Subject: [PATCH 03/68] Analysis and optimization passes --- spada/cli/compiler.py | 5 +- spada/lowering/spatial_ir_to_csl.py | 14 +- spada/syntax/spatial_ir/canonicalization.py | 21 +- spada/syntax/spatial_ir/stream_lifetime.py | 644 ++++++++++++++++++ tests/spatial_ir/samples/two_phase_split.sptl | 5 +- tests/spatial_ir/test_stream_lifetime.py | 462 +++++++++++++ 6 files changed, 1147 insertions(+), 4 deletions(-) create mode 100644 spada/syntax/spatial_ir/stream_lifetime.py create mode 100644 tests/spatial_ir/test_stream_lifetime.py diff --git a/spada/cli/compiler.py b/spada/cli/compiler.py index 7bb19a37..360307c7 100644 --- a/spada/cli/compiler.py +++ b/spada/cli/compiler.py @@ -22,10 +22,12 @@ @click.option('--disable-task-fusion', is_flag=True, help='Disable task fusion optimization') @click.option('--disable-task-recycling', is_flag=True, help='Disable task ID recycling') @click.option('--disable-copy-elision', is_flag=True, help='Disable copy elimination optimization pass') +@click.option('--disable-close-elision', is_flag=True, help='Disable elision of unnecessary stream closes') def compile_spatial_ir(input_file: str, output_folder: str, param: list[str], offset_x: int, offset_y: int, generate_only: bool, disable_benchmarking: bool, disable_asynchronous: bool, disable_dsd: bool, disable_map: bool, - disable_task_fusion: bool, disable_task_recycling: bool, disable_copy_elision: bool): + disable_task_fusion: bool, disable_task_recycling: bool, disable_copy_elision: bool, + disable_close_elision: bool): # Parse parameters into dictionary kernel_parameters = {} for p in param: @@ -95,6 +97,7 @@ def compile_spatial_ir(input_file: str, output_folder: str, param: list[str], of task_fusion=not disable_task_fusion, copy_elision=not disable_copy_elision, task_id_recycling=not disable_task_recycling, + close_elision=not disable_close_elision, ) # Create output folder if it doesn't exist diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 3c36712c..35b59f50 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -11,6 +11,7 @@ from spada.syntax.spatial_ir import irnodes as spir, canonicalization, analysis, passes from spada.syntax.spatial_ir import copy_elimination from spada.syntax.spatial_ir import canonical_subgrids +from spada.syntax.spatial_ir import stream_lifetime from spada.syntax.spatial_ir.canonicalization import PEBlock, Rectangle from spada.syntax.csl import constants as csl, preprocessing, tasks as tdag, statements as cslstmt, dsd_ops from spada.syntax.csl import benchmarking as cslbench @@ -37,6 +38,7 @@ def canonicalize_kernel(kernel: spir.Kernel) -> spir.Kernel: """ kernel = canonicalization.inline_metaprogramming(kernel) kernel = canonicalization.canonicalize_phases(kernel) + kernel = stream_lifetime.insert_implicit_closes(kernel) kernel = canonicalization.reduce_streams(kernel) kernel = canonical_subgrids.canonicalize_subgrids(kernel) kernel = canonicalization.resolve_auto_hops(kernel) @@ -52,7 +54,8 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, task_fusion: bool = True, copy_elision: bool = True, prune_memory: bool = True, - task_id_recycling: bool = True) -> list[CodeFile]: + task_id_recycling: bool = True, + close_elision: bool = True) -> list[CodeFile]: """ Lowers a routed Spatial IR kernel into Cerebras CSL code. @@ -66,6 +69,7 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, :param copy_elision: If True, enables copy elision optimization pass. :param prune_memory: If True, enables unused field pruning optimization pass. :param task_id_recycling: If True, enables task ID recycling pass. + :param close_elision: If True, removes stream closes that no router has to act on. :return: List of code-file objects that can be written to files. See ``write_code_to_files``. """ # PRECONDITION: Rectangles of dataflow/compute/place do not intersect (comes from Spatial IR) @@ -129,7 +133,15 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, routing_instructions: list[str] = [] color_maps = [] + # Verify stream lifetimes. This runs after channels have been resolved by + # ``_collect_colors_globally``, and before ``elide_redundant_closes`` so that no diagnostic can + # be hidden by the elision. channel_to_color = _collect_colors_globally(kernel, rectangles, use_memcpy_mode) + stream_lifetime.verify_stream_bounds(rectangles) + stream_lifetime.check_use_after_close(rectangles) + stream_lifetime.check_channel_conflicts(rectangles) + if close_elision: + stream_lifetime.elide_redundant_closes(rectangles) for rect in rectangles: # Create a unique CSL code file based on rectangle offset diff --git a/spada/syntax/spatial_ir/canonicalization.py b/spada/syntax/spatial_ir/canonicalization.py index 998a88b7..4997a827 100644 --- a/spada/syntax/spatial_ir/canonicalization.py +++ b/spada/syntax/spatial_ir/canonicalization.py @@ -211,6 +211,24 @@ def _rewrite_stream_declarations( return replacements, appended_statements +def _ends_with_phase_barrier(statements: list[spir.Statement]) -> bool: + """ + Returns whether a compute block already ends with a phase barrier, so that appending another + ``awaitall`` would be redundant. + + ``insert_implicit_closes`` ends every phase with an ``awaitall`` followed by the implicit + ``close`` statements, which are awaited and therefore leave nothing outstanding. Scanning back + over those closes finds the barrier they belong to. + """ + for statement in reversed(statements): + if isinstance(statement, spir.AwaitAllStatement): + return True + if isinstance(statement, spir.CloseStatement) and statement.completion_name is None: + continue + return False + return False + + def inline_phases(kernel: spir.Kernel) -> spir.Kernel: """ Inlines phases into their constituent computation and dataflow blocks by adding waits and appending all streams, @@ -295,7 +313,8 @@ def inline_phases(kernel: spir.Kernel) -> spir.Kernel: for compute in block.compute: rect = compute.get_grid_rect() if rect in rect_compute: - rect_compute[rect].statements.append(spir.AwaitAllStatement()) + if not _ends_with_phase_barrier(rect_compute[rect].statements): + rect_compute[rect].statements.append(spir.AwaitAllStatement()) replacements = dict(phase_replacements[rect]) replacements.update({ oldv.identifier: newv.identifier diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py new file mode 100644 index 00000000..6d54cdd3 --- /dev/null +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -0,0 +1,644 @@ +""" +Stream lifetime analysis, verification, and optimization passes for Spatial IR. + +A stream is *open* on a PE from its first use (``send``/``receive``) until it is *closed*, either +explicitly with ``await s.close()`` or implicitly at the end of the phase in which the PE last uses +it. That interval is the stream's *epoch*. Closing a stream releases its ``channel``, which may then +be taken over by another stream; see ``irspec/docs/spatial/routing.md``. + +The passes in this module are deliberately separate so that each can be tested on its own: + +* :func:`insert_implicit_closes` -- materializes the implicit end-of-phase closes. +* :func:`collect_stream_uses` -- the shared analysis every other pass builds on. +* :func:`verify_stream_bounds` -- checks ``stream`` against the transferred element count. +* :func:`check_use_after_close` -- rejects any use of a stream past its close. +* :func:`check_channel_conflicts` -- rejects concurrent use of a channel. +* :func:`elide_redundant_closes` -- drops closes whose channel is never reused. +""" +from collections import defaultdict +import copy +from dataclasses import dataclass, field +from typing import Optional, Union + +from spada.syntax.spatial_ir import irnodes as spir +from spada.syntax.spatial_ir.grid_geometry import Rectangle + +StreamExpression = Union[spir.Identifier, spir.ArraySlice] + + +@dataclass +class StreamUse: + """ + The lifetime of one stream within one compute block (i.e., on one PE equivalence class). + """ + #: The stream expression as it was first written (e.g. ``eastwards`` or ``inp[i]``). + expression: StreamExpression + #: Index of the first statement that sends on or receives from the stream. + first_use: int + #: Index of the last statement that sends on or receives from the stream. + last_use: int + #: Index of the ``close`` statement, or ``None`` if the stream is never closed here. + close: Optional[int] = None + #: Indices of every statement that uses the stream, in order. + uses: list[int] = field(default_factory=list) + #: Whether the stream is sent on / received from in this compute block. + sent: bool = False + received: bool = False + + @property + def name(self) -> spir.Identifier: + return _underlying_stream(self.expression) + + +def _underlying_stream(expression: StreamExpression) -> spir.Identifier: + """ + Returns the stream identifier behind a stream expression, unwrapping array slices. + """ + if isinstance(expression, spir.ArraySlice): + return expression.array + return expression + + +def _stream_references(statement: spir.Statement) -> list[tuple[str, StreamExpression]]: + """ + Returns every stream reference in a statement (including nested ones) as ``(kind, expression)`` + pairs, where ``kind`` is one of ``'send'``, ``'receive'``, or ``'close'``. + + References nested inside the same top-level statement are not ordered relative to each other; + all of them are attributed to the enclosing top-level statement. + """ + references: list[tuple[str, StreamExpression]] = [] + for node in statement.walk(): + if isinstance(node, spir.SendStatement): + references.append(('send', node.stream_name)) + elif isinstance(node, (spir.ReceiveStatement, spir.ReceiveGenerator)): + references.append(('receive', node.stream_name)) + elif isinstance(node, spir.CloseStatement): + references.append(('close', node.stream_name)) + return references + + +def collect_stream_uses(compute: spir.ComputeBlock) -> dict[spir.Identifier, StreamUse]: + """ + Collects the lifetime of every stream used in a compute block, keyed by stream identifier. + + :param compute: The compute block to analyze. + :return: A dictionary mapping each stream identifier to its :class:`StreamUse`. + """ + result: dict[spir.Identifier, StreamUse] = {} + for index, statement in enumerate(compute.statements): + for kind, expression in _stream_references(statement): + name = _underlying_stream(expression) + use = result.get(name) + if use is None: + use = StreamUse(expression=expression, first_use=index, last_use=index) + result[name] = use + + if kind == 'close': + # Keep the first close: a second one is reported by ``check_use_after_close``. + if use.close is None: + use.close = index + continue + + use.uses.append(index) + use.last_use = index + if kind == 'send': + use.sent = True + else: + use.received = True + + # A stream that is only closed has no data use; ``first_use`` then points at the close itself. + return result + + +def _declared_stream_names(kernel: spir.Kernel) -> set[spir.Identifier]: + """ + Returns the names of every stream declared in a ``dataflow`` block of the kernel. + + Streams that are kernel arguments are not included: they are lowered to extern streams or extern + fields later in the pipeline and carry no on-chip channel of their own at this point. + """ + return { + statement.stream_name + for node in kernel.walk() + if isinstance(node, spir.DataflowBlock) + for statement in node.statements + } + + +def insert_implicit_closes(kernel: spir.Kernel) -> spir.Kernel: + """ + Materializes the implicit close of every stream at the end of its scope. + + For each compute block, an ``awaitall`` followed by ``await .close()`` is appended for + every stream the block uses and does not already close, in the phase in which that block last + uses it. The closes are emitted *after* the barrier because the phase's implicit awaits may be + waiting on operations that are still using those very streams. + + Must run after ``canonicalize_phases`` (so that ``kernel.body`` contains only phases and place + blocks) and before ``inline_phases``. + + :param kernel: The kernel to transform, modified in place. + :return: The transformed kernel. + """ + declared = _declared_stream_names(kernel) + if not declared: + return kernel + + phases = [block for block in kernel.body if isinstance(block, spir.Phase)] + + # Find, for each (rectangle, stream), the last phase in which the rectangle uses the stream. + # A stream declared at kernel level stays in scope across phases, so it may only be closed + # after its final use. + last_phase: dict[tuple[tuple[int, int, int, int], spir.Identifier], int] = {} + uses_per_block: list[list[tuple[spir.ComputeBlock, dict[spir.Identifier, StreamUse]]]] = [] + for phase_index, phase in enumerate(phases): + blocks = [] + for compute in phase.compute: + uses = collect_stream_uses(compute) + blocks.append((compute, uses)) + rect = compute.get_grid_rect() + for name, use in uses.items(): + if name in declared and use.uses: + last_phase[(rect, name)] = phase_index + uses_per_block.append(blocks) + + for phase_index, blocks in enumerate(uses_per_block): + for compute, uses in blocks: + rect = compute.get_grid_rect() + to_close = [ + use for name, use in uses.items() + if name in declared and use.uses and use.close is None and last_phase[(rect, name)] == phase_index + ] + if not to_close: + continue + + compute.statements.append(spir.AwaitAllStatement()) + for use in to_close: + close = spir.CloseStatement(copy.deepcopy(use.expression)) + close.lineinfo = getattr(use.expression, 'lineinfo', None) + compute.statements.append(close) + + return kernel + + +### +# Verification passes +### + + +def _location(node: spir.SpatialNode) -> str: + lineinfo = getattr(node, 'lineinfo', None) + return f' at {lineinfo}' if lineinfo else '' + + +def check_use_after_close(rectangles: list[Rectangle]) -> None: + """ + Raises a ``SyntaxError`` if a stream is used after it has been closed on the same PE. + + A closed stream is dead: reusing its channel requires declaring another stream. Every + ``send``, ``receive``, ``foreach`` generator, and second ``close`` past the close is rejected. + + :param rectangles: The consolidated PE rectangles of the kernel. + """ + for rect in rectangles: + compute = rect.metadata.compute + closed: dict[spir.Identifier, spir.CloseStatement] = {} + for statement in compute.statements: + for kind, expression in _stream_references(statement): + name = _underlying_stream(expression) + if name in closed: + what = 'closed again' if kind == 'close' else f'used in a `{kind}`' + raise SyntaxError( + f"Stream '{name.as_ir()}' is {what} after it was closed.\n" + f" closed{_location(closed[name])}\n" + f" used{_location(statement)}\n" + " note: a closed stream cannot be reopened; declare another stream on the " + "same channel to reuse it") + if kind == 'close': + closed[name] = statement + + +def _stream_declarations(rect: Rectangle) -> dict[spir.Identifier, spir.StreamDeclaration]: + return {declaration.stream_name: declaration for declaration in rect.metadata.dataflow.statements} + + +def stream_group_key(declaration: spir.StreamDeclaration) -> str: + """ + Returns the identity of a stream for routing purposes: its channel together with its routing + pattern. + + Stream *names* cannot be used for this. ``inline_phases`` freshens colliding names per + rectangle, so one logical stream that is declared in several dataflow blocks shows up as + ``bcast``, ``bcast#1``, ... Two declarations that route identically on the same channel occupy + every PE in the same way and therefore never conflict, which is exactly what this key captures. + """ + return ' '.join(declaration.stream.as_ir().split()) + + +def _constant_range_length(rng: spir.RangeExpression) -> Optional[int]: + """ + Returns the number of iterations of a range expression, or ``None`` if it is not a constant. + """ + try: + start = rng.start.eval() + stop = rng.stop.eval() if rng.stop is not None else None + step = rng.step.eval() if rng.step is not None else 1 + except Exception: # pragma: no cover - defensive: any non-constant expression + return None + if stop is None: + return 1 + if not all(isinstance(value, int) for value in (start, stop, step)) or step == 0: + return None + return max(0, -(-(stop - start) // step)) + + +def _transferred_elements(statement: spir.Statement, stream: spir.Identifier, + identifier_sizes: dict[spir.Identifier, list[int]]) -> Optional[int]: + """ + Returns how many elements a top-level statement transfers over a stream, or ``None`` if that + cannot be determined statically. + """ + if isinstance(statement, (spir.SendStatement, spir.ReceiveStatement)): + if _underlying_stream(statement.stream_name) != stream: + return 0 + try: + dimensions = statement.get_size(identifier_sizes) + except Exception: + return None + count = 1 + for dimension in dimensions: + count *= dimension + return count + + if isinstance(statement, spir.ForeachStatement): + if _underlying_stream(statement.receive_stream.stream_name) != stream: + return None if _uses_stream(statement, stream) else 0 + if not statement.parameter_range: + return None # Receives until the sender is done + count = 1 + for rng in statement.parameter_range: + length = _constant_range_length(rng) + if length is None: + return None + count *= length + return count + + if isinstance(statement, (spir.ForStatement, spir.MapStatement)): + if not _uses_stream(statement, stream): + return 0 + trips = 1 + for rng in statement.range_expression: + length = _constant_range_length(rng) + if length is None: + return None + trips *= length + inner = 0 + for inner_statement in statement.body: + count = _transferred_elements(inner_statement, stream, identifier_sizes) + if count is None: + return None + inner += count + return trips * inner + + if isinstance(statement, spir.AsyncBlock): + if not _uses_stream(statement, stream): + return 0 + total = 0 + for inner_statement in statement.body: + count = _transferred_elements(inner_statement, stream, identifier_sizes) + if count is None: + return None + total += count + return total + + return None if _uses_stream(statement, stream) else 0 + + +def _uses_stream(statement: spir.Statement, stream: spir.Identifier) -> bool: + return any(_underlying_stream(expression) == stream for kind, expression in _stream_references(statement) + if kind != 'close') + + +def verify_stream_bounds(rectangles: list[Rectangle]) -> None: + """ + Raises a ``SyntaxError`` if the number of elements transferred over a bounded stream can be + determined statically and does not match the stream's bound. + + Streams whose element count cannot be analyzed are silently accepted. + + :param rectangles: The consolidated PE rectangles of the kernel. + """ + for rect in rectangles: + declarations = _stream_declarations(rect) + identifier_sizes = _identifier_sizes(rect.metadata.place) + for name, use in collect_stream_uses(rect.metadata.compute).items(): + declaration = declarations.get(name) + if declaration is None or declaration.dtype.bound is None or not use.uses: + continue + try: + bound = declaration.dtype.bound.eval() + except Exception: + continue + if not isinstance(bound, int): + continue + + total = 0 + for index in use.uses: + count = _transferred_elements(rect.metadata.compute.statements[index], name, identifier_sizes) + if count is None: + total = None + break + total += count + if total is None or total == bound: + continue + + raise SyntaxError(f"Stream '{name.as_ir()}' is declared with bound {bound}, but {total} " + f"element(s) are transferred over it{_location(declaration)}.\n" + " note: the bound of a stream is the exact number of elements it carries " + "before it closes itself") + + +def _identifier_sizes(place: spir.PlaceBlock) -> dict[spir.Identifier, list[int]]: + """ + Like ``analysis.get_identifier_sizes``, but tolerant of shapes that are not yet constant. + """ + result: dict[spir.Identifier, list[int]] = {} + for declaration in place.statements: + if isinstance(declaration.dtype, spir.ScalarType): + result[declaration.field_name] = [] + continue + shape = [] + for dimension in declaration.dtype.shape: + if isinstance(dimension, int): + shape.append(dimension) + else: + try: + shape.append(dimension.eval()) + except Exception: + break + else: + result[declaration.field_name] = shape + return result + + +### +# Channel occupancy +### + + +def _stream_path_offsets(declaration: spir.StreamDeclaration) -> Optional[list[tuple[int, int]]]: + """ + Returns the PE offsets, relative to the sending PE, that a stream occupies -- that is, every PE + whose router carries the stream. Returns ``None`` for streams that have no on-chip routing. + """ + stream = declaration.stream + if isinstance(stream, spir.ExternStreamDeclaration): + return None + + if isinstance(stream, spir.MulticastRangeStreamDeclaration): + multicast_range = stream.multicast_range + try: + start = int(multicast_range.start.eval()) + stop = int(multicast_range.stop.eval()) + except Exception: + return None + step = 1 if stop > start else -1 + covered = list(range(0, stop, step)) + if stream.multicast_axis == 'x': + return [(offset, 0) for offset in covered] + return [(0, offset) for offset in covered] + + if stream.routing is None or stream.routing.hops == 'auto': + return None + + offsets = [(0, 0)] + x, y = 0, 0 + for hop in stream.routing.hops: + x += hop.offset[0] + y += hop.offset[1] + offsets.append((x, y)) + return offsets + + +def _rect_points(rect: Rectangle) -> list[tuple[int, int]]: + return [(x, y) + for x in range(rect.x_range[0], rect.x_range[1], rect.x_range[2]) + for y in range(rect.y_range[0], rect.y_range[1], rect.y_range[2])] + + +def channel_occupancy(rectangles: list[Rectangle]) -> dict[tuple[int, tuple[int, int]], set[tuple[int, str]]]: + """ + Builds the parametric routing graph's occupancy map: which streams occupy which PE on which + channel. + + Only channels that carry more than one stream declaration are enumerated, since a channel with a + single stream can never conflict with itself. + + :param rectangles: The consolidated PE rectangles of the kernel. + :return: A dictionary mapping ``(channel, (x, y))`` to the set of ``(rectangle index, stream + group key)`` pairs that occupy it. See :func:`stream_group_key`. + """ + streams_per_channel = groups_per_channel(rectangles) + candidates: list[tuple[int, spir.StreamDeclaration, int]] = [] + for rect_index, rect in enumerate(rectangles): + declarations = _stream_declarations(rect) + for name, use in collect_stream_uses(rect.metadata.compute).items(): + declaration = declarations.get(name) + if declaration is None or declaration.stream.routing is None or not use.sent: + continue # Paths are enumerated from the sending PE, as in the routing emission + channel = declaration.stream.routing.resolved_channel + if channel == 'auto': + continue + candidates.append((channel, declaration, rect_index)) + + bounds = _kernel_bounds(rectangles) + occupancy: dict[tuple[int, tuple[int, int]], set[tuple[int, str]]] = defaultdict(set) + for channel, declaration, rect_index in candidates: + if len(streams_per_channel[channel]) < 2: + continue # A channel with a single stream cannot conflict with itself + offsets = _stream_path_offsets(declaration) + if offsets is None: + continue + rect = rectangles[rect_index] + key = (rect_index, stream_group_key(declaration)) + for x, y in _rect_points(rect): + for dx, dy in offsets: + pe = (x + dx, y + dy) + if _in_bounds(pe, bounds): + occupancy[(channel, pe)].add(key) + + return occupancy + + +def _kernel_bounds(rectangles: list[Rectangle]) -> tuple[int, int, int, int]: + """ + Returns the bounding box of every PE of the kernel, as ``(min_x, max_x, min_y, max_y)``, with + the maxima exclusive. Paths that leave it belong to senders that have no receiver. + """ + if not rectangles: + return (0, 0, 0, 0) + return (min(rect.x_range[0] for rect in rectangles), max(rect.x_range[1] for rect in rectangles), + min(rect.y_range[0] for rect in rectangles), max(rect.y_range[1] for rect in rectangles)) + + +def _in_bounds(pe: tuple[int, int], bounds: tuple[int, int, int, int]) -> bool: + return bounds[0] <= pe[0] < bounds[1] and bounds[2] <= pe[1] < bounds[3] + + +def groups_per_channel(rectangles: list[Rectangle]) -> dict[int, set[str]]: + """ + Returns, for each resolved channel, the set of distinct stream groups that use it. + """ + result: dict[int, set[str]] = defaultdict(set) + for rect in rectangles: + for declaration in rect.metadata.dataflow.statements: + if declaration.stream.routing is None: + continue + channel = declaration.stream.routing.resolved_channel + if channel != 'auto': + result[channel].add(stream_group_key(declaration)) + return result + + +def check_channel_conflicts(rectangles: list[Rectangle]) -> None: + """ + Raises a ``SyntaxError`` if two streams may use the same channel concurrently, that is, if two + streams share a channel and a PE without being ordered by empties-before. The ordering is + established by closing the earlier stream everywhere it is used, and, where both streams are + used on the same PE, by closing it before the later stream's first use. + + The check is best-effort in the sense of the specification: it establishes the ordering + statically where it can, and does not reject what it cannot decide. + + Note that a *single* stream that is both sent and received on the same PE (a systolic forwarding + pattern, as in ``laplacian_routed.sptl``) is not a channel conflict: the two directions are + ordered by the local order of the receive and the send. They do require two router + configurations on one color, which is handled by switch planning during CSL lowering. + + :param rectangles: The consolidated PE rectangles of the kernel. + """ + uses_per_rect = [collect_stream_uses(rect.metadata.compute) for rect in rectangles] + # Per rectangle, the local name each stream group goes by + groups_per_rect: list[dict[str, spir.Identifier]] = [] + for rect, uses in zip(rectangles, uses_per_rect): + declarations = _stream_declarations(rect) + groups_per_rect.append({ + stream_group_key(declarations[name]): name + for name in uses if name in declarations + }) + + reported: set[tuple[str, str]] = set() + for (channel, pe), occupants in sorted(channel_occupancy(rectangles).items()): + groups = sorted({group for _, group in occupants}) + if len(groups) < 2: + continue + + for first, second in _ordered_pairs(groups): + if (first, second) in reported: + continue + if _empties_before(first, second, uses_per_rect, groups_per_rect) or \ + _empties_before(second, first, uses_per_rect, groups_per_rect): + continue + reported.add((first, second)) + first_decl = _find_declaration(rectangles, first) + second_decl = _find_declaration(rectangles, second) + first_name = first_decl.stream_name.as_ir() if first_decl else first + second_name = second_decl.stream_name.as_ir() if second_decl else second + raise SyntaxError( + f"Streams '{first_name}' and '{second_name}' both use channel {channel} and share " + f"PE ({pe[0]}, {pe[1]}), but are not ordered by empties-before.\n" + f" '{first_name}' declared{_location(first_decl)}\n" + f" '{second_name}' declared{_location(second_decl)}\n" + f" note: close one of them on every PE that uses it before the other is used, " + f"e.g. `await {first_name}.close()`") + + +def _ordered_pairs(names: list[str]): + for i, first in enumerate(names): + for second in names[i + 1:]: + yield first, second + + +def _find_declaration(rectangles: list[Rectangle], group: str) -> Optional[spir.StreamDeclaration]: + for rect in rectangles: + for declaration in rect.metadata.dataflow.statements: + if stream_group_key(declaration) == group: + return declaration + return None + + +def _empties_before(first: str, second: str, uses_per_rect: list[dict[spir.Identifier, StreamUse]], + groups_per_rect: list[dict[str, spir.Identifier]]) -> bool: + """ + Returns whether the stream group ``first`` provably empties before the stream group ``second``. + + This requires ``first`` to be closed on every PE that uses it, and, wherever both streams are + used on the same PE, for that close to precede the first use of ``second`` in local order. + """ + used_anywhere = False + for uses, groups in zip(uses_per_rect, groups_per_rect): + first_name = groups.get(first) + if first_name is None: + continue + first_use = uses[first_name] + if not first_use.uses: + continue + used_anywhere = True + if first_use.close is None: + return False + + second_name = groups.get(second) + if second_name is None: + continue + second_use = uses[second_name] + if second_use.uses and first_use.close > second_use.first_use: + return False + + return used_anywhere + + +### +# Optimization passes +### + + +def elide_redundant_closes(rectangles: list[Rectangle]) -> int: + """ + Removes every close whose channel is never taken over by another stream, since no router has to + advance in that case. + + :param rectangles: The consolidated PE rectangles of the kernel, modified in place. + :return: The number of closes that were removed. + """ + streams_per_channel = groups_per_channel(rectangles) + + removed = 0 + for rect in rectangles: + declarations = _stream_declarations(rect) + keep = [] + for statement in rect.metadata.compute.statements: + if isinstance(statement, spir.CloseStatement): + name = _underlying_stream(statement.stream_name) + declaration = declarations.get(name) + if _is_redundant_close(declaration, streams_per_channel): + removed += 1 + continue + keep.append(statement) + rect.metadata.compute.statements = keep + + return removed + + +def _is_redundant_close(declaration: Optional[spir.StreamDeclaration], + streams_per_channel: dict[int, set[str]]) -> bool: + if declaration is None: + return True # Not a routed stream (e.g. a kernel argument): nothing to release + if isinstance(declaration.stream, spir.ExternStreamDeclaration): + return True # No on-chip routing + if declaration.stream.routing is None: + return True + channel = declaration.stream.routing.resolved_channel + if channel == 'auto': + return True + return len(streams_per_channel[channel]) < 2 diff --git a/tests/spatial_ir/samples/two_phase_split.sptl b/tests/spatial_ir/samples/two_phase_split.sptl index ed3e6282..c950c154 100644 --- a/tests/spatial_ir/samples/two_phase_split.sptl +++ b/tests/spatial_ir/samples/two_phase_split.sptl @@ -55,6 +55,7 @@ kernel @two_phase (stream[4] readonly in, await foreach i32 k, f32 x in [0:K], receive(hop1) { a[k] = a[k] + x } + await hop1.close() await foreach i32 k, f32 x in [0:K], receive(hop2) { a[k] = a[k] + x } @@ -63,18 +64,20 @@ kernel @two_phase (stream[4] readonly in, compute i32 i, i32 j in [1, 0] { await receive(a, in[i]) await send(a, hop1) + await hop1.close() } compute i32 i, i32 j in [2, 0] { await receive(a, in[i]) await foreach i32 k, f32 x in [0:K], receive(hop1) { a[k] = a[k] + x } + await hop1.close() await send(a, hop2) } compute i32 i, i32 j in [3, 0] { await receive(a, in[i]) await send(a, hop1) - + await hop1.close() } } diff --git a/tests/spatial_ir/test_stream_lifetime.py b/tests/spatial_ir/test_stream_lifetime.py new file mode 100644 index 00000000..7535a72d --- /dev/null +++ b/tests/spatial_ir/test_stream_lifetime.py @@ -0,0 +1,462 @@ +""" +Tests for the stream lifetime passes (implicit closes, bound verification, use-after-close, +channel conflicts, and close elision). +""" +import os + +import pytest + +from spada.lowering import spatial_ir_to_csl as s2c +from spada.syntax.spatial_ir import canonicalization, irnodes as spir, parser, passes, stream_lifetime + +SAMPLES = os.path.join(os.path.dirname(__file__), 'samples') + + +def _canonicalized(code: str, **parameters) -> spir.Kernel: + kernel = parser.parse_string(code, 'test.sptl') + if parameters: + kernel = passes.concretize_parameters(kernel, **parameters) + kernel = passes.constexpr_propagation(kernel) + return s2c.canonicalize_kernel(kernel) + + +def _rectangles(code: str, **parameters): + kernel = _canonicalized(code, **parameters) + return canonicalization.consolidate_rectangles_to_equivalence_classes(kernel) + + +def _rectangles_from_file(filename: str, **parameters): + with open(os.path.join(SAMPLES, filename)) as fp: + return _rectangles(fp.read(), **parameters) + + +def _closed_streams(compute: spir.ComputeBlock) -> list[str]: + """ + Names of the streams closed in a compute block, without the version suffix that + ``inline_phases`` adds when it freshens colliding stream names per rectangle. + """ + return [ + statement.stream_name.as_ir().split('#')[0] for statement in compute.statements + if isinstance(statement, spir.CloseStatement) + ] + + +### +# insert_implicit_closes +### + +_TWO_STREAM_KERNEL = """ +kernel @test(stream[2] readonly inp) { + place u16 i, u16 j in [0:2, 0:1] { + f32[K] a + } + dataflow u16 i, u16 j in [0:2, 0:1] { + stream eastwards = relative_stream(1, 0) { + hops = [(1, 0)], + channel = 0 + } + } + compute u16 i, u16 j in [0:1, 0:1] { + await receive(a, inp[i]) + await send(a, eastwards) + } + compute u16 i, u16 j in [1:2, 0:1] { + await foreach i32 k, f32 x in [0:K], receive(eastwards) { + a[k] = x + } + } +} +""" + + +def test_implicit_close_is_inserted_for_every_participant(): + """Both the sending and the receiving PE close the stream: closing is collective.""" + rects = _rectangles(_TWO_STREAM_KERNEL, K=4) + assert len(rects) == 2 + for rect in rects: + assert _closed_streams(rect.metadata.compute) == ['eastwards'] + + +def test_implicit_close_comes_after_the_phase_barrier(): + """ + The close must follow the implicit awaits: they may be waiting on operations that are still + using the stream. + """ + rects = _rectangles(_TWO_STREAM_KERNEL, K=4) + statements = rects[0].metadata.compute.statements + close_index = next(i for i, s in enumerate(statements) if isinstance(s, spir.CloseStatement)) + barrier_index = next(i for i, s in enumerate(statements) if isinstance(s, spir.AwaitAllStatement)) + assert barrier_index < close_index + + +def test_phase_barrier_is_not_duplicated(): + """ + ``insert_implicit_closes`` ends a phase with ``awaitall`` + closes, and ``inline_phases`` would + otherwise append a second barrier right behind it. The closes are awaited, so nothing is + outstanding and the second barrier is redundant. + """ + code = """ + kernel @test(stream[2] readonly inp) { + place u16 i, u16 j in [0:2, 0:1] { f32[K] a } + phase { + dataflow u16 i, u16 j in [0:2, 0:1] { + stream e1 = relative_stream(1, 0) { hops = [(1, 0)], channel = 0 } + } + compute u16 i, u16 j in [0:1, 0:1] { await send(a, e1) } + } + phase { + dataflow u16 i, u16 j in [0:2, 0:1] { + stream e2 = relative_stream(1, 0) { hops = [(1, 0)], channel = 1 } + } + compute u16 i, u16 j in [0:1, 0:1] { await send(a, e2) } + } + } + """ + rects = _rectangles(code, K=4) + sender = next(rect for rect in rects if rect.x_range[0] == 0) + statements = sender.metadata.compute.statements + barriers = [i for i, s in enumerate(statements) if isinstance(s, spir.AwaitAllStatement)] + assert len(barriers) == 2, sender.metadata.compute.as_ir() + assert all(second - first > 1 for first, second in zip(barriers, barriers[1:])) + + +def test_implicit_close_is_not_duplicated(): + """A stream the user already closed is not closed a second time.""" + code = _TWO_STREAM_KERNEL.replace('await send(a, eastwards)', 'await send(a, eastwards)\n' + ' await eastwards.close()') + rects = _rectangles(code, K=4) + sender = next(rect for rect in rects if rect.x_range[0] == 0) + assert _closed_streams(sender.metadata.compute) == ['eastwards'] + + +def test_stream_used_in_two_phases_is_closed_after_its_last_use(): + """ + A stream declared at kernel level stays in scope across phases, so it is closed only in the + phase that uses it last. + """ + code = """ + kernel @test(stream[2] readonly inp) { + place u16 i, u16 j in [0:2, 0:1] { + f32[K] a + } + dataflow u16 i, u16 j in [0:2, 0:1] { + stream eastwards = relative_stream(1, 0) { + hops = [(1, 0)], + channel = 0 + } + } + phase { + compute u16 i, u16 j in [0:1, 0:1] { + await receive(a, inp[i]) + await send(a, eastwards) + } + } + phase { + compute u16 i, u16 j in [0:1, 0:1] { + await send(a, eastwards) + } + } + } + """ + rects = _rectangles(code, K=4) + sender = next(rect for rect in rects if rect.x_range[0] == 0) + # Exactly one close, and it comes after the second send + statements = sender.metadata.compute.statements + closes = [i for i, s in enumerate(statements) if isinstance(s, spir.CloseStatement)] + sends = [i for i, s in enumerate(statements) if isinstance(s, spir.SendStatement)] + assert len(closes) == 1 + assert closes[0] > max(sends) + + +### +# verify_stream_bounds +### + + +def _bounded_kernel(bound: str, elements: str) -> str: + return f""" + kernel @test(stream[2] readonly inp) {{ + place u16 i, u16 j in [0:2, 0:1] {{ + f32[{elements}] a + }} + dataflow u16 i, u16 j in [0:2, 0:1] {{ + stream eastwards = relative_stream(1, 0) {{ + hops = [(1, 0)], + channel = 0 + }} + }} + compute u16 i, u16 j in [0:1, 0:1] {{ + await send(a, eastwards) + }} + compute u16 i, u16 j in [1:2, 0:1] {{ + await foreach i32 k, f32 x in [0:{elements}], receive(eastwards) {{ + a[k] = x + }} + }} + }} + """ + + +def test_matching_bound_is_accepted(): + stream_lifetime.verify_stream_bounds(_rectangles(_bounded_kernel('4', '4'), K=4)) + + +def test_mismatched_bound_is_rejected(): + with pytest.raises(SyntaxError, match='declared with bound 8, but 4 element'): + stream_lifetime.verify_stream_bounds(_rectangles(_bounded_kernel('8', '4'), K=4)) + + +def test_unanalyzable_bound_is_silently_accepted(): + """A ``foreach`` without a range receives until the sender is done, so nothing can be checked.""" + code = """ + kernel @test(stream[2] readonly inp) { + place u16 i, u16 j in [0:2, 0:1] { + f32[K] a + } + dataflow u16 i, u16 j in [0:2, 0:1] { + stream eastwards = relative_stream(1, 0) { + hops = [(1, 0)], + channel = 0 + } + } + compute u16 i, u16 j in [1:2, 0:1] { + await foreach f32 x in receive(eastwards) { + a[0] = x + } + } + } + """ + stream_lifetime.verify_stream_bounds(_rectangles(code, K=4)) + + +def test_bound_counts_enclosing_loop_trips(): + """A send inside a ``for`` loop transfers ``trips * size`` elements.""" + code = """ + kernel @test(stream[2] readonly inp) { + place u16 i, u16 j in [0:2, 0:1] { + f32[4] a + } + dataflow u16 i, u16 j in [0:2, 0:1] { + stream eastwards = relative_stream(1, 0) { + hops = [(1, 0)], + channel = 0 + } + } + compute u16 i, u16 j in [0:1, 0:1] { + for i32 t in [0:3] { + await send(a, eastwards) + } + } + } + """ + with pytest.raises(SyntaxError, match='bound 8, but 12 element'): + stream_lifetime.verify_stream_bounds(_rectangles(code, K=4)) + + +### +# check_use_after_close +### + + +def _use_after_close_kernel(body: str) -> str: + return f""" + kernel @test(stream[2] readonly inp) {{ + place u16 i, u16 j in [0:2, 0:1] {{ + f32[K] a + }} + dataflow u16 i, u16 j in [0:2, 0:1] {{ + stream eastwards = relative_stream(1, 0) {{ + hops = [(1, 0)], + channel = 0 + }} + }} + compute u16 i, u16 j in [0:1, 0:1] {{ +{body} + }} + compute u16 i, u16 j in [1:2, 0:1] {{ + await foreach i32 k, f32 x in [0:K], receive(eastwards) {{ + a[k] = x + }} + }} + }} + """ + + +def test_send_after_close_is_rejected(): + code = _use_after_close_kernel(""" + await send(a, eastwards) + await eastwards.close() + await send(a, eastwards)""") + with pytest.raises(SyntaxError, match='used in a `send` after it was closed'): + stream_lifetime.check_use_after_close(_rectangles(code, K=4)) + + +def test_receive_after_close_is_rejected(): + code = _use_after_close_kernel(""" + await send(a, eastwards) + await eastwards.close() + await foreach i32 k, f32 x in [0:K], receive(eastwards) { + a[k] = x + }""") + with pytest.raises(SyntaxError, match='used in a `receive` after it was closed'): + stream_lifetime.check_use_after_close(_rectangles(code, K=4)) + + +def test_double_close_is_rejected(): + code = _use_after_close_kernel(""" + await send(a, eastwards) + await eastwards.close() + await eastwards.close()""") + with pytest.raises(SyntaxError, match='closed again after it was closed'): + stream_lifetime.check_use_after_close(_rectangles(code, K=4)) + + +def test_close_after_last_use_is_accepted(): + code = _use_after_close_kernel(""" + await send(a, eastwards) + await eastwards.close()""") + stream_lifetime.check_use_after_close(_rectangles(code, K=4)) + + +### +# check_channel_conflicts +### + +_SHARED_CHANNEL_KERNEL = """ +kernel @test(stream[4] readonly inp) {{ + place i16 i, i16 j in [0:4, 0:1] {{ + f32[K] a + }} + dataflow i32 i, i32 j in [0:4, 0:1] {{ + stream hop1 = relative_stream(-1, 0) {{ + hops = [(-1, 0)], + channel = 0 + }} + stream hop2 = relative_stream(-2, 0) {{ + hops = [(-1, 0), (-1, 0)], + channel = {second_channel} + }} + }} + compute i32 i, i32 j in [3:4, 0:1] {{ + await receive(a, inp[i]) + await send(a, hop1) + }} + compute i32 i, i32 j in [2:3, 0:1] {{ + await foreach i32 k, f32 x in [0:K], receive(hop1) {{ + a[k] = x + }} + {close} + await send(a, hop2) + }} + compute i32 i, i32 j in [0:1, 0:1] {{ + await foreach i32 k, f32 x in [0:K], receive(hop2) {{ + a[k] = x + }} + }} +}} +""" + + +def test_unordered_channel_reuse_is_rejected(): + code = _SHARED_CHANNEL_KERNEL.format(second_channel=0, close='') + with pytest.raises(SyntaxError, match="both use channel 0 and share PE"): + stream_lifetime.check_channel_conflicts(_rectangles(code, K=4)) + + +def test_channel_reuse_after_close_is_accepted(): + code = _SHARED_CHANNEL_KERNEL.format(second_channel=0, close='await hop1.close()') + stream_lifetime.check_channel_conflicts(_rectangles(code, K=4)) + + +def test_distinct_channels_never_conflict(): + code = _SHARED_CHANNEL_KERNEL.format(second_channel=1, close='') + stream_lifetime.check_channel_conflicts(_rectangles(code, K=4)) + + +def test_channel_reuse_across_phases_is_accepted(): + """``two_phase.sptl`` reuses channel 0 in two phases; the phase barrier orders the streams.""" + rects = _rectangles_from_file('two_phase.sptl', K=4) + stream_lifetime.check_channel_conflicts(rects) + + +def test_streams_on_one_channel_with_disjoint_paths_are_accepted(): + """Two streams may share a channel without ever meeting on a PE.""" + code = """ + kernel @test(stream[4] readonly inp) { + place i16 i, i16 j in [0:4, 0:1] { + f32[K] a + } + dataflow i32 i, i32 j in [0:2, 0:1] { + stream left = relative_stream(1, 0) { + hops = [(1, 0)], + channel = 0 + } + } + dataflow i32 i, i32 j in [2:4, 0:1] { + stream right = relative_stream(1, 0) { + hops = [(1, 0)], + channel = 0 + } + } + compute i32 i, i32 j in [0:1, 0:1] { + await receive(a, inp[i]) + await send(a, left) + } + compute i32 i, i32 j in [1:2, 0:1] { + await foreach i32 k, f32 x in [0:K], receive(left) { + a[k] = x + } + } + compute i32 i, i32 j in [2:3, 0:1] { + await receive(a, inp[i]) + await send(a, right) + } + compute i32 i, i32 j in [3:4, 0:1] { + await foreach i32 k, f32 x in [0:K], receive(right) { + a[k] = x + } + } + } + """ + stream_lifetime.check_channel_conflicts(_rectangles(code, K=4)) + + +def test_systolic_forwarding_is_not_a_channel_conflict(): + """ + One stream that is received and then forwarded on the same PE needs two router configurations, + but is ordered by local order and therefore not a conflict. + """ + rects = _rectangles_from_file('multihop.sptl', K=4) + stream_lifetime.check_channel_conflicts(rects) + + +### +# elide_redundant_closes +### + + +def test_closes_are_elided_when_the_channel_is_never_reused(): + rects = _rectangles(_TWO_STREAM_KERNEL, K=4) + assert stream_lifetime.elide_redundant_closes(rects) == 2 + assert all(not _closed_streams(rect.metadata.compute) for rect in rects) + + +def test_closes_are_kept_when_the_channel_is_reused(): + code = _SHARED_CHANNEL_KERNEL.format(second_channel=0, close='await hop1.close()') + rects = _rectangles(code, K=4) + assert stream_lifetime.elide_redundant_closes(rects) == 0 + assert any('hop1' in _closed_streams(rect.metadata.compute) for rect in rects) + + +def test_auto_channel_kernels_emit_no_closes(): + """ + Kernels produced by the GT4Py path use ``auto`` channels, one per declaration, so every close is + elided and code generation is unaffected by this feature. + """ + rects = _rectangles_from_file('two_phase_unrouted.sptl', K=4) + stream_lifetime.elide_redundant_closes(rects) + assert all(not _closed_streams(rect.metadata.compute) for rect in rects) + + +if __name__ == '__main__': + pytest.main([__file__]) From f5e99b0e333cc1007938f33ca268085e299720fd Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Thu, 13 Aug 2026 13:50:58 -0700 Subject: [PATCH 04/68] Implement code generation and switch planning/collapsing --- spada/cli/compiler.py | 4 +- spada/lowering/spatial_ir_to_csl.py | 585 +++++++++++++++------ spada/syntax/csl/constants.py | 16 + spada/syntax/csl/statements.py | 11 +- spada/syntax/csl/switching.py | 142 +++++ spada/syntax/spatial_ir/irnodes.py | 7 + spada/syntax/spatial_ir/stream_lifetime.py | 45 +- tests/spatial_ir/test_switching.py | 317 +++++++++++ 8 files changed, 935 insertions(+), 192 deletions(-) create mode 100644 spada/syntax/csl/switching.py create mode 100644 tests/spatial_ir/test_switching.py diff --git a/spada/cli/compiler.py b/spada/cli/compiler.py index 360307c7..1fda636e 100644 --- a/spada/cli/compiler.py +++ b/spada/cli/compiler.py @@ -23,11 +23,12 @@ @click.option('--disable-task-recycling', is_flag=True, help='Disable task ID recycling') @click.option('--disable-copy-elision', is_flag=True, help='Disable copy elimination optimization pass') @click.option('--disable-close-elision', is_flag=True, help='Disable elision of unnecessary stream closes') +@click.option('--disable-switching', is_flag=True, help='Disable router switch positions for shared channels') def compile_spatial_ir(input_file: str, output_folder: str, param: list[str], offset_x: int, offset_y: int, generate_only: bool, disable_benchmarking: bool, disable_asynchronous: bool, disable_dsd: bool, disable_map: bool, disable_task_fusion: bool, disable_task_recycling: bool, disable_copy_elision: bool, - disable_close_elision: bool): + disable_close_elision: bool, disable_switching: bool): # Parse parameters into dictionary kernel_parameters = {} for p in param: @@ -98,6 +99,7 @@ def compile_spatial_ir(input_file: str, output_folder: str, param: list[str], of copy_elision=not disable_copy_elision, task_id_recycling=not disable_task_recycling, close_elision=not disable_close_elision, + disable_switching=disable_switching, ) # Create output folder if it doesn't exist diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 35b59f50..79f3cac0 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -4,6 +4,7 @@ from collections import defaultdict import copy +from dataclasses import dataclass import functools from io import StringIO import textwrap @@ -16,6 +17,7 @@ from spada.syntax.csl import constants as csl, preprocessing, tasks as tdag, statements as cslstmt, dsd_ops from spada.syntax.csl import benchmarking as cslbench from spada.syntax.csl import structures as cslstruct +from spada.syntax.csl import switching as cslswitch from spada.syntax.csl import task_recycling, prune_unused_fields as csl_pruning from spada.syntax.csl.codefile import CodeFile from spada.syntax.csl.statements import name_to_csl, dtype_as_csl, expr_to_csl @@ -55,7 +57,8 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, copy_elision: bool = True, prune_memory: bool = True, task_id_recycling: bool = True, - close_elision: bool = True) -> list[CodeFile]: + close_elision: bool = True, + disable_switching: bool = False) -> list[CodeFile]: """ Lowers a routed Spatial IR kernel into Cerebras CSL code. @@ -70,6 +73,8 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, :param prune_memory: If True, enables unused field pruning optimization pass. :param task_id_recycling: If True, enables task ID recycling pass. :param close_elision: If True, removes stream closes that no router has to act on. + :param disable_switching: If True, emits one route configuration per stream instead of merging + them into router switch positions. :return: List of code-file objects that can be written to files. See ``write_code_to_files``. """ # PRECONDITION: Rectangles of dataflow/compute/place do not intersect (comes from Spatial IR) @@ -140,8 +145,11 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, stream_lifetime.verify_stream_bounds(rectangles) stream_lifetime.check_use_after_close(rectangles) stream_lifetime.check_channel_conflicts(rectangles) + + # Plan the router switch advances, then drop every close no router has to act on + plan_switch_advances(rectangles) if close_elision: - stream_lifetime.elide_redundant_closes(rectangles) + stream_lifetime.elide_redundant_closes(rectangles, needs_advance=lambda stmt: bool(stmt.switch_advance)) for rect in rectangles: # Create a unique CSL code file based on rectangle offset @@ -167,7 +175,7 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, rect_size = x1 - x0, y1 - y0 # Collect unique routes for all rectangles - routes_per_rectangle = _collect_routes(rectangles, color_maps) + routes_per_rectangle = _collect_routes(rectangles, color_maps, disable_switching) if use_memcpy_mode: layout_code.write(f''' @@ -337,6 +345,7 @@ def generate_rectangle(kernel: spir.Kernel, # * Generate routing instructions from dataflow blocks # * Make unique colors out of streams, reduce number of streams color_map = _allocate_colors(rect, header, kernel, use_memcpy_mode, stream_extents, channel_to_color) + _declare_switch_advances(rect, header, color_map) dtypes = _collect_identifier_types(rect.metadata, kernel.arguments) # Preprocess potential data tasks to convert to loops if possible @@ -661,6 +670,32 @@ def _collect_colors_globally(kernel: spir.Kernel, rectangles: list[Rectangle[PEB return channel_to_color +def _declare_switch_advances(rect: Rectangle[PEBlock], header: StringIO, color_map: dict[str, int]) -> None: + """ + Declares the fabric output descriptors that carry a stream's switch-advance control wavelet. + + The wavelet itself is emitted by ``statements.generate_csl_statement`` from the close's + ``switch_advance`` field; this only has to provide the descriptor it is sent through, because + that is where the color is known. + """ + kept = [ + statement for statement in rect.metadata.compute.statements + if isinstance(statement, spir.CloseStatement) and statement.switch_advance + ] + if not kept: + return + + header.write('\nconst ctrl = @import_module("");\n') + queue = csl.OUTPUT_QUEUE_IDS[0] + for statement in kept: + name = name_to_csl(stream_lifetime.underlying_stream(statement.stream_name)) + dsd_name = f'{name}_switch_dsd' + if f'const {dsd_name}' not in header.getvalue(): + header.write(f'const {dsd_name} = @get_dsd(fabout_dsd, .{{ .extent = 1, ' + f'.fabric_color = @get_color({color_map[name + "_OUT"]}), .control = true, ' + f'.output_queue = @get_output_queue({queue}) }});\n') + + def _allocate_colors(rect: Rectangle[PEBlock], header: StringIO, kernel: spir.Kernel, use_memcpy_mode: bool, stream_extents: analysis.StreamExtents, channel_to_color: dict[int, int]) -> dict[str, int]: """ @@ -1214,192 +1249,398 @@ def _route_dir(dx: int, dy: int): return ('NORTH', 'SOUTH') -def _collect_routes(rectangles: list[Rectangle[PEBlock]], color_maps: list[dict[str, - int]]) -> dict[tuple[int, int], str]: +@dataclass(frozen=True) +class _RouteSite: + """ + One ``@set_color_config`` target: the rectangle of PEs (already shifted by the relay offset) that + receive a route configuration for one color. + """ + color: int + x_range: tuple[int, int, int] + y_range: tuple[int, int, int] + + def as_rectangle(self) -> Rectangle: + return Rectangle(self.x_range, self.y_range, None) + + def describe(self) -> str: + return (f'PEs [{self.x_range[0]}:{self.x_range[1]}, {self.y_range[0]}:{self.y_range[1]}]') + + +@dataclass +class _RouteEntry: + """ + One route configuration contributed to a site, with the key that orders it against the other + configurations of the same site. + """ + config: cslswitch.RouteConfig + order: tuple[int, int] + origin_rect: int + origin_offset: tuple[int, int] + stream: spir.Identifier + #: Routing identity of the stream; stable across the per-rectangle renaming of ``inline_phases`` + group: str = '' + + +def _stream_use_order(compute: spir.ComputeBlock) -> dict[spir.Identifier, dict[str, tuple[int, int]]]: + """ + Returns, per stream, the order key of its first receive and of its first send in a compute block. + + The key is ``(barrier index, statement index)``: statements are ordered first by how many phase + barriers precede them, then by their position in the block. This is what sequences the route + configurations of a router into switch positions. + """ + result: dict[spir.Identifier, dict[str, tuple[int, int]]] = {} + barrier = 0 + for index, statement in enumerate(compute.statements): + if isinstance(statement, spir.AwaitAllStatement): + barrier += 1 + continue + for kind, expression in stream_lifetime.stream_references(statement): + if kind == 'close': + continue + name = stream_lifetime.underlying_stream(expression) + orders = result.setdefault(name, {}) + orders.setdefault(kind, (barrier, index)) + return result + + +def _offset_expression(axis: str, offset: int) -> str: + return axis if offset == 0 else f'{axis} + {offset}' + + +def _collect_routes(rectangles: list[Rectangle[PEBlock]], + color_maps: list[dict[str, int]], + disable_switching: bool = False) -> dict[tuple[int, int], str]: """ Creates a parametric version of the Routing Graph (see the Spatial IR specification for more information) and returns a dictionary of code segements to add to the layout CSL file based on the streams. + Route configurations are collected per *site* -- a rectangle of PEs and a color -- and merged + across rectangles, because a multi-hop stream configures its relay PEs from the sending + rectangle's loop. A site that ends up with more than one configuration is lowered to router + switch positions, ordered by the local order of the statements that use the streams. + :param rectangles: All rectangles involved in this kernel. + :param color_maps: Per-rectangle mapping of stream names to colors. + :param disable_switching: If True, emit each configuration as its own ``@set_color_config`` + instead of merging them into switch positions. :return: A dictionary mapping the starting point of each rectangle to a string representing the layout instructions. """ INDENT = 12 * ' ' - result = {} - # Create a routing graph - for rect, color_map in zip(rectangles, color_maps): - # Test whether a receive/send statement are called for creating inbound/outbound routes - sends_recvs = analysis.sends_and_receives(rect.metadata.compute) - inst = '' + entries: dict[_RouteSite, list[_RouteEntry]] = {} + for rect_index, (rect, color_map) in enumerate(zip(rectangles, color_maps)): + for site, entry in _rectangle_route_entries(rect_index, rect, color_map): + entries.setdefault(site, []).append(entry) + + _check_site_overlap(entries) + + result = {(rect.x_range[0], rect.y_range[0]): '' for rect in rectangles} + for site, site_entries in entries.items(): + site_entries.sort(key=lambda entry: entry.order) + + # The site is configured from the loop of one rectangle: the one that owns these PEs if + # there is one, otherwise the first relay that reaches them. + owner = min(site_entries, key=lambda entry: (entry.origin_offset != (0, 0), entry.origin_rect)) + owner_rect = rectangles[owner.origin_rect] + key = (owner_rect.x_range[0], owner_rect.y_range[0]) + x = _offset_expression('pe_x', owner.origin_offset[0]) + y = _offset_expression('pe_y', owner.origin_offset[1]) + color = f'@get_color({site.color})' + + if disable_switching: + for entry in site_entries: + plan = cslswitch.ColorSwitchPlan([entry.config]) + text = cslswitch.set_color_config(x, y, color, plan, INDENT) + if text not in result[key]: + result[key] += text + continue + + plan = cslswitch.ColorSwitchPlan() + for entry in site_entries: + plan.add(entry.config) + plan.validate(site.color, site.describe()) + result[key] += cslswitch.set_color_config(x, y, color, plan, INDENT) + + return result + + +def _route_sites(rectangles: list[Rectangle[PEBlock]], + color_maps: list[dict[str, int]]) -> dict['_RouteSite', list['_RouteEntry']]: + """ + Collects the route configurations of every site, sorted into switch-position order. + """ + entries: dict[_RouteSite, list[_RouteEntry]] = {} + for rect_index, (rect, color_map) in enumerate(zip(rectangles, color_maps)): + for site, entry in _rectangle_route_entries(rect_index, rect, color_map): + entries.setdefault(site, []).append(entry) + for site_entries in entries.values(): + site_entries.sort(key=lambda entry: entry.order) + return entries - # Make routing instructions unique - routing_instructions: set[str] = set() - # For each hop, make a color WEST-EAST/NORTH-SOUTH pair. For the first and last hop, pair with RAMP - for stream in rect.metadata.dataflow.statements: - if stream.stream_name not in sends_recvs: # Skip unused streams +def _channel_color_maps(rectangles: list[Rectangle[PEBlock]]) -> list[dict[str, int]]: + """ + Builds color maps that use the stream's *channel* in place of its color. + + Switch planning has to run before colors are allocated per rectangle, and the shape of a + router's configuration sequence only depends on the channel (a channel maps to exactly one + color). + """ + color_maps = [] + for rect in rectangles: + color_map = {} + for declaration in rect.metadata.dataflow.statements: + if declaration.stream.routing is None: continue - sent, received = sends_recvs[stream.stream_name] - if received: - color_name_inbound = f'@get_color({color_map[name_to_csl(stream.stream_name) + "_IN"]})' - if sent: - color_name_outbound = f'@get_color({color_map[name_to_csl(stream.stream_name) + "_OUT"]})' - - if isinstance(stream.stream, spir.ExternStreamDeclaration): - continue # Extern streams do not have on-chip routing - - if isinstance(stream.stream, spir.MulticastRangeStreamDeclaration): - if sent and received: - raise ValueError( - f"Multicast stream '{stream.stream_name.as_ir()}' is both sent and received " - f"within the same compute rectangle [{rect.x_range[0]}:{rect.x_range[1]}, " - f"{rect.y_range[0]}:{rect.y_range[1]}]. " - "Sender and receiver compute blocks must be in separate rectangles for multicast streams.") - if not sent: - # All multicast routing is emitted by the rectangle that sends this stream. + channel = declaration.stream.routing.resolved_channel + if channel == 'auto': + continue + name = name_to_csl(declaration.stream_name) + color_map[name + '_IN'] = channel + color_map[name + '_OUT'] = channel + color_maps.append(color_map) + return color_maps + + +def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: + """ + Determines, for every ``close`` statement, which routers along the stream's path must advance + their switch. + + A close only produces code on a PE that *sends* the stream: the control message it emits travels + the path being retired and advances each router it traverses. Routers whose configuration does + not change are given a no-op so that they stay where they are. A close on a receiving PE emits + nothing; its router is advanced by the sender's message. + + The result is recorded on each ``CloseStatement`` as its ``switch_advance`` field; a close that + needs no advance keeps ``None`` there and generates no code. + + :param rectangles: The consolidated PE rectangles of the kernel, annotated in place. + :return: The number of closes that retire a route configuration. + """ + sites = _route_sites(rectangles, _channel_color_maps(rectangles)) + + # Per site, the switch position of every configuration, and how many positions the site holds + position_of: dict[_RouteSite, list[tuple[_RouteEntry, int]]] = {} + positions_at: dict[_RouteSite, int] = {} + for site, site_entries in sites.items(): + configs: list[cslswitch.RouteConfig] = [] + placed = [] + for entry in site_entries: + if not configs or configs[-1] != entry.config: + configs.append(entry.config) + placed.append((entry, len(configs) - 1)) + position_of[site] = placed + positions_at[site] = len(configs) + + planned = 0 + for rect in rectangles: + declarations = {d.stream_name: d for d in rect.metadata.dataflow.statements} + uses = stream_lifetime.collect_stream_uses(rect.metadata.compute) + for statement in rect.metadata.compute.statements: + if not isinstance(statement, spir.CloseStatement): + continue + statement.switch_advance = None + name = stream_lifetime.underlying_stream(statement.stream_name) + declaration = declarations.get(name) + if declaration is None or name not in uses or not uses[name].sent: + continue # Not sent here: the sending PE's control message advances this router + channel = declaration.stream.routing.resolved_channel if declaration.stream.routing else 'auto' + offsets = stream_lifetime._stream_path_offsets(declaration) + if channel == 'auto' or offsets is None: + continue + group = stream_lifetime.stream_group_key(declaration) + + commands = [] + for dx, dy in offsets: + site = _find_site(position_of, channel, + (rect.x_range[0] + dx, rect.x_range[1] + dx, rect.x_range[2]), + (rect.y_range[0] + dy, rect.y_range[1] + dy, rect.y_range[2])) + if site is None: + commands.append(False) continue - rng = stream.stream.multicast_range - start = int(rng.start.eval()) - stop = int(rng.stop.eval()) - axis = stream.stream.multicast_axis - is_negative = start < 0 - - if axis == 'y': - if is_negative: - tx_dir, rx_dir = 'NORTH', 'SOUTH' - else: - tx_dir, rx_dir = 'SOUTH', 'NORTH' + position = _traffic_position(position_of[site], group, is_sender=(dx == 0 and dy == 0)) + commands.append(position is not None and position + 1 < positions_at.get(site, 0)) - def _coord(k): # noqa: E731 - if k >= 0: - return 'pe_x', f'pe_y + {k}' - return 'pe_x', f'pe_y - {-k}' - else: - if is_negative: - tx_dir, rx_dir = 'WEST', 'EAST' - else: - tx_dir, rx_dir = 'EAST', 'WEST' - - def _coord(k): # noqa: E731 - if k >= 0: - return f'pe_x + {k}', 'pe_y' - return f'pe_x - {-k}', 'pe_y' - - # Sender: inject into fabric toward receivers. - routing_inst = INDENT + '@set_color_config(pe_x, pe_y, %s, .{ .routes = .{ .rx = .{RAMP}, .tx = .{%s} } });\n' % ( - color_name_outbound, tx_dir) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - - if is_negative: - # Negative multicast: receivers at start, start-1, …, stop+1 (stop exclusive). - k_last = stop + 1 # farthest receiver - - # Gap relay-only PEs between sender and first receiver (when start < -1). - for k in range(-1, start, -1): - cx, cy = _coord(k) - routing_inst = INDENT + '@set_color_config(%s, %s, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{%s} } });\n' % ( - cx, cy, color_name_outbound, rx_dir, tx_dir) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - - # Intermediate receivers: forward toward farthest and deliver to RAMP. - for k in range(start, k_last, -1): - cx, cy = _coord(k) - routing_inst = INDENT + '@set_color_config(%s, %s, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{%s, RAMP} } });\n' % ( - cx, cy, color_name_outbound, rx_dir, tx_dir) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - - # Last (farthest) receiver: deliver to RAMP only, no forwarding. - cx, cy = _coord(k_last) - routing_inst = INDENT + '@set_color_config(%s, %s, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{RAMP} } });\n' % ( - cx, cy, color_name_outbound, rx_dir) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - else: - # Positive multicast: receivers at start, start+1, …, stop-1 (stop exclusive). - # Gap relay-only PEs between sender and first receiver (when start > 1). - for k in range(1, start): - cx, cy = _coord(k) - routing_inst = INDENT + '@set_color_config(%s, %s, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{%s} } });\n' % ( - cx, cy, color_name_outbound, rx_dir, tx_dir) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - - # Intermediate receivers: forward and simultaneously deliver to RAMP. - for k in range(start, stop - 1): - cx, cy = _coord(k) - routing_inst = INDENT + '@set_color_config(%s, %s, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{%s, RAMP} } });\n' % ( - cx, cy, color_name_outbound, rx_dir, tx_dir) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - - # Last receiver: deliver to RAMP only, no forwarding. - k_last = stop - 1 - cx, cy = _coord(k_last) - routing_inst = INDENT + '@set_color_config(%s, %s, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{RAMP} } });\n' % ( - cx, cy, color_name_outbound, rx_dir) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) + if any(commands): + statement.switch_advance = commands + planned += 1 + + return planned + + +def _find_site(sites, color: int, x_range: tuple[int, int, int], y_range: tuple[int, int, int]): + """ + Finds the site that configures a shifted rectangle of PEs for a color. + + An exact match is the common case, but a stream whose sender and receiver live in the *same* + rectangle shifts that rectangle onto itself: PE ``(i, j)`` sends to ``(i, j+1)``, which the same + parametric loop configures. Those lookups are resolved by intersection. + """ + exact = _RouteSite(color=color, x_range=x_range, y_range=y_range) + if exact in sites: + return exact + shifted = Rectangle(x_range, y_range, None) + for site in sites: + if site.color == color and site.as_rectangle().intersects(shifted): + return site + return None + + +def _traffic_position(placed: list[tuple['_RouteEntry', int]], group: str, is_sender: bool): + """ + Returns the switch position that a stream's traffic occupies at one router. + + The sending PE injects from its own ramp, so its configuration is the one with ``rx = RAMP``; + every router further along the path forwards traffic that arrives from the fabric. Picking the + right one matters when a PE both receives and sends the same stream, as in a systolic chain: + the message that retires the incoming configuration has to advance that router onto the + outgoing one. + """ + for entry, position in placed: + if entry.group != group: + continue + if is_sender == (entry.config.rx == ('RAMP', )): + return position + return None + + +def _check_site_overlap(entries: dict['_RouteSite', list['_RouteEntry']]) -> None: + """ + Raises a ``SyntaxError`` if two sites of the same color cover overlapping but different sets of + PEs, which cannot be expressed as a single parametric ``@set_color_config`` loop. + """ + sites = list(entries) + for index, first in enumerate(sites): + for second in sites[index + 1:]: + if first.color != second.color: continue + if not first.as_rectangle().intersects(second.as_rectangle()): + continue + raise SyntaxError( + f'Color {first.color} is configured differently on overlapping but distinct PE ' + f'regions {first.describe()} and {second.describe()}.\n' + ' note: the two regions would need separate switch sequences; split the compute ' + 'blocks so that the regions coincide or are disjoint') - if len(stream.stream.routing.hops) == 1: # Inbound and outbound generated together - route = _route_dir(*stream.stream.routing.hops[0].offset) - if sent: - routing_inst = INDENT + '@set_color_config(pe_x, pe_y, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{%s} } });\n' % ( - color_name_outbound, 'RAMP', route[1]) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - if received: - routing_inst = INDENT + '@set_color_config(pe_x, pe_y, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{%s} } });\n' % ( - color_name_inbound, route[0], 'RAMP') - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - else: # Multi-hop - if sent: - first_hop = stream.stream.routing.hops[0] - route = ('RAMP', _route_dir(*first_hop.offset)[1]) - routing_inst = INDENT + '@set_color_config(pe_x, pe_y, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{%s} } });\n' % ( - color_name_outbound, route[0], route[1]) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - cur_offx = 0 - cur_offy = 0 - for hop in stream.stream.routing.hops[1:]: - route = _route_dir(*hop.offset) - cur_offx += hop.offset[0] - cur_offy += hop.offset[1] - routing_inst = INDENT + '@set_color_config(pe_x + %d, pe_y + %d, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{%s} } });\n' % ( - cur_offx, cur_offy, color_name_outbound, route[0], route[1]) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - if received: - # The receiver only configures itself (pe_x + 0, pe_y + 0). - # Intermediate PEs are configured by the sender block above, - # which walks forward through hops[1:] relative to the sender PE. - last_hop = stream.stream.routing.hops[-1] - route = (_route_dir(*last_hop.offset)[0], 'RAMP') - routing_inst = INDENT + '@set_color_config(pe_x, pe_y, %s, .{ .routes = .{ .rx = .{%s}, .tx = .{%s} } });\n' % ( - color_name_inbound, route[0], route[1]) - if routing_inst not in routing_instructions: - inst += routing_inst - routing_instructions.add(routing_inst) - - result[(rect.x_range[0], rect.y_range[0])] = inst - return result +def _rectangle_route_entries(rect_index: int, rect: Rectangle[PEBlock], + color_map: dict[str, int]) -> list[tuple['_RouteSite', '_RouteEntry']]: + """ + Collects every route configuration a single rectangle contributes, as ``(site, entry)`` pairs. + """ + # Test whether a receive/send statement are called for creating inbound/outbound routes + sends_recvs = analysis.sends_and_receives(rect.metadata.compute) + use_order = _stream_use_order(rect.metadata.compute) + collected: list[tuple[_RouteSite, _RouteEntry]] = [] + + def add(offset: tuple[int, int], color: int, rx: tuple[str, ...], tx: tuple[str, ...], + order: tuple[int, int], stream_name: spir.Identifier, group: str) -> None: + site = _RouteSite( + color=color, + x_range=(rect.x_range[0] + offset[0], rect.x_range[1] + offset[0], rect.x_range[2]), + y_range=(rect.y_range[0] + offset[1], rect.y_range[1] + offset[1], rect.y_range[2])) + collected.append((site, + _RouteEntry(cslswitch.RouteConfig(rx, tx), order, rect_index, offset, stream_name, group))) + + # For each hop, make a color WEST-EAST/NORTH-SOUTH pair. For the first and last hop, pair with RAMP + for stream in rect.metadata.dataflow.statements: + if stream.stream_name not in sends_recvs: # Skip unused streams + continue + sent, received = sends_recvs[stream.stream_name] + group = stream_lifetime.stream_group_key(stream) + orders = use_order.get(stream.stream_name, {}) + receive_order = orders.get('receive', (0, 0)) + send_order = orders.get('send', (0, 0)) + if received: + color_inbound = color_map[name_to_csl(stream.stream_name) + "_IN"] + if sent: + color_outbound = color_map[name_to_csl(stream.stream_name) + "_OUT"] + + if isinstance(stream.stream, spir.ExternStreamDeclaration): + continue # Extern streams do not have on-chip routing + + if isinstance(stream.stream, spir.MulticastRangeStreamDeclaration): + if sent and received: + raise ValueError( + f"Multicast stream '{stream.stream_name.as_ir()}' is both sent and received " + f"within the same compute rectangle [{rect.x_range[0]}:{rect.x_range[1]}, " + f"{rect.y_range[0]}:{rect.y_range[1]}]. " + "Sender and receiver compute blocks must be in separate rectangles for multicast streams.") + if not sent: + # All multicast routing is emitted by the rectangle that sends this stream. + continue + rng = stream.stream.multicast_range + start = int(rng.start.eval()) + stop = int(rng.stop.eval()) + axis = stream.stream.multicast_axis + is_negative = start < 0 + + if axis == 'y': + tx_dir, rx_dir = ('NORTH', 'SOUTH') if is_negative else ('SOUTH', 'NORTH') + + def _coord(k): # noqa: E731 + return (0, k) + else: + tx_dir, rx_dir = ('WEST', 'EAST') if is_negative else ('EAST', 'WEST') + + def _coord(k): # noqa: E731 + return (k, 0) + + # Sender: inject into fabric toward receivers. + add((0, 0), color_outbound, ('RAMP', ), (tx_dir, ), send_order, stream.stream_name, group) + + if is_negative: + # Negative multicast: receivers at start, start-1, …, stop+1 (stop exclusive). + k_last = stop + 1 # farthest receiver + gap = range(-1, start, -1) + intermediate = range(start, k_last, -1) + else: + # Positive multicast: receivers at start, start+1, …, stop-1 (stop exclusive). + k_last = stop - 1 + gap = range(1, start) + intermediate = range(start, stop - 1) + + # Gap relay-only PEs between the sender and the first receiver. + for k in gap: + add(_coord(k), color_outbound, (rx_dir, ), (tx_dir, ), send_order, stream.stream_name, group) + + # Intermediate receivers: forward toward the farthest one and deliver to RAMP. + for k in intermediate: + add(_coord(k), color_outbound, (rx_dir, ), (tx_dir, 'RAMP'), send_order, stream.stream_name, group) + + # Last (farthest) receiver: deliver to RAMP only, no forwarding. + add(_coord(k_last), color_outbound, (rx_dir, ), ('RAMP', ), send_order, stream.stream_name, group) + continue + + if len(stream.stream.routing.hops) == 1: # Inbound and outbound generated together + route = _route_dir(*stream.stream.routing.hops[0].offset) + if sent: + add((0, 0), color_outbound, ('RAMP', ), (route[1], ), send_order, stream.stream_name, group) + if received: + add((0, 0), color_inbound, (route[0], ), ('RAMP', ), receive_order, stream.stream_name, group) + else: # Multi-hop + if sent: + first_hop = stream.stream.routing.hops[0] + add((0, 0), color_outbound, ('RAMP', ), (_route_dir(*first_hop.offset)[1], ), send_order, + stream.stream_name, group) + cur_offx = 0 + cur_offy = 0 + for hop in stream.stream.routing.hops[1:]: + route = _route_dir(*hop.offset) + cur_offx += hop.offset[0] + cur_offy += hop.offset[1] + add((cur_offx, cur_offy), color_outbound, (route[0], ), (route[1], ), send_order, + stream.stream_name, group) + if received: + # The receiver only configures itself. Intermediate PEs are configured by the sender + # block above, which walks forward through hops[1:] relative to the sender PE. + last_hop = stream.stream.routing.hops[-1] + add((0, 0), color_inbound, (_route_dir(*last_hop.offset)[0], ), ('RAMP', ), receive_order, + stream.stream_name, group) + + return collected def _write_indented_block(current_code: StringIO, block: str, indent: str) -> None: diff --git a/spada/syntax/csl/constants.py b/spada/syntax/csl/constants.py index cb83062a..fccaa017 100644 --- a/spada/syntax/csl/constants.py +++ b/spada/syntax/csl/constants.py @@ -53,3 +53,19 @@ 'wse3': (762, 1172), } HARDWARE_FABRIC_DIMS = _HARDWARE_FABRIC_DIMS[ARCH] + +# Router switches. Each router holds one base route configuration (``.routes``) plus up to three +# switch positions (``.pos1``, ``.pos2``, ``.pos3``) per color. +# See https://sdk.cerebras.ai/csl/language/builtins#switching-configuration-semantics +SWITCH_POSITIONS = 4 + +# Maximum number of switching commands that fit in one control wavelet (````'s MAX_CMDS). +# One command is consumed per router the wavelet traverses. +MAX_CONTROL_COMMANDS = 8 + +# Colors whose routers support switches. WSE-3 only implements switches on a subset of colors. +_SWITCHABLE_COLORS = { + 'wse2': list(range(0, 21)), + 'wse3': [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 12, 13, 16, 17, 20], +} +SWITCHABLE_COLORS = [color for color in _SWITCHABLE_COLORS[ARCH] if color in COLORS] diff --git a/spada/syntax/csl/statements.py b/spada/syntax/csl/statements.py index 8cb93140..acb24974 100644 --- a/spada/syntax/csl/statements.py +++ b/spada/syntax/csl/statements.py @@ -2,6 +2,7 @@ from typing import Optional from spada.syntax.csl.structures import DataStructureDescriptor from spada.syntax.csl import dsd_ops +from spada.syntax.csl import switching from spada.syntax.spatial_ir import irnodes as spir UniqueDSDDict = dict[str, list[tuple[str, DataStructureDescriptor]]] @@ -50,9 +51,13 @@ def generate_csl_statement(statement: spir.Statement, # Skip (taken care of when tasks are defined) return "" elif isinstance(statement, spir.CloseStatement): - # TODO(switching): Lower to a switch advance on the stream's channel. - raise NotImplementedError('Closing a stream is not yet supported by the CSL backend.\n' - f' In line {statement.lineinfo}') + # Retiring a route configuration is a control wavelet that advances the routers along the + # stream's path, or nothing at all when none of them has to move. + if not statement.switch_advance: + return "" + stream = statement.stream_name + name = name_to_csl(stream.array if isinstance(stream, spir.ArraySlice) else stream) + return '@mov32(%s_switch_dsd, %s);' % (name, switching.switch_advance_payload(statement.switch_advance)) if op is None: return f'// TODO: Convert {statement} to CSL' diff --git a/spada/syntax/csl/switching.py b/spada/syntax/csl/switching.py new file mode 100644 index 00000000..ea28e712 --- /dev/null +++ b/spada/syntax/csl/switching.py @@ -0,0 +1,142 @@ +""" +Router route configurations and switch planning for CSL code generation. + +Each router holds, per color, one base route configuration plus up to three *switch positions*. A +router advances from one position to the next when a switch-advance control message passes through +it, which is what a stream ``close`` lowers to. This module owns: + +* :class:`RouteConfig` -- one ``.rx``/``.tx`` pair, the configuration of one router for one color. +* :class:`ColorSwitchPlan` -- the ordered sequence of configurations a router cycles through. +* :func:`set_color_config` -- the single place where ``@set_color_config`` text is produced. +* :func:`switch_advance_payload` -- the ```` payload that advances a path's routers. + +See ``irspec/docs/spatial/routing.md`` ("Lowering to Switches") for the semantics. +""" +from dataclasses import dataclass, field + +from spada.syntax.csl import constants + + +@dataclass(frozen=True) +class RouteConfig: + """ + The configuration of one router for one color: the directions it receives from and transmits to. + + ``RAMP`` denotes the PE's own compute element. + """ + rx: tuple[str, ...] + tx: tuple[str, ...] + + def as_csl(self) -> str: + return '.{ .rx = .{%s}, .tx = .{%s} }' % (', '.join(self.rx), ', '.join(self.tx)) + + def as_switch_position(self) -> str: + """ + Renders this configuration as a ``.posN`` struct. The ``rx`` field of a switch position only + accepts a single direction, unlike the base configuration. + """ + if len(self.rx) != 1: + raise ValueError(f'A switch position can only receive from a single direction, got {self.rx}') + return '.{ .rx = %s, .tx = .{%s} }' % (self.rx[0], ', '.join(self.tx)) + + +@dataclass +class ColorSwitchPlan: + """ + The ordered route configurations a router cycles through for one color. + + ``positions[0]`` is the base configuration, and every further entry becomes a switch position. + Consecutive identical configurations are collapsed by :meth:`add`, so a router that keeps the + same configuration across an epoch boundary consumes no switch position and needs no advance. + """ + positions: list[RouteConfig] = field(default_factory=list) + ring_mode: bool = False + + def add(self, config: RouteConfig) -> None: + if self.positions and self.positions[-1] == config: + return + self.positions.append(config) + + @property + def uses_switches(self) -> bool: + return len(self.positions) > 1 + + def index_of_epoch(self, epoch: int) -> int: + """ + Returns the switch position an epoch maps to, given that identical configurations collapse. + """ + return min(epoch, len(self.positions) - 1) + + def as_csl(self) -> str: + base = self.positions[0].as_csl() + if not self.uses_switches: + return '.{ .routes = %s }' % base + + switches = [ + '.pos%d = %s' % (index, config.as_switch_position()) + for index, config in enumerate(self.positions[1:], start=1) + ] + if self.ring_mode: + switches.append('.ring_mode = true') + return '.{ .routes = %s, .switches = .{ %s } }' % (base, ', '.join(switches)) + + def validate(self, color: int, location: str) -> None: + """ + Raises a ``SyntaxError`` if the plan exceeds what a router can hold. + + :param color: The color the plan is for, used for the diagnostic and for the WSE-3 check. + :param location: A human-readable description of the PE the plan belongs to. + """ + if len(self.positions) > constants.SWITCH_POSITIONS: + raise SyntaxError( + f'Color {color} at {location} requires {len(self.positions)} route configurations, ' + f'but a router holds at most {constants.SWITCH_POSITIONS} per color on ' + f'{constants.ARCH}.\n' + ' note: assign a different channel to some of the streams, at the cost of an ' + 'additional color') + if self.uses_switches and color not in constants.SWITCHABLE_COLORS: + raise SyntaxError( + f'Color {color} at {location} needs router switches, but {constants.ARCH} only ' + f'supports switches on colors {constants.SWITCHABLE_COLORS}.\n' + ' note: assign the channel to a switchable color') + + +def set_color_config(x: str, y: str, color: str, plan: ColorSwitchPlan, indent: str = '') -> str: + """ + Renders a single ``@set_color_config`` call. This is the only place that produces such text. + + :param x: The PE x coordinate expression (e.g. ``pe_x`` or ``pe_x + -1``). + :param y: The PE y coordinate expression. + :param color: The color expression (e.g. ``@get_color(0)``). + :param plan: The route configurations the router cycles through. + :param indent: Indentation to prefix the line with. + """ + return indent + '@set_color_config(%s, %s, %s, %s);\n' % (x, y, color, plan.as_csl()) + + +def switch_advance_payload(commands: list[bool]) -> str: + """ + Returns the ```` expression for a switch-advance control wavelet. + + One command is consumed per router the wavelet traverses, in order, so ``commands[i]`` says + whether the ``i``-th router on the path advances. Routers whose configuration does not change + are given a ``NOP`` so that they stay on their current position. + + :param commands: Per-router advance flags, starting at the sending PE's own router. + """ + if not commands: + raise ValueError('A switch advance needs at least one router command') + if len(commands) > constants.MAX_CONTROL_COMMANDS: + raise SyntaxError( + f'A switch advance along this path needs {len(commands)} router commands, but a control ' + f'wavelet carries at most {constants.MAX_CONTROL_COMMANDS}.\n' + ' note: shorten the routing path, or split it across two channels') + + if len(commands) == 1: + opcode = 'ctrl.opcode.SWITCH_ADV' if commands[0] else 'ctrl.opcode.NOP' + return f'ctrl.encode_single_payload({opcode}, true, {{}}, 0)' + + opcodes = ', '.join('ctrl.opcode.SWITCH_ADV' if advance else 'ctrl.opcode.NOP' for advance in commands) + ce_ignore = ', '.join('true' for _ in commands) + return ('ctrl.encode_payload(.{ .opcodes = .{%s}, .ce_ignore = .{%s}, ' + '.ce_ignore_remaining = true })' % (opcodes, ce_ignore)) diff --git a/spada/syntax/spatial_ir/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index bc77eb64..276ad4fe 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -924,11 +924,18 @@ class CloseStatement(Statement): """ stream_name: Union[Identifier, ArraySlice] completion_name: Optional[Completion] = None + #: Which routers along the stream's path advance their switch when this close retires the + #: stream's route configuration, starting at the sending PE. Filled in during lowering by + #: ``spatial_ir_to_csl.plan_switch_advances``; ``None`` means no router has to act, in which + #: case the close generates no code. Not part of the surface syntax. + switch_advance: Optional[list[bool]] = None def validate(self) -> None: assert isinstance(self.stream_name, (Identifier, ArraySlice)) if self.completion_name: assert isinstance(self.completion_name, Completion) + if self.switch_advance is not None: + assert all(isinstance(command, bool) for command in self.switch_advance) def as_ir(self, indent: int = 0) -> str: indent_str = ' ' * indent diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py index 6d54cdd3..b5318a7f 100644 --- a/spada/syntax/spatial_ir/stream_lifetime.py +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -47,10 +47,10 @@ class StreamUse: @property def name(self) -> spir.Identifier: - return _underlying_stream(self.expression) + return underlying_stream(self.expression) -def _underlying_stream(expression: StreamExpression) -> spir.Identifier: +def underlying_stream(expression: StreamExpression) -> spir.Identifier: """ Returns the stream identifier behind a stream expression, unwrapping array slices. """ @@ -59,7 +59,7 @@ def _underlying_stream(expression: StreamExpression) -> spir.Identifier: return expression -def _stream_references(statement: spir.Statement) -> list[tuple[str, StreamExpression]]: +def stream_references(statement: spir.Statement) -> list[tuple[str, StreamExpression]]: """ Returns every stream reference in a statement (including nested ones) as ``(kind, expression)`` pairs, where ``kind`` is one of ``'send'``, ``'receive'``, or ``'close'``. @@ -87,8 +87,8 @@ def collect_stream_uses(compute: spir.ComputeBlock) -> dict[spir.Identifier, Str """ result: dict[spir.Identifier, StreamUse] = {} for index, statement in enumerate(compute.statements): - for kind, expression in _stream_references(statement): - name = _underlying_stream(expression) + for kind, expression in stream_references(statement): + name = underlying_stream(expression) use = result.get(name) if use is None: use = StreamUse(expression=expression, first_use=index, last_use=index) @@ -205,8 +205,8 @@ def check_use_after_close(rectangles: list[Rectangle]) -> None: compute = rect.metadata.compute closed: dict[spir.Identifier, spir.CloseStatement] = {} for statement in compute.statements: - for kind, expression in _stream_references(statement): - name = _underlying_stream(expression) + for kind, expression in stream_references(statement): + name = underlying_stream(expression) if name in closed: what = 'closed again' if kind == 'close' else f'used in a `{kind}`' raise SyntaxError( @@ -260,7 +260,7 @@ def _transferred_elements(statement: spir.Statement, stream: spir.Identifier, cannot be determined statically. """ if isinstance(statement, (spir.SendStatement, spir.ReceiveStatement)): - if _underlying_stream(statement.stream_name) != stream: + if underlying_stream(statement.stream_name) != stream: return 0 try: dimensions = statement.get_size(identifier_sizes) @@ -272,7 +272,7 @@ def _transferred_elements(statement: spir.Statement, stream: spir.Identifier, return count if isinstance(statement, spir.ForeachStatement): - if _underlying_stream(statement.receive_stream.stream_name) != stream: + if underlying_stream(statement.receive_stream.stream_name) != stream: return None if _uses_stream(statement, stream) else 0 if not statement.parameter_range: return None # Receives until the sender is done @@ -316,7 +316,7 @@ def _transferred_elements(statement: spir.Statement, stream: spir.Identifier, def _uses_stream(statement: spir.Statement, stream: spir.Identifier) -> bool: - return any(_underlying_stream(expression) == stream for kind, expression in _stream_references(statement) + return any(underlying_stream(expression) == stream for kind, expression in stream_references(statement) if kind != 'close') @@ -603,12 +603,22 @@ def _empties_before(first: str, second: str, uses_per_rect: list[dict[spir.Ident ### -def elide_redundant_closes(rectangles: list[Rectangle]) -> int: +def elide_redundant_closes(rectangles: list[Rectangle], needs_advance=None) -> int: """ - Removes every close whose channel is never taken over by another stream, since no router has to - advance in that case. + Removes every close that no router has to act on. + + Without ``needs_advance``, a close survives only when its channel carries more than one stream + group, which is the cheapest sound approximation and is what makes kernels with ``auto`` + channels come out exactly as they did before this feature existed. + + Code generation passes the precise predicate instead: a close survives only if some router on + the stream's path actually changes configuration at that epoch boundary. That also covers a + single stream whose own configuration changes at a PE, as in a systolic forwarding chain, which + the channel-level approximation would wrongly drop. :param rectangles: The consolidated PE rectangles of the kernel, modified in place. + :param needs_advance: Optional predicate taking a ``CloseStatement`` and returning whether it + has to be kept. :return: The number of closes that were removed. """ streams_per_channel = groups_per_channel(rectangles) @@ -619,9 +629,12 @@ def elide_redundant_closes(rectangles: list[Rectangle]) -> int: keep = [] for statement in rect.metadata.compute.statements: if isinstance(statement, spir.CloseStatement): - name = _underlying_stream(statement.stream_name) - declaration = declarations.get(name) - if _is_redundant_close(declaration, streams_per_channel): + if needs_advance is not None: + redundant = not needs_advance(statement) + else: + name = underlying_stream(statement.stream_name) + redundant = _is_redundant_close(declarations.get(name), streams_per_channel) + if redundant: removed += 1 continue keep.append(statement) diff --git a/tests/spatial_ir/test_switching.py b/tests/spatial_ir/test_switching.py new file mode 100644 index 00000000..930ddce8 --- /dev/null +++ b/tests/spatial_ir/test_switching.py @@ -0,0 +1,317 @@ +""" +Tests for router switch planning: route configuration merging, switch positions, the control +wavelets that retire a configuration, and the hardware capacity limits. +""" +import os + +import pytest + +from spada.lowering import spatial_ir_to_csl as s2c +from spada.syntax.csl import constants as csl, switching as cslswitch +from spada.syntax.spatial_ir import canonicalization, parser, passes + +SAMPLES = os.path.join(os.path.dirname(__file__), 'samples') + + +def _lower(filename: str, **parameters) -> dict[str, str]: + kernel = parser.parse_file(os.path.join(SAMPLES, filename)) + if parameters: + kernel = passes.concretize_parameters(kernel, **parameters) + kernel = passes.constexpr_propagation(kernel) + return {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel)} + + +def _lower_string(code: str, **parameters) -> dict[str, str]: + kernel = parser.parse_string(code, 'test.sptl') + if parameters: + kernel = passes.concretize_parameters(kernel, **parameters) + kernel = passes.constexpr_propagation(kernel) + return {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel)} + + +def _rectangles(code: str, **parameters): + kernel = parser.parse_string(code, 'test.sptl') + if parameters: + kernel = passes.concretize_parameters(kernel, **parameters) + kernel = passes.constexpr_propagation(kernel) + kernel = s2c.canonicalize_kernel(kernel) + return canonicalization.consolidate_rectangles_to_equivalence_classes(kernel) + + +### +# RouteConfig / ColorSwitchPlan +### + + +def test_single_configuration_emits_no_switches(): + plan = cslswitch.ColorSwitchPlan() + plan.add(cslswitch.RouteConfig(('EAST', ), ('RAMP', ))) + assert plan.as_csl() == '.{ .routes = .{ .rx = .{EAST}, .tx = .{RAMP} } }' + assert not plan.uses_switches + + +def test_identical_configurations_collapse(): + """Two streams that route identically through a PE consume a single switch position.""" + plan = cslswitch.ColorSwitchPlan() + config = cslswitch.RouteConfig(('EAST', ), ('RAMP', )) + plan.add(config) + plan.add(cslswitch.RouteConfig(('EAST', ), ('RAMP', ))) + assert len(plan.positions) == 1 + assert not plan.uses_switches + + +def test_switch_positions_are_emitted(): + plan = cslswitch.ColorSwitchPlan() + plan.add(cslswitch.RouteConfig(('RAMP', ), ('WEST', ))) + plan.add(cslswitch.RouteConfig(('EAST', ), ('WEST', ))) + assert plan.as_csl() == ('.{ .routes = .{ .rx = .{RAMP}, .tx = .{WEST} }, ' + '.switches = .{ .pos1 = .{ .rx = EAST, .tx = .{WEST} } } }') + + +def test_too_many_configurations_is_rejected(): + plan = cslswitch.ColorSwitchPlan() + for direction in ('NORTH', 'SOUTH', 'EAST', 'WEST', 'RAMP'): + plan.add(cslswitch.RouteConfig((direction, ), ('RAMP', ))) + with pytest.raises(SyntaxError, match='requires 5 route configurations'): + plan.validate(0, 'PEs [2:3, 0:1]') + + +def test_exactly_four_configurations_is_accepted(): + plan = cslswitch.ColorSwitchPlan() + for direction in ('NORTH', 'SOUTH', 'EAST', 'WEST'): + plan.add(cslswitch.RouteConfig((direction, ), ('RAMP', ))) + plan.validate(0, 'PEs [2:3, 0:1]') + assert plan.as_csl().count('.pos') == 3 + + +def test_non_switchable_color_is_rejected(monkeypatch): + """WSE-3 only implements switches on a subset of colors.""" + monkeypatch.setattr(csl, 'SWITCHABLE_COLORS', [0, 1, 2]) + plan = cslswitch.ColorSwitchPlan() + plan.add(cslswitch.RouteConfig(('RAMP', ), ('WEST', ))) + plan.add(cslswitch.RouteConfig(('EAST', ), ('WEST', ))) + plan.validate(1, 'PEs [0:1, 0:1]') + with pytest.raises(SyntaxError, match='needs router switches'): + plan.validate(7, 'PEs [0:1, 0:1]') + + +### +# Control wavelets +### + + +def test_single_router_advance_uses_the_single_payload_helper(): + assert cslswitch.switch_advance_payload([True]) == \ + 'ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0)' + + +def test_routers_that_keep_their_configuration_get_a_nop(): + payload = cslswitch.switch_advance_payload([False, True]) + assert '.opcodes = .{ctrl.opcode.NOP, ctrl.opcode.SWITCH_ADV}' in payload + + +def test_path_longer_than_the_control_wavelet_is_rejected(): + commands = [True] * (csl.MAX_CONTROL_COMMANDS + 1) + with pytest.raises(SyntaxError, match='at most 8'): + cslswitch.switch_advance_payload(commands) + + +### +# two_phase_split: per-PE elision +### + + +def test_two_phase_split_switch_plans(): + """ + ``hop1`` and ``hop2`` share channel 0. Only the PEs whose configuration actually changes get a + switch position: + + * PE 0 receives both streams identically -> one configuration, no switches. + * PE 1 sends ``hop1`` and relays ``hop2`` -> two positions. The relay configuration comes from + rectangle 2's loop, so this PE also proves that configurations are merged across rectangles. + * PE 2 receives ``hop1`` and sends ``hop2`` -> two positions. + * PE 3 only sends ``hop1`` -> one configuration, no switches. + """ + files = _lower('two_phase_split.sptl', K=32) + layout = files['layout.csl'] + configs = [line.strip() for line in layout.splitlines() if '@set_color_config' in line] + assert len(configs) == 4 + + with_switches = [line for line in configs if '.switches' in line] + assert len(with_switches) == 2 + assert any('.rx = .{RAMP}, .tx = .{WEST} }, .switches = .{ .pos1 = .{ .rx = EAST, .tx = .{WEST} } }' in line + for line in with_switches), layout + assert any('.rx = .{EAST}, .tx = .{RAMP} }, .switches = .{ .pos1 = .{ .rx = RAMP, .tx = .{WEST} } }' in line + for line in with_switches), layout + + +def test_two_phase_split_emits_two_control_wavelets(): + """ + Of the six closes in the sample, only the two on PEs that *send* a stream whose path contains a + router that must advance survive elision. + """ + files = _lower('two_phase_split.sptl', K=32) + emitting = {name: code for name, code in files.items() if 'switch_dsd' in code} + assert sorted(emitting) == ['code_1_0.csl', 'code_3_0.csl'] + + # PE 1 advances its own router, PE 3 advances the router of its receiver + assert '.opcodes = .{ctrl.opcode.SWITCH_ADV, ctrl.opcode.NOP}' in emitting['code_1_0.csl'] + assert '.opcodes = .{ctrl.opcode.NOP, ctrl.opcode.SWITCH_ADV}' in emitting['code_3_0.csl'] + for code in emitting.values(): + assert 'const ctrl = @import_module("");' in code + assert '.control = true' in code + + +def test_two_phase_split_without_switching_falls_back(): + """``--disable-switching`` emits one configuration per stream, as before switches existed.""" + kernel = parser.parse_file(os.path.join(SAMPLES, 'two_phase_split.sptl')) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, K=32)) + layout = next(f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_switching=True) + if f.filename == 'layout.csl') + assert '.switches' not in layout + # One call per configuration instead of one per PE: the later call silently wins, which is the + # behaviour switches replace. + assert len([line for line in layout.splitlines() if '@set_color_config' in line]) == 6 + + +### +# Kernels that need no switching are unaffected +### + + +def test_auto_channel_kernels_emit_no_switches_and_no_control_wavelets(): + """ + Kernels from the GT4Py path use ``auto`` channels, one per stream declaration, so nothing is + ever shared and code generation is exactly as it was before this feature. + """ + code = """ + kernel @auto_channels(stream[2] readonly inp) { + place u16 i, u16 j in [0:2, 0:1] { + f32[K] a + } + dataflow u16 i, u16 j in [0:2, 0:1] { + stream eastwards = relative_stream(1, 0) { + hops = auto, + channel = auto + } + } + compute u16 i, u16 j in [0:1, 0:1] { + await receive(a, inp[i]) + await send(a, eastwards) + } + compute u16 i, u16 j in [1:2, 0:1] { + await foreach i32 k, f32 x in [0:K], receive(eastwards) { + a[k] = x + } + } + } + """ + files = _lower_string(code, K=4) + for name, code in files.items(): + assert '.switches' not in code, name + assert 'switch_dsd' not in code, name + assert '' not in code, name + + +def test_systolic_forwarding_gets_two_positions(): + """ + A PE that receives a stream and forwards it on the same channel needs two configurations, and + the upstream PE's close is what advances it onto the second. + """ + code = """ + kernel @chain(stream[3] readonly inp, stream writeonly out) { + place u16 i, u16 j in [0:3, 0:1] { + f32[K] a + } + dataflow u16 i, u16 j in [0:3, 0:1] { + stream eastwards = relative_stream(1, 0) { + hops = [(1, 0)], + channel = 0 + } + } + compute u16 i, u16 j in [0:1, 0:1] { + await receive(a, inp[i]) + await send(a, eastwards) + } + compute u16 i, u16 j in [1:2, 0:1] { + await foreach i32 k, f32 x in [0:K], receive(eastwards) { + a[k] = x + } + await send(a, eastwards) + } + compute u16 i, u16 j in [2:3, 0:1] { + await foreach i32 k, f32 x in [0:K], receive(eastwards) { + a[k] = x + } + await send(a, out) + } + } + """ + files = _lower_string(code, K=4) + layout = files['layout.csl'] + middle = [line for line in layout.splitlines() if '@set_color_config' in line and '.switches' in line] + assert len(middle) == 1, layout + assert '.rx = .{WEST}, .tx = .{RAMP} }, .switches = .{ .pos1 = .{ .rx = RAMP, .tx = .{EAST} } }' in middle[0] + + # The first PE retires its outgoing configuration so that the middle PE advances to sending + assert any('switch_dsd' in code for name, code in files.items() if name == 'code_0_0.csl') + + +### +# Capacity stress +### + +_STRESS_PHASES = [ + # (declaration, sender subgrid, receiver subgrid) -- all on channel 0, all crossing PE 2 + ('relative_stream(-1, 0) {{ hops = [(-1, 0)], channel = 0 }}', '3:4', '2:3'), + ('relative_stream(-1, 0) {{ hops = [(-1, 0)], channel = 0 }}', '2:3', '1:2'), + ('relative_stream(1, 0) {{ hops = [(1, 0)], channel = 0 }}', '1:2', '2:3'), + ('relative_stream(1, 0) {{ hops = [(1, 0)], channel = 0 }}', '2:3', '3:4'), + ('relative_stream(-2, 0) {{ hops = [(-1, 0), (-1, 0)], channel = 0 }}', '3:4', '1:2'), +] + + +def _stress_kernel(phases: int) -> str: + body = '' + for index, (declaration, sender, receiver) in enumerate(_STRESS_PHASES[:phases]): + body += f""" + phase {{ + dataflow u16 i, u16 j in [0:5, 0:1] {{ + stream s{index} = {declaration.format()} + }} + compute u16 i, u16 j in [{sender}, 0:1] {{ + await send(a, s{index}) + await s{index}.close() + }} + compute u16 i, u16 j in [{receiver}, 0:1] {{ + await foreach i32 k, f32 x in [0:K], receive(s{index}) {{ + a[k] = x + }} + await s{index}.close() + }} + }}""" + return f""" + kernel @stress(stream[5] readonly inp) {{ + place u16 i, u16 j in [0:5, 0:1] {{ + f32[K] a + }} +{body} + }} + """ + + +@pytest.mark.parametrize('phases', [2, 3, 4]) +def test_switch_positions_within_capacity(phases): + """Up to four route configurations fit in a router.""" + files = _lower_string(_stress_kernel(phases), K=4) + assert any('.switches' in code for code in files.values()) + + +def test_switch_positions_beyond_capacity_are_rejected(): + """A fifth distinct configuration on one color has nowhere to go.""" + with pytest.raises(SyntaxError, match='route configurations'): + _lower_string(_stress_kernel(5), K=4) + + +if __name__ == '__main__': + pytest.main([__file__]) From 8d1326e24d63a1c304c1eb971328e311589915a3 Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Thu, 13 Aug 2026 13:59:54 -0700 Subject: [PATCH 05/68] Refactor routing into a separate file --- spada/lowering/spatial_ir_to_csl.py | 444 +------------ spada/syntax/csl/routing.py | 590 ++++++++++++++++++ spada/syntax/csl/statements.py | 4 +- spada/syntax/csl/switching.py | 142 ----- spada/syntax/spatial_ir/irnodes.py | 2 +- .../{test_switching.py => test_routing.py} | 38 +- 6 files changed, 616 insertions(+), 604 deletions(-) create mode 100644 spada/syntax/csl/routing.py delete mode 100644 spada/syntax/csl/switching.py rename tests/spatial_ir/{test_switching.py => test_routing.py} (91%) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 79f3cac0..2d2556f6 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -4,7 +4,6 @@ from collections import defaultdict import copy -from dataclasses import dataclass import functools from io import StringIO import textwrap @@ -17,7 +16,7 @@ from spada.syntax.csl import constants as csl, preprocessing, tasks as tdag, statements as cslstmt, dsd_ops from spada.syntax.csl import benchmarking as cslbench from spada.syntax.csl import structures as cslstruct -from spada.syntax.csl import switching as cslswitch +from spada.syntax.csl import routing as cslrouting from spada.syntax.csl import task_recycling, prune_unused_fields as csl_pruning from spada.syntax.csl.codefile import CodeFile from spada.syntax.csl.statements import name_to_csl, dtype_as_csl, expr_to_csl @@ -147,7 +146,7 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, stream_lifetime.check_channel_conflicts(rectangles) # Plan the router switch advances, then drop every close no router has to act on - plan_switch_advances(rectangles) + cslrouting.plan_switch_advances(rectangles) if close_elision: stream_lifetime.elide_redundant_closes(rectangles, needs_advance=lambda stmt: bool(stmt.switch_advance)) @@ -175,7 +174,7 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, rect_size = x1 - x0, y1 - y0 # Collect unique routes for all rectangles - routes_per_rectangle = _collect_routes(rectangles, color_maps, disable_switching) + routes_per_rectangle = cslrouting.collect_routes(rectangles, color_maps, disable_switching) if use_memcpy_mode: layout_code.write(f''' @@ -345,7 +344,7 @@ def generate_rectangle(kernel: spir.Kernel, # * Generate routing instructions from dataflow blocks # * Make unique colors out of streams, reduce number of streams color_map = _allocate_colors(rect, header, kernel, use_memcpy_mode, stream_extents, channel_to_color) - _declare_switch_advances(rect, header, color_map) + cslrouting.declare_switch_advances(rect, header, color_map) dtypes = _collect_identifier_types(rect.metadata, kernel.arguments) # Preprocess potential data tasks to convert to loops if possible @@ -670,32 +669,6 @@ def _collect_colors_globally(kernel: spir.Kernel, rectangles: list[Rectangle[PEB return channel_to_color -def _declare_switch_advances(rect: Rectangle[PEBlock], header: StringIO, color_map: dict[str, int]) -> None: - """ - Declares the fabric output descriptors that carry a stream's switch-advance control wavelet. - - The wavelet itself is emitted by ``statements.generate_csl_statement`` from the close's - ``switch_advance`` field; this only has to provide the descriptor it is sent through, because - that is where the color is known. - """ - kept = [ - statement for statement in rect.metadata.compute.statements - if isinstance(statement, spir.CloseStatement) and statement.switch_advance - ] - if not kept: - return - - header.write('\nconst ctrl = @import_module("");\n') - queue = csl.OUTPUT_QUEUE_IDS[0] - for statement in kept: - name = name_to_csl(stream_lifetime.underlying_stream(statement.stream_name)) - dsd_name = f'{name}_switch_dsd' - if f'const {dsd_name}' not in header.getvalue(): - header.write(f'const {dsd_name} = @get_dsd(fabout_dsd, .{{ .extent = 1, ' - f'.fabric_color = @get_color({color_map[name + "_OUT"]}), .control = true, ' - f'.output_queue = @get_output_queue({queue}) }});\n') - - def _allocate_colors(rect: Rectangle[PEBlock], header: StringIO, kernel: spir.Kernel, use_memcpy_mode: bool, stream_extents: analysis.StreamExtents, channel_to_color: dict[int, int]) -> dict[str, int]: """ @@ -1234,415 +1207,6 @@ def visit_AssignmentStatement(self, node: spir.AssignmentStatement): return -def _route_dir(dx: int, dy: int): - """ - Helper function that returns directions for routing: (source, target). - """ - assert abs(dx + dy) == 1 - if dx == -1: - return ('EAST', 'WEST') - elif dx == 1: - return ('WEST', 'EAST') - elif dy == -1: - return ('SOUTH', 'NORTH') - elif dy == 1: - return ('NORTH', 'SOUTH') - - -@dataclass(frozen=True) -class _RouteSite: - """ - One ``@set_color_config`` target: the rectangle of PEs (already shifted by the relay offset) that - receive a route configuration for one color. - """ - color: int - x_range: tuple[int, int, int] - y_range: tuple[int, int, int] - - def as_rectangle(self) -> Rectangle: - return Rectangle(self.x_range, self.y_range, None) - - def describe(self) -> str: - return (f'PEs [{self.x_range[0]}:{self.x_range[1]}, {self.y_range[0]}:{self.y_range[1]}]') - - -@dataclass -class _RouteEntry: - """ - One route configuration contributed to a site, with the key that orders it against the other - configurations of the same site. - """ - config: cslswitch.RouteConfig - order: tuple[int, int] - origin_rect: int - origin_offset: tuple[int, int] - stream: spir.Identifier - #: Routing identity of the stream; stable across the per-rectangle renaming of ``inline_phases`` - group: str = '' - - -def _stream_use_order(compute: spir.ComputeBlock) -> dict[spir.Identifier, dict[str, tuple[int, int]]]: - """ - Returns, per stream, the order key of its first receive and of its first send in a compute block. - - The key is ``(barrier index, statement index)``: statements are ordered first by how many phase - barriers precede them, then by their position in the block. This is what sequences the route - configurations of a router into switch positions. - """ - result: dict[spir.Identifier, dict[str, tuple[int, int]]] = {} - barrier = 0 - for index, statement in enumerate(compute.statements): - if isinstance(statement, spir.AwaitAllStatement): - barrier += 1 - continue - for kind, expression in stream_lifetime.stream_references(statement): - if kind == 'close': - continue - name = stream_lifetime.underlying_stream(expression) - orders = result.setdefault(name, {}) - orders.setdefault(kind, (barrier, index)) - return result - - -def _offset_expression(axis: str, offset: int) -> str: - return axis if offset == 0 else f'{axis} + {offset}' - - -def _collect_routes(rectangles: list[Rectangle[PEBlock]], - color_maps: list[dict[str, int]], - disable_switching: bool = False) -> dict[tuple[int, int], str]: - """ - Creates a parametric version of the Routing Graph (see the Spatial IR specification for more information) and - returns a dictionary of code segements to add to the layout CSL file based on the streams. - - Route configurations are collected per *site* -- a rectangle of PEs and a color -- and merged - across rectangles, because a multi-hop stream configures its relay PEs from the sending - rectangle's loop. A site that ends up with more than one configuration is lowered to router - switch positions, ordered by the local order of the statements that use the streams. - - :param rectangles: All rectangles involved in this kernel. - :param color_maps: Per-rectangle mapping of stream names to colors. - :param disable_switching: If True, emit each configuration as its own ``@set_color_config`` - instead of merging them into switch positions. - :return: A dictionary mapping the starting point of each rectangle to a string representing the layout instructions. - """ - INDENT = 12 * ' ' - - entries: dict[_RouteSite, list[_RouteEntry]] = {} - for rect_index, (rect, color_map) in enumerate(zip(rectangles, color_maps)): - for site, entry in _rectangle_route_entries(rect_index, rect, color_map): - entries.setdefault(site, []).append(entry) - - _check_site_overlap(entries) - - result = {(rect.x_range[0], rect.y_range[0]): '' for rect in rectangles} - for site, site_entries in entries.items(): - site_entries.sort(key=lambda entry: entry.order) - - # The site is configured from the loop of one rectangle: the one that owns these PEs if - # there is one, otherwise the first relay that reaches them. - owner = min(site_entries, key=lambda entry: (entry.origin_offset != (0, 0), entry.origin_rect)) - owner_rect = rectangles[owner.origin_rect] - key = (owner_rect.x_range[0], owner_rect.y_range[0]) - x = _offset_expression('pe_x', owner.origin_offset[0]) - y = _offset_expression('pe_y', owner.origin_offset[1]) - color = f'@get_color({site.color})' - - if disable_switching: - for entry in site_entries: - plan = cslswitch.ColorSwitchPlan([entry.config]) - text = cslswitch.set_color_config(x, y, color, plan, INDENT) - if text not in result[key]: - result[key] += text - continue - - plan = cslswitch.ColorSwitchPlan() - for entry in site_entries: - plan.add(entry.config) - plan.validate(site.color, site.describe()) - result[key] += cslswitch.set_color_config(x, y, color, plan, INDENT) - - return result - - -def _route_sites(rectangles: list[Rectangle[PEBlock]], - color_maps: list[dict[str, int]]) -> dict['_RouteSite', list['_RouteEntry']]: - """ - Collects the route configurations of every site, sorted into switch-position order. - """ - entries: dict[_RouteSite, list[_RouteEntry]] = {} - for rect_index, (rect, color_map) in enumerate(zip(rectangles, color_maps)): - for site, entry in _rectangle_route_entries(rect_index, rect, color_map): - entries.setdefault(site, []).append(entry) - for site_entries in entries.values(): - site_entries.sort(key=lambda entry: entry.order) - return entries - - -def _channel_color_maps(rectangles: list[Rectangle[PEBlock]]) -> list[dict[str, int]]: - """ - Builds color maps that use the stream's *channel* in place of its color. - - Switch planning has to run before colors are allocated per rectangle, and the shape of a - router's configuration sequence only depends on the channel (a channel maps to exactly one - color). - """ - color_maps = [] - for rect in rectangles: - color_map = {} - for declaration in rect.metadata.dataflow.statements: - if declaration.stream.routing is None: - continue - channel = declaration.stream.routing.resolved_channel - if channel == 'auto': - continue - name = name_to_csl(declaration.stream_name) - color_map[name + '_IN'] = channel - color_map[name + '_OUT'] = channel - color_maps.append(color_map) - return color_maps - - -def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: - """ - Determines, for every ``close`` statement, which routers along the stream's path must advance - their switch. - - A close only produces code on a PE that *sends* the stream: the control message it emits travels - the path being retired and advances each router it traverses. Routers whose configuration does - not change are given a no-op so that they stay where they are. A close on a receiving PE emits - nothing; its router is advanced by the sender's message. - - The result is recorded on each ``CloseStatement`` as its ``switch_advance`` field; a close that - needs no advance keeps ``None`` there and generates no code. - - :param rectangles: The consolidated PE rectangles of the kernel, annotated in place. - :return: The number of closes that retire a route configuration. - """ - sites = _route_sites(rectangles, _channel_color_maps(rectangles)) - - # Per site, the switch position of every configuration, and how many positions the site holds - position_of: dict[_RouteSite, list[tuple[_RouteEntry, int]]] = {} - positions_at: dict[_RouteSite, int] = {} - for site, site_entries in sites.items(): - configs: list[cslswitch.RouteConfig] = [] - placed = [] - for entry in site_entries: - if not configs or configs[-1] != entry.config: - configs.append(entry.config) - placed.append((entry, len(configs) - 1)) - position_of[site] = placed - positions_at[site] = len(configs) - - planned = 0 - for rect in rectangles: - declarations = {d.stream_name: d for d in rect.metadata.dataflow.statements} - uses = stream_lifetime.collect_stream_uses(rect.metadata.compute) - for statement in rect.metadata.compute.statements: - if not isinstance(statement, spir.CloseStatement): - continue - statement.switch_advance = None - name = stream_lifetime.underlying_stream(statement.stream_name) - declaration = declarations.get(name) - if declaration is None or name not in uses or not uses[name].sent: - continue # Not sent here: the sending PE's control message advances this router - channel = declaration.stream.routing.resolved_channel if declaration.stream.routing else 'auto' - offsets = stream_lifetime._stream_path_offsets(declaration) - if channel == 'auto' or offsets is None: - continue - group = stream_lifetime.stream_group_key(declaration) - - commands = [] - for dx, dy in offsets: - site = _find_site(position_of, channel, - (rect.x_range[0] + dx, rect.x_range[1] + dx, rect.x_range[2]), - (rect.y_range[0] + dy, rect.y_range[1] + dy, rect.y_range[2])) - if site is None: - commands.append(False) - continue - position = _traffic_position(position_of[site], group, is_sender=(dx == 0 and dy == 0)) - commands.append(position is not None and position + 1 < positions_at.get(site, 0)) - - if any(commands): - statement.switch_advance = commands - planned += 1 - - return planned - - -def _find_site(sites, color: int, x_range: tuple[int, int, int], y_range: tuple[int, int, int]): - """ - Finds the site that configures a shifted rectangle of PEs for a color. - - An exact match is the common case, but a stream whose sender and receiver live in the *same* - rectangle shifts that rectangle onto itself: PE ``(i, j)`` sends to ``(i, j+1)``, which the same - parametric loop configures. Those lookups are resolved by intersection. - """ - exact = _RouteSite(color=color, x_range=x_range, y_range=y_range) - if exact in sites: - return exact - shifted = Rectangle(x_range, y_range, None) - for site in sites: - if site.color == color and site.as_rectangle().intersects(shifted): - return site - return None - - -def _traffic_position(placed: list[tuple['_RouteEntry', int]], group: str, is_sender: bool): - """ - Returns the switch position that a stream's traffic occupies at one router. - - The sending PE injects from its own ramp, so its configuration is the one with ``rx = RAMP``; - every router further along the path forwards traffic that arrives from the fabric. Picking the - right one matters when a PE both receives and sends the same stream, as in a systolic chain: - the message that retires the incoming configuration has to advance that router onto the - outgoing one. - """ - for entry, position in placed: - if entry.group != group: - continue - if is_sender == (entry.config.rx == ('RAMP', )): - return position - return None - - -def _check_site_overlap(entries: dict['_RouteSite', list['_RouteEntry']]) -> None: - """ - Raises a ``SyntaxError`` if two sites of the same color cover overlapping but different sets of - PEs, which cannot be expressed as a single parametric ``@set_color_config`` loop. - """ - sites = list(entries) - for index, first in enumerate(sites): - for second in sites[index + 1:]: - if first.color != second.color: - continue - if not first.as_rectangle().intersects(second.as_rectangle()): - continue - raise SyntaxError( - f'Color {first.color} is configured differently on overlapping but distinct PE ' - f'regions {first.describe()} and {second.describe()}.\n' - ' note: the two regions would need separate switch sequences; split the compute ' - 'blocks so that the regions coincide or are disjoint') - - -def _rectangle_route_entries(rect_index: int, rect: Rectangle[PEBlock], - color_map: dict[str, int]) -> list[tuple['_RouteSite', '_RouteEntry']]: - """ - Collects every route configuration a single rectangle contributes, as ``(site, entry)`` pairs. - """ - # Test whether a receive/send statement are called for creating inbound/outbound routes - sends_recvs = analysis.sends_and_receives(rect.metadata.compute) - use_order = _stream_use_order(rect.metadata.compute) - collected: list[tuple[_RouteSite, _RouteEntry]] = [] - - def add(offset: tuple[int, int], color: int, rx: tuple[str, ...], tx: tuple[str, ...], - order: tuple[int, int], stream_name: spir.Identifier, group: str) -> None: - site = _RouteSite( - color=color, - x_range=(rect.x_range[0] + offset[0], rect.x_range[1] + offset[0], rect.x_range[2]), - y_range=(rect.y_range[0] + offset[1], rect.y_range[1] + offset[1], rect.y_range[2])) - collected.append((site, - _RouteEntry(cslswitch.RouteConfig(rx, tx), order, rect_index, offset, stream_name, group))) - - # For each hop, make a color WEST-EAST/NORTH-SOUTH pair. For the first and last hop, pair with RAMP - for stream in rect.metadata.dataflow.statements: - if stream.stream_name not in sends_recvs: # Skip unused streams - continue - sent, received = sends_recvs[stream.stream_name] - group = stream_lifetime.stream_group_key(stream) - orders = use_order.get(stream.stream_name, {}) - receive_order = orders.get('receive', (0, 0)) - send_order = orders.get('send', (0, 0)) - if received: - color_inbound = color_map[name_to_csl(stream.stream_name) + "_IN"] - if sent: - color_outbound = color_map[name_to_csl(stream.stream_name) + "_OUT"] - - if isinstance(stream.stream, spir.ExternStreamDeclaration): - continue # Extern streams do not have on-chip routing - - if isinstance(stream.stream, spir.MulticastRangeStreamDeclaration): - if sent and received: - raise ValueError( - f"Multicast stream '{stream.stream_name.as_ir()}' is both sent and received " - f"within the same compute rectangle [{rect.x_range[0]}:{rect.x_range[1]}, " - f"{rect.y_range[0]}:{rect.y_range[1]}]. " - "Sender and receiver compute blocks must be in separate rectangles for multicast streams.") - if not sent: - # All multicast routing is emitted by the rectangle that sends this stream. - continue - rng = stream.stream.multicast_range - start = int(rng.start.eval()) - stop = int(rng.stop.eval()) - axis = stream.stream.multicast_axis - is_negative = start < 0 - - if axis == 'y': - tx_dir, rx_dir = ('NORTH', 'SOUTH') if is_negative else ('SOUTH', 'NORTH') - - def _coord(k): # noqa: E731 - return (0, k) - else: - tx_dir, rx_dir = ('WEST', 'EAST') if is_negative else ('EAST', 'WEST') - - def _coord(k): # noqa: E731 - return (k, 0) - - # Sender: inject into fabric toward receivers. - add((0, 0), color_outbound, ('RAMP', ), (tx_dir, ), send_order, stream.stream_name, group) - - if is_negative: - # Negative multicast: receivers at start, start-1, …, stop+1 (stop exclusive). - k_last = stop + 1 # farthest receiver - gap = range(-1, start, -1) - intermediate = range(start, k_last, -1) - else: - # Positive multicast: receivers at start, start+1, …, stop-1 (stop exclusive). - k_last = stop - 1 - gap = range(1, start) - intermediate = range(start, stop - 1) - - # Gap relay-only PEs between the sender and the first receiver. - for k in gap: - add(_coord(k), color_outbound, (rx_dir, ), (tx_dir, ), send_order, stream.stream_name, group) - - # Intermediate receivers: forward toward the farthest one and deliver to RAMP. - for k in intermediate: - add(_coord(k), color_outbound, (rx_dir, ), (tx_dir, 'RAMP'), send_order, stream.stream_name, group) - - # Last (farthest) receiver: deliver to RAMP only, no forwarding. - add(_coord(k_last), color_outbound, (rx_dir, ), ('RAMP', ), send_order, stream.stream_name, group) - continue - - if len(stream.stream.routing.hops) == 1: # Inbound and outbound generated together - route = _route_dir(*stream.stream.routing.hops[0].offset) - if sent: - add((0, 0), color_outbound, ('RAMP', ), (route[1], ), send_order, stream.stream_name, group) - if received: - add((0, 0), color_inbound, (route[0], ), ('RAMP', ), receive_order, stream.stream_name, group) - else: # Multi-hop - if sent: - first_hop = stream.stream.routing.hops[0] - add((0, 0), color_outbound, ('RAMP', ), (_route_dir(*first_hop.offset)[1], ), send_order, - stream.stream_name, group) - cur_offx = 0 - cur_offy = 0 - for hop in stream.stream.routing.hops[1:]: - route = _route_dir(*hop.offset) - cur_offx += hop.offset[0] - cur_offy += hop.offset[1] - add((cur_offx, cur_offy), color_outbound, (route[0], ), (route[1], ), send_order, - stream.stream_name, group) - if received: - # The receiver only configures itself. Intermediate PEs are configured by the sender - # block above, which walks forward through hops[1:] relative to the sender PE. - last_hop = stream.stream.routing.hops[-1] - add((0, 0), color_inbound, (_route_dir(*last_hop.offset)[0], ), ('RAMP', ), receive_order, - stream.stream_name, group) - - return collected - - def _write_indented_block(current_code: StringIO, block: str, indent: str) -> None: block = textwrap.dedent(block).strip('\n') if not block: diff --git a/spada/syntax/csl/routing.py b/spada/syntax/csl/routing.py new file mode 100644 index 00000000..e6ae2a70 --- /dev/null +++ b/spada/syntax/csl/routing.py @@ -0,0 +1,590 @@ +""" +Routing and switch planning for CSL code generation. + +A stream's routing declaration says which PEs its data traverses; this module turns that into the +``@set_color_config`` calls of ``layout.csl``. Each router holds, per color, one base route +configuration plus up to three *switch positions*, and advances between them when a switch-advance +control message passes through. Streams that share a channel, and streams that a PE both receives +and forwards, are what make a router need more than one configuration. + +The pieces, in the order they appear below: + +* :class:`RouteConfig` -- one ``.rx``/``.tx`` pair, the configuration of one router for one color. +* :class:`ColorSwitchPlan` -- the ordered configurations a router cycles through. +* :func:`set_color_config` / :func:`switch_advance_payload` -- the only places that produce + ``@set_color_config`` and ```` text. +* :func:`declare_switch_advances` -- the fabric descriptors those control messages are sent through. +* :func:`collect_routes` -- the parametric routing graph, as per-rectangle layout code. +* :func:`plan_switch_advances` -- which routers each ``close`` has to advance. + +See ``irspec/docs/spatial/routing.md`` for the semantics this implements. +""" +from dataclasses import dataclass, field +from io import StringIO + +from spada.syntax.csl import constants +from spada.syntax.csl import statements as cslstmt +from spada.syntax.spatial_ir import analysis, stream_lifetime +from spada.syntax.spatial_ir import irnodes as spir +from spada.syntax.spatial_ir.canonicalization import PEBlock +from spada.syntax.spatial_ir.grid_geometry import Rectangle + + +@dataclass(frozen=True) +class RouteConfig: + """ + The configuration of one router for one color: the directions it receives from and transmits to. + + ``RAMP`` denotes the PE's own compute element. + """ + rx: tuple[str, ...] + tx: tuple[str, ...] + + def as_csl(self) -> str: + return '.{ .rx = .{%s}, .tx = .{%s} }' % (', '.join(self.rx), ', '.join(self.tx)) + + def as_switch_position(self) -> str: + """ + Renders this configuration as a ``.posN`` struct. The ``rx`` field of a switch position only + accepts a single direction, unlike the base configuration. + """ + if len(self.rx) != 1: + raise ValueError(f'A switch position can only receive from a single direction, got {self.rx}') + return '.{ .rx = %s, .tx = .{%s} }' % (self.rx[0], ', '.join(self.tx)) + + +@dataclass +class ColorSwitchPlan: + """ + The ordered route configurations a router cycles through for one color. + + ``positions[0]`` is the base configuration, and every further entry becomes a switch position. + Consecutive identical configurations are collapsed by :meth:`add`, so a router that keeps the + same configuration across an epoch boundary consumes no switch position and needs no advance. + """ + positions: list[RouteConfig] = field(default_factory=list) + ring_mode: bool = False + + def add(self, config: RouteConfig) -> None: + if self.positions and self.positions[-1] == config: + return + self.positions.append(config) + + @property + def uses_switches(self) -> bool: + return len(self.positions) > 1 + + def index_of_epoch(self, epoch: int) -> int: + """ + Returns the switch position an epoch maps to, given that identical configurations collapse. + """ + return min(epoch, len(self.positions) - 1) + + def as_csl(self) -> str: + base = self.positions[0].as_csl() + if not self.uses_switches: + return '.{ .routes = %s }' % base + + switches = [ + '.pos%d = %s' % (index, config.as_switch_position()) + for index, config in enumerate(self.positions[1:], start=1) + ] + if self.ring_mode: + switches.append('.ring_mode = true') + return '.{ .routes = %s, .switches = .{ %s } }' % (base, ', '.join(switches)) + + def validate(self, color: int, location: str) -> None: + """ + Raises a ``SyntaxError`` if the plan exceeds what a router can hold. + + :param color: The color the plan is for, used for the diagnostic and for the WSE-3 check. + :param location: A human-readable description of the PE the plan belongs to. + """ + if len(self.positions) > constants.SWITCH_POSITIONS: + raise SyntaxError( + f'Color {color} at {location} requires {len(self.positions)} route configurations, ' + f'but a router holds at most {constants.SWITCH_POSITIONS} per color on ' + f'{constants.ARCH}.\n' + ' note: assign a different channel to some of the streams, at the cost of an ' + 'additional color') + if self.uses_switches and color not in constants.SWITCHABLE_COLORS: + raise SyntaxError( + f'Color {color} at {location} needs router switches, but {constants.ARCH} only ' + f'supports switches on colors {constants.SWITCHABLE_COLORS}.\n' + ' note: assign the channel to a switchable color') + + +def set_color_config(x: str, y: str, color: str, plan: ColorSwitchPlan, indent: str = '') -> str: + """ + Renders a single ``@set_color_config`` call. This is the only place that produces such text. + + :param x: The PE x coordinate expression (e.g. ``pe_x`` or ``pe_x + -1``). + :param y: The PE y coordinate expression. + :param color: The color expression (e.g. ``@get_color(0)``). + :param plan: The route configurations the router cycles through. + :param indent: Indentation to prefix the line with. + """ + return indent + '@set_color_config(%s, %s, %s, %s);\n' % (x, y, color, plan.as_csl()) + + +def switch_advance_payload(commands: list[bool]) -> str: + """ + Returns the ```` expression for a switch-advance control wavelet. + + One command is consumed per router the wavelet traverses, in order, so ``commands[i]`` says + whether the ``i``-th router on the path advances. Routers whose configuration does not change + are given a ``NOP`` so that they stay on their current position. + + :param commands: Per-router advance flags, starting at the sending PE's own router. + """ + if not commands: + raise ValueError('A switch advance needs at least one router command') + if len(commands) > constants.MAX_CONTROL_COMMANDS: + raise SyntaxError( + f'A switch advance along this path needs {len(commands)} router commands, but a control ' + f'wavelet carries at most {constants.MAX_CONTROL_COMMANDS}.\n' + ' note: shorten the routing path, or split it across two channels') + + if len(commands) == 1: + opcode = 'ctrl.opcode.SWITCH_ADV' if commands[0] else 'ctrl.opcode.NOP' + return f'ctrl.encode_single_payload({opcode}, true, {{}}, 0)' + + opcodes = ', '.join('ctrl.opcode.SWITCH_ADV' if advance else 'ctrl.opcode.NOP' for advance in commands) + ce_ignore = ', '.join('true' for _ in commands) + return ('ctrl.encode_payload(.{ .opcodes = .{%s}, .ce_ignore = .{%s}, ' + '.ce_ignore_remaining = true })' % (opcodes, ce_ignore)) + + +def declare_switch_advances(rect: Rectangle[PEBlock], header: StringIO, color_map: dict[str, int]) -> None: + """ + Declares the fabric output descriptors that carry a stream's switch-advance control wavelet. + + The wavelet itself is emitted by ``statements.generate_csl_statement`` from the close's + ``switch_advance`` field; this only has to provide the descriptor it is sent through, because + that is where the color is known. + """ + kept = [ + statement for statement in rect.metadata.compute.statements + if isinstance(statement, spir.CloseStatement) and statement.switch_advance + ] + if not kept: + return + + header.write('\nconst ctrl = @import_module("");\n') + queue = constants.OUTPUT_QUEUE_IDS[0] + for statement in kept: + name = cslstmt.name_to_csl(stream_lifetime.underlying_stream(statement.stream_name)) + dsd_name = f'{name}_switch_dsd' + if f'const {dsd_name}' not in header.getvalue(): + header.write(f'const {dsd_name} = @get_dsd(fabout_dsd, .{{ .extent = 1, ' + f'.fabric_color = @get_color({color_map[name + "_OUT"]}), .control = true, ' + f'.output_queue = @get_output_queue({queue}) }});\n') + + +def route_dir(dx: int, dy: int): + """ + Helper function that returns directions for routing: (source, target). + """ + assert abs(dx + dy) == 1 + if dx == -1: + return ('EAST', 'WEST') + elif dx == 1: + return ('WEST', 'EAST') + elif dy == -1: + return ('SOUTH', 'NORTH') + elif dy == 1: + return ('NORTH', 'SOUTH') + + +@dataclass(frozen=True) +class _RouteSite: + """ + One ``@set_color_config`` target: the rectangle of PEs (already shifted by the relay offset) that + receive a route configuration for one color. + """ + color: int + x_range: tuple[int, int, int] + y_range: tuple[int, int, int] + + def as_rectangle(self) -> Rectangle: + return Rectangle(self.x_range, self.y_range, None) + + def describe(self) -> str: + return (f'PEs [{self.x_range[0]}:{self.x_range[1]}, {self.y_range[0]}:{self.y_range[1]}]') + + +@dataclass +class _RouteEntry: + """ + One route configuration contributed to a site, with the key that orders it against the other + configurations of the same site. + """ + config: RouteConfig + order: tuple[int, int] + origin_rect: int + origin_offset: tuple[int, int] + stream: spir.Identifier + #: Routing identity of the stream; stable across the per-rectangle renaming of ``inline_phases`` + group: str = '' + + +def _stream_use_order(compute: spir.ComputeBlock) -> dict[spir.Identifier, dict[str, tuple[int, int]]]: + """ + Returns, per stream, the order key of its first receive and of its first send in a compute block. + + The key is ``(barrier index, statement index)``: statements are ordered first by how many phase + barriers precede them, then by their position in the block. This is what sequences the route + configurations of a router into switch positions. + """ + result: dict[spir.Identifier, dict[str, tuple[int, int]]] = {} + barrier = 0 + for index, statement in enumerate(compute.statements): + if isinstance(statement, spir.AwaitAllStatement): + barrier += 1 + continue + for kind, expression in stream_lifetime.stream_references(statement): + if kind == 'close': + continue + name = stream_lifetime.underlying_stream(expression) + orders = result.setdefault(name, {}) + orders.setdefault(kind, (barrier, index)) + return result + + +def _offset_expression(axis: str, offset: int) -> str: + return axis if offset == 0 else f'{axis} + {offset}' + + +def collect_routes(rectangles: list[Rectangle[PEBlock]], + color_maps: list[dict[str, int]], + disable_switching: bool = False) -> dict[tuple[int, int], str]: + """ + Creates a parametric version of the Routing Graph (see the Spatial IR specification for more information) and + returns a dictionary of code segements to add to the layout CSL file based on the streams. + + Route configurations are collected per *site* -- a rectangle of PEs and a color -- and merged + across rectangles, because a multi-hop stream configures its relay PEs from the sending + rectangle's loop. A site that ends up with more than one configuration is lowered to router + switch positions, ordered by the local order of the statements that use the streams. + + :param rectangles: All rectangles involved in this kernel. + :param color_maps: Per-rectangle mapping of stream names to colors. + :param disable_switching: If True, emit each configuration as its own ``@set_color_config`` + instead of merging them into switch positions. + :return: A dictionary mapping the starting point of each rectangle to a string representing the layout instructions. + """ + INDENT = 12 * ' ' + + entries: dict[_RouteSite, list[_RouteEntry]] = {} + for rect_index, (rect, color_map) in enumerate(zip(rectangles, color_maps)): + for site, entry in _rectangle_route_entries(rect_index, rect, color_map): + entries.setdefault(site, []).append(entry) + + _check_site_overlap(entries) + + result = {(rect.x_range[0], rect.y_range[0]): '' for rect in rectangles} + for site, site_entries in entries.items(): + site_entries.sort(key=lambda entry: entry.order) + + # The site is configured from the loop of one rectangle: the one that owns these PEs if + # there is one, otherwise the first relay that reaches them. + owner = min(site_entries, key=lambda entry: (entry.origin_offset != (0, 0), entry.origin_rect)) + owner_rect = rectangles[owner.origin_rect] + key = (owner_rect.x_range[0], owner_rect.y_range[0]) + x = _offset_expression('pe_x', owner.origin_offset[0]) + y = _offset_expression('pe_y', owner.origin_offset[1]) + color = f'@get_color({site.color})' + + if disable_switching: + for entry in site_entries: + plan = ColorSwitchPlan([entry.config]) + text = set_color_config(x, y, color, plan, INDENT) + if text not in result[key]: + result[key] += text + continue + + plan = ColorSwitchPlan() + for entry in site_entries: + plan.add(entry.config) + plan.validate(site.color, site.describe()) + result[key] += set_color_config(x, y, color, plan, INDENT) + + return result + + +def _route_sites(rectangles: list[Rectangle[PEBlock]], + color_maps: list[dict[str, int]]) -> dict['_RouteSite', list['_RouteEntry']]: + """ + Collects the route configurations of every site, sorted into switch-position order. + """ + entries: dict[_RouteSite, list[_RouteEntry]] = {} + for rect_index, (rect, color_map) in enumerate(zip(rectangles, color_maps)): + for site, entry in _rectangle_route_entries(rect_index, rect, color_map): + entries.setdefault(site, []).append(entry) + for site_entries in entries.values(): + site_entries.sort(key=lambda entry: entry.order) + return entries + + +def _channel_color_maps(rectangles: list[Rectangle[PEBlock]]) -> list[dict[str, int]]: + """ + Builds color maps that use the stream's *channel* in place of its color. + + Switch planning has to run before colors are allocated per rectangle, and the shape of a + router's configuration sequence only depends on the channel (a channel maps to exactly one + color). + """ + color_maps = [] + for rect in rectangles: + color_map = {} + for declaration in rect.metadata.dataflow.statements: + if declaration.stream.routing is None: + continue + channel = declaration.stream.routing.resolved_channel + if channel == 'auto': + continue + name = cslstmt.name_to_csl(declaration.stream_name) + color_map[name + '_IN'] = channel + color_map[name + '_OUT'] = channel + color_maps.append(color_map) + return color_maps + + +def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: + """ + Determines, for every ``close`` statement, which routers along the stream's path must advance + their switch. + + A close only produces code on a PE that *sends* the stream: the control message it emits travels + the path being retired and advances each router it traverses. Routers whose configuration does + not change are given a no-op so that they stay where they are. A close on a receiving PE emits + nothing; its router is advanced by the sender's message. + + The result is recorded on each ``CloseStatement`` as its ``switch_advance`` field; a close that + needs no advance keeps ``None`` there and generates no code. + + :param rectangles: The consolidated PE rectangles of the kernel, annotated in place. + :return: The number of closes that retire a route configuration. + """ + sites = _route_sites(rectangles, _channel_color_maps(rectangles)) + + # Per site, the switch position of every configuration, and how many positions the site holds + position_of: dict[_RouteSite, list[tuple[_RouteEntry, int]]] = {} + positions_at: dict[_RouteSite, int] = {} + for site, site_entries in sites.items(): + configs: list[RouteConfig] = [] + placed = [] + for entry in site_entries: + if not configs or configs[-1] != entry.config: + configs.append(entry.config) + placed.append((entry, len(configs) - 1)) + position_of[site] = placed + positions_at[site] = len(configs) + + planned = 0 + for rect in rectangles: + declarations = {d.stream_name: d for d in rect.metadata.dataflow.statements} + uses = stream_lifetime.collect_stream_uses(rect.metadata.compute) + for statement in rect.metadata.compute.statements: + if not isinstance(statement, spir.CloseStatement): + continue + statement.switch_advance = None + name = stream_lifetime.underlying_stream(statement.stream_name) + declaration = declarations.get(name) + if declaration is None or name not in uses or not uses[name].sent: + continue # Not sent here: the sending PE's control message advances this router + channel = declaration.stream.routing.resolved_channel if declaration.stream.routing else 'auto' + offsets = stream_lifetime._stream_path_offsets(declaration) + if channel == 'auto' or offsets is None: + continue + group = stream_lifetime.stream_group_key(declaration) + + commands = [] + for dx, dy in offsets: + site = _find_site(position_of, channel, + (rect.x_range[0] + dx, rect.x_range[1] + dx, rect.x_range[2]), + (rect.y_range[0] + dy, rect.y_range[1] + dy, rect.y_range[2])) + if site is None: + commands.append(False) + continue + position = _traffic_position(position_of[site], group, is_sender=(dx == 0 and dy == 0)) + commands.append(position is not None and position + 1 < positions_at.get(site, 0)) + + if any(commands): + statement.switch_advance = commands + planned += 1 + + return planned + + +def _find_site(sites, color: int, x_range: tuple[int, int, int], y_range: tuple[int, int, int]): + """ + Finds the site that configures a shifted rectangle of PEs for a color. + + An exact match is the common case, but a stream whose sender and receiver live in the *same* + rectangle shifts that rectangle onto itself: PE ``(i, j)`` sends to ``(i, j+1)``, which the same + parametric loop configures. Those lookups are resolved by intersection. + """ + exact = _RouteSite(color=color, x_range=x_range, y_range=y_range) + if exact in sites: + return exact + shifted = Rectangle(x_range, y_range, None) + for site in sites: + if site.color == color and site.as_rectangle().intersects(shifted): + return site + return None + + +def _traffic_position(placed: list[tuple['_RouteEntry', int]], group: str, is_sender: bool): + """ + Returns the switch position that a stream's traffic occupies at one router. + + The sending PE injects from its own ramp, so its configuration is the one with ``rx = RAMP``; + every router further along the path forwards traffic that arrives from the fabric. Picking the + right one matters when a PE both receives and sends the same stream, as in a systolic chain: + the message that retires the incoming configuration has to advance that router onto the + outgoing one. + """ + for entry, position in placed: + if entry.group != group: + continue + if is_sender == (entry.config.rx == ('RAMP', )): + return position + return None + + +def _check_site_overlap(entries: dict['_RouteSite', list['_RouteEntry']]) -> None: + """ + Raises a ``SyntaxError`` if two sites of the same color cover overlapping but different sets of + PEs, which cannot be expressed as a single parametric ``@set_color_config`` loop. + """ + sites = list(entries) + for index, first in enumerate(sites): + for second in sites[index + 1:]: + if first.color != second.color: + continue + if not first.as_rectangle().intersects(second.as_rectangle()): + continue + raise SyntaxError( + f'Color {first.color} is configured differently on overlapping but distinct PE ' + f'regions {first.describe()} and {second.describe()}.\n' + ' note: the two regions would need separate switch sequences; split the compute ' + 'blocks so that the regions coincide or are disjoint') + + +def _rectangle_route_entries(rect_index: int, rect: Rectangle[PEBlock], + color_map: dict[str, int]) -> list[tuple['_RouteSite', '_RouteEntry']]: + """ + Collects every route configuration a single rectangle contributes, as ``(site, entry)`` pairs. + """ + # Test whether a receive/send statement are called for creating inbound/outbound routes + sends_recvs = analysis.sends_and_receives(rect.metadata.compute) + use_order = _stream_use_order(rect.metadata.compute) + collected: list[tuple[_RouteSite, _RouteEntry]] = [] + + def add(offset: tuple[int, int], color: int, rx: tuple[str, ...], tx: tuple[str, ...], + order: tuple[int, int], stream_name: spir.Identifier, group: str) -> None: + site = _RouteSite( + color=color, + x_range=(rect.x_range[0] + offset[0], rect.x_range[1] + offset[0], rect.x_range[2]), + y_range=(rect.y_range[0] + offset[1], rect.y_range[1] + offset[1], rect.y_range[2])) + collected.append((site, + _RouteEntry(RouteConfig(rx, tx), order, rect_index, offset, stream_name, group))) + + # For each hop, make a color WEST-EAST/NORTH-SOUTH pair. For the first and last hop, pair with RAMP + for stream in rect.metadata.dataflow.statements: + if stream.stream_name not in sends_recvs: # Skip unused streams + continue + sent, received = sends_recvs[stream.stream_name] + group = stream_lifetime.stream_group_key(stream) + orders = use_order.get(stream.stream_name, {}) + receive_order = orders.get('receive', (0, 0)) + send_order = orders.get('send', (0, 0)) + if received: + color_inbound = color_map[cslstmt.name_to_csl(stream.stream_name) + "_IN"] + if sent: + color_outbound = color_map[cslstmt.name_to_csl(stream.stream_name) + "_OUT"] + + if isinstance(stream.stream, spir.ExternStreamDeclaration): + continue # Extern streams do not have on-chip routing + + if isinstance(stream.stream, spir.MulticastRangeStreamDeclaration): + if sent and received: + raise ValueError( + f"Multicast stream '{stream.stream_name.as_ir()}' is both sent and received " + f"within the same compute rectangle [{rect.x_range[0]}:{rect.x_range[1]}, " + f"{rect.y_range[0]}:{rect.y_range[1]}]. " + "Sender and receiver compute blocks must be in separate rectangles for multicast streams.") + if not sent: + # All multicast routing is emitted by the rectangle that sends this stream. + continue + rng = stream.stream.multicast_range + start = int(rng.start.eval()) + stop = int(rng.stop.eval()) + axis = stream.stream.multicast_axis + is_negative = start < 0 + + if axis == 'y': + tx_dir, rx_dir = ('NORTH', 'SOUTH') if is_negative else ('SOUTH', 'NORTH') + + def _coord(k): # noqa: E731 + return (0, k) + else: + tx_dir, rx_dir = ('WEST', 'EAST') if is_negative else ('EAST', 'WEST') + + def _coord(k): # noqa: E731 + return (k, 0) + + # Sender: inject into fabric toward receivers. + add((0, 0), color_outbound, ('RAMP', ), (tx_dir, ), send_order, stream.stream_name, group) + + if is_negative: + # Negative multicast: receivers at start, start-1, …, stop+1 (stop exclusive). + k_last = stop + 1 # farthest receiver + gap = range(-1, start, -1) + intermediate = range(start, k_last, -1) + else: + # Positive multicast: receivers at start, start+1, …, stop-1 (stop exclusive). + k_last = stop - 1 + gap = range(1, start) + intermediate = range(start, stop - 1) + + # Gap relay-only PEs between the sender and the first receiver. + for k in gap: + add(_coord(k), color_outbound, (rx_dir, ), (tx_dir, ), send_order, stream.stream_name, group) + + # Intermediate receivers: forward toward the farthest one and deliver to RAMP. + for k in intermediate: + add(_coord(k), color_outbound, (rx_dir, ), (tx_dir, 'RAMP'), send_order, stream.stream_name, group) + + # Last (farthest) receiver: deliver to RAMP only, no forwarding. + add(_coord(k_last), color_outbound, (rx_dir, ), ('RAMP', ), send_order, stream.stream_name, group) + continue + + if len(stream.stream.routing.hops) == 1: # Inbound and outbound generated together + route = route_dir(*stream.stream.routing.hops[0].offset) + if sent: + add((0, 0), color_outbound, ('RAMP', ), (route[1], ), send_order, stream.stream_name, group) + if received: + add((0, 0), color_inbound, (route[0], ), ('RAMP', ), receive_order, stream.stream_name, group) + else: # Multi-hop + if sent: + first_hop = stream.stream.routing.hops[0] + add((0, 0), color_outbound, ('RAMP', ), (route_dir(*first_hop.offset)[1], ), send_order, + stream.stream_name, group) + cur_offx = 0 + cur_offy = 0 + for hop in stream.stream.routing.hops[1:]: + route = route_dir(*hop.offset) + cur_offx += hop.offset[0] + cur_offy += hop.offset[1] + add((cur_offx, cur_offy), color_outbound, (route[0], ), (route[1], ), send_order, + stream.stream_name, group) + if received: + # The receiver only configures itself. Intermediate PEs are configured by the sender + # block above, which walks forward through hops[1:] relative to the sender PE. + last_hop = stream.stream.routing.hops[-1] + add((0, 0), color_inbound, (route_dir(*last_hop.offset)[0], ), ('RAMP', ), receive_order, + stream.stream_name, group) + + return collected diff --git a/spada/syntax/csl/statements.py b/spada/syntax/csl/statements.py index acb24974..3c460c70 100644 --- a/spada/syntax/csl/statements.py +++ b/spada/syntax/csl/statements.py @@ -2,7 +2,7 @@ from typing import Optional from spada.syntax.csl.structures import DataStructureDescriptor from spada.syntax.csl import dsd_ops -from spada.syntax.csl import switching +from spada.syntax.csl import routing from spada.syntax.spatial_ir import irnodes as spir UniqueDSDDict = dict[str, list[tuple[str, DataStructureDescriptor]]] @@ -57,7 +57,7 @@ def generate_csl_statement(statement: spir.Statement, return "" stream = statement.stream_name name = name_to_csl(stream.array if isinstance(stream, spir.ArraySlice) else stream) - return '@mov32(%s_switch_dsd, %s);' % (name, switching.switch_advance_payload(statement.switch_advance)) + return '@mov32(%s_switch_dsd, %s);' % (name, routing.switch_advance_payload(statement.switch_advance)) if op is None: return f'// TODO: Convert {statement} to CSL' diff --git a/spada/syntax/csl/switching.py b/spada/syntax/csl/switching.py deleted file mode 100644 index ea28e712..00000000 --- a/spada/syntax/csl/switching.py +++ /dev/null @@ -1,142 +0,0 @@ -""" -Router route configurations and switch planning for CSL code generation. - -Each router holds, per color, one base route configuration plus up to three *switch positions*. A -router advances from one position to the next when a switch-advance control message passes through -it, which is what a stream ``close`` lowers to. This module owns: - -* :class:`RouteConfig` -- one ``.rx``/``.tx`` pair, the configuration of one router for one color. -* :class:`ColorSwitchPlan` -- the ordered sequence of configurations a router cycles through. -* :func:`set_color_config` -- the single place where ``@set_color_config`` text is produced. -* :func:`switch_advance_payload` -- the ```` payload that advances a path's routers. - -See ``irspec/docs/spatial/routing.md`` ("Lowering to Switches") for the semantics. -""" -from dataclasses import dataclass, field - -from spada.syntax.csl import constants - - -@dataclass(frozen=True) -class RouteConfig: - """ - The configuration of one router for one color: the directions it receives from and transmits to. - - ``RAMP`` denotes the PE's own compute element. - """ - rx: tuple[str, ...] - tx: tuple[str, ...] - - def as_csl(self) -> str: - return '.{ .rx = .{%s}, .tx = .{%s} }' % (', '.join(self.rx), ', '.join(self.tx)) - - def as_switch_position(self) -> str: - """ - Renders this configuration as a ``.posN`` struct. The ``rx`` field of a switch position only - accepts a single direction, unlike the base configuration. - """ - if len(self.rx) != 1: - raise ValueError(f'A switch position can only receive from a single direction, got {self.rx}') - return '.{ .rx = %s, .tx = .{%s} }' % (self.rx[0], ', '.join(self.tx)) - - -@dataclass -class ColorSwitchPlan: - """ - The ordered route configurations a router cycles through for one color. - - ``positions[0]`` is the base configuration, and every further entry becomes a switch position. - Consecutive identical configurations are collapsed by :meth:`add`, so a router that keeps the - same configuration across an epoch boundary consumes no switch position and needs no advance. - """ - positions: list[RouteConfig] = field(default_factory=list) - ring_mode: bool = False - - def add(self, config: RouteConfig) -> None: - if self.positions and self.positions[-1] == config: - return - self.positions.append(config) - - @property - def uses_switches(self) -> bool: - return len(self.positions) > 1 - - def index_of_epoch(self, epoch: int) -> int: - """ - Returns the switch position an epoch maps to, given that identical configurations collapse. - """ - return min(epoch, len(self.positions) - 1) - - def as_csl(self) -> str: - base = self.positions[0].as_csl() - if not self.uses_switches: - return '.{ .routes = %s }' % base - - switches = [ - '.pos%d = %s' % (index, config.as_switch_position()) - for index, config in enumerate(self.positions[1:], start=1) - ] - if self.ring_mode: - switches.append('.ring_mode = true') - return '.{ .routes = %s, .switches = .{ %s } }' % (base, ', '.join(switches)) - - def validate(self, color: int, location: str) -> None: - """ - Raises a ``SyntaxError`` if the plan exceeds what a router can hold. - - :param color: The color the plan is for, used for the diagnostic and for the WSE-3 check. - :param location: A human-readable description of the PE the plan belongs to. - """ - if len(self.positions) > constants.SWITCH_POSITIONS: - raise SyntaxError( - f'Color {color} at {location} requires {len(self.positions)} route configurations, ' - f'but a router holds at most {constants.SWITCH_POSITIONS} per color on ' - f'{constants.ARCH}.\n' - ' note: assign a different channel to some of the streams, at the cost of an ' - 'additional color') - if self.uses_switches and color not in constants.SWITCHABLE_COLORS: - raise SyntaxError( - f'Color {color} at {location} needs router switches, but {constants.ARCH} only ' - f'supports switches on colors {constants.SWITCHABLE_COLORS}.\n' - ' note: assign the channel to a switchable color') - - -def set_color_config(x: str, y: str, color: str, plan: ColorSwitchPlan, indent: str = '') -> str: - """ - Renders a single ``@set_color_config`` call. This is the only place that produces such text. - - :param x: The PE x coordinate expression (e.g. ``pe_x`` or ``pe_x + -1``). - :param y: The PE y coordinate expression. - :param color: The color expression (e.g. ``@get_color(0)``). - :param plan: The route configurations the router cycles through. - :param indent: Indentation to prefix the line with. - """ - return indent + '@set_color_config(%s, %s, %s, %s);\n' % (x, y, color, plan.as_csl()) - - -def switch_advance_payload(commands: list[bool]) -> str: - """ - Returns the ```` expression for a switch-advance control wavelet. - - One command is consumed per router the wavelet traverses, in order, so ``commands[i]`` says - whether the ``i``-th router on the path advances. Routers whose configuration does not change - are given a ``NOP`` so that they stay on their current position. - - :param commands: Per-router advance flags, starting at the sending PE's own router. - """ - if not commands: - raise ValueError('A switch advance needs at least one router command') - if len(commands) > constants.MAX_CONTROL_COMMANDS: - raise SyntaxError( - f'A switch advance along this path needs {len(commands)} router commands, but a control ' - f'wavelet carries at most {constants.MAX_CONTROL_COMMANDS}.\n' - ' note: shorten the routing path, or split it across two channels') - - if len(commands) == 1: - opcode = 'ctrl.opcode.SWITCH_ADV' if commands[0] else 'ctrl.opcode.NOP' - return f'ctrl.encode_single_payload({opcode}, true, {{}}, 0)' - - opcodes = ', '.join('ctrl.opcode.SWITCH_ADV' if advance else 'ctrl.opcode.NOP' for advance in commands) - ce_ignore = ', '.join('true' for _ in commands) - return ('ctrl.encode_payload(.{ .opcodes = .{%s}, .ce_ignore = .{%s}, ' - '.ce_ignore_remaining = true })' % (opcodes, ce_ignore)) diff --git a/spada/syntax/spatial_ir/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index 276ad4fe..2e8b4c0e 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -926,7 +926,7 @@ class CloseStatement(Statement): completion_name: Optional[Completion] = None #: Which routers along the stream's path advance their switch when this close retires the #: stream's route configuration, starting at the sending PE. Filled in during lowering by - #: ``spatial_ir_to_csl.plan_switch_advances``; ``None`` means no router has to act, in which + #: ``csl.routing.plan_switch_advances``; ``None`` means no router has to act, in which #: case the close generates no code. Not part of the surface syntax. switch_advance: Optional[list[bool]] = None diff --git a/tests/spatial_ir/test_switching.py b/tests/spatial_ir/test_routing.py similarity index 91% rename from tests/spatial_ir/test_switching.py rename to tests/spatial_ir/test_routing.py index 930ddce8..88524804 100644 --- a/tests/spatial_ir/test_switching.py +++ b/tests/spatial_ir/test_routing.py @@ -7,7 +7,7 @@ import pytest from spada.lowering import spatial_ir_to_csl as s2c -from spada.syntax.csl import constants as csl, switching as cslswitch +from spada.syntax.csl import constants as csl, routing as cslrouting from spada.syntax.spatial_ir import canonicalization, parser, passes SAMPLES = os.path.join(os.path.dirname(__file__), 'samples') @@ -44,42 +44,42 @@ def _rectangles(code: str, **parameters): def test_single_configuration_emits_no_switches(): - plan = cslswitch.ColorSwitchPlan() - plan.add(cslswitch.RouteConfig(('EAST', ), ('RAMP', ))) + plan = cslrouting.ColorSwitchPlan() + plan.add(cslrouting.RouteConfig(('EAST', ), ('RAMP', ))) assert plan.as_csl() == '.{ .routes = .{ .rx = .{EAST}, .tx = .{RAMP} } }' assert not plan.uses_switches def test_identical_configurations_collapse(): """Two streams that route identically through a PE consume a single switch position.""" - plan = cslswitch.ColorSwitchPlan() - config = cslswitch.RouteConfig(('EAST', ), ('RAMP', )) + plan = cslrouting.ColorSwitchPlan() + config = cslrouting.RouteConfig(('EAST', ), ('RAMP', )) plan.add(config) - plan.add(cslswitch.RouteConfig(('EAST', ), ('RAMP', ))) + plan.add(cslrouting.RouteConfig(('EAST', ), ('RAMP', ))) assert len(plan.positions) == 1 assert not plan.uses_switches def test_switch_positions_are_emitted(): - plan = cslswitch.ColorSwitchPlan() - plan.add(cslswitch.RouteConfig(('RAMP', ), ('WEST', ))) - plan.add(cslswitch.RouteConfig(('EAST', ), ('WEST', ))) + plan = cslrouting.ColorSwitchPlan() + plan.add(cslrouting.RouteConfig(('RAMP', ), ('WEST', ))) + plan.add(cslrouting.RouteConfig(('EAST', ), ('WEST', ))) assert plan.as_csl() == ('.{ .routes = .{ .rx = .{RAMP}, .tx = .{WEST} }, ' '.switches = .{ .pos1 = .{ .rx = EAST, .tx = .{WEST} } } }') def test_too_many_configurations_is_rejected(): - plan = cslswitch.ColorSwitchPlan() + plan = cslrouting.ColorSwitchPlan() for direction in ('NORTH', 'SOUTH', 'EAST', 'WEST', 'RAMP'): - plan.add(cslswitch.RouteConfig((direction, ), ('RAMP', ))) + plan.add(cslrouting.RouteConfig((direction, ), ('RAMP', ))) with pytest.raises(SyntaxError, match='requires 5 route configurations'): plan.validate(0, 'PEs [2:3, 0:1]') def test_exactly_four_configurations_is_accepted(): - plan = cslswitch.ColorSwitchPlan() + plan = cslrouting.ColorSwitchPlan() for direction in ('NORTH', 'SOUTH', 'EAST', 'WEST'): - plan.add(cslswitch.RouteConfig((direction, ), ('RAMP', ))) + plan.add(cslrouting.RouteConfig((direction, ), ('RAMP', ))) plan.validate(0, 'PEs [2:3, 0:1]') assert plan.as_csl().count('.pos') == 3 @@ -87,9 +87,9 @@ def test_exactly_four_configurations_is_accepted(): def test_non_switchable_color_is_rejected(monkeypatch): """WSE-3 only implements switches on a subset of colors.""" monkeypatch.setattr(csl, 'SWITCHABLE_COLORS', [0, 1, 2]) - plan = cslswitch.ColorSwitchPlan() - plan.add(cslswitch.RouteConfig(('RAMP', ), ('WEST', ))) - plan.add(cslswitch.RouteConfig(('EAST', ), ('WEST', ))) + plan = cslrouting.ColorSwitchPlan() + plan.add(cslrouting.RouteConfig(('RAMP', ), ('WEST', ))) + plan.add(cslrouting.RouteConfig(('EAST', ), ('WEST', ))) plan.validate(1, 'PEs [0:1, 0:1]') with pytest.raises(SyntaxError, match='needs router switches'): plan.validate(7, 'PEs [0:1, 0:1]') @@ -101,19 +101,19 @@ def test_non_switchable_color_is_rejected(monkeypatch): def test_single_router_advance_uses_the_single_payload_helper(): - assert cslswitch.switch_advance_payload([True]) == \ + assert cslrouting.switch_advance_payload([True]) == \ 'ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0)' def test_routers_that_keep_their_configuration_get_a_nop(): - payload = cslswitch.switch_advance_payload([False, True]) + payload = cslrouting.switch_advance_payload([False, True]) assert '.opcodes = .{ctrl.opcode.NOP, ctrl.opcode.SWITCH_ADV}' in payload def test_path_longer_than_the_control_wavelet_is_rejected(): commands = [True] * (csl.MAX_CONTROL_COMMANDS + 1) with pytest.raises(SyntaxError, match='at most 8'): - cslswitch.switch_advance_payload(commands) + cslrouting.switch_advance_payload(commands) ### From 99a8e7f61149d5dca07616ccd80c4fc54d09a920 Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Thu, 13 Aug 2026 14:19:01 -0700 Subject: [PATCH 06/68] Add end-to-end tests --- .../spatial/collectives/scalar_reduce_1D.sptl | 3 +- spada/syntax/spatial_ir/stream_lifetime.py | 77 ++++++++++++------- tests/csl_runtime/test_bounded_chain.sh | 51 ++++++++++++ tests/csl_runtime/test_scalar_reduce_1d.sh | 45 +++++++++++ tests/csl_runtime/test_two_phase_switch.sh | 53 +++++++++++++ tests/spatial_ir/samples/bounded_chain.sptl | 48 ++++++++++++ tests/spatial_ir/test_routing.py | 29 +++++++ tests/spatial_ir/test_stream_lifetime.py | 2 +- 8 files changed, 279 insertions(+), 29 deletions(-) create mode 100755 tests/csl_runtime/test_bounded_chain.sh create mode 100755 tests/csl_runtime/test_scalar_reduce_1d.sh create mode 100755 tests/csl_runtime/test_two_phase_switch.sh create mode 100644 tests/spatial_ir/samples/bounded_chain.sptl diff --git a/samples/spatial/collectives/scalar_reduce_1D.sptl b/samples/spatial/collectives/scalar_reduce_1D.sptl index 4784d176..33143977 100644 --- a/samples/spatial/collectives/scalar_reduce_1D.sptl +++ b/samples/spatial/collectives/scalar_reduce_1D.sptl @@ -2,7 +2,8 @@ * Simple 1D scalar chain reduction * N is the number of PEs in the first row. * Root is 0,0 - * WARNING: Unsupported example by CSL backend -- Requires switching configuration between receive and send! + * Receiving and sending share channel 0, so every middle PE's router switches between the + * two configurations. See tests/csl_runtime/test_scalar_reduce_1d.sh **/ kernel @reduce(stream[N] readonly inp, stream writeonly out) { diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py index b5318a7f..c8ef3be0 100644 --- a/spada/syntax/spatial_ir/stream_lifetime.py +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -253,14 +253,21 @@ def _constant_range_length(rng: spir.RangeExpression) -> Optional[int]: return max(0, -(-(stop - start) // step)) -def _transferred_elements(statement: spir.Statement, stream: spir.Identifier, +def _transferred_elements(statement: spir.Statement, stream: spir.Identifier, kind: str, identifier_sizes: dict[spir.Identifier, list[int]]) -> Optional[int]: """ - Returns how many elements a top-level statement transfers over a stream, or ``None`` if that - cannot be determined statically. + Returns how many elements a top-level statement transfers over a stream in one direction, or + ``None`` if that cannot be determined statically. + + Sends and receives are counted separately because they are different stream edges: a PE in a + systolic chain receives a stream's ``BOUND`` elements from upstream and sends ``BOUND`` elements + downstream, and neither edge carries more than the bound. + + :param kind: Either ``'send'`` or ``'receive'``. """ if isinstance(statement, (spir.SendStatement, spir.ReceiveStatement)): - if underlying_stream(statement.stream_name) != stream: + wanted = spir.SendStatement if kind == 'send' else spir.ReceiveStatement + if not isinstance(statement, wanted) or underlying_stream(statement.stream_name) != stream: return 0 try: dimensions = statement.get_size(identifier_sizes) @@ -272,8 +279,15 @@ def _transferred_elements(statement: spir.Statement, stream: spir.Identifier, return count if isinstance(statement, spir.ForeachStatement): - if underlying_stream(statement.receive_stream.stream_name) != stream: - return None if _uses_stream(statement, stream) else 0 + receives_here = underlying_stream(statement.receive_stream.stream_name) == stream + if kind == 'receive' and not receives_here: + return 0 if not _uses_stream(statement, stream, kind) else None + if kind == 'send': + # A nested send repeats once per received element, which is only known when the loop + # carries an explicit range. + if not _uses_stream(statement, stream, 'send'): + return 0 + return None if not statement.parameter_range: return None # Receives until the sender is done count = 1 @@ -285,7 +299,7 @@ def _transferred_elements(statement: spir.Statement, stream: spir.Identifier, return count if isinstance(statement, (spir.ForStatement, spir.MapStatement)): - if not _uses_stream(statement, stream): + if not _uses_stream(statement, stream, kind): return 0 trips = 1 for rng in statement.range_expression: @@ -295,29 +309,30 @@ def _transferred_elements(statement: spir.Statement, stream: spir.Identifier, trips *= length inner = 0 for inner_statement in statement.body: - count = _transferred_elements(inner_statement, stream, identifier_sizes) + count = _transferred_elements(inner_statement, stream, kind, identifier_sizes) if count is None: return None inner += count return trips * inner if isinstance(statement, spir.AsyncBlock): - if not _uses_stream(statement, stream): + if not _uses_stream(statement, stream, kind): return 0 total = 0 for inner_statement in statement.body: - count = _transferred_elements(inner_statement, stream, identifier_sizes) + count = _transferred_elements(inner_statement, stream, kind, identifier_sizes) if count is None: return None total += count return total - return None if _uses_stream(statement, stream) else 0 + return None if _uses_stream(statement, stream, kind) else 0 -def _uses_stream(statement: spir.Statement, stream: spir.Identifier) -> bool: - return any(underlying_stream(expression) == stream for kind, expression in stream_references(statement) - if kind != 'close') +def _uses_stream(statement: spir.Statement, stream: spir.Identifier, kind: Optional[str] = None) -> bool: + return any(underlying_stream(expression) == stream + for reference_kind, expression in stream_references(statement) + if reference_kind != 'close' and (kind is None or reference_kind == kind)) def verify_stream_bounds(rectangles: list[Rectangle]) -> None: @@ -325,6 +340,9 @@ def verify_stream_bounds(rectangles: list[Rectangle]) -> None: Raises a ``SyntaxError`` if the number of elements transferred over a bounded stream can be determined statically and does not match the stream's bound. + Each direction is checked on its own: a PE that forwards a stream receives its bound from + upstream and sends its bound downstream, which are two stream edges of the same size. + Streams whose element count cannot be analyzed are silently accepted. :param rectangles: The consolidated PE rectangles of the kernel. @@ -343,20 +361,25 @@ def verify_stream_bounds(rectangles: list[Rectangle]) -> None: if not isinstance(bound, int): continue - total = 0 - for index in use.uses: - count = _transferred_elements(rect.metadata.compute.statements[index], name, identifier_sizes) - if count is None: - total = None - break - total += count - if total is None or total == bound: - continue + for kind in ('send', 'receive'): + if not (use.sent if kind == 'send' else use.received): + continue + total = 0 + for index in use.uses: + count = _transferred_elements(rect.metadata.compute.statements[index], name, kind, + identifier_sizes) + if count is None: + total = None + break + total += count + if total is None or total == bound: + continue - raise SyntaxError(f"Stream '{name.as_ir()}' is declared with bound {bound}, but {total} " - f"element(s) are transferred over it{_location(declaration)}.\n" - " note: the bound of a stream is the exact number of elements it carries " - "before it closes itself") + direction = 'sent over' if kind == 'send' else 'received from' + raise SyntaxError(f"Stream '{name.as_ir()}' is declared with bound {bound}, but {total} " + f"element(s) are {direction} it{_location(declaration)}.\n" + " note: the bound of a stream is the exact number of elements each of " + "its stream edges carries before it closes itself") def _identifier_sizes(place: spir.PlaceBlock) -> dict[spir.Identifier, list[int]]: diff --git a/tests/csl_runtime/test_bounded_chain.sh b/tests/csl_runtime/test_bounded_chain.sh new file mode 100755 index 00000000..366700cc --- /dev/null +++ b/tests/csl_runtime/test_bounded_chain.sh @@ -0,0 +1,51 @@ +#!/bin/sh +# E2E test: a bounded stream closing itself (bounded_chain.sptl). +# +# `stream eastwards` carries exactly K elements per stream edge, so it closes itself without +# any explicit `close` in the source. That implicit close is what advances the middle PE's router +# from its receiving configuration to its sending one -- both directions share channel 0. +# +# Reference: OUT_out == inp[0] + inp[1] (the third input row is never read). + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +. "$SCRIPT_DIR/_lib.sh" + +K=4 +FOLDER="bounded_chain_sptl" +SAMPLES_DIR="$(cd "$SCRIPT_DIR/../spatial_ir/samples" && pwd)" + +sptlc "$SAMPLES_DIR/bounded_chain.sptl" "$FOLDER" -p K=$K + +# The forwarding PE needs two configurations, and the head PE retires the first one for it +grep -q '\.switches' "$FOLDER/layout.csl" || { + echo "Test failed: no router switch configuration was generated." + exit 1 +} +grep -q 'SWITCH_ADV' "$FOLDER"/code_0_0.csl || { + echo "Test failed: the bounded stream did not emit a switch advance." + exit 1 +} + +python3 - < one configuration, no switch +# PE 1: sends hop1, then relays hop2 -> two positions, advanced by its own close +# PE 2: receives hop1, then sends hop2 -> two positions, advanced by PE 3's close +# PE 3: sends hop1 only -> one configuration, no switch +# +# Reference: OUT_out == sum over all four input rows. + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +. "$SCRIPT_DIR/_lib.sh" + +K=32 +FOLDER="two_phase_switch_sptl" +SAMPLES_DIR="$(cd "$SCRIPT_DIR/../spatial_ir/samples" && pwd)" + +sptlc "$SAMPLES_DIR/two_phase_split.sptl" "$FOLDER" -p K=$K + +# The channel must be reused through switch positions rather than a second color +grep -q '\.switches' "$FOLDER/layout.csl" || { + echo "Test failed: no router switch configuration was generated." + exit 1 +} + +python3 - <(stream[3] readonly inp, stream writeonly out) { + + place u16 i, u16 j in [0:3, 0:1] { + f32[K] a + } + + dataflow u16 i, u16 j in [0:3, 0:1] { + stream eastwards = relative_stream(1, 0) { + hops = [(1, 0)], + channel = 0 + } + } + + // Head of the chain + compute u16 i, u16 j in [0:1, 0:1] { + await receive(a, inp[i]) + await send(a, eastwards) + } + + // Middle: receive, accumulate, forward on the same stream + compute u16 i, u16 j in [1:2, 0:1] { + await receive(a, inp[i]) + await foreach i32 k, f32 x in [0:K], receive(eastwards) { + a[k] = a[k] + x + } + await send(a, eastwards) + } + + // Tail of the chain + compute u16 i, u16 j in [2:3, 0:1] { + await foreach i32 k, f32 x in [0:K], receive(eastwards) { + a[k] = x + } + await send(a, out) + } +} diff --git a/tests/spatial_ir/test_routing.py b/tests/spatial_ir/test_routing.py index 88524804..4e6af7db 100644 --- a/tests/spatial_ir/test_routing.py +++ b/tests/spatial_ir/test_routing.py @@ -257,6 +257,35 @@ def test_systolic_forwarding_gets_two_positions(): assert any('switch_dsd' in code for name, code in files.items() if name == 'code_0_0.csl') +def test_bounded_chain_sample_lowers_with_switches(): + """ + ``bounded_chain.sptl`` backs ``tests/csl_runtime/test_bounded_chain.sh``: a bounded stream that + closes itself after its bound, which is what advances the forwarding PE's router. + """ + files = _lower('bounded_chain.sptl', K=4) + layout = files['layout.csl'] + switched = [line for line in layout.splitlines() if '@set_color_config' in line and '.switches' in line] + assert len(switched) == 1, layout + assert '.rx = .{WEST}, .tx = .{RAMP} }, .switches = .{ .pos1 = .{ .rx = RAMP, .tx = .{EAST} } }' in switched[0] + + # The head of the chain retires the incoming configuration for the PE that forwards + assert 'ctrl.opcode.NOP, ctrl.opcode.SWITCH_ADV' in files['code_0_0.csl'] + assert not any('switch_dsd' in code for name, code in files.items() if name == 'code_2_0.csl') + + +def test_scalar_reduce_1d_sample_lowers_with_switches(): + """ + ``scalar_reduce_1D.sptl`` used to carry a warning that the CSL backend could not lower it: every + middle PE receives and sends on one channel. It backs + ``tests/csl_runtime/test_scalar_reduce_1d.sh``. + """ + path = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'collectives') + kernel = parser.parse_file(os.path.join(path, 'scalar_reduce_1D.sptl')) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, N=4)) + files = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel)} + assert '.switches' in files['layout.csl'] + + ### # Capacity stress ### diff --git a/tests/spatial_ir/test_stream_lifetime.py b/tests/spatial_ir/test_stream_lifetime.py index 7535a72d..0b363c6e 100644 --- a/tests/spatial_ir/test_stream_lifetime.py +++ b/tests/spatial_ir/test_stream_lifetime.py @@ -202,7 +202,7 @@ def test_matching_bound_is_accepted(): def test_mismatched_bound_is_rejected(): - with pytest.raises(SyntaxError, match='declared with bound 8, but 4 element'): + with pytest.raises(SyntaxError, match='bound 8, but 4 element'): stream_lifetime.verify_stream_bounds(_rectangles(_bounded_kernel('8', '4'), K=4)) From 586dc948b0c17f3e4d0e44a6ca6f54dae3ffab8b Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Thu, 13 Aug 2026 14:24:31 -0700 Subject: [PATCH 07/68] Type hint fixes --- spada/lowering/spatial_ir_to_csl.py | 11 ++++++---- spada/syntax/common/basenode.py | 33 +++++++++++++++++++---------- spada/syntax/spatial_ir/irnodes.py | 26 ++++++----------------- 3 files changed, 36 insertions(+), 34 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 2d2556f6..a63d060d 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -603,7 +603,7 @@ def generate_rectangle(kernel: spir.Kernel, def _collect_colors_globally(kernel: spir.Kernel, rectangles: list[Rectangle[PEBlock]], - use_memcpy_mode: bool) -> dict[str, int]: + use_memcpy_mode: bool) -> dict[int, int]: """ Returns a mapping of each channel to a CSL color. @@ -623,6 +623,7 @@ def _collect_colors_globally(kernel: spir.Kernel, rectangles: list[Rectangle[PEB for stream_decl in rect.metadata.dataflow.statements: if stream_decl.stream_name not in sends_recvs: continue # Unused stream + assert stream_decl.stream.routing is not None outbound, inbound = sends_recvs[stream_decl.stream_name] if stream_decl.stream.routing.resolved_channel == "auto": if outbound: @@ -640,6 +641,7 @@ def _collect_colors_globally(kernel: spir.Kernel, rectangles: list[Rectangle[PEB # Assign all "auto" channels for rect in rectangles: for stream_decl in rect.metadata.dataflow.statements: + assert stream_decl.stream.routing is not None if stream_decl.stream.routing.resolved_channel == "auto": stream_decl.stream.routing.channel = max_channel + 1 if stream_decl.stream_name in auto_stream_is_written: @@ -814,8 +816,8 @@ def _dsd_from_array(array_candidates: dict[str, tuple[spir.FieldDeclaration, lis extents = [str(s) if isinstance(s, int) else s.as_ir() for s in shape] # Find the index in the array - def _find_index(ind: spir.Expression) -> spir.Identifier: - candidates = [] + def _find_index(ind: spir.Expression | spir.RangeExpression) -> spir.Identifier | None: + candidates: list[spir.Identifier] = [] for n in ind.walk(): if isinstance(n, spir.Identifier): candidates.append(n) @@ -828,7 +830,8 @@ def _find_index(ind: spir.Expression) -> spir.Identifier: if isinstance(node, spir.ArraySlice): # Find and replace index with __index - idxvars = [_find_index(ind) for ind in node.indices if _find_index(ind) is not None] + idxvars = [_find_index(ind) for ind in node.indices] + idxvars = [ind for ind in idxvars if ind is not None] if use_index: assert len( idxvars) == 1, f'Expected one index variable for 1D array, got {idxvars}.\n In line {node.lineinfo}' diff --git a/spada/syntax/common/basenode.py b/spada/syntax/common/basenode.py index a07dd311..8c5254f3 100644 --- a/spada/syntax/common/basenode.py +++ b/spada/syntax/common/basenode.py @@ -9,9 +9,21 @@ from collections import deque import pprint from enum import Enum -from typing import Generic +from typing import Generic, Optional +@dataclass +class LineInfo: + """ + Represents source line information for a node in the IR. + """ + filename: str + line: int + column: int + + def __str__(self) -> str: + return f"{self.filename}:{self.line}:{self.column}" + @dataclass class BaseNode: @@ -36,7 +48,7 @@ class BaseNode: """ @classmethod - def validate_schema(cls, visited: set[type['BaseNode']] = None): + def validate_schema(cls, visited: Optional[set[type['BaseNode']]] = None): """ Validates that the node type and all its child node types abide by the rules defined on ``BaseNode``. @@ -62,7 +74,7 @@ def _check_sequence(sequence, f_name): _check_union(item, field_name) elif issubclass(item, BaseNode): item.validate_schema(visited) - elif not isinstance(item, type) or not issubclass(item, (int, float, str, type(None), Enum)): + elif not isinstance(item, type) or not issubclass(item, (int, float, str, type(None), Enum, LineInfo)): raise TypeError(f'Unsupported sequence content {item} for field {f_name} of {cls}') def _check_union(union, f_name): @@ -80,7 +92,7 @@ def _check_union(union, f_name): subtype.validate_schema(visited) # Raise error for unsupported types elif not isinstance(subtype, type) or not issubclass(subtype, - (BaseNode, int, float, str, type(None), Enum)): + (BaseNode, int, float, str, type(None), Enum, LineInfo)): raise TypeError(f'Unsupported union type {subtype} for field {f_name} of {cls}') # Use get_type_hints to resolve forward references @@ -101,7 +113,7 @@ def _check_union(union, f_name): # Check contents of sequences _check_sequence(field_type, field_name) else: - if not isinstance(field_type, type) or not issubclass(field_type, (int, float, str, type(None), Enum)): + if not isinstance(field_type, type) or not issubclass(field_type, (int, float, str, type(None), Enum, LineInfo)): raise TypeError(f'Unsupported terminator type {field_type} for field {field_name} of {cls}') return True @@ -154,7 +166,7 @@ def walk(self): (including the node itself), in breadth-first order. This function is based on ``ast.walk``. """ - todo = deque([self]) + todo: deque[BaseNode] = deque([self]) while todo: node = todo.popleft() todo.extend(node.iter_child_nodes()) @@ -203,8 +215,7 @@ def get_type(self): :return: The type restriction of the wildcard. """ - - if hasattr(self, '__orig_class__'): - return self.__orig_class__.__args__[0] - else: - return typing.Any # Fallback to Any if no type argument is provided + orig_class = getattr(self, '__orig_class__', None) + if orig_class: + return orig_class.__args__[0] + return typing.Any # Fallback to Any if no type argument is provided diff --git a/spada/syntax/spatial_ir/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index 2e8b4c0e..3744d19e 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -2,9 +2,9 @@ import copy from dataclasses import dataclass, field -from typing import Union, Tuple, Optional, Literal +from typing import Union, Optional, Literal from spada.syntax.common import visitor -from spada.syntax.common.basenode import BaseNode +from spada.syntax.common.basenode import BaseNode, LineInfo from spada.syntax.common.types import ScalarType, IRType from spada.syntax.spatial_ir.grid_geometry import Rectangle @@ -14,6 +14,7 @@ class SpatialNode(BaseNode): """ Base class for all spatial IR nodes. """ + lineinfo: Optional[LineInfo] = field(default=None, init=False, repr=False, compare=False) @classmethod def from_lark(cls, args): @@ -27,19 +28,6 @@ def as_ir(self, indent: int = 0) -> str: raise NotImplementedError() -@dataclass -class LineInfo: - """ - Represents source line information for a node in the IR. - """ - filename: str - line: int - column: int - - def __str__(self) -> str: - return f"{self.filename}:{self.line}:{self.column}" - - # Constant Literals @dataclass class ConstantLiteral(SpatialNode): @@ -378,8 +366,8 @@ class RangeExpression(SpatialNode): A range expression (start:stop or start:stop:step). """ start: Expression - stop: Expression = None - step: Expression = None + stop: Optional[Expression] = None + step: Optional[Expression] = None def validate(self) -> None: assert isinstance(self.start, Expression) @@ -390,7 +378,7 @@ def validate(self) -> None: assert isinstance(self.step, Expression) def as_ir(self, indent: int = 0) -> str: - if self.step: + if self.step and self.stop: return f'{self.start.as_ir()}:{self.stop.as_ir()}:{self.step.as_ir()}' elif self.stop: return f'{self.start.as_ir()}:{self.stop.as_ir()}' @@ -398,7 +386,7 @@ def as_ir(self, indent: int = 0) -> str: return self.start.as_ir() @staticmethod - def from_args(start: int, stop: int, step: int = None) -> 'RangeExpression': + def from_args(start: int, stop: int, step: Optional[int] = None) -> 'RangeExpression': start_expr = Expression(ConstantLiteral(start, ScalarType.i32)) stop_expr = Expression(ConstantLiteral(stop, ScalarType.i32)) step_expr = Expression(ConstantLiteral(step if step else 1, ScalarType.i32)) From 49d243d17ce8e28fb3e7d238df2f659a90b914d6 Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Thu, 13 Aug 2026 17:43:59 -0700 Subject: [PATCH 08/68] Fix switch advance for WSE2 and test on simulator --- irspec/docs/spatial/routing.md | 68 +++-- .../spatial/collectives/scalar_reduce_1D.sptl | 8 +- spada/lowering/spatial_ir_to_csl.py | 54 ++-- spada/syntax/csl/constants.py | 18 +- spada/syntax/csl/routing.py | 235 ++++++++++++++---- spada/syntax/csl/statements.py | 6 +- spada/syntax/spatial_ir/irnodes.py | 13 +- ...duce_1d.sh => pending_scalar_reduce_1d.sh} | 7 + tests/csl_runtime/test_bounded_chain.sh | 4 +- tests/csl_runtime/test_two_phase_switch.sh | 4 +- tests/spatial_ir/samples/bounded_chain.sptl | 8 +- tests/spatial_ir/samples/two_phase_split.sptl | 14 +- tests/spatial_ir/test_routing.py | 122 +++++++-- 13 files changed, 426 insertions(+), 135 deletions(-) rename tests/csl_runtime/{test_scalar_reduce_1d.sh => pending_scalar_reduce_1d.sh} (73%) diff --git a/irspec/docs/spatial/routing.md b/irspec/docs/spatial/routing.md index bea1f166..6c9b5caa 100644 --- a/irspec/docs/spatial/routing.md +++ b/irspec/docs/spatial/routing.md @@ -195,6 +195,14 @@ to receive. ## Lowering to Switches +!!! note "Note: Scope" + Everything in this section describes how the **Cerebras WSE / CSL backend** realizes epochs, not + the semantics of the Spatial IR itself. Colors, routers, switch positions and control messages + are properties of that target; a different backend may implement channel reuse by any means that + preserves the correctness conditions of [Undefined Behavior](#undefined-behavior) above. Where + the two WSE generations differ, the text says so; the compiler selects between them on + `WSE_ARCH`. + Channels are a scarce resource: each channel that is live at a PE occupies one of the hardware's routing colors. Epochs are what makes it possible to reuse a channel, and hence a color, for several streams. @@ -206,33 +214,59 @@ at $(i, j)$ by their epochs yields a sequence of route configurations $R_0, R_1, \dotsc, R_{n-1}$, which is realized by the PE's *switch* for the color assigned to $C$: $R_0$ is the initial configuration and the router *advances* to $R_{k+1}$ at the epoch boundary. +Whichever side a switch position leaves unspecified keeps the value it currently has, so positions +compose incrementally. + +!!! warning "WSE-2: A Switch Position Carries One Direction" + On WSE-2 a switch position records *either* the input the router receives from or the output it + transmits to, never both — `cslc` rejects a position naming both with *"cannot have both an + input and an output in the same switch position"*. A transition that changes both sides — a PE + that stops receiving on a channel and starts sending on it, as in a systolic chain — therefore + occupies **two** positions, passing through an intermediate configuration that keeps the old + input and takes the new output. The intermediate is a pure relay, occupied only between the two + advances that retire the configuration, and it must keep the old input so that the second + advance still reaches the router. + + WSE-3 accepts both directions in one position, so the same transition costs one position and one + advance there. + !!! danger "Error: Too Many Route Configurations" - A router holds a bounded number of route configurations per color (four on both WSE-2 and - WSE-3). *If the streams sharing a channel require more configurations than that at a single PE, - a compile error is raised.* Assigning a different channel to some of the streams resolves it, - at the cost of an additional color. + A router holds a bounded number of switch positions per color (four on both WSE-2 and WSE-3). + *If the streams sharing a channel require more positions than that at a single PE, a compile + error is raised.* Assigning a different channel to some of the streams resolves it, at the cost + of an additional color. On WSE-2 a configuration which changes both the input and the output + direction costs two positions, so four positions is fewer than four turnarounds there. + +When the sequence of configurations at a router is periodic — a halo exchange that alternates +between sending and receiving across phases produces $R_0, R_1, R_0, R_1$ — only one period is +stored and the switch wraps around from the last position back to the base one (`ring_mode`). Consecutive configurations that are equal do not consume a position and do not require an advance. This is a common case: two streams declared as `relative_stream(-2, 0)` in successive phases induce the same configuration at every PE of their paths, so their shared channel needs no switching at all. -An advance is driven by the `close` that ends the epoch, and is emitted only at those PEs whose -next configuration differs. Two lowerings are available: - -- The sending PE marks the last transfer of a bounded stream so that its router advances once the - stream's `BOUND` elements have left the fabric. This adds no traffic to the channel. -- The sending PE emits a *switch-advance control message* on the channel. It follows the stream's - path using the configuration that is being retired, and advances the router of each PE it - traverses, after all data of the epoch. PEs on the path whose configuration does not change are - skipped. +An advance is driven by the `close` that ends the epoch. The sending PE emits a *switch-advance +control message* on the channel, one per position to be traversed. It follows the stream's path +using the configuration that is being retired, and advances the router of each PE it traverses, +after all data of the epoch. + +!!! warning "WSE: Advances Are Not Selective" + A CSL control wavelet nominally carries up to eight per-router switching commands + (``'s `MAX_CMDS`), which would let one message advance some routers on a path and leave + others alone. **On the WSE hardware it does not work that way.** Measured on the simulator, only + command slot 0 is ever executed, and **every** switch-configured router the wavelet reaches + applies it; slots 1–7 had no effect in any topology tested — the sender's own router, one hop, + two hops through a plain relay, and two switch-configured routers in sequence. The compiler + therefore emits `encode_single_payload`, which writes slot 0 only. + + The consequence is that a message cannot advance one router while leaving another on the same + path where it is. *If the routers along one path would have to advance by different amounts, a + compile error is raised.* A router that is already on its last configuration is exempt: it never + routes anything again, so a message passing through may over-advance it harmlessly. Because the control message travels the path of the retired configuration in order behind the data, a receiving PE needs to emit nothing to advance its own router: the ordering required by the [lemma above](#undefined-behavior) is provided by the fabric. A receiver's `close` therefore has no runtime effect; it exists so that the lifetime of the stream — and hence the number of elements it carries — is stated by every participant and can be checked. - -!!! note "Note: Number of Control Messages" - A single control message can carry advance commands for a bounded number of consecutive routers - (eight on both WSE-2 and WSE-3). *A path that would require more raises a compile error.* diff --git a/samples/spatial/collectives/scalar_reduce_1D.sptl b/samples/spatial/collectives/scalar_reduce_1D.sptl index 33143977..1f51ab3f 100644 --- a/samples/spatial/collectives/scalar_reduce_1D.sptl +++ b/samples/spatial/collectives/scalar_reduce_1D.sptl @@ -3,7 +3,13 @@ * N is the number of PEs in the first row. * Root is 0,0 * Receiving and sending share channel 0, so every middle PE's router switches between the - * two configurations. See tests/csl_runtime/test_scalar_reduce_1d.sh + * two configurations. See tests/csl_runtime/pending_scalar_reduce_1d.sh + * + * WARNING: this sample does not compile yet, for a reason unrelated to routing. `await + * receive(rcv_val, westwards)` targets a scalar, and scalars in place blocks get no DSD, so + * `emit_copy` falls through to a plain assignment and emits `rcv_val = westwards;` -- a reference + * to an undeclared identifier. Scalar receives need to lower either to a data task or to a + * one-element DSD before this runs. The routing itself lowers correctly. **/ kernel @reduce(stream[N] readonly inp, stream writeonly out) { diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index a63d060d..22184b76 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -931,6 +931,39 @@ def _collect_unique_dsds( # TODO: Infer input/output queue ID based on concurrency input_queue_id_ctr = 0 output_queue_id_ctr = 0 + + # Streams that share a channel share a color, and a color binds to exactly one fabric queue per + # PE -- the hardware rejects "two master input queues for the same color". Queues are therefore + # handed out per channel; streams on ``auto`` channels get a color to themselves, so they key on + # their own name. + channel_of_stream = { + declaration.stream_name.as_ir(): declaration.stream.routing.resolved_channel + for declaration in rect.dataflow.statements + if getattr(declaration.stream, 'routing', None) is not None + } + + def queue_key(stream: spir.Identifier) -> str: + channel = channel_of_stream.get(stream.as_ir(), 'auto') + return stream.as_ir() if channel == 'auto' else f'channel {channel}' + + input_queue_of: dict[str, int] = {} + output_queue_of: dict[str, int] = {} + + def allocate_input_queue(stream: spir.Identifier) -> int: + nonlocal input_queue_id_ctr + key = queue_key(stream) + if key not in input_queue_of: + input_queue_of[key] = csl.INPUT_QUEUE_IDS[input_queue_id_ctr % len(csl.INPUT_QUEUE_IDS)] + input_queue_id_ctr += 1 + return input_queue_of[key] + + def allocate_output_queue(stream: spir.Identifier) -> int: + nonlocal output_queue_id_ctr + key = queue_key(stream) + if key not in output_queue_of: + output_queue_of[key] = csl.OUTPUT_QUEUE_IDS[output_queue_id_ctr % len(csl.OUTPUT_QUEUE_IDS)] + output_queue_id_ctr += 1 + return output_queue_of[key] for stmt in rect.compute.statements: # Find out if compute block uses this stream for receive/send if isinstance(stmt, (spir.ReceiveStatement, spir.SendStatement)): @@ -955,9 +988,7 @@ def _collect_unique_dsds( lambda a, b: a * b, [s.eval() if not isinstance(s, int) else s for s in dtypes[stmt.local_array].shape], 1) fabric_color = f'{name_to_csl(stream_name)}_color' - dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, - csl.INPUT_QUEUE_IDS[input_queue_id_ctr % len(csl.INPUT_QUEUE_IDS)]) - input_queue_id_ctr += 1 + dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, allocate_input_queue(stream_name)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) elif isinstance(stmt, spir.SendStatement) and stream_name.as_ir() in stream_candidates: dsd_type = cslstruct.DSDType.fabout @@ -975,9 +1006,7 @@ def _collect_unique_dsds( lambda a, b: a * b, [s.eval() if not isinstance(s, int) else s for s in dtypes[stmt.local_array].shape], 1) fabric_color = f'{name_to_csl(stream_name)}_color' - dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, - csl.OUTPUT_QUEUE_IDS[output_queue_id_ctr % len(csl.OUTPUT_QUEUE_IDS)]) - output_queue_id_ctr += 1 + dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, allocate_output_queue(stream_name)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) if isinstance(stmt, spir.SendStatement): @@ -1038,9 +1067,8 @@ def _collect_unique_dsds( extents = (end.eval() - start.eval()) // (step.eval() if step is not None else 1) fabric_color = f'{name_to_csl(stream_name)}_color' dsd = cslstruct.FabricDSD(cslstruct.DSDType.fabin, fabric_color, extents, - csl.INPUT_QUEUE_IDS[input_queue_id_ctr % len(csl.INPUT_QUEUE_IDS)]) + allocate_input_queue(stream_name)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) - input_queue_id_ctr += 1 def _visit_nested_send(substmt: spir.SendStatement): if substmt.stream_name.as_ir() not in stream_candidates: @@ -1062,10 +1090,7 @@ def _visit_nested_send(substmt: spir.SendStatement): lambda a, b: a * b, [s.eval() if not isinstance(s, int) else s for s in dtypes[substmt.local_array].shape], 1) fabric_color = f'{name_to_csl(stream_name)}_color' - nonlocal output_queue_id_ctr - dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, - csl.OUTPUT_QUEUE_IDS[output_queue_id_ctr % len(csl.OUTPUT_QUEUE_IDS)]) - output_queue_id_ctr += 1 + dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, allocate_output_queue(stream_name)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) def _visit_nested_receive(substmt: spir.ReceiveStatement): @@ -1091,10 +1116,7 @@ def _visit_nested_receive(substmt: spir.ReceiveStatement): lambda a, b: a * b, [s.eval() if not isinstance(s, int) else s for s in dtypes[local_array].shape], 1) fabric_color = f'{name_to_csl(stream_name)}_color' - nonlocal input_queue_id_ctr - dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, - csl.INPUT_QUEUE_IDS[input_queue_id_ctr % len(csl.INPUT_QUEUE_IDS)]) - input_queue_id_ctr += 1 + dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, allocate_input_queue(stream_name)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) def _visit_dsd(substmt, in_scope, in_assignment): diff --git a/spada/syntax/csl/constants.py b/spada/syntax/csl/constants.py index fccaa017..47dd270a 100644 --- a/spada/syntax/csl/constants.py +++ b/spada/syntax/csl/constants.py @@ -59,8 +59,16 @@ # See https://sdk.cerebras.ai/csl/language/builtins#switching-configuration-semantics SWITCH_POSITIONS = 4 -# Maximum number of switching commands that fit in one control wavelet (````'s MAX_CMDS). -# One command is consumed per router the wavelet traverses. +# Number of switching command slots a control wavelet carries (````'s MAX_CMDS). +# +# NOTE: only slot 0 is ever executed. Measured on the simulator, every switch-configured router a +# wavelet reaches applies the command in slot 0; slots 1-7 had no effect in any topology tested +# (the sender's own router, one hop, two hops through a plain relay, and two switch-configured +# routers in sequence). A wavelet therefore cannot advance one router while skipping another on its +# path, which is why ``routing.plan_switch_advances`` requires the routers along a path to agree. +# ````'s ``encode_payload`` also loops over all eight slots regardless of the array length +# it is given, so it must be passed exactly eight; ``encode_single_payload`` writes slot 0 only and +# is what the compiler emits. MAX_CONTROL_COMMANDS = 8 # Colors whose routers support switches. WSE-3 only implements switches on a subset of colors. @@ -69,3 +77,9 @@ 'wse3': [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 12, 13, 16, 17, 20], } SWITCHABLE_COLORS = [color for color in _SWITCHABLE_COLORS[ARCH] if color in COLORS] + +# Whether one switch position may change both the receiving and the transmitting direction. +# WSE-2 rejects it ("cannot have both an input and an output in the same switch position"), so a PE +# that receives and then sends on one color cannot be expressed there with a single advance. +_SWITCH_POSITION_ALLOWS_BOTH = {'wse2': False, 'wse3': True} +SWITCH_POSITION_ALLOWS_BOTH = _SWITCH_POSITION_ALLOWS_BOTH[ARCH] diff --git a/spada/syntax/csl/routing.py b/spada/syntax/csl/routing.py index e6ae2a70..5f3b4451 100644 --- a/spada/syntax/csl/routing.py +++ b/spada/syntax/csl/routing.py @@ -43,14 +43,88 @@ class RouteConfig: def as_csl(self) -> str: return '.{ .rx = .{%s}, .tx = .{%s} }' % (', '.join(self.rx), ', '.join(self.tx)) - def as_switch_position(self) -> str: + def as_switch_position(self, previous: 'RouteConfig') -> str: """ - Renders this configuration as a ``.posN`` struct. The ``rx`` field of a switch position only - accepts a single direction, unlike the base configuration. + Renders this configuration as a ``.posN`` struct, relative to the configuration it replaces. + + Only the side that changes is written: whatever a position leaves out keeps the value it + currently has, so positions compose incrementally. The ``rx`` of a switch position also + accepts a single direction only, unlike the base configuration. """ - if len(self.rx) != 1: - raise ValueError(f'A switch position can only receive from a single direction, got {self.rx}') - return '.{ .rx = %s, .tx = .{%s} }' % (self.rx[0], ', '.join(self.tx)) + parts = [] + if self.rx != previous.rx: + if len(self.rx) != 1: + raise ValueError(f'A switch position can only receive from a single direction, got {self.rx}') + parts.append('.rx = %s' % self.rx[0]) + if self.tx != previous.tx or not parts: + parts.append('.tx = .{%s}' % ', '.join(self.tx)) + return '.{ %s }' % ', '.join(parts) + + def changes_both_sides(self, previous: 'RouteConfig') -> bool: + return self.rx != previous.rx and self.tx != previous.tx + + +def logical_positions(configs: list[RouteConfig]) -> tuple[list[RouteConfig], bool]: + """ + Reduces a router's configuration sequence to one period, reporting whether it repeats. + + A PE that alternates between sending and receiving on one channel -- a halo exchange reusing a + channel across phases, for instance -- produces ``A, B, A, B``. Storing all four would exhaust + the router; storing ``A, B`` and letting the switch wrap around from the last position back to + the base one expresses the same thing, which is what ``.ring_mode`` is for. + + :param configs: The configurations in the order the router takes them. + :return: ``(period, ring_mode)``. + """ + count = len(configs) + for period in range(1, count): + if count % period: + continue + if all(configs[index] == configs[index % period] for index in range(count)): + return configs[:period], True + return configs, False + + +def expand_positions(configs: list[RouteConfig], ring: bool = False) -> tuple[list[RouteConfig], list[int]]: + """ + Turns a router's logical sequence of route configurations into the switch positions it holds. + + On WSE-2 a switch position carries either an input or an output, never both: the compiler + rejects ``.pos1 = .{ .rx = RAMP, .tx = .{EAST} }`` outright. A transition that changes both sides + is therefore split into two positions -- first the new output while the old input is kept, then + the new input -- and costs two switch advances instead of one. The intermediate configuration is + a pure relay, occupied only between the two wavelets that a ``close`` emits back to back. + Architectures that accept both sides in one position (see + ``constants.SWITCH_POSITION_ALLOWS_BOTH``) keep the transition as a single position. + + The intermediate keeps the *old* input direction, so a switch-advance wavelet arriving from the + same neighbour as before is still accepted once the router has taken the intermediate position; + the second wavelet would never reach the router otherwise. + + :param configs: The logical configurations, in the order the router takes them. + :param ring: Whether the router wraps from the last configuration back to the first, which may + need a trailing intermediate of its own. + :return: ``(positions, index_of)``, where ``positions`` are the hardware switch positions and + ``index_of[i]`` is the position that ``configs[i]`` ends up at. + """ + if not configs: + return [], [] + + def split(previous: RouteConfig, config: RouteConfig) -> bool: + return not constants.SWITCH_POSITION_ALLOWS_BOTH and config.changes_both_sides(previous) + + positions = [configs[0]] + index_of = [0] + for config in configs[1:]: + previous = positions[-1] + if split(previous, config): + positions.append(RouteConfig(previous.rx, config.tx)) + positions.append(config) + index_of.append(len(positions) - 1) + + if ring and len(configs) > 1 and split(positions[-1], configs[0]): + positions.append(RouteConfig(positions[-1].rx, configs[0].tx)) + return positions, index_of @dataclass @@ -71,25 +145,36 @@ def add(self, config: RouteConfig) -> None: self.positions.append(config) @property - def uses_switches(self) -> bool: - return len(self.positions) > 1 + def cycle(self) -> tuple[list[RouteConfig], bool]: + """ + The router's configurations reduced to one period, and whether the switch wraps around. + """ + configs, ring = logical_positions(self.positions) + return configs, ring or self.ring_mode - def index_of_epoch(self, epoch: int) -> int: + @property + def hardware_positions(self) -> list[RouteConfig]: """ - Returns the switch position an epoch maps to, given that identical configurations collapse. + The switch positions the router actually holds, with both-sided transitions split in two. """ - return min(epoch, len(self.positions) - 1) + configs, ring = self.cycle + return expand_positions(configs, ring)[0] + + @property + def uses_switches(self) -> bool: + return len(self.cycle[0]) > 1 def as_csl(self) -> str: base = self.positions[0].as_csl() if not self.uses_switches: return '.{ .routes = %s }' % base + hardware = self.hardware_positions switches = [ - '.pos%d = %s' % (index, config.as_switch_position()) - for index, config in enumerate(self.positions[1:], start=1) + '.pos%d = %s' % (index, config.as_switch_position(hardware[index - 1])) + for index, config in enumerate(hardware[1:], start=1) ] - if self.ring_mode: + if self.cycle[1]: switches.append('.ring_mode = true') return '.{ .routes = %s, .switches = .{ %s } }' % (base, ', '.join(switches)) @@ -100,13 +185,20 @@ def validate(self, color: int, location: str) -> None: :param color: The color the plan is for, used for the diagnostic and for the WSE-3 check. :param location: A human-readable description of the PE the plan belongs to. """ - if len(self.positions) > constants.SWITCH_POSITIONS: + hardware = self.hardware_positions + if len(hardware) > constants.SWITCH_POSITIONS: + extra = '' + if len(hardware) > len(self.positions): + extra = (f' ({len(self.positions)} route configurations, {len(hardware) - len(self.positions)} ' + 'of which change both the input and the output direction and so take two ' + 'positions each)') raise SyntaxError( - f'Color {color} at {location} requires {len(self.positions)} route configurations, ' + f'Color {color} at {location} requires {len(hardware)} switch positions{extra}, ' f'but a router holds at most {constants.SWITCH_POSITIONS} per color on ' f'{constants.ARCH}.\n' ' note: assign a different channel to some of the streams, at the cost of an ' 'additional color') + if self.uses_switches and color not in constants.SWITCHABLE_COLORS: raise SyntaxError( f'Color {color} at {location} needs router switches, but {constants.ARCH} only ' @@ -127,32 +219,33 @@ def set_color_config(x: str, y: str, color: str, plan: ColorSwitchPlan, indent: return indent + '@set_color_config(%s, %s, %s, %s);\n' % (x, y, color, plan.as_csl()) -def switch_advance_payload(commands: list[bool]) -> str: +def switch_advance_payload() -> str: """ Returns the ```` expression for a switch-advance control wavelet. - One command is consumed per router the wavelet traverses, in order, so ``commands[i]`` says - whether the ``i``-th router on the path advances. Routers whose configuration does not change - are given a ``NOP`` so that they stay on their current position. + The wavelet carries a single command, and every switch-configured router it passes through + applies it: the hardware does not index the command array by hop, so a wavelet cannot advance + one router while leaving another on the path where it is. ``ce_ignore`` keeps the wavelet from + reaching any compute element, so no task fires when it is delivered. + """ + return 'ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0)' + - :param commands: Per-router advance flags, starting at the sending PE's own router. +def switch_advance_statements(dsd_name: str, advances: int) -> str: """ - if not commands: - raise ValueError('A switch advance needs at least one router command') - if len(commands) > constants.MAX_CONTROL_COMMANDS: - raise SyntaxError( - f'A switch advance along this path needs {len(commands)} router commands, but a control ' - f'wavelet carries at most {constants.MAX_CONTROL_COMMANDS}.\n' - ' note: shorten the routing path, or split it across two channels') + Returns the statements that retire a route configuration by advancing switches ``advances`` times. - if len(commands) == 1: - opcode = 'ctrl.opcode.SWITCH_ADV' if commands[0] else 'ctrl.opcode.NOP' - return f'ctrl.encode_single_payload({opcode}, true, {{}}, 0)' + A transition that changes both the input and the output direction of a router occupies two + switch positions (see :func:`expand_positions`), and therefore needs two wavelets sent back to + back; the router is a pure relay in between. - opcodes = ', '.join('ctrl.opcode.SWITCH_ADV' if advance else 'ctrl.opcode.NOP' for advance in commands) - ce_ignore = ', '.join('true' for _ in commands) - return ('ctrl.encode_payload(.{ .opcodes = .{%s}, .ce_ignore = .{%s}, ' - '.ce_ignore_remaining = true })' % (opcodes, ce_ignore)) + :param dsd_name: The fabric output descriptor the wavelets are sent through. + :param advances: How many switch positions the routers on the path move forward. + """ + if advances <= 0: + raise ValueError(f'A switch advance must move at least one position, got {advances}') + line = '@mov32(%s, %s);' % (dsd_name, switch_advance_payload()) + return '\n'.join(line for _ in range(advances)) def declare_switch_advances(rect: Rectangle[PEBlock], header: StringIO, color_map: dict[str, int]) -> None: @@ -352,25 +445,30 @@ def _channel_color_maps(rectangles: list[Rectangle[PEBlock]]) -> list[dict[str, def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: """ - Determines, for every ``close`` statement, which routers along the stream's path must advance - their switch. + Determines, for every ``close`` statement, how many switch advances it has to emit. - A close only produces code on a PE that *sends* the stream: the control message it emits travels - the path being retired and advances each router it traverses. Routers whose configuration does - not change are given a no-op so that they stay where they are. A close on a receiving PE emits - nothing; its router is advanced by the sender's message. + A close only produces code on a PE that *sends* the stream: the control wavelets it emits travel + the path being retired. Every switch-configured router such a wavelet reaches advances -- the + hardware applies the wavelet's single command at each of them rather than indexing a per-router + command array -- so a close cannot move one router while leaving another on its path behind. + All routers on the path that hold switch positions must therefore advance by the same amount, + and that amount is how many wavelets are sent. A close on a receiving PE emits nothing; its + router is advanced by the sending PE's wavelets. The result is recorded on each ``CloseStatement`` as its ``switch_advance`` field; a close that needs no advance keeps ``None`` there and generates no code. :param rectangles: The consolidated PE rectangles of the kernel, annotated in place. :return: The number of closes that retire a route configuration. + :raises SyntaxError: If the routers along one path would have to advance by different amounts. """ sites = _route_sites(rectangles, _channel_color_maps(rectangles)) - # Per site, the switch position of every configuration, and how many positions the site holds + # Per site, where each contributed configuration sits in the router's logical sequence, and the + # hardware switch position each of those configurations maps to. position_of: dict[_RouteSite, list[tuple[_RouteEntry, int]]] = {} - positions_at: dict[_RouteSite, int] = {} + hardware_index: dict[_RouteSite, list[int]] = {} + wraps: dict[_RouteSite, tuple[bool, int]] = {} for site, site_entries in sites.items(): configs: list[RouteConfig] = [] placed = [] @@ -379,7 +477,14 @@ def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: configs.append(entry.config) placed.append((entry, len(configs) - 1)) position_of[site] = placed - positions_at[site] = len(configs) + cycle, ring = logical_positions(configs) + positions, indices = expand_positions(cycle, ring) + if len(positions) > constants.SWITCH_POSITIONS: + continue # Over capacity: ``collect_routes`` reports it, planning advances is moot + # A router that wraps around returns to its base position, so every configuration has a + # successor; otherwise the last one is final and never advances again. + hardware_index[site] = indices + wraps[site] = (ring, len(positions)) planned = 0 for rect in rectangles: @@ -392,27 +497,53 @@ def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: name = stream_lifetime.underlying_stream(statement.stream_name) declaration = declarations.get(name) if declaration is None or name not in uses or not uses[name].sent: - continue # Not sent here: the sending PE's control message advances this router + continue # Not sent here: the sending PE's wavelets advance this router channel = declaration.stream.routing.resolved_channel if declaration.stream.routing else 'auto' offsets = stream_lifetime._stream_path_offsets(declaration) if channel == 'auto' or offsets is None: continue group = stream_lifetime.stream_group_key(declaration) - commands = [] + # How far each switch-configured router on the path has to move, keyed by the router so + # that a disagreement can name it. + advances: dict[str, int] = {} for dx, dy in offsets: site = _find_site(position_of, channel, (rect.x_range[0] + dx, rect.x_range[1] + dx, rect.x_range[2]), (rect.y_range[0] + dy, rect.y_range[1] + dy, rect.y_range[2])) if site is None: - commands.append(False) continue + indices = hardware_index.get(site) + if indices is None or len(indices) < 2: + continue # One configuration only: no switch positions, nothing to advance position = _traffic_position(position_of[site], group, is_sender=(dx == 0 and dy == 0)) - commands.append(position is not None and position + 1 < positions_at.get(site, 0)) - - if any(commands): - statement.switch_advance = commands - planned += 1 + if position is None: + continue + ring, total = wraps[site] + position %= len(indices) + if position + 1 < len(indices): + advances[site.describe()] = indices[position + 1] - indices[position] + elif ring: + advances[site.describe()] = total - indices[position] + # Otherwise this router is on its last configuration for this color and never routes + # anything again, so wavelets passing through may over-advance it harmlessly. + + distinct = set(advances.values()) + if not distinct: + continue + if len(distinct) > 1: + detail = ', '.join(f'{where} by {count}' for where, count in sorted(advances.items())) + raise SyntaxError( + f"Closing '{name.as_ir()}' has to advance the routers along its path by " + f'different amounts ({detail}), but one switch-advance wavelet moves every ' + f'switch-configured router it reaches by one position.\n' + ' note: a control wavelet carries a single command that each router applies; ' + 'it cannot skip a router on its path\n' + ' note: give the streams that disagree separate channels, at the cost of an ' + 'additional color') + + statement.switch_advance = distinct.pop() + planned += 1 return planned diff --git a/spada/syntax/csl/statements.py b/spada/syntax/csl/statements.py index 3c460c70..c3c7a8d2 100644 --- a/spada/syntax/csl/statements.py +++ b/spada/syntax/csl/statements.py @@ -51,13 +51,13 @@ def generate_csl_statement(statement: spir.Statement, # Skip (taken care of when tasks are defined) return "" elif isinstance(statement, spir.CloseStatement): - # Retiring a route configuration is a control wavelet that advances the routers along the - # stream's path, or nothing at all when none of them has to move. + # Retiring a route configuration means advancing the switches along the stream's path, one + # control wavelet per position, or nothing at all when no router has to move. if not statement.switch_advance: return "" stream = statement.stream_name name = name_to_csl(stream.array if isinstance(stream, spir.ArraySlice) else stream) - return '@mov32(%s_switch_dsd, %s);' % (name, routing.switch_advance_payload(statement.switch_advance)) + return routing.switch_advance_statements(f'{name}_switch_dsd', statement.switch_advance) if op is None: return f'// TODO: Convert {statement} to CSL' diff --git a/spada/syntax/spatial_ir/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index 3744d19e..88d719f8 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -912,18 +912,19 @@ class CloseStatement(Statement): """ stream_name: Union[Identifier, ArraySlice] completion_name: Optional[Completion] = None - #: Which routers along the stream's path advance their switch when this close retires the - #: stream's route configuration, starting at the sending PE. Filled in during lowering by - #: ``csl.routing.plan_switch_advances``; ``None`` means no router has to act, in which - #: case the close generates no code. Not part of the surface syntax. - switch_advance: Optional[list[bool]] = None + #: How many switch positions the routers along the stream's path move forward when this close + #: retires the stream's route configuration. One wavelet is emitted per position, and a + #: transition that changes both a router's input and its output direction takes two. Filled in + #: during lowering by ``csl.routing.plan_switch_advances``; ``None`` means no router has to act, + #: in which case the close generates no code. Not part of the surface syntax. + switch_advance: Optional[int] = None def validate(self) -> None: assert isinstance(self.stream_name, (Identifier, ArraySlice)) if self.completion_name: assert isinstance(self.completion_name, Completion) if self.switch_advance is not None: - assert all(isinstance(command, bool) for command in self.switch_advance) + assert isinstance(self.switch_advance, int) and self.switch_advance > 0 def as_ir(self, indent: int = 0) -> str: indent_str = ' ' * indent diff --git a/tests/csl_runtime/test_scalar_reduce_1d.sh b/tests/csl_runtime/pending_scalar_reduce_1d.sh similarity index 73% rename from tests/csl_runtime/test_scalar_reduce_1d.sh rename to tests/csl_runtime/pending_scalar_reduce_1d.sh index b92169dc..54b8e4f7 100755 --- a/tests/csl_runtime/test_scalar_reduce_1d.sh +++ b/tests/csl_runtime/pending_scalar_reduce_1d.sh @@ -1,4 +1,11 @@ #!/bin/sh +# NOT RUN YET (named "pending_" so run_tests.sh does not collect it). +# +# The routing this exercises lowers correctly -- every middle PE turns its router around on one +# channel, which now compiles and runs. What blocks the sample is unrelated: `await +# receive(rcv_val, westwards)` targets a scalar, and scalars in place blocks get no DSD, so +# emit_copy falls through to a plain assignment and emits `rcv_val = westwards;`, which cslc +# rejects as an undeclared identifier. Rename back to test_*.sh once scalar receives lower. # E2E test: 1D scalar chain reduction over a single channel (scalar_reduce_1D.sptl). # # Every PE but the last receives a partial sum from the east and sends the accumulated value west, diff --git a/tests/csl_runtime/test_bounded_chain.sh b/tests/csl_runtime/test_bounded_chain.sh index 366700cc..f42c22a7 100755 --- a/tests/csl_runtime/test_bounded_chain.sh +++ b/tests/csl_runtime/test_bounded_chain.sh @@ -29,7 +29,7 @@ grep -q 'SWITCH_ADV' "$FOLDER"/code_0_0.csl || { python3 - <(stream[3] readonly inp, stream writeonly out) { +kernel @bounded_chain(stream[3,1] readonly inp, stream[1,1] writeonly out) { place u16 i, u16 j in [0:3, 0:1] { f32[K] a @@ -25,13 +25,13 @@ kernel @bounded_chain(stream[3] readonly inp, stream writeonl // Head of the chain compute u16 i, u16 j in [0:1, 0:1] { - await receive(a, inp[i]) + await receive(a, inp[i, j]) await send(a, eastwards) } // Middle: receive, accumulate, forward on the same stream compute u16 i, u16 j in [1:2, 0:1] { - await receive(a, inp[i]) + await receive(a, inp[i, j]) await foreach i32 k, f32 x in [0:K], receive(eastwards) { a[k] = a[k] + x } @@ -43,6 +43,6 @@ kernel @bounded_chain(stream[3] readonly inp, stream writeonl await foreach i32 k, f32 x in [0:K], receive(eastwards) { a[k] = x } - await send(a, out) + await send(a, out[i, j]) } } diff --git a/tests/spatial_ir/samples/two_phase_split.sptl b/tests/spatial_ir/samples/two_phase_split.sptl index c950c154..723e2958 100644 --- a/tests/spatial_ir/samples/two_phase_split.sptl +++ b/tests/spatial_ir/samples/two_phase_split.sptl @@ -1,5 +1,5 @@ -kernel @two_phase (stream[4] readonly in, - stream writeonly out) { +kernel @two_phase (stream[4,1] readonly in, + stream[1,1] writeonly out) { place i16 i, i16 j in [0, 0] { f32[K] a @@ -51,7 +51,7 @@ kernel @two_phase (stream[4] readonly in, } } compute i32 i, i32 j in [0, 0] { - await receive(a, in[i]) + await receive(a, in[i, j]) await foreach i32 k, f32 x in [0:K], receive(hop1) { a[k] = a[k] + x } @@ -59,15 +59,15 @@ kernel @two_phase (stream[4] readonly in, await foreach i32 k, f32 x in [0:K], receive(hop2) { a[k] = a[k] + x } - await send(a, out) + await send(a, out[i, j]) } compute i32 i, i32 j in [1, 0] { - await receive(a, in[i]) + await receive(a, in[i, j]) await send(a, hop1) await hop1.close() } compute i32 i, i32 j in [2, 0] { - await receive(a, in[i]) + await receive(a, in[i, j]) await foreach i32 k, f32 x in [0:K], receive(hop1) { a[k] = a[k] + x } @@ -75,7 +75,7 @@ kernel @two_phase (stream[4] readonly in, await send(a, hop2) } compute i32 i, i32 j in [3, 0] { - await receive(a, in[i]) + await receive(a, in[i, j]) await send(a, hop1) await hop1.close() } diff --git a/tests/spatial_ir/test_routing.py b/tests/spatial_ir/test_routing.py index 4e6af7db..44c27840 100644 --- a/tests/spatial_ir/test_routing.py +++ b/tests/spatial_ir/test_routing.py @@ -38,6 +38,18 @@ def _rectangles(code: str, **parameters): return canonicalization.consolidate_rectangles_to_equivalence_classes(kernel) +def _turnaround(direction: str) -> str: + """ + The ``.switches`` text for a router that stops receiving and starts sending in ``direction``. + + Changing both the input and the output direction takes two switch positions wherever a position + carries only one of them. + """ + if csl.SWITCH_POSITION_ALLOWS_BOTH: + return '.pos1 = .{ .rx = RAMP, .tx = .{%s} }' % direction + return '.pos1 = .{ .tx = .{%s} }, .pos2 = .{ .rx = RAMP }' % direction + + ### # RouteConfig / ColorSwitchPlan ### @@ -64,18 +76,58 @@ def test_switch_positions_are_emitted(): plan = cslrouting.ColorSwitchPlan() plan.add(cslrouting.RouteConfig(('RAMP', ), ('WEST', ))) plan.add(cslrouting.RouteConfig(('EAST', ), ('WEST', ))) + # Only the side that changes is written: a switch position carries an input or an output assert plan.as_csl() == ('.{ .routes = .{ .rx = .{RAMP}, .tx = .{WEST} }, ' - '.switches = .{ .pos1 = .{ .rx = EAST, .tx = .{WEST} } } }') + '.switches = .{ .pos1 = .{ .rx = EAST } } }') def test_too_many_configurations_is_rejected(): plan = cslrouting.ColorSwitchPlan() for direction in ('NORTH', 'SOUTH', 'EAST', 'WEST', 'RAMP'): plan.add(cslrouting.RouteConfig((direction, ), ('RAMP', ))) - with pytest.raises(SyntaxError, match='requires 5 route configurations'): + with pytest.raises(SyntaxError, match='requires 5 switch positions'): plan.validate(0, 'PEs [2:3, 0:1]') +def test_both_sided_transition_becomes_two_positions(): + """ + A router that changes its input *and* its output direction cannot express that as one switch + position on WSE-2, so it goes through an intermediate pure-relay configuration. + """ + receive = cslrouting.RouteConfig(('EAST', ), ('RAMP', )) + send = cslrouting.RouteConfig(('RAMP', ), ('WEST', )) + positions, index_of = cslrouting.expand_positions([receive, send]) + + if csl.SWITCH_POSITION_ALLOWS_BOTH: + assert positions == [receive, send] + assert index_of == [0, 1] + else: + assert positions == [receive, cslrouting.RouteConfig(('EAST', ), ('WEST', )), send] + assert index_of == [0, 2] + + plan = cslrouting.ColorSwitchPlan() + plan.add(receive) + plan.add(send) + # Each position names only the side it changes, and they compose incrementally. + assert plan.as_csl() == ('.{ .routes = .{ .rx = .{EAST}, .tx = .{RAMP} }, ' + '.switches = .{ .pos1 = .{ .tx = .{WEST} }, .pos2 = .{ .rx = RAMP } } }') + + +def test_both_sided_transition_counts_against_capacity(): + """Two both-sided transitions take four positions, which exactly fills a router.""" + plan = cslrouting.ColorSwitchPlan() + plan.add(cslrouting.RouteConfig(('EAST', ), ('RAMP', ))) + plan.add(cslrouting.RouteConfig(('RAMP', ), ('WEST', ))) + plan.add(cslrouting.RouteConfig(('NORTH', ), ('RAMP', ))) + + if csl.SWITCH_POSITION_ALLOWS_BOTH: + plan.validate(0, 'PEs [1:2, 0:1]') + else: + assert len(plan.hardware_positions) == 5 + with pytest.raises(SyntaxError, match='requires 5 switch positions'): + plan.validate(0, 'PEs [1:2, 0:1]') + + def test_exactly_four_configurations_is_accepted(): plan = cslrouting.ColorSwitchPlan() for direction in ('NORTH', 'SOUTH', 'EAST', 'WEST'): @@ -100,20 +152,29 @@ def test_non_switchable_color_is_rejected(monkeypatch): ### -def test_single_router_advance_uses_the_single_payload_helper(): - assert cslrouting.switch_advance_payload([True]) == \ +def test_switch_advance_uses_the_single_command_payload(): + """ + The hardware applies a control wavelet's single command at every switch-configured router it + reaches, so there is nothing to index per router. + """ + assert cslrouting.switch_advance_payload() == \ 'ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0)' -def test_routers_that_keep_their_configuration_get_a_nop(): - payload = cslrouting.switch_advance_payload([False, True]) - assert '.opcodes = .{ctrl.opcode.NOP, ctrl.opcode.SWITCH_ADV}' in payload +def test_one_advance_emits_one_wavelet(): + assert cslrouting.switch_advance_statements('s_switch_dsd', 1) == \ + '@mov32(s_switch_dsd, ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0));' -def test_path_longer_than_the_control_wavelet_is_rejected(): - commands = [True] * (csl.MAX_CONTROL_COMMANDS + 1) - with pytest.raises(SyntaxError, match='at most 8'): - cslrouting.switch_advance_payload(commands) +def test_both_sided_transition_emits_two_wavelets(): + text = cslrouting.switch_advance_statements('s_switch_dsd', 2) + assert text.count('@mov32(s_switch_dsd,') == 2 + assert text.count('\n') == 1 + + +def test_empty_advance_is_rejected(): + with pytest.raises(ValueError, match='at least one position'): + cslrouting.switch_advance_statements('s_switch_dsd', 0) ### @@ -139,9 +200,15 @@ def test_two_phase_split_switch_plans(): with_switches = [line for line in configs if '.switches' in line] assert len(with_switches) == 2 - assert any('.rx = .{RAMP}, .tx = .{WEST} }, .switches = .{ .pos1 = .{ .rx = EAST, .tx = .{WEST} } }' in line + # PE 1 only changes where it receives from, which is one position on every architecture. + assert any('.rx = .{RAMP}, .tx = .{WEST} }, .switches = .{ .pos1 = .{ .rx = EAST } }' in line for line in with_switches), layout - assert any('.rx = .{EAST}, .tx = .{RAMP} }, .switches = .{ .pos1 = .{ .rx = RAMP, .tx = .{WEST} } }' in line + # PE 2 turns around from receiving to sending, changing both sides at once. + if csl.SWITCH_POSITION_ALLOWS_BOTH: + turnaround = '.switches = .{ .pos1 = .{ .rx = RAMP, .tx = .{WEST} } }' + else: + turnaround = '.switches = .{ .pos1 = .{ .tx = .{WEST} }, .pos2 = .{ .rx = RAMP } }' + assert any('.rx = .{EAST}, .tx = .{RAMP} }, ' + turnaround in line for line in with_switches), layout @@ -154,9 +221,11 @@ def test_two_phase_split_emits_two_control_wavelets(): emitting = {name: code for name, code in files.items() if 'switch_dsd' in code} assert sorted(emitting) == ['code_1_0.csl', 'code_3_0.csl'] - # PE 1 advances its own router, PE 3 advances the router of its receiver - assert '.opcodes = .{ctrl.opcode.SWITCH_ADV, ctrl.opcode.NOP}' in emitting['code_1_0.csl'] - assert '.opcodes = .{ctrl.opcode.NOP, ctrl.opcode.SWITCH_ADV}' in emitting['code_3_0.csl'] + payload = 'ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0)' + # PE 1 moves its own router one position; PE 3 retires PE 2's incoming configuration, and PE 2 + # has to turn around, which takes two positions where a switch carries only one direction. + assert emitting['code_1_0.csl'].count(payload) == 1 + assert emitting['code_3_0.csl'].count(payload) == (1 if csl.SWITCH_POSITION_ALLOWS_BOTH else 2) for code in emitting.values(): assert 'const ctrl = @import_module("");' in code assert '.control = true' in code @@ -251,7 +320,7 @@ def test_systolic_forwarding_gets_two_positions(): layout = files['layout.csl'] middle = [line for line in layout.splitlines() if '@set_color_config' in line and '.switches' in line] assert len(middle) == 1, layout - assert '.rx = .{WEST}, .tx = .{RAMP} }, .switches = .{ .pos1 = .{ .rx = RAMP, .tx = .{EAST} } }' in middle[0] + assert _turnaround('EAST') in middle[0], middle[0] # The first PE retires its outgoing configuration so that the middle PE advances to sending assert any('switch_dsd' in code for name, code in files.items() if name == 'code_0_0.csl') @@ -266,10 +335,10 @@ def test_bounded_chain_sample_lowers_with_switches(): layout = files['layout.csl'] switched = [line for line in layout.splitlines() if '@set_color_config' in line and '.switches' in line] assert len(switched) == 1, layout - assert '.rx = .{WEST}, .tx = .{RAMP} }, .switches = .{ .pos1 = .{ .rx = RAMP, .tx = .{EAST} } }' in switched[0] + assert _turnaround('EAST') in switched[0], switched[0] # The head of the chain retires the incoming configuration for the PE that forwards - assert 'ctrl.opcode.NOP, ctrl.opcode.SWITCH_ADV' in files['code_0_0.csl'] + assert 'ctrl.opcode.SWITCH_ADV' in files['code_0_0.csl'] assert not any('switch_dsd' in code for name, code in files.items() if name == 'code_2_0.csl') @@ -329,17 +398,24 @@ def _stress_kernel(phases: int) -> str: """ +# Every phase of the stress kernel turns PE 2 around between receiving and sending, so on an +# architecture where a switch position carries a single direction each phase costs two positions. +_MAX_STRESS_PHASES = 4 if csl.SWITCH_POSITION_ALLOWS_BOTH else 2 + + @pytest.mark.parametrize('phases', [2, 3, 4]) def test_switch_positions_within_capacity(phases): - """Up to four route configurations fit in a router.""" + """A router holds four switch positions; how many phases that is depends on the architecture.""" + if phases > _MAX_STRESS_PHASES: + pytest.skip(f'{csl.ARCH} fits at most {_MAX_STRESS_PHASES} turnarounds in one router') files = _lower_string(_stress_kernel(phases), K=4) assert any('.switches' in code for code in files.values()) def test_switch_positions_beyond_capacity_are_rejected(): - """A fifth distinct configuration on one color has nowhere to go.""" - with pytest.raises(SyntaxError, match='route configurations'): - _lower_string(_stress_kernel(5), K=4) + """One configuration too many on a color has nowhere to go.""" + with pytest.raises(SyntaxError, match='switch positions'): + _lower_string(_stress_kernel(_MAX_STRESS_PHASES + 1), K=4) if __name__ == '__main__': From 505b88f47b1d2377c1327e48067745327e9e0a21 Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Fri, 14 Aug 2026 00:16:45 -0700 Subject: [PATCH 09/68] 1D bitonic sorting network sample and inter-phase stream numbering passes --- samples/spatial/sorting/bitonic_sort_1D.sptl | 158 +++ spada/lowering/spatial_ir_to_csl.py | 47 +- spada/syntax/csl/constants.py | 6 +- spada/syntax/csl/dsd_ops.py | 2 +- spada/syntax/csl/routing.py | 52 +- spada/syntax/csl/statements.py | 4 +- spada/syntax/csl/structures.py | 5 + spada/syntax/csl/tasks.py | 1128 +++++++++--------- spada/syntax/spatial_ir/canonicalization.py | 91 ++ spada/syntax/spatial_ir/irnodes.py | 10 + spada/syntax/spatial_ir/stream_lifetime.py | 87 +- tests/csl_runtime/test_bitonic_sort_1d.sh | 77 ++ tests/spatial_ir/test_routing.py | 51 + 13 files changed, 1098 insertions(+), 620 deletions(-) create mode 100644 samples/spatial/sorting/bitonic_sort_1D.sptl create mode 100644 tests/csl_runtime/test_bitonic_sort_1d.sh diff --git a/samples/spatial/sorting/bitonic_sort_1D.sptl b/samples/spatial/sorting/bitonic_sort_1D.sptl new file mode 100644 index 00000000..08b8e0be --- /dev/null +++ b/samples/spatial/sorting/bitonic_sort_1D.sptl @@ -0,0 +1,158 @@ +/** + * Batcher's bitonic sorting network over N = 2^L PEs in a row. + * + * Each PE holds K keys, and the network sorts K independent sequences at once: sequence k is made + * of element k of every PE. A compare-exchange is therefore an elementwise min/max over K values, + * which streams as one K-element transfer per epoch. + * + * The point of this sample is *channel economy*. A bitonic network on N keys performs L(L+1)/2 + * compare-exchange steps at distances 1, 2, 4, ..., N/2, and every PE takes part in every step. + * Giving each step its own channel would need one per step; giving each concurrently exchanging + * pair its own channel would need O(N). This kernel instead uses + * + * L = log2(N) channels + * + * -- exactly one per exchange distance -- and reuses each of them across every step, every lane and + * both directions of travel. That reuse is what the router switches pay for. + * + * Why the network needs lanes + * --------------------------- + * At distance J = 2^d every PE with bit d clear is a "low" PE and exchanges with the PE J to its + * east. Those paths overlap: 0 -> 4 and 1 -> 5 both cross PEs 1..4, so they cannot share a channel + * concurrently (see the channel-conflict rule in the routing specification). The step is therefore + * split into J *lanes*: lane c handles the PEs congruent to c modulo 2J, whose paths are exactly + * disjoint. + * + * Why each lane is two phases + * --------------------------- + * A compare-exchange has to move a key in each direction, and both directions use the same channel. + * They cannot be concurrent, so the lane is split into an eastward phase and a westward one; the + * phase boundary closes the stream, which is what frees the channel. Reversing the direction of + * travel is what makes this kernel demanding: every PE on the path -- sender, relay and receiver + * alike -- has to change both its router's input and its output between the two phases. + * + * Structure, generated entirely by compile-time `for` blocks: + * + * for s in [0, L) stage: groups of size 2^(s+1) are sorted + * for e in [0, s] step: exchange distance J = 2^(s-e), halving + * for c in [0, J) lane: two phases, two epochs on channel s-e + * for g in ... the ascending groups, then the descending ones + * + * A group of size 2^(s+1) sorts ascending when bit s+1 of its index is clear and descending + * otherwise, so the ascending groups start at multiples of 2^(s+2) and the descending ones are + * offset by 2^(s+1). In the final stage the descending range is empty, which is what leaves the + * whole array sorted ascending. + * + * One lane (L=3, s=2, e=0, J=4, c=1) on channel 2, as two epochs: + * + * PE: 0 1 2 3 4 5 6 7 + * east *---->-----------------* 1 -> 5, relays 2,3,4 + * west *-----------------<----* 5 -> 1, relays 4,3,2 + * + * (!) Assumes L >= 1. Requires WSE-3: reversing a router between the two phases changes both its + * input and its output direction, and enough of those accumulate on one color to exceed the + * four switch positions a WSE-2 router can hold (where each such reversal costs two). (!) + * + * (!) The router budget caps this at L = 2. At distance 2^d the channel is reused by 2^d lanes in + * two directions each, so an interior router cycles through 2^(d+1) configurations; four + * switch positions run out at d = 2. Sorting more keys needs a channel per direction, which + * doubles the channel count to 2*log2(N) and, since no router then reverses, also runs on + * WSE-2. (!) + * + * Note on syntax: `+`/`-` are right-associative in this grammar, so `s-e+1` would parse as + * `s-(e+1)`. The step expressions below write `(s-e)+1` explicitly. + **/ +kernel @bitonic_sort_1d(stream[1<[1< east = relative_stream(1<<(s-e), 0) { + hops = auto, + channel = s-e + } + } + for i16 g in [0 : 1< west = relative_stream(-(1<<(s-e)), 0) { + hops = auto, + channel = s-e + } + } + for i16 g in [0 : 1< val[k] else val[k]) + } + } + } + for i16 g in [1<<(s+1) : 1< val[k] else val[k]) + } + } + compute i16 i, i16 j in [g + c + (1<<(s-e)) : g + (1<<(s+1)) : 1<<((s-e)+1), 0] { + await send(val, west) + await map i32 k in [0:K] { + val[k] = (other[k] if other[k] < val[k] else val[k]) + } + } + } + } + } + } + } + + // Store the sorted keys. + phase { + compute i16 i, i16 j in [0:1< spir.Kernel: @@ -39,7 +39,9 @@ def canonicalize_kernel(kernel: spir.Kernel) -> spir.Kernel: """ kernel = canonicalization.inline_metaprogramming(kernel) kernel = canonicalization.canonicalize_phases(kernel) + kernel = canonicalization.uniquify_stream_names(kernel) kernel = stream_lifetime.insert_implicit_closes(kernel) + kernel = canonicalization.number_stream_phases(kernel) kernel = canonicalization.reduce_streams(kernel) kernel = canonical_subgrids.canonicalize_subgrids(kernel) kernel = canonicalization.resolve_auto_hops(kernel) @@ -344,7 +346,6 @@ def generate_rectangle(kernel: spir.Kernel, # * Generate routing instructions from dataflow blocks # * Make unique colors out of streams, reduce number of streams color_map = _allocate_colors(rect, header, kernel, use_memcpy_mode, stream_extents, channel_to_color) - cslrouting.declare_switch_advances(rect, header, color_map) dtypes = _collect_identifier_types(rect.metadata, kernel.arguments) # Preprocess potential data tasks to convert to loops if possible @@ -383,6 +384,9 @@ def generate_rectangle(kernel: spir.Kernel, raise ValueError(f"Error in {e.args[0].lineinfo}. Undefined identifier \"{e.args[0].as_ir()}\".") raise + cslrouting.declare_switch_advances(rect, header, color_map, dsds) + _declare_queue_initialization(dsds, rect, footer, color_map) + # Fuse tasks as much as possible to reduce number of resources if task_fusion: orig_len = 0 @@ -881,6 +885,45 @@ def _dsd_from_stream(stream_candidates: dict[str, tuple[spir.StreamDeclaration | return cslstruct.MemoryDSD(dsd_type, name, extents, idxvars, indices) +def _declare_queue_initialization(dsds: UniqueDSDDict, rect: PEBlock, footer: StringIO, + color_map: dict[str, int]) -> None: + """ + Binds every fabric queue this PE uses to its color, which WSE-3 requires. + + On WSE-2 a fabric queue picks its color up from the descriptor that uses it. WSE-3 does not: + a queue must be tied to a color with ``@initialize_queue`` before any transfer over it will + proceed, and a program that omits it simply hangs. Queues are handed out per channel (see + ``_collect_unique_dsds``), so each one is named by exactly one color here. + + :param dsds: The descriptors collected for this rectangle. + :param rect: The PE block being generated, used for the switch-advance descriptors. + :param footer: The ``comptime`` block to write the bindings into. + :param color_map: Stream name to color number, for the switch-advance descriptors. + """ + if not csl.ARCH == 'wse3': + return + + # (queue kind, queue id) -> color expression. Both the data descriptors and the control + # descriptors that carry switch advances need their queue bound. + bindings: dict[tuple[str, int], str] = {} + for entries in dsds.values(): + for _, dsd in entries: + if not isinstance(dsd, cslstruct.FabricDSD) or not dsd.color: + continue + direction = 'in' if dsd.dsd_type == cslstruct.DSDType.fabin else 'out' + kind = 'input_queue' if dsd.dsd_type == cslstruct.DSDType.fabin else 'output_queue' + bindings.setdefault((kind, dsd.queue), f'{dsd.color}_{direction}') + + for statement in rect.metadata.compute.statements: + if isinstance(statement, spir.CloseStatement) and statement.switch_advance: + name = cslstmt.name_to_csl(stream_lifetime.underlying_stream(statement.stream_name)) + queue = csl.OUTPUT_QUEUE_IDS[0] + bindings.setdefault(('output_queue', queue), f'@get_color({color_map[name + "_OUT"]})') + + for (kind, queue), color in sorted(bindings.items()): + footer.write(f' @initialize_queue(@get_{kind}({queue}), .{{ .color = {color} }});\n') + + def _collect_unique_dsds( tasks: list[tdag.CSLTask], rect: PEBlock, diff --git a/spada/syntax/csl/constants.py b/spada/syntax/csl/constants.py index 47dd270a..6c8e7292 100644 --- a/spada/syntax/csl/constants.py +++ b/spada/syntax/csl/constants.py @@ -38,13 +38,15 @@ # See https://sdk.cerebras.net/csl/language/dsds#fabric-queues _INPUT_QUEUE_IDS = { 'wse2': list(range(0, 2)), # Ignoring 2-7 as they are smaller in capacity - 'wse3': list(range(0, 8)), # 0 is better than 1-7 + # On WSE-3 a data task's ID *is* its input queue, and memcpy takes 0 and 1 for its own; binding + # either of them with ``@initialize_queue`` is rejected as "already been set". + 'wse3': list(range(2, 8)), } INPUT_QUEUE_IDS = _INPUT_QUEUE_IDS[ARCH] _OUTPUT_QUEUE_IDS = { 'wse2': list(range(2, 4)), # Ignoring 0-1,4-5 as they are smaller in capacity - 'wse3': list(range(0, 8)), # All queues are equivalent + 'wse3': list(range(2, 8)), # All queues are equivalent, but memcpy reserves 0 and 1 } OUTPUT_QUEUE_IDS = _OUTPUT_QUEUE_IDS[ARCH] diff --git a/spada/syntax/csl/dsd_ops.py b/spada/syntax/csl/dsd_ops.py index 0123754c..f1e76ee0 100644 --- a/spada/syntax/csl/dsd_ops.py +++ b/spada/syntax/csl/dsd_ops.py @@ -7,7 +7,7 @@ from spada.syntax.spatial_ir import irnodes as spir from spada.syntax.csl import structures as cslstruct -UniqueDSDDict = dict[str, list[tuple[str, cslstruct.DataStructureDescriptor]]] +UniqueDSDDict = cslstruct.UniqueDSDDict @dataclass diff --git a/spada/syntax/csl/routing.py b/spada/syntax/csl/routing.py index 5f3b4451..1e49a966 100644 --- a/spada/syntax/csl/routing.py +++ b/spada/syntax/csl/routing.py @@ -24,6 +24,7 @@ from spada.syntax.csl import constants from spada.syntax.csl import statements as cslstmt +from spada.syntax.csl import structures as cslstruct from spada.syntax.spatial_ir import analysis, stream_lifetime from spada.syntax.spatial_ir import irnodes as spir from spada.syntax.spatial_ir.canonicalization import PEBlock @@ -248,13 +249,21 @@ def switch_advance_statements(dsd_name: str, advances: int) -> str: return '\n'.join(line for _ in range(advances)) -def declare_switch_advances(rect: Rectangle[PEBlock], header: StringIO, color_map: dict[str, int]) -> None: +def declare_switch_advances(rect: Rectangle[PEBlock], header: StringIO, color_map: dict[str, int], + dsds: cslstruct.UniqueDSDDict) -> None: """ Declares the fabric output descriptors that carry a stream's switch-advance control wavelet. The wavelet itself is emitted by ``statements.generate_csl_statement`` from the close's - ``switch_advance`` field; this only has to provide the descriptor it is sent through, because - that is where the color is known. + ``switch_advance`` field; this only has to provide the descriptor it is sent through. + + The descriptor reuses the *same* output queue as the stream's data, which is mandatory rather + than tidy: a queue is bound to one color on WSE-3 (see ``_declare_queue_initialization``), and + queues are handed out per channel, so sending a control wavelet for one color through the queue + that belongs to another silently fails to advance anything. A close only ever runs on a PE that + sends the stream, so the outgoing descriptor always exists. + + :param dsds: The descriptors collected for this rectangle, to take the stream's queue from. """ kept = [ statement for statement in rect.metadata.compute.statements @@ -264,14 +273,22 @@ def declare_switch_advances(rect: Rectangle[PEBlock], header: StringIO, color_ma return header.write('\nconst ctrl = @import_module("");\n') - queue = constants.OUTPUT_QUEUE_IDS[0] for statement in kept: - name = cslstmt.name_to_csl(stream_lifetime.underlying_stream(statement.stream_name)) + stream = stream_lifetime.underlying_stream(statement.stream_name) + name = cslstmt.name_to_csl(stream) dsd_name = f'{name}_switch_dsd' - if f'const {dsd_name}' not in header.getvalue(): - header.write(f'const {dsd_name} = @get_dsd(fabout_dsd, .{{ .extent = 1, ' - f'.fabric_color = @get_color({color_map[name + "_OUT"]}), .control = true, ' - f'.output_queue = @get_output_queue({queue}) }});\n') + if f'const {dsd_name}' in header.getvalue(): + continue + queue = None + for _, dsd in dsds.get(stream.as_ir(), ()): + if isinstance(dsd, cslstruct.FabricDSD) and dsd.dsd_type == cslstruct.DSDType.fabout: + queue = dsd.queue + break + if queue is None: + queue = constants.OUTPUT_QUEUE_IDS[0] + header.write(f'const {dsd_name} = @get_dsd(fabout_dsd, .{{ .extent = 1, ' + f'.fabric_color = @get_color({color_map[name + "_OUT"]}), .control = true, ' + f'.output_queue = @get_output_queue({queue}) }});\n') def route_dir(dx: int, dy: int): @@ -313,7 +330,12 @@ class _RouteEntry: configurations of the same site. """ config: RouteConfig - order: tuple[int, int] + #: ``(phase index, barrier index, statement index)``. The phase index is kernel-wide, which is + #: what makes entries from different rectangles -- notably a relay configuration contributed by + #: the sending rectangle -- comparable with each other. The other two come from + #: :func:`_stream_use_order` and only separate uses *within* one rectangle, which is what + #: sequences a receive before the send that forwards it. + order: tuple[int, int, int] origin_rect: int origin_offset: tuple[int, int] stream: spir.Identifier @@ -614,7 +636,7 @@ def _rectangle_route_entries(rect_index: int, rect: Rectangle[PEBlock], collected: list[tuple[_RouteSite, _RouteEntry]] = [] def add(offset: tuple[int, int], color: int, rx: tuple[str, ...], tx: tuple[str, ...], - order: tuple[int, int], stream_name: spir.Identifier, group: str) -> None: + order: tuple[int, int, int], stream_name: spir.Identifier, group: str) -> None: site = _RouteSite( color=color, x_range=(rect.x_range[0] + offset[0], rect.x_range[1] + offset[0], rect.x_range[2]), @@ -629,8 +651,12 @@ def add(offset: tuple[int, int], color: int, rx: tuple[str, ...], tx: tuple[str, sent, received = sends_recvs[stream.stream_name] group = stream_lifetime.stream_group_key(stream) orders = use_order.get(stream.stream_name, {}) - receive_order = orders.get('receive', (0, 0)) - send_order = orders.get('send', (0, 0)) + # A stream's epoch is fixed kernel-wide by the phase it is declared in; within that phase + # the local statement order separates a receive from a send of the same stream, which is + # what sequences a systolic forward. + epoch = stream.phase if stream.phase is not None else 0 + receive_order = (epoch, ) + orders.get('receive', (0, 0)) + send_order = (epoch, ) + orders.get('send', (0, 0)) if received: color_inbound = color_map[cslstmt.name_to_csl(stream.stream_name) + "_IN"] if sent: diff --git a/spada/syntax/csl/statements.py b/spada/syntax/csl/statements.py index c3c7a8d2..31761b9a 100644 --- a/spada/syntax/csl/statements.py +++ b/spada/syntax/csl/statements.py @@ -1,12 +1,10 @@ from io import StringIO from typing import Optional -from spada.syntax.csl.structures import DataStructureDescriptor +from spada.syntax.csl.structures import UniqueDSDDict from spada.syntax.csl import dsd_ops from spada.syntax.csl import routing from spada.syntax.spatial_ir import irnodes as spir -UniqueDSDDict = dict[str, list[tuple[str, DataStructureDescriptor]]] - def generate_csl_statement(statement: spir.Statement, dsds: UniqueDSDDict, diff --git a/spada/syntax/csl/structures.py b/spada/syntax/csl/structures.py index 1855895c..693c01b2 100644 --- a/spada/syntax/csl/structures.py +++ b/spada/syntax/csl/structures.py @@ -21,6 +21,11 @@ def as_csl(self) -> str: raise NotImplementedError +#: The descriptors collected for one PE, keyed by the Spatial IR name they belong to, each entry +#: pairing the CSL identifier the descriptor is declared under with the descriptor itself. +UniqueDSDDict = dict[str, list[tuple[str, DataStructureDescriptor]]] + + @dataclass class MemoryDSD(DataStructureDescriptor): """ diff --git a/spada/syntax/csl/tasks.py b/spada/syntax/csl/tasks.py index 838ae46b..7bd7981f 100644 --- a/spada/syntax/csl/tasks.py +++ b/spada/syntax/csl/tasks.py @@ -1,564 +1,564 @@ -""" -Contains a CSL task DAG representation and creation methods. -""" - -from copy import deepcopy -from dataclasses import dataclass -from enum import Enum, auto -import networkx as nx # TODO: Switch to igraph -from typing import Any, Literal, Optional -from spada.syntax.spatial_ir import irnodes as spir, analysis -from spada.syntax.csl import constants, dsd_ops, structures as cslstruct - -UniqueDSDDict = dict[str, list[tuple[str, cslstruct.DataStructureDescriptor]]] - - -class TaskCreationBehavior(Enum): - """ - Enumeration prescribing how tasks should be created. - """ - - NO_TASKS = auto() # All statements in one task - FAIL_ON_OVERRUN = auto() # Error if too many tasks - STATE_MACHINE_ON_OVERRUN = auto() # Recycle task IDs with a state machine - SYNCHRONOUS_ON_OVERRUN = auto() # Run tasks synchronously if too many - - -class InterTaskEdge(Enum): - """ - Enumeration representing a task dependency edge type. - """ - - UNSET = auto() - SEQUENCE = auto() - ACTIVATE = auto() - UNBLOCK = auto() - - -@dataclass -class CSLTask: - """ - Object representing a task DAG node. - """ - - task_id: int - task_type: Literal["local", "data"] # We do not generate control tasks at the moment - statements: list[int] # Index is the statement's index from the completion DAG - outgoing: list[tuple[int, InterTaskEdge]] # For each statement, the next task ID and the dependency type - blocked: bool # Whether there is an unblock edge leading to this task - - -def should_be_asynchronous(dtypes: dict[spir.Identifier, spir.IRType], stmt: spir.Statement) -> bool: - """ - Returns True if a statement can and should be executed asynchronously in CSL. - The only statements that apply are DSD operations that have to do with fabric DSDs (e.g., send, receive). - """ - if isinstance(stmt, (spir.SendStatement, spir.ReceiveStatement)): - return True - if isinstance(stmt, spir.ForeachStatement) and stmt.receive_stream: - return dsd_ops.get_dsd_op(dtypes, stmt) is not None - - return False - - -# _DEBUG_i = 1 - - -def create_csl_tasks( - completion_dag: nx.DiGraph, - block: spir.ComputeBlock, - dtypes: dict[spir.Identifier, spir.IRType], - task_creation_behavior: TaskCreationBehavior = TaskCreationBehavior.FAIL_ON_OVERRUN, -) -> list[CSLTask]: - """ - Creates a list of CSL tasks. The nodes are tasks that contain a unique ID - and the list of statements to include; and the edges are the type of dependency across tasks. - The algorithm operates as follows. - - Statements can take on different task types, based on the statement type and its contents: - - * Foreach statements may take the form of a CSL data task, if they cannot trivially be represented by a - single DSD operation (@mov, @fadd*, etc.) - * Send and receive statements that can be lowered to a ``FabricDSD`` operation, in turn can (and should) - be nonblocking, or ``async`` in CSL terms. In this lowering pipeline, these live in CSL local tasks. - * Other statements (e.g. free assignments) are blocking and also live in local tasks. - - Given that tasks can ``@activate`` and ``@unblock`` other tasks, and that ``FabricDSD`` operations can also - do the same, both a task terminator and a nonblocking statement can trigger other tasks. Given that there are - no other options to trigger tasks, we run a preprocessing pass on the graph to convert nodes with in-degree over 2 - to a series of ``wait`` nodes. - - Subsequently, we traverse the Completion DAG topologically (to ensure proper local order). We then decide to create - new tasks based on a set of necessary rules in which a new task must be formed: - - 1. A node with no predecessors creates a new activated and unblocked task - 2. A node with more than one incoming edge must start a new task - (the conditions below thus apply to the case where a node has one predecessor) - 3. If a node's predecessor represents one kind of CSL task (e.g., data) and this node represents another - 4. Node pairs with ``wait->wait`` edges create a new task (this also fulfills the condition for the above - preprocessing pass) - 5. ``post->wait`` node pairs where the post is a nonblocking operation creates a new task for the ``wait`` node - and sets the nonblocking DSD to ``.activate`` the waiting task, or ``.unblock`` it if there is another edge - - The last (i.e., sink) task is called ``exit_task`` and is built into the generation of rectangle code. - - This means that post->post nodes of nonblocking operations can coexist in the same task. - """ - result: list[CSLTask] = [] - num_local_tasks = 0 - num_data_tasks = 0 - - completion_dag = _canonicalize_dag(completion_dag) - # global _DEBUG_i - # nx.nx_pydot.write_dot(completion_dag, f'canon{_DEBUG_i}.dot') - # _DEBUG_i += 1 - - # Mappings between IR statements and tasks - cnode: analysis.CompletionDAGNode - current_task: CSLTask = None - statement_id_to_task_id: dict[int, int] = {} - cnode_to_task_id: dict[analysis.CompletionDAGNode, int] = {} - - # Loop over completion DAG to coarsen completions to tasks - for cnode in nx.topological_sort(completion_dag): - node = block.statements[cnode.statement_id] - # Figure out whether this task type is a local task or a data task - if ( - isinstance(node, spir.ForeachStatement) - and dsd_ops.get_dsd_op(dtypes, node) is None - and cnode.optype == "post" - ): - # Only if it is a complex task (i.e., not a DSD operation) - this_task_type = "data" - else: - this_task_type = "local" - - task_id = None - # Look at incoming edges: - indeg = completion_dag.in_degree(cnode) - if indeg == 1: # A node with zero or more than one incoming edge has to start a new task - pred: analysis.CompletionDAGNode - pred, _ = next(iter(completion_dag.in_edges(cnode))) - # If {wait,post}->post and there is one edge, and the previous task is a local task, inherit task ID - if result[statement_id_to_task_id[pred.statement_id]].task_type == "local" and this_task_type == "local": - if pred.optype == "post" and cnode.optype == "post": - task_id = statement_id_to_task_id[pred.statement_id] - elif pred.optype == "wait" and cnode.optype == "post": - task_id = statement_id_to_task_id[pred.statement_id] - elif pred.optype == "post" and cnode.optype == "wait": - # ``post->wait`` node pairs where the post is a nonblocking operation creates a new task for the - # ``wait`` node, depending on task creation behavior - should_create_task = True - if task_creation_behavior == TaskCreationBehavior.NO_TASKS: - should_create_task = False # Always inherit prior task - elif task_creation_behavior == TaskCreationBehavior.SYNCHRONOUS_ON_OVERRUN: - if num_local_tasks >= len(constants.LOCAL_TASK_IDS): - should_create_task = False - if should_create_task and not should_be_asynchronous(dtypes, block.statements[pred.statement_id]): - should_create_task = False - if not should_create_task: - task_id = statement_id_to_task_id[pred.statement_id] - # wait->wait will create a new task - elif result[statement_id_to_task_id[pred.statement_id]].task_type == "local" and this_task_type == "data": - # An empty wait task before a data task can be contracted - if not result[statement_id_to_task_id[pred.statement_id]].statements: - task_id = statement_id_to_task_id[pred.statement_id] - - # Otherwise, we need a new task - - # The one condition in which a wait->wait edge can be contracted is if there is a (post,post)->wait->wait, - # which can be represented by two edges with unblock and activate. - # TODO(later): this is a performance optimization that can be done later - - # If task ID is not None, append statement to prior task - if task_id is not None: - previous_task: CSLTask = result[task_id] - cnode_to_task_id[cnode] = task_id - - # Modify task type - if previous_task.task_type != this_task_type: - previous_task.task_type = this_task_type - - if cnode.statement_id not in statement_id_to_task_id: - previous_task.statements.append(cnode.statement_id) - previous_task.outgoing.append((-1, InterTaskEdge.UNSET)) - statement_id_to_task_id[cnode.statement_id] = task_id - - continue - - # Create a new task - task_id = len(result) - cnode_to_task_id[cnode] = task_id - current_task = CSLTask( - task_id, this_task_type, [], [], blocked=((indeg > 1) or (this_task_type == "data" and indeg > 0)) - ) - result.append(current_task) - if this_task_type == "local": - num_local_tasks += 1 - else: - num_data_tasks += 1 - statement_id_to_task_id[cnode.statement_id] = task_id - - if cnode.optype == "wait": - # Nothing to do within the task - pass - else: # 'post' - current_task.statements.append(cnode.statement_id) - current_task.outgoing.append((-1, InterTaskEdge.UNSET)) - - # For edge type detection - task_has_activate: set[int] = set() - - # Determine edge types between task statements - for cnode in nx.topological_sort(completion_dag): - stmt_task = cnode_to_task_id[cnode] - if cnode.optype == "post": - # Find matching "wait" successor - succ_task = None - for succ in completion_dag.successors(cnode): - if succ.optype == "wait": - succ_task = cnode_to_task_id[succ] - break - - assert succ_task is not None # An asynchronous statement must have a unique successor - - # Find outgoing index within task - ind = next(i for i, s in enumerate(result[stmt_task].statements) if s == cnode.statement_id) - - elif cnode.optype == "wait": # Set the next task after the await to begin sequentially - ind = next((i for i, s in enumerate(result[stmt_task].statements) if s == cnode.statement_id), None) - if ind is None: # Wait already omitted from task - continue - # After canonicalization, there must be one successor for each wait node - num_successors = len(list(completion_dag.successors(cnode))) - if num_successors == 1: - succ_task = next(succ for succ in completion_dag.successors(cnode)) - succ_task = cnode_to_task_id[succ_task] - elif num_successors > 1: - node = block.statements[cnode.statement_id] - raise ValueError( - "Multiple successors for a wait task should not appear after canonicalization.\n In " - f"line {node.lineinfo}" - ) - else: # No successors - continue - - # Determine edge type and assign outgoing edge - # Successor lives within same task, make sequence - if stmt_task == succ_task: - etype = InterTaskEdge.SEQUENCE - else: - # Check if task already has an activate edge - if succ_task in task_has_activate or result[succ_task].task_type == "data": - etype = InterTaskEdge.UNBLOCK - else: - etype = InterTaskEdge.ACTIVATE - task_has_activate.add(succ_task) - result[stmt_task].outgoing[ind] = (succ_task, etype) - - # If the last task is local and empty, we can contract it with our exit task - if len(result) > 0 and not result[-1].statements: - result = result[:-1] - - # Determine terminators: if a task has a predecessor but no matching activator (outgoing statement), - # add a terminator statement (@activate or @unblock, depending on other dependencies). - # We define a terminator as a statement with ID "TERMINATOR" - for cnode in nx.topological_sort(completion_dag): - preds = completion_dag.predecessors(cnode) - stmt_task = cnode_to_task_id[cnode] - for pred in preds: - pred_task = cnode_to_task_id[pred] - if pred_task == stmt_task: # Skip sequential edges - continue - # Predecessor lived on the contracted empty last task; no task slot remains for it. - if pred_task >= len(result): - continue - - has_edge = any(e == stmt_task for e, _ in result[pred_task].outgoing) - if not has_edge: - task = result[pred_task] - if task.task_type == "local": - task.statements.append("TERMINATOR") - # After contracting an empty trailing task, stmt_task may equal len(result), meaning exit. - succ_is_data = stmt_task < len(result) and result[stmt_task].task_type == "data" - if stmt_task in task_has_activate or succ_is_data: - task.outgoing.append((stmt_task, InterTaskEdge.UNBLOCK)) - else: - task.outgoing.append((stmt_task, InterTaskEdge.ACTIVATE)) - task_has_activate.add(stmt_task) - - # Assign task IDs for local and data tasks - current_local_task_id = -1 - current_data_task_id = -1 - task_id_to_local_id: dict[int, int] = {} - task_id_to_data_id: dict[int, int] = {} - for task_id, task in enumerate(result): - # Increment the current task ID and add a new task with the specified type - if task.task_type == "local": - current_local_task_id += 1 - task_id_to_local_id[task_id] = current_local_task_id - else: # 'data' - current_data_task_id += 1 - task_id_to_data_id[task_id] = current_data_task_id - - # Re-number task IDs and outgoing connections based on CSL IDs - for task_id, task in enumerate(result): - # TODO(later): Task IDs can be recycled with a global ``var`` that can be set prior to activating - # a task, like a state machine - - # NOTE: This check happens after task fusion, so it is less likely to trigger - # if task.task_type == 'data': - # tid = task_id_to_data_id[task_id] - # if tid >= len(constants.DATA_TASK_IDS) and task_creation_behavior == TaskCreationBehavior.FAIL_ON_OVERRUN: - # raise ValueError('Too many data tasks') - # task.task_id = constants.DATA_TASK_IDS[tid] - # elif task.task_type == 'local': - # tid = task_id_to_local_id[task_id] - # if tid >= len(constants.LOCAL_TASK_IDS) and task_creation_behavior == TaskCreationBehavior.FAIL_ON_OVERRUN: - # raise ValueError('Too many local tasks') - # task.task_id = constants.LOCAL_TASK_IDS[tid] - - for i, (target, e) in enumerate(task.outgoing): - # Mark exit task explicitly - if target == len(result): - task.outgoing[i] = (-1, e) - continue - - # Do not modify outgoing task IDs - # if result[target].task_type == 'local': - # target_id = constants.LOCAL_TASK_IDS[task_id_to_local_id[target]] - # else: - # target_id = constants.DATA_TASK_IDS[task_id_to_data_id[target]] - # task.outgoing[i] = (target_id, e) - - # Add explicit terminator statements for tasks that have no successors - sink_tasks = [ - t for i, t in enumerate(result) if not any(n != i for n, _ in t.outgoing) or -1 in set(n for n, _ in t.outgoing) - ] - if len(sink_tasks) > 2: - raise ValueError("Too many sink tasks") - for i, task in enumerate(sink_tasks): - edge_type = InterTaskEdge.ACTIVATE if i == 0 else InterTaskEdge.UNBLOCK - if task_creation_behavior == TaskCreationBehavior.SYNCHRONOUS_ON_OVERRUN and num_local_tasks >= len( - constants.LOCAL_TASK_IDS - ): - edge_type = InterTaskEdge.SEQUENCE - elif task_creation_behavior == TaskCreationBehavior.NO_TASKS: - edge_type = InterTaskEdge.SEQUENCE - elif len(sink_tasks) == 1: # Save on task IDs if there is only one sink task - edge_type = InterTaskEdge.SEQUENCE - - if task.statements[-1] != "TERMINATOR" and task.outgoing[-1][0] != -1 and task.task_type == "local": - task.statements.append("TERMINATOR") - task.outgoing.append((-1, edge_type)) - elif task.outgoing[-1][0] == -1: # Modify existing UNBLOCK edge if there is more than one sink - task.outgoing[-1] = (-1, edge_type) - - # Inject local tasks in front of data tasks with more than one input - to_append = [] - for i, task in enumerate(result): - if task.task_type != "data": - continue - predecessors = set(j for j, t in enumerate(result) if any(o == i for o, _ in t.outgoing)) - if len(predecessors) > 1: - new_id = len(result) + len(to_append) - new_task = CSLTask(i, "local", ["TERMINATOR"], [(new_id, InterTaskEdge.UNBLOCK)], blocked=True) - to_append.append(task) - result[i] = new_task - # Make one of the predecessors into an ACTIVATE edge - tpred = next(iter(predecessors)) - result[tpred].outgoing = [ - (o, InterTaskEdge.ACTIVATE) if o == i else (o, e) for o, e in result[tpred].outgoing - ] - - result.extend(to_append) - - return result - - -def _contract_node(g: nx.DiGraph, n: Any): - if g.out_degree(n) == 0: # Keep sink node - return - for u, _ in g.in_edges(n): - for _, v in g.out_edges(n): - g.add_edge(u, v) - g.remove_node(n) - - -def _canonicalize_dag(completion_dag: nx.DiGraph) -> nx.DiGraph: - completion_dag = deepcopy(completion_dag) - - # Reduce in-degree of nodes to up to 2 - _limit_indegree(completion_dag) - - return completion_dag - - -def _limit_indegree(dag: nx.DiGraph): - """ - Injects extra wait nodes to completion DAGs where the in-degree of a node is larger than two. - - :param dag: The completion DAG. - """ - counter = -1 - for node in list(dag.nodes): # Copy nodes to a list - if dag.in_degree(node) > 2: - edges = list(dag.in_edges(node)) - current_node = node - # Create intermediate wait nodes (the counter changes the statement ID because it has to be unique) - for u, _ in edges[1:]: - new_node = analysis.CompletionDAGNode("wait", counter) - counter -= 1 - dag.remove_edge(u, node) - dag.add_edge(new_node, current_node) - dag.add_edge(u, new_node) - current_node = new_node - - -def fuse_tasks( - tasks: list[CSLTask], - dsds: UniqueDSDDict, - dtypes: dict[spir.Identifier, spir.IRType], - rect, - use_memcpy_mode: bool, - compute: spir.ComputeBlock, -) -> list[CSLTask]: - """ - Fuses tasks where possible to reduce the number of tasks. - - :param tasks: The list of CSL tasks. - :param dsds: The unique DSD dictionary. - :param dtypes: The dictionary of identifier types. - :param kernel: The spatial IR kernel. - :param use_memcpy_mode: Whether memcpy mode is used. - :return: The fused list of CSL tasks. - """ - fused: set[int] = set() - removed: set[int] = set() - redirect: dict[int, int] = {} - for i, task in enumerate(tasks): - if i in fused or i in removed: - continue - if task.task_type == "data": - continue - if not task.statements: - # Remove task - removed.add(i) - continue - outgoing_id, _ = task.outgoing[-1] - if outgoing_id == -1: - # This is a sink task, nothing to fuse with - continue - if outgoing_id in fused or outgoing_id in removed: - continue - if any(et == InterTaskEdge.UNBLOCK for t in tasks for n, et in t.outgoing if n == outgoing_id): - # Cannot fuse if there are multiple predecessors to next task - continue - if tasks[outgoing_id].task_type != "local": - # Cannot fuse with data tasks - continue - - last_stmt = task.statements[-1] - if not isinstance(last_stmt, int): - # Last statement is a terminator, cannot fuse - continue - - # Identify fusion opportunities - next_task = tasks[outgoing_id] if outgoing_id < len(tasks) else None - if next_task is None: - continue - stmt = compute.statements[last_stmt] - if isinstance(stmt, (spir.SendStatement, spir.ReceiveStatement)): - stream = dsd_ops._get_id(stmt.stream_name) - localarr = dsd_ops._get_id(stmt.local_array) - if stream.as_ir() in dsds and isinstance(dsds[stream.as_ir()][0][1], cslstruct.FabricDSD): - # Cannot fuse if the last statement is a fabric DSD operation - continue - if localarr.as_ir() in dsds and isinstance(dsds[localarr.as_ir()][0][1], cslstruct.FabricDSD): - # Cannot fuse if the last statement is a fabric DSD operation - continue - elif isinstance(stmt, spir.ForeachStatement): - # Try to fuse synchronous DSD operations - dsd_op: type[dsd_ops.DSDOp] = dsd_ops.DSD_ASSIGNMENT_MAPPING.get(dsd_ops.get_dsd_op(dtypes, stmt), None) - if dsd_op is None: - continue - # Pure memory DSD operations are synchronous and can be fused - dsd_stmt = dsd_ops.get_dsd_statement(dtypes, stmt) - if dsd_stmt is None: - continue - dsd_objects = dsd_op().used_dsd_objects(dsd_stmt, dsds) - if any(isinstance(dsd, cslstruct.FabricDSD) for dsd in dsd_objects): - continue - # Also check the foreach receive generator - stream = dsd_ops._get_id(stmt.receive_stream.stream_name) - if stream.as_ir() in dsds and isinstance(dsds[stream.as_ir()][0][1], cslstruct.FabricDSD): - continue - else: - # Fusible pattern not found - continue - - # If we can fuse with the next task, do so - fused.add(outgoing_id) - redirect[outgoing_id] = i - task.outgoing[-1] = (i, InterTaskEdge.SEQUENCE) # Remove outgoing edge - task.statements.extend(next_task.statements) - task.outgoing.extend(next_task.outgoing) - - new_tasks: list[CSLTask] = [] - old_to_new: dict[int, int] = {} - for idx, task in enumerate(tasks): - if idx in fused or idx in removed: - continue - new_idx = len(new_tasks) - old_to_new[idx] = new_idx - new_tasks.append(task) - - def resolve_target(target: int) -> int: - while target in redirect: - target = redirect[target] - return target - - for task in new_tasks: - for j, (target, et) in enumerate(task.outgoing): - if target == -1: - continue - resolved = resolve_target(target) - if resolved not in old_to_new: - raise ValueError(f"Dangling task reference {resolved} after fusion") - task.outgoing[j] = (old_to_new[resolved], et) - - return new_tasks - - -def renumber_tasks(tasks: list[CSLTask], task_creation_behavior: TaskCreationBehavior) -> None: - """ - Renumbers tasks to map to hardware task IDs, based on the task creation behavior. - - :param tasks: The list of CSL tasks to operate in-place on. - :param task_creation_behavior: The task creation behavior. - """ - current_local_task_id = -1 - current_data_task_id = -1 - task_id_to_local_id: dict[int, int] = {} - task_id_to_data_id: dict[int, int] = {} - for task_id, task in enumerate(tasks): - # Increment the current task ID and add a new task with the specified type - if task.task_type == "local": - current_local_task_id += 1 - task_id_to_local_id[task_id] = current_local_task_id - else: # 'data' - current_data_task_id += 1 - task_id_to_data_id[task_id] = current_data_task_id - - # Re-number task IDs and outgoing connections based on CSL IDs - for task_id, task in enumerate(tasks): - if task.task_type == "data": - tid = task_id_to_data_id[task_id] - if tid >= len(constants.DATA_TASK_IDS) and task_creation_behavior == TaskCreationBehavior.FAIL_ON_OVERRUN: - raise ValueError("Too many data tasks") - task.task_id = constants.DATA_TASK_IDS[tid] - elif task.task_type == "local": - tid = task_id_to_local_id[task_id] - if tid >= len(constants.LOCAL_TASK_IDS) and task_creation_behavior == TaskCreationBehavior.FAIL_ON_OVERRUN: - raise ValueError("Too many local tasks") - task.task_id = constants.LOCAL_TASK_IDS[tid] +""" +Contains a CSL task DAG representation and creation methods. +""" + +from copy import deepcopy +from dataclasses import dataclass +from enum import Enum, auto +import networkx as nx # TODO: Switch to igraph +from typing import Any, Literal, Optional +from spada.syntax.spatial_ir import irnodes as spir, analysis +from spada.syntax.csl import constants, dsd_ops, structures as cslstruct + +UniqueDSDDict = cslstruct.UniqueDSDDict + + +class TaskCreationBehavior(Enum): + """ + Enumeration prescribing how tasks should be created. + """ + + NO_TASKS = auto() # All statements in one task + FAIL_ON_OVERRUN = auto() # Error if too many tasks + STATE_MACHINE_ON_OVERRUN = auto() # Recycle task IDs with a state machine + SYNCHRONOUS_ON_OVERRUN = auto() # Run tasks synchronously if too many + + +class InterTaskEdge(Enum): + """ + Enumeration representing a task dependency edge type. + """ + + UNSET = auto() + SEQUENCE = auto() + ACTIVATE = auto() + UNBLOCK = auto() + + +@dataclass +class CSLTask: + """ + Object representing a task DAG node. + """ + + task_id: int + task_type: Literal["local", "data"] # We do not generate control tasks at the moment + statements: list[int] # Index is the statement's index from the completion DAG + outgoing: list[tuple[int, InterTaskEdge]] # For each statement, the next task ID and the dependency type + blocked: bool # Whether there is an unblock edge leading to this task + + +def should_be_asynchronous(dtypes: dict[spir.Identifier, spir.IRType], stmt: spir.Statement) -> bool: + """ + Returns True if a statement can and should be executed asynchronously in CSL. + The only statements that apply are DSD operations that have to do with fabric DSDs (e.g., send, receive). + """ + if isinstance(stmt, (spir.SendStatement, spir.ReceiveStatement)): + return True + if isinstance(stmt, spir.ForeachStatement) and stmt.receive_stream: + return dsd_ops.get_dsd_op(dtypes, stmt) is not None + + return False + + +# _DEBUG_i = 1 + + +def create_csl_tasks( + completion_dag: nx.DiGraph, + block: spir.ComputeBlock, + dtypes: dict[spir.Identifier, spir.IRType], + task_creation_behavior: TaskCreationBehavior = TaskCreationBehavior.FAIL_ON_OVERRUN, +) -> list[CSLTask]: + """ + Creates a list of CSL tasks. The nodes are tasks that contain a unique ID + and the list of statements to include; and the edges are the type of dependency across tasks. + The algorithm operates as follows. + + Statements can take on different task types, based on the statement type and its contents: + + * Foreach statements may take the form of a CSL data task, if they cannot trivially be represented by a + single DSD operation (@mov, @fadd*, etc.) + * Send and receive statements that can be lowered to a ``FabricDSD`` operation, in turn can (and should) + be nonblocking, or ``async`` in CSL terms. In this lowering pipeline, these live in CSL local tasks. + * Other statements (e.g. free assignments) are blocking and also live in local tasks. + + Given that tasks can ``@activate`` and ``@unblock`` other tasks, and that ``FabricDSD`` operations can also + do the same, both a task terminator and a nonblocking statement can trigger other tasks. Given that there are + no other options to trigger tasks, we run a preprocessing pass on the graph to convert nodes with in-degree over 2 + to a series of ``wait`` nodes. + + Subsequently, we traverse the Completion DAG topologically (to ensure proper local order). We then decide to create + new tasks based on a set of necessary rules in which a new task must be formed: + + 1. A node with no predecessors creates a new activated and unblocked task + 2. A node with more than one incoming edge must start a new task + (the conditions below thus apply to the case where a node has one predecessor) + 3. If a node's predecessor represents one kind of CSL task (e.g., data) and this node represents another + 4. Node pairs with ``wait->wait`` edges create a new task (this also fulfills the condition for the above + preprocessing pass) + 5. ``post->wait`` node pairs where the post is a nonblocking operation creates a new task for the ``wait`` node + and sets the nonblocking DSD to ``.activate`` the waiting task, or ``.unblock`` it if there is another edge + + The last (i.e., sink) task is called ``exit_task`` and is built into the generation of rectangle code. + + This means that post->post nodes of nonblocking operations can coexist in the same task. + """ + result: list[CSLTask] = [] + num_local_tasks = 0 + num_data_tasks = 0 + + completion_dag = _canonicalize_dag(completion_dag) + # global _DEBUG_i + # nx.nx_pydot.write_dot(completion_dag, f'canon{_DEBUG_i}.dot') + # _DEBUG_i += 1 + + # Mappings between IR statements and tasks + cnode: analysis.CompletionDAGNode + current_task: CSLTask = None + statement_id_to_task_id: dict[int, int] = {} + cnode_to_task_id: dict[analysis.CompletionDAGNode, int] = {} + + # Loop over completion DAG to coarsen completions to tasks + for cnode in nx.topological_sort(completion_dag): + node = block.statements[cnode.statement_id] + # Figure out whether this task type is a local task or a data task + if ( + isinstance(node, spir.ForeachStatement) + and dsd_ops.get_dsd_op(dtypes, node) is None + and cnode.optype == "post" + ): + # Only if it is a complex task (i.e., not a DSD operation) + this_task_type = "data" + else: + this_task_type = "local" + + task_id = None + # Look at incoming edges: + indeg = completion_dag.in_degree(cnode) + if indeg == 1: # A node with zero or more than one incoming edge has to start a new task + pred: analysis.CompletionDAGNode + pred, _ = next(iter(completion_dag.in_edges(cnode))) + # If {wait,post}->post and there is one edge, and the previous task is a local task, inherit task ID + if result[statement_id_to_task_id[pred.statement_id]].task_type == "local" and this_task_type == "local": + if pred.optype == "post" and cnode.optype == "post": + task_id = statement_id_to_task_id[pred.statement_id] + elif pred.optype == "wait" and cnode.optype == "post": + task_id = statement_id_to_task_id[pred.statement_id] + elif pred.optype == "post" and cnode.optype == "wait": + # ``post->wait`` node pairs where the post is a nonblocking operation creates a new task for the + # ``wait`` node, depending on task creation behavior + should_create_task = True + if task_creation_behavior == TaskCreationBehavior.NO_TASKS: + should_create_task = False # Always inherit prior task + elif task_creation_behavior == TaskCreationBehavior.SYNCHRONOUS_ON_OVERRUN: + if num_local_tasks >= len(constants.LOCAL_TASK_IDS): + should_create_task = False + if should_create_task and not should_be_asynchronous(dtypes, block.statements[pred.statement_id]): + should_create_task = False + if not should_create_task: + task_id = statement_id_to_task_id[pred.statement_id] + # wait->wait will create a new task + elif result[statement_id_to_task_id[pred.statement_id]].task_type == "local" and this_task_type == "data": + # An empty wait task before a data task can be contracted + if not result[statement_id_to_task_id[pred.statement_id]].statements: + task_id = statement_id_to_task_id[pred.statement_id] + + # Otherwise, we need a new task + + # The one condition in which a wait->wait edge can be contracted is if there is a (post,post)->wait->wait, + # which can be represented by two edges with unblock and activate. + # TODO(later): this is a performance optimization that can be done later + + # If task ID is not None, append statement to prior task + if task_id is not None: + previous_task: CSLTask = result[task_id] + cnode_to_task_id[cnode] = task_id + + # Modify task type + if previous_task.task_type != this_task_type: + previous_task.task_type = this_task_type + + if cnode.statement_id not in statement_id_to_task_id: + previous_task.statements.append(cnode.statement_id) + previous_task.outgoing.append((-1, InterTaskEdge.UNSET)) + statement_id_to_task_id[cnode.statement_id] = task_id + + continue + + # Create a new task + task_id = len(result) + cnode_to_task_id[cnode] = task_id + current_task = CSLTask( + task_id, this_task_type, [], [], blocked=((indeg > 1) or (this_task_type == "data" and indeg > 0)) + ) + result.append(current_task) + if this_task_type == "local": + num_local_tasks += 1 + else: + num_data_tasks += 1 + statement_id_to_task_id[cnode.statement_id] = task_id + + if cnode.optype == "wait": + # Nothing to do within the task + pass + else: # 'post' + current_task.statements.append(cnode.statement_id) + current_task.outgoing.append((-1, InterTaskEdge.UNSET)) + + # For edge type detection + task_has_activate: set[int] = set() + + # Determine edge types between task statements + for cnode in nx.topological_sort(completion_dag): + stmt_task = cnode_to_task_id[cnode] + if cnode.optype == "post": + # Find matching "wait" successor + succ_task = None + for succ in completion_dag.successors(cnode): + if succ.optype == "wait": + succ_task = cnode_to_task_id[succ] + break + + assert succ_task is not None # An asynchronous statement must have a unique successor + + # Find outgoing index within task + ind = next(i for i, s in enumerate(result[stmt_task].statements) if s == cnode.statement_id) + + elif cnode.optype == "wait": # Set the next task after the await to begin sequentially + ind = next((i for i, s in enumerate(result[stmt_task].statements) if s == cnode.statement_id), None) + if ind is None: # Wait already omitted from task + continue + # After canonicalization, there must be one successor for each wait node + num_successors = len(list(completion_dag.successors(cnode))) + if num_successors == 1: + succ_task = next(succ for succ in completion_dag.successors(cnode)) + succ_task = cnode_to_task_id[succ_task] + elif num_successors > 1: + node = block.statements[cnode.statement_id] + raise ValueError( + "Multiple successors for a wait task should not appear after canonicalization.\n In " + f"line {node.lineinfo}" + ) + else: # No successors + continue + + # Determine edge type and assign outgoing edge + # Successor lives within same task, make sequence + if stmt_task == succ_task: + etype = InterTaskEdge.SEQUENCE + else: + # Check if task already has an activate edge + if succ_task in task_has_activate or result[succ_task].task_type == "data": + etype = InterTaskEdge.UNBLOCK + else: + etype = InterTaskEdge.ACTIVATE + task_has_activate.add(succ_task) + result[stmt_task].outgoing[ind] = (succ_task, etype) + + # If the last task is local and empty, we can contract it with our exit task + if len(result) > 0 and not result[-1].statements: + result = result[:-1] + + # Determine terminators: if a task has a predecessor but no matching activator (outgoing statement), + # add a terminator statement (@activate or @unblock, depending on other dependencies). + # We define a terminator as a statement with ID "TERMINATOR" + for cnode in nx.topological_sort(completion_dag): + preds = completion_dag.predecessors(cnode) + stmt_task = cnode_to_task_id[cnode] + for pred in preds: + pred_task = cnode_to_task_id[pred] + if pred_task == stmt_task: # Skip sequential edges + continue + # Predecessor lived on the contracted empty last task; no task slot remains for it. + if pred_task >= len(result): + continue + + has_edge = any(e == stmt_task for e, _ in result[pred_task].outgoing) + if not has_edge: + task = result[pred_task] + if task.task_type == "local": + task.statements.append("TERMINATOR") + # After contracting an empty trailing task, stmt_task may equal len(result), meaning exit. + succ_is_data = stmt_task < len(result) and result[stmt_task].task_type == "data" + if stmt_task in task_has_activate or succ_is_data: + task.outgoing.append((stmt_task, InterTaskEdge.UNBLOCK)) + else: + task.outgoing.append((stmt_task, InterTaskEdge.ACTIVATE)) + task_has_activate.add(stmt_task) + + # Assign task IDs for local and data tasks + current_local_task_id = -1 + current_data_task_id = -1 + task_id_to_local_id: dict[int, int] = {} + task_id_to_data_id: dict[int, int] = {} + for task_id, task in enumerate(result): + # Increment the current task ID and add a new task with the specified type + if task.task_type == "local": + current_local_task_id += 1 + task_id_to_local_id[task_id] = current_local_task_id + else: # 'data' + current_data_task_id += 1 + task_id_to_data_id[task_id] = current_data_task_id + + # Re-number task IDs and outgoing connections based on CSL IDs + for task_id, task in enumerate(result): + # TODO(later): Task IDs can be recycled with a global ``var`` that can be set prior to activating + # a task, like a state machine + + # NOTE: This check happens after task fusion, so it is less likely to trigger + # if task.task_type == 'data': + # tid = task_id_to_data_id[task_id] + # if tid >= len(constants.DATA_TASK_IDS) and task_creation_behavior == TaskCreationBehavior.FAIL_ON_OVERRUN: + # raise ValueError('Too many data tasks') + # task.task_id = constants.DATA_TASK_IDS[tid] + # elif task.task_type == 'local': + # tid = task_id_to_local_id[task_id] + # if tid >= len(constants.LOCAL_TASK_IDS) and task_creation_behavior == TaskCreationBehavior.FAIL_ON_OVERRUN: + # raise ValueError('Too many local tasks') + # task.task_id = constants.LOCAL_TASK_IDS[tid] + + for i, (target, e) in enumerate(task.outgoing): + # Mark exit task explicitly + if target == len(result): + task.outgoing[i] = (-1, e) + continue + + # Do not modify outgoing task IDs + # if result[target].task_type == 'local': + # target_id = constants.LOCAL_TASK_IDS[task_id_to_local_id[target]] + # else: + # target_id = constants.DATA_TASK_IDS[task_id_to_data_id[target]] + # task.outgoing[i] = (target_id, e) + + # Add explicit terminator statements for tasks that have no successors + sink_tasks = [ + t for i, t in enumerate(result) if not any(n != i for n, _ in t.outgoing) or -1 in set(n for n, _ in t.outgoing) + ] + if len(sink_tasks) > 2: + raise ValueError("Too many sink tasks") + for i, task in enumerate(sink_tasks): + edge_type = InterTaskEdge.ACTIVATE if i == 0 else InterTaskEdge.UNBLOCK + if task_creation_behavior == TaskCreationBehavior.SYNCHRONOUS_ON_OVERRUN and num_local_tasks >= len( + constants.LOCAL_TASK_IDS + ): + edge_type = InterTaskEdge.SEQUENCE + elif task_creation_behavior == TaskCreationBehavior.NO_TASKS: + edge_type = InterTaskEdge.SEQUENCE + elif len(sink_tasks) == 1: # Save on task IDs if there is only one sink task + edge_type = InterTaskEdge.SEQUENCE + + if task.statements[-1] != "TERMINATOR" and task.outgoing[-1][0] != -1 and task.task_type == "local": + task.statements.append("TERMINATOR") + task.outgoing.append((-1, edge_type)) + elif task.outgoing[-1][0] == -1: # Modify existing UNBLOCK edge if there is more than one sink + task.outgoing[-1] = (-1, edge_type) + + # Inject local tasks in front of data tasks with more than one input + to_append = [] + for i, task in enumerate(result): + if task.task_type != "data": + continue + predecessors = set(j for j, t in enumerate(result) if any(o == i for o, _ in t.outgoing)) + if len(predecessors) > 1: + new_id = len(result) + len(to_append) + new_task = CSLTask(i, "local", ["TERMINATOR"], [(new_id, InterTaskEdge.UNBLOCK)], blocked=True) + to_append.append(task) + result[i] = new_task + # Make one of the predecessors into an ACTIVATE edge + tpred = next(iter(predecessors)) + result[tpred].outgoing = [ + (o, InterTaskEdge.ACTIVATE) if o == i else (o, e) for o, e in result[tpred].outgoing + ] + + result.extend(to_append) + + return result + + +def _contract_node(g: nx.DiGraph, n: Any): + if g.out_degree(n) == 0: # Keep sink node + return + for u, _ in g.in_edges(n): + for _, v in g.out_edges(n): + g.add_edge(u, v) + g.remove_node(n) + + +def _canonicalize_dag(completion_dag: nx.DiGraph) -> nx.DiGraph: + completion_dag = deepcopy(completion_dag) + + # Reduce in-degree of nodes to up to 2 + _limit_indegree(completion_dag) + + return completion_dag + + +def _limit_indegree(dag: nx.DiGraph): + """ + Injects extra wait nodes to completion DAGs where the in-degree of a node is larger than two. + + :param dag: The completion DAG. + """ + counter = -1 + for node in list(dag.nodes): # Copy nodes to a list + if dag.in_degree(node) > 2: + edges = list(dag.in_edges(node)) + current_node = node + # Create intermediate wait nodes (the counter changes the statement ID because it has to be unique) + for u, _ in edges[1:]: + new_node = analysis.CompletionDAGNode("wait", counter) + counter -= 1 + dag.remove_edge(u, node) + dag.add_edge(new_node, current_node) + dag.add_edge(u, new_node) + current_node = new_node + + +def fuse_tasks( + tasks: list[CSLTask], + dsds: UniqueDSDDict, + dtypes: dict[spir.Identifier, spir.IRType], + rect, + use_memcpy_mode: bool, + compute: spir.ComputeBlock, +) -> list[CSLTask]: + """ + Fuses tasks where possible to reduce the number of tasks. + + :param tasks: The list of CSL tasks. + :param dsds: The unique DSD dictionary. + :param dtypes: The dictionary of identifier types. + :param kernel: The spatial IR kernel. + :param use_memcpy_mode: Whether memcpy mode is used. + :return: The fused list of CSL tasks. + """ + fused: set[int] = set() + removed: set[int] = set() + redirect: dict[int, int] = {} + for i, task in enumerate(tasks): + if i in fused or i in removed: + continue + if task.task_type == "data": + continue + if not task.statements: + # Remove task + removed.add(i) + continue + outgoing_id, _ = task.outgoing[-1] + if outgoing_id == -1: + # This is a sink task, nothing to fuse with + continue + if outgoing_id in fused or outgoing_id in removed: + continue + if any(et == InterTaskEdge.UNBLOCK for t in tasks for n, et in t.outgoing if n == outgoing_id): + # Cannot fuse if there are multiple predecessors to next task + continue + if tasks[outgoing_id].task_type != "local": + # Cannot fuse with data tasks + continue + + last_stmt = task.statements[-1] + if not isinstance(last_stmt, int): + # Last statement is a terminator, cannot fuse + continue + + # Identify fusion opportunities + next_task = tasks[outgoing_id] if outgoing_id < len(tasks) else None + if next_task is None: + continue + stmt = compute.statements[last_stmt] + if isinstance(stmt, (spir.SendStatement, spir.ReceiveStatement)): + stream = dsd_ops._get_id(stmt.stream_name) + localarr = dsd_ops._get_id(stmt.local_array) + if stream.as_ir() in dsds and isinstance(dsds[stream.as_ir()][0][1], cslstruct.FabricDSD): + # Cannot fuse if the last statement is a fabric DSD operation + continue + if localarr.as_ir() in dsds and isinstance(dsds[localarr.as_ir()][0][1], cslstruct.FabricDSD): + # Cannot fuse if the last statement is a fabric DSD operation + continue + elif isinstance(stmt, spir.ForeachStatement): + # Try to fuse synchronous DSD operations + dsd_op: type[dsd_ops.DSDOp] = dsd_ops.DSD_ASSIGNMENT_MAPPING.get(dsd_ops.get_dsd_op(dtypes, stmt), None) + if dsd_op is None: + continue + # Pure memory DSD operations are synchronous and can be fused + dsd_stmt = dsd_ops.get_dsd_statement(dtypes, stmt) + if dsd_stmt is None: + continue + dsd_objects = dsd_op().used_dsd_objects(dsd_stmt, dsds) + if any(isinstance(dsd, cslstruct.FabricDSD) for dsd in dsd_objects): + continue + # Also check the foreach receive generator + stream = dsd_ops._get_id(stmt.receive_stream.stream_name) + if stream.as_ir() in dsds and isinstance(dsds[stream.as_ir()][0][1], cslstruct.FabricDSD): + continue + else: + # Fusible pattern not found + continue + + # If we can fuse with the next task, do so + fused.add(outgoing_id) + redirect[outgoing_id] = i + task.outgoing[-1] = (i, InterTaskEdge.SEQUENCE) # Remove outgoing edge + task.statements.extend(next_task.statements) + task.outgoing.extend(next_task.outgoing) + + new_tasks: list[CSLTask] = [] + old_to_new: dict[int, int] = {} + for idx, task in enumerate(tasks): + if idx in fused or idx in removed: + continue + new_idx = len(new_tasks) + old_to_new[idx] = new_idx + new_tasks.append(task) + + def resolve_target(target: int) -> int: + while target in redirect: + target = redirect[target] + return target + + for task in new_tasks: + for j, (target, et) in enumerate(task.outgoing): + if target == -1: + continue + resolved = resolve_target(target) + if resolved not in old_to_new: + raise ValueError(f"Dangling task reference {resolved} after fusion") + task.outgoing[j] = (old_to_new[resolved], et) + + return new_tasks + + +def renumber_tasks(tasks: list[CSLTask], task_creation_behavior: TaskCreationBehavior) -> None: + """ + Renumbers tasks to map to hardware task IDs, based on the task creation behavior. + + :param tasks: The list of CSL tasks to operate in-place on. + :param task_creation_behavior: The task creation behavior. + """ + current_local_task_id = -1 + current_data_task_id = -1 + task_id_to_local_id: dict[int, int] = {} + task_id_to_data_id: dict[int, int] = {} + for task_id, task in enumerate(tasks): + # Increment the current task ID and add a new task with the specified type + if task.task_type == "local": + current_local_task_id += 1 + task_id_to_local_id[task_id] = current_local_task_id + else: # 'data' + current_data_task_id += 1 + task_id_to_data_id[task_id] = current_data_task_id + + # Re-number task IDs and outgoing connections based on CSL IDs + for task_id, task in enumerate(tasks): + if task.task_type == "data": + tid = task_id_to_data_id[task_id] + if tid >= len(constants.DATA_TASK_IDS) and task_creation_behavior == TaskCreationBehavior.FAIL_ON_OVERRUN: + raise ValueError("Too many data tasks") + task.task_id = constants.DATA_TASK_IDS[tid] + elif task.task_type == "local": + tid = task_id_to_local_id[task_id] + if tid >= len(constants.LOCAL_TASK_IDS) and task_creation_behavior == TaskCreationBehavior.FAIL_ON_OVERRUN: + raise ValueError("Too many local tasks") + task.task_id = constants.LOCAL_TASK_IDS[tid] diff --git a/spada/syntax/spatial_ir/canonicalization.py b/spada/syntax/spatial_ir/canonicalization.py index 4997a827..b5bd68cf 100644 --- a/spada/syntax/spatial_ir/canonicalization.py +++ b/spada/syntax/spatial_ir/canonicalization.py @@ -211,6 +211,69 @@ def _rewrite_stream_declarations( return replacements, appended_statements +def uniquify_stream_names(kernel: spir.Kernel) -> spir.Kernel: + """ + Gives every stream declaration in the kernel a name of its own. + + Reusing a name across phases is normal in metaprogrammed kernels: each phase a compile-time + ``for`` block generates repeats the same declaration verbatim, so a kernel of a hundred phases + can hold a hundred streams all called ``east``. Those are distinct streams -- each is scoped to + its phase and dies at the phase boundary -- but nothing in the IR says so, and a pass that keys + on the stream name sees one long-lived stream where there are many. + + This pass settles the scoping once, up front: a redeclaration is renamed to a fresh version and + the uses that follow are rewritten to match, up to the next redeclaration of that name. + Afterwards a stream name identifies exactly one declaration kernel-wide, so keying on it is + safe -- which is what lets ``insert_implicit_closes`` decide where a stream is last used by + name alone. ``inline_phases`` performs the same rewriting for the names it merges into one + rectangle; running it here as well leaves it nothing to rename. + + Must run after ``canonicalize_phases`` (so that every dataflow block sits in a phase, which is + what puts the declarations in order) and before ``insert_implicit_closes``. + + :param kernel: The kernel to transform, modified in place. + :return: The transformed kernel. + """ + used_versions: dict[str, set[int]] = defaultdict(set) + # The streams currently in scope, mapping the name a declaration was written with to the name it + # carries now. A phase's own declarations are folded in before its uses are rewritten, so a + # redeclaration shadows the stream it replaces from that phase onwards. + in_scope: dict[spir.Identifier, spir.Identifier] = {} + + for phase in kernel.body: + if not isinstance(phase, spir.Phase): + continue + + # The phase is the unit of scope, not the declaration: one stream is declared once per + # subgrid that takes part in it, so the repeats *within* a phase are the same stream and + # have to keep sharing a name. Only a redeclaration in a later phase is a new stream. + phase_names: dict[spir.Identifier, spir.Identifier] = {} + for dataflow in phase.dataflow: + statements = [] + for statement in dataflow.statements: + stream_name = statement.stream_name + if stream_name not in phase_names: + if stream_name.version in used_versions[stream_name.name]: + phase_names[stream_name] = _make_fresh_identifier(used_versions, stream_name) + else: + phase_names[stream_name] = stream_name + _register_identifier(used_versions, phase_names[stream_name]) + + rewritten_statement = copy.deepcopy(statement) + rewritten_statement.stream_name = copy.deepcopy(phase_names[stream_name]) + statements.append(rewritten_statement) + dataflow.statements = statements + + in_scope.update({old: new for old, new in phase_names.items() if old != new}) + if not in_scope: + continue + + replacer = passes.FindAndReplace(in_scope) + phase.compute = [replacer.visit(compute) for compute in phase.compute] + + return kernel + + def _ends_with_phase_barrier(statements: list[spir.Statement]) -> bool: """ Returns whether a compute block already ends with a phase barrier, so that appending another @@ -229,6 +292,34 @@ def _ends_with_phase_barrier(statements: list[spir.Statement]) -> bool: return False +def number_stream_phases(kernel: spir.Kernel) -> spir.Kernel: + """ + Stamps every stream declaration with the index of the phase it belongs to. + + Router configurations have to be ordered by the epoch they serve, and after ``inline_phases`` + that order is no longer recoverable: a compute block only carries barriers for the phases it + takes part in, so counting them gives a per-rectangle numbering that cannot be compared across + rectangles -- and a PE that merely relays a stream has no statements at all, its configuration + being contributed by the sending rectangle. Numbering the declarations here, while phases are + still explicit, gives routing a kernel-wide order to sort by. + + Must run after ``canonicalize_phases`` and before ``inline_phases``. + + :param kernel: The kernel to annotate, modified in place. + :return: The annotated kernel. + """ + phase_index = 0 + for block in kernel.body: + if not isinstance(block, spir.Phase): + continue + for dataflow in block.dataflow: + for statement in dataflow.statements: + if isinstance(statement, spir.StreamDeclaration): + statement.phase = phase_index + phase_index += 1 + return kernel + + def inline_phases(kernel: spir.Kernel) -> spir.Kernel: """ Inlines phases into their constituent computation and dataflow blocks by adding waits and appending all streams, diff --git a/spada/syntax/spatial_ir/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index 88d719f8..11396838 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -735,11 +735,21 @@ class StreamDeclaration(SpatialNode): dtype: StreamType stream_name: Identifier stream: RelativeStreamDeclaration | MulticastRangeStreamDeclaration | ExternStreamDeclaration + #: Index of the phase this stream is declared in, counted over the whole kernel. Filled in by + #: ``canonicalization.number_stream_phases`` while phases are still explicit, and used + #: afterwards to order a router's configurations: once phases are inlined, a compute block only + #: carries barriers for the phases *it* takes part in, so its local barrier count is not + #: comparable with another block's. A relay PE has no statements at all, and its configuration + #: is contributed by the sending rectangle, so only a kernel-wide index orders the two. + #: ``None`` on streams that never went through the pass. Not part of the surface syntax. + phase: Optional[int] = None def validate(self) -> None: assert isinstance(self.dtype, StreamType) assert isinstance(self.stream_name, Identifier) assert isinstance(self.stream, (RelativeStreamDeclaration, MulticastRangeStreamDeclaration, ExternStreamDeclaration)) + if self.phase is not None: + assert isinstance(self.phase, int) and self.phase >= 0 def as_ir(self, indent: int = 0) -> str: indent_str = ' ' * indent diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py index c8ef3be0..0f75f843 100644 --- a/spada/syntax/spatial_ir/stream_lifetime.py +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -136,7 +136,8 @@ def insert_implicit_closes(kernel: spir.Kernel) -> spir.Kernel: waiting on operations that are still using those very streams. Must run after ``canonicalize_phases`` (so that ``kernel.body`` contains only phases and place - blocks) and before ``inline_phases``. + blocks) and ``uniquify_stream_names`` (so that a stream name means one stream), and before + ``inline_phases``. :param kernel: The kernel to transform, modified in place. :return: The transformed kernel. @@ -147,9 +148,10 @@ def insert_implicit_closes(kernel: spir.Kernel) -> spir.Kernel: phases = [block for block in kernel.body if isinstance(block, spir.Phase)] - # Find, for each (rectangle, stream), the last phase in which the rectangle uses the stream. - # A stream declared at kernel level stays in scope across phases, so it may only be closed - # after its final use. + # Find, for each (rectangle, stream), the last phase in which the rectangle uses the stream: a + # stream that outlives the phase it was declared in may only be closed after its final use. + # ``uniquify_stream_names`` has already given every declaration a name of its own, so a name + # that recurs across phases really is one stream and not a redeclaration wearing the same name. last_phase: dict[tuple[tuple[int, int, int, int], spir.Identifier], int] = {} uses_per_block: list[list[tuple[spir.ComputeBlock, dict[spir.Identifier, StreamUse]]]] = [] for phase_index, phase in enumerate(phases): @@ -167,8 +169,8 @@ def insert_implicit_closes(kernel: spir.Kernel) -> spir.Kernel: for compute, uses in blocks: rect = compute.get_grid_rect() to_close = [ - use for name, use in uses.items() - if name in declared and use.uses and use.close is None and last_phase[(rect, name)] == phase_index + use for name, use in uses.items() if name in declared and use.uses and use.close is None and + last_phase.get((rect, name), phase_index) == phase_index ] if not to_close: continue @@ -542,14 +544,17 @@ def check_channel_conflicts(rectangles: list[Rectangle]) -> None: :param rectangles: The consolidated PE rectangles of the kernel. """ uses_per_rect = [collect_stream_uses(rect.metadata.compute) for rect in rectangles] - # Per rectangle, the local name each stream group goes by - groups_per_rect: list[dict[str, spir.Identifier]] = [] + # Per rectangle, every local name a stream group goes by. A group generally has more than one: + # ``inline_phases`` freshens colliding names, so a channel reused by the same routing signature + # in several phases shows up as ``east``, ``east#1``, ... within one compute block. + groups_per_rect: list[dict[str, list[spir.Identifier]]] = [] for rect, uses in zip(rectangles, uses_per_rect): declarations = _stream_declarations(rect) - groups_per_rect.append({ - stream_group_key(declarations[name]): name - for name in uses if name in declarations - }) + names_by_group: dict[str, list[spir.Identifier]] = defaultdict(list) + for name in uses: + if name in declarations: + names_by_group[stream_group_key(declarations[name])].append(name) + groups_per_rect.append(dict(names_by_group)) reported: set[tuple[str, str]] = set() for (channel, pe), occupants in sorted(channel_occupancy(rectangles).items()): @@ -560,8 +565,7 @@ def check_channel_conflicts(rectangles: list[Rectangle]) -> None: for first, second in _ordered_pairs(groups): if (first, second) in reported: continue - if _empties_before(first, second, uses_per_rect, groups_per_rect) or \ - _empties_before(second, first, uses_per_rect, groups_per_rect): + if _never_concurrent(first, second, uses_per_rect, groups_per_rect): continue reported.add((first, second)) first_decl = _find_declaration(rectangles, first) @@ -591,32 +595,45 @@ def _find_declaration(rectangles: list[Rectangle], group: str) -> Optional[spir. return None -def _empties_before(first: str, second: str, uses_per_rect: list[dict[spir.Identifier, StreamUse]], - groups_per_rect: list[dict[str, spir.Identifier]]) -> bool: +def _never_concurrent(first: str, second: str, uses_per_rect: list[dict[spir.Identifier, StreamUse]], + groups_per_rect: list[dict[str, list[spir.Identifier]]]) -> bool: """ - Returns whether the stream group ``first`` provably empties before the stream group ``second``. + Returns whether two stream groups provably never occupy their shared channel at the same time. - This requires ``first`` to be closed on every PE that uses it, and, wherever both streams are - used on the same PE, for that close to precede the first use of ``second`` in local order. + Each use of a stream opens an epoch that runs until its close, so a group contributes one + interval ``[first use, close]`` per local name it goes by. The groups are safely ordered when + + * every one of those epochs is closed -- an unclosed stream holds the channel indefinitely, so + nothing can be said about what follows it, on this PE or on any relay further along its path; + and + * no epoch of one group overlaps an epoch of the other in local order. + + The two groups may alternate any number of times, which is what a channel reused by successive + phases does: ``east``, close, ``west``, close, ``east#1``, close, ... + + This is best-effort in the sense of the specification: it establishes the ordering where it can + and does not reject what it cannot decide. """ used_anywhere = False for uses, groups in zip(uses_per_rect, groups_per_rect): - first_name = groups.get(first) - if first_name is None: - continue - first_use = uses[first_name] - if not first_use.uses: - continue - used_anywhere = True - if first_use.close is None: - return False - - second_name = groups.get(second) - if second_name is None: - continue - second_use = uses[second_name] - if second_use.uses and first_use.close > second_use.first_use: - return False + epochs: list[tuple[int, int, str]] = [] + for group in (first, second): + for name in groups.get(group, ()): + use = uses[name] + if not use.uses: + continue # Declared and closed but never actually used here + if use.close is None: + return False + used_anywhere = True + epochs.append((use.first_use, use.close, group)) + + epochs.sort() + for index, (start, end, group) in enumerate(epochs): + for other_start, _, other_group in epochs[index + 1:]: + if other_group == group: + continue + if other_start < end: + return False return used_anywhere diff --git a/tests/csl_runtime/test_bitonic_sort_1d.sh b/tests/csl_runtime/test_bitonic_sort_1d.sh new file mode 100644 index 00000000..311b8f43 --- /dev/null +++ b/tests/csl_runtime/test_bitonic_sort_1d.sh @@ -0,0 +1,77 @@ +#!/bin/sh +# E2E test: Batcher's bitonic sorting network on 2^L PEs (bitonic_sort_1D.sptl). +# +# This is the heaviest user of channel reuse in the suite. The network runs L(L+1)/2 +# compare-exchange steps at distances 1, 2, ..., 2^(L-1), split into lanes so that concurrent +# paths stay disjoint, and every lane runs two epochs -- one eastward, one westward -- over a +# single channel per distance. That is log2(N) channels for N keys, and it makes each router +# cycle through sender, relay and receiver configurations, reversing direction every epoch. +# +# WSE-3 only: reversing a router changes both its input and its output direction. On WSE-2 a +# switch position carries only one of the two, so each reversal costs two positions and the +# interior routers need far more than the four a router holds. The compiler rejects it there with +# a capacity error, which `test_bitonic_sort_1d_rejected_on_wse2` in +# tests/spatial_ir/test_routing.py asserts. +# +# Reference: OUT_out == sorted(a_in). + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +. "$SCRIPT_DIR/_lib.sh" + +if [ "${WSE_ARCH:-wse2}" != "wse3" ]; then + echo "Skipping: bitonic_sort_1D needs switch positions that carry both directions (WSE-3)." + echo " re-run with WSE_ARCH=wse3 to exercise it." + exit 0 +fi + +L=2 +N=4 +K=4 +FOLDER="bitonic_sort_1d_sptl" +SAMPLES_DIR="$(cd "$SCRIPT_DIR/../../samples/spatial/sorting" && pwd)" + +sptlc "$SAMPLES_DIR/bitonic_sort_1D.sptl" "$FOLDER" -p L=$L -p K=$K + +# One channel per exchange distance, so exactly L colors carry the network. +colors=$(grep -o '@get_color([0-9]*)' "$FOLDER/layout.csl" | sort -u | wc -l) +if [ "$colors" -gt "$L" ]; then + echo "Test failed: expected at most $L colors for the network, found $colors." + exit 1 +fi + +# Reuse must be realized by switch positions, and the periodic lane structure by ring mode. +grep -q '\.switches' "$FOLDER/layout.csl" || { + echo "Test failed: no router switch configuration was generated." + exit 1 +} +grep -q 'ring_mode' "$FOLDER/layout.csl" || { + echo "Test failed: the repeating lane pattern did not collapse into a switch ring." + exit 1 +} + +python3 - < dict[str, str]: + kernel = parser.parse_file(_BITONIC) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, L=L, K=K)) + return {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel)} + + +@pytest.mark.skipif(not csl.SWITCH_POSITION_ALLOWS_BOTH, + reason=f'{csl.ARCH} cannot reverse a router within four switch positions') +def test_bitonic_sort_uses_one_channel_per_distance(): + """ + A bitonic network on 2^L keys needs L(L+1)/2 exchange steps but only L channels: one per + exchange distance, reused by every lane, every stage and both directions of travel. + + L is 2 here because that is what the router budget allows: at distance 2^d the channel is + reused by 2^d lanes in two directions each, so an interior router cycles through 2^(d+1) + configurations, and four positions run out at d = 2. + """ + files = _lower_bitonic(2) + colors = set(re.findall(r'@get_color\((\d+)\)', files['layout.csl'])) + assert len(colors) <= 2, sorted(colors) + + layout = files['layout.csl'] + assert '.switches' in layout + # The lane pattern repeats, so the configuration sequence closes into a ring. + assert 'ring_mode' in layout + # Every router stays inside its four positions. + for line in layout.splitlines(): + if '.switches' in line: + assert len(re.findall(r'\.pos\d', line)) < csl.SWITCH_POSITIONS, line + + +@pytest.mark.skipif(csl.SWITCH_POSITION_ALLOWS_BOTH, + reason='this architecture can reverse a router in a single switch position') +def test_bitonic_sort_is_rejected_on_wse2(): + """ + Reversing a router costs two positions where a position carries one direction, and the interior + routers of the network reverse often enough to exhaust them. + """ + with pytest.raises(SyntaxError, match='switch positions'): + _lower_bitonic(2) + + if __name__ == '__main__': pytest.main([__file__]) From 7a13d6133e3da210c00ad35becaf53acc3f3246a Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 17 Aug 2026 14:07:38 +0200 Subject: [PATCH 10/68] Initial spmv, 1D batcher sort --- README.md | 3 +- samples/spatial/blas/spmv.sptl | 140 +++++++++ samples/spatial/sort/batcher_oddeven_1D.sptl | 133 ++++++++ samples/spatial/sort/plot_batcher_routing.py | 313 +++++++++++++++++++ tests/csl_runtime/test_batcher_oddeven_1d.sh | 49 +++ tests/csl_runtime/test_spmv.sh | 100 ++++++ 6 files changed, 737 insertions(+), 1 deletion(-) create mode 100644 samples/spatial/blas/spmv.sptl create mode 100644 samples/spatial/sort/batcher_oddeven_1D.sptl create mode 100644 samples/spatial/sort/plot_batcher_routing.py create mode 100644 tests/csl_runtime/test_batcher_oddeven_1d.sh create mode 100644 tests/csl_runtime/test_spmv.sh diff --git a/README.md b/README.md index b40ff77a..f9ba38d8 100644 --- a/README.md +++ b/README.md @@ -102,8 +102,9 @@ Sample SPADA programs are in `samples/`: | `samples/advanced_stencils.py` | GT4Py definitions for horizontal diffusion kernels | | `samples/benchmarks/` | Pre-compiled `.spst`/`.sptl` pairs for five kernels at five domain sizes | | `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy` | -| `samples/spatial/blas/` | Dense linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase` | +| `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | +| `samples/spatial/sort/` | Sorting networks: `batcher_oddeven_1D` | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | diff --git a/samples/spatial/blas/spmv.sptl b/samples/spatial/blas/spmv.sptl new file mode 100644 index 00000000..fc7994ba --- /dev/null +++ b/samples/spatial/blas/spmv.sptl @@ -0,0 +1,140 @@ +/** + * Distributed sparse GEMV: y = alpha * A * x + beta * y + * + * Grid [0:PX, 0:PY] — PX PEs along the i/x axis, PY PEs along the j/y axis. + * + * Data layout (same 1.5D blocking as gemv.sptl, except A blocks are COO): + * A: (PY*K) rows × (PX*K) cols — PE(i,j) holds a K×K block of A in padded COO + * A_val[NZ], A_row[NZ], A_col[NZ], with unused slots (val=0, row=col=0) + * x: PX*K elements — PE(i,0) initially holds x[i*K:(i+1)*K] + * y: PY*K elements — PE(0,j) initially holds y[j*K:(j+1)*K] + * + * Algorithm: + * Phase 1: All PEs load their padded COO block of A from the host. + * Phase 2: The j=0 column loads x; the i=0 column loads y from the host. + * Phase 3: Multicast x in the Y direction (j=0 → j=PY-1) using native hardware multicast. + * Phase 4: Each PE computes its local contribution z = A_block @ x as a + * compile-time COO loop over NZ entries. + * Phase 5: Pipelined chain reduction of z in the X direction (i=PX-1 → i=0). + * The root PE(0,j) applies alpha*z + beta*y and writes to the host. + * + * Constraints: PX >= 2, PY >= 2, K >= 1, NZ >= 1 + * A_row[p], A_col[p] must lie in [0, K). + **/ +kernel @spmv( + stream[PX, 1] readonly inp_x, // x blocks: PE(i,0) for i=0..PX-1 + stream[PX, PY] readonly inp_A_val, // COO values, padded to NZ + stream[PX, PY] readonly inp_A_row, // COO row indices in [0, K) + stream[PX, PY] readonly inp_A_col, // COO column indices in [0, K) + stream[1, PY] readonly inp_y, // y blocks: PE(0,j) for j=0..PY-1 + f32 alpha, // scalar multiplier for A*x + f32 beta, // scalar multiplier for y + stream[1, PY] writeonly out // result: PE(0,j) for j=0..PY-1 +) { + place i16 i, i16 j in [0:PX, 0:PY] { + f32[NZ] A_val // COO nonzero values (zero-padded) + i16[NZ] A_row // COO row indices + i16[NZ] A_col // COO column indices + f32[K] x // x chunk for this PE column (populated by multicast) + f32[K] z // local partial result; accumulated during reduction + f32[K] y_block // y chunk; read from host, used only at i=0 + } + + // Phase 1: Load COO A blocks on every PE. + phase { + compute i16 i, i16 j in [0:PX, 0:PY] { + await receive(A_val, inp_A_val[i, j]) + await receive(A_row, inp_A_row[i, j]) + await receive(A_col, inp_A_col[i, j]) + } + } + + // Phase 2: Load x on the j=0 column and y on the i=0 column. + phase { + compute i16 i, i16 j in [0:PX, 0] { + await receive(x, inp_x[i, j]) + } + compute i16 i, i16 j in [0, 0:PY] { + await receive(y_block, inp_y[i, j]) + } + } + + // Phase 3: Multicast x in the Y direction (j=0 → j=PY-1). + // Uses CSL native multicast: each intermediate PE forwards to both RAMP and + // the next PE southward, so all PY-1 receivers are served in a single phase. + phase { + dataflow i16 i, i16 j in [0:PX, 0:PY] { + stream bcast = relative_stream(0, [1:PY]) { + hops = auto, + channel = 0 + } + } + + compute i16 i, i16 j in [0:PX, 0:1] { + await send(x, bcast) + } + + compute i16 i, i16 j in [0:PX, 1:PY] { + await receive(x, bcast) + } + } + + // Phase 4: Local COO SpMV: z[A_row[p]] += A_val[p] * x[A_col[p]]. + phase { + compute i16 i, i16 j in [0:PX, 0:PY] { + for i16 k in [0:K] { + z[k] = 0.0 + } + for i16 p in [0:NZ] { + z[A_row[p]] = z[A_row[p]] + A_val[p] * x[A_col[p]] + } + } + } + + // Phase 5: Pipelined chain reduction of z in the X direction (i=PX-1 → i=0). + // Root PE(0,j) applies y[k] = alpha*z[k] + beta*y_block[k] and writes to host. + phase { + dataflow i16 i, i16 j in [0:PX, 0:PY] { + stream yellow = relative_stream(-1, 0) { + hops = [(-1, 0)], + channel = 1 + } + stream green = relative_stream(-1, 0) { + hops = [(-1, 0)], + channel = 2 + } + } + + // East column (i=PX-1): start the reduction. + compute i16 i, i16 j in [PX-1, 0:PY] { + await send(z, yellow if (PX-1) % 2 == 0 else green) + } + + // Odd i: receive yellow, accumulate, forward on green. + compute i16 i, i16 j in [1:PX-1:2, 0:PY] { + await foreach i16 k, f32 v in [0:K], receive(yellow) { + z[k] = z[k] + v + await send(z[k], green) + } + } + + // Even i (middle): receive green, accumulate, forward on yellow. + compute i16 i, i16 j in [2:PX-1:2, 0:PY] { + await foreach i16 k, f32 v in [0:K], receive(green) { + z[k] = z[k] + v + await send(z[k], yellow) + } + } + + // Root i=0: accumulate, apply alpha/beta scaling, output result. + compute i16 i, i16 j in [0, 0:PY] { + await foreach i16 k, f32 v in [0:K], receive(green) { + z[k] = z[k] + v + } + for i16 k in [0:K] { + z[k] = alpha * z[k] + beta * y_block[k] + } + await send(z, out[i, j]) + } + } +} diff --git a/samples/spatial/sort/batcher_oddeven_1D.sptl b/samples/spatial/sort/batcher_oddeven_1D.sptl new file mode 100644 index 00000000..7d05268c --- /dev/null +++ b/samples/spatial/sort/batcher_oddeven_1D.sptl @@ -0,0 +1,133 @@ +/** + * 1D Batcher odd-even mergesort over N = 2^L PEs, one f32 key per PE. + * Ascending: after the network, PE i holds the i-th smallest input. + * + * Merge stages l = 1 .. L, each of width 2^l. Within stage l: + * p = 1: every PE in each 2^l box compares at dist = 2^{l-1} + * p = 2 .. l: skip box endpoints, compare at dist = 2^{l-p} + * + * Each drawn comparator is two messages (low -> high, then high -> low). + * Lower index keeps min; higher index keeps max. + * + * Static channel assignment (no router reconfiguration): + * At distance d, the d interleaved matchings (offset r = 0 .. d-1) overlap + * on the 1D mesh, so each (d, r) pair gets its own colors: + * fwd (east, +d) : channel 2*(((d - 1) + r)) + * bwd (west, -d) : channel 2*(((d - 1) + r)) + 1 + * Stages that share the same d reuse those colors (identical hop table). + * Total colors = 2*(N - 1). Intended for small N (e.g. L <= 3). + * + * Example L=3 (n=8): + * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) + * (l=2,p=1) dist 2: (0,2)(1,3)(4,6)(5,7) + * (l=2,p=2) dist 1: (1,2)(5,6) + * (l=3,p=1) dist 4: (0,4)(1,5)(2,6)(3,7) + * (l=3,p=2) dist 2: (2,4)(3,5) + * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) + * + * Constraints: L >= 1 + **/ +kernel @batcher_oddeven_1d( + stream[1<[1< fwd = relative_stream(1<<(l-1), 0) { + hops = auto, + channel = 2 * (((1<<(l-1)) - 1) + r) + } + stream bwd = relative_stream(-(1<<(l-1)), 0) { + hops = auto, + channel = (2 * (((1<<(l-1)) - 1) + r)) + 1 + } + } + dataflow i16 i, i16 j in [(r + (1<<(l-1))):1< fwd = relative_stream(1<<(l-1), 0) { + hops = auto, + channel = 2 * (((1<<(l-1)) - 1) + r) + } + stream bwd = relative_stream(-(1<<(l-1)), 0) { + hops = auto, + channel = (2 * (((1<<(l-1)) - 1) + r)) + 1 + } + } + + compute i16 i, i16 j in [r:1< val else val + } + } + } + + // p = 2 .. l: skip first/last dist PEs of each 2^l box, dist = 1<<(l-p). + for i16 p in [2:l+1] { + phase { + for i16 r in [0:1<<(l-p)] { + for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { + hops = auto, + channel = 2 * (((1<<(l-p)) - 1) + r) + } + stream bwd = relative_stream(-(1<<(l-p)), 0) { + hops = auto, + channel = (2 * (((1<<(l-p)) - 1) + r)) + 1 + } + } + dataflow i16 i, i16 j in [(b + ((2 * (1<<(l-p))) + r)):(b + (1< fwd = relative_stream(1<<(l-p), 0) { + hops = auto, + channel = 2 * (((1<<(l-p)) - 1) + r) + } + stream bwd = relative_stream(-(1<<(l-p)), 0) { + hops = auto, + channel = (2 * (((1<<(l-p)) - 1) + r)) + 1 + } + } + + compute i16 i, i16 j in [(b + ((1<<(l-p)) + r)):((b + (1< val else val + } + } + } + } + } + } + + // Write sorted keys back to the host. + phase { + compute i16 i, i16 j in [0:1< int: + return 2 * ((dist - 1) + offset) + + +def bwd_channel(dist: int, offset: int) -> int: + return fwd_channel(dist, offset) + 1 + + +def batcher_phases(n: int) -> list[Phase]: + if n < 2 or n & (n - 1): + raise ValueError(f"n must be a power of two >= 2, got {n}") + log_n = int(math.log2(n)) + phases: list[Phase] = [] + index = 0 + for l in range(1, log_n + 1): + dist = 1 << (l - 1) + matchings = [] + for r in range(dist): + pairs = tuple((i, i + dist) for i in range(r, n, 1 << l)) + matchings.append( + Matching(l, 1, dist, r, pairs, fwd_channel(dist, r), bwd_channel(dist, r)) + ) + phases.append(Phase(index, l, 1, dist, tuple(matchings))) + index += 1 + for p in range(2, l + 1): + dist = 1 << (l - p) + matchings = [] + for r in range(dist): + pairs = [] + for b in range(0, n, 1 << l): + start = b + dist + r + stop = b + (1 << l) - dist + r + for lo in range(start, stop, 2 * dist): + pairs.append((lo, lo + dist)) + matchings.append( + Matching(l, p, dist, r, tuple(pairs), fwd_channel(dist, r), bwd_channel(dist, r)) + ) + phases.append(Phase(index, l, p, dist, tuple(matchings))) + index += 1 + return phases + + +def channel_color(channel: int, n_channels: int, cmap_name: str = "tab20"): + cmap = plt.get_cmap(cmap_name) + if n_channels <= 20: + return cmap(channel % 20) + return plt.get_cmap("gist_ncar")(channel / max(n_channels - 1, 1)) + + +def _dir_arrow(ax, x: float, y_from: float, y_to: float, color) -> None: + ax.annotate( + "", + xy=(x, y_to), + xytext=(x, y_from), + arrowprops=dict(arrowstyle="-|>", color=color, lw=1.6, mutation_scale=9, shrinkA=1.5, shrinkB=1.5), + zorder=2, + ) + + +def _draw_network(ax, phases: list[Phase], n: int) -> None: + """Draw concurrent matchings as sub-columns; fwd and bwd as offset arrows.""" + n_channels = 2 * (n - 1) + slot = 0.38 + gap = 0.55 + dx = 0.06 + origins = [] + x = 0.0 + for ph in phases: + origins.append(x) + x += max(len(ph.matchings), 1) * slot + gap + x_end = x - gap + + ax.set_xlim(-0.7, x_end + 0.4) + ax.set_ylim(n - 0.5, -0.5) + ax.set_yticks(range(n)) + ax.set_yticklabels([str(i) for i in range(n)]) + ax.set_ylabel("PE index") + ax.set_xlabel("Phase (left arrow ↓ fwd, right arrow ↑ bwd)") + ax.set_title(f"Batcher network, n={n}: downward = fwd (east), upward = bwd (west)") + + tick_pos = [] + tick_lab = [] + for ph, x0 in zip(phases, origins): + n_m = max(len(ph.matchings), 1) + width = n_m * slot + ax.axvspan(x0 - 0.10, x0 + width - slot + 0.10, color="0.93", zorder=0) + tick_pos.append(x0 + (n_m - 1) * slot / 2) + tick_lab.append(f"l={ph.l}\np={ph.p}\nd={ph.dist}") + + ax.set_xticks(tick_pos) + ax.set_xticklabels(tick_lab, fontsize=8) + + for pe in range(n): + ax.plot([-0.5, x_end + 0.2], [pe, pe], color="0.78", lw=0.8, zorder=1) + + cap = 0.05 + for ph, x0 in zip(phases, origins): + for k, m in enumerate(ph.matchings): + x = x0 + k * slot + fwd_c = channel_color(m.fwd, n_channels) + bwd_c = channel_color(m.bwd, n_channels) + for lo, hi in m.pairs: + x_fwd = x - dx + x_bwd = x + dx + _dir_arrow(ax, x_fwd, lo, hi, fwd_c) + _dir_arrow(ax, x_bwd, hi, lo, bwd_c) + ax.plot([x_fwd - cap, x_fwd + cap], [lo, lo], color=fwd_c, lw=1.6, zorder=3) + ax.plot([x_fwd - cap, x_fwd + cap], [hi, hi], color=fwd_c, lw=1.6, zorder=3) + ax.plot([x_bwd - cap, x_bwd + cap], [lo, lo], color=bwd_c, lw=1.6, zorder=3) + ax.plot([x_bwd - cap, x_bwd + cap], [hi, hi], color=bwd_c, lw=1.6, zorder=3) + + handles = [] + for ch in range(n_channels): + arrow = "↓" if ch % 2 == 0 else "↑" + kind = "fwd" if ch % 2 == 0 else "bwd" + handles.append( + Line2D( + [0], + [0], + color=channel_color(ch, n_channels), + lw=2.0, + label=f"{arrow} ch {ch} ({kind})", + ) + ) + ax.legend( + handles=handles, + title="channel", + loc="upper left", + bbox_to_anchor=(1.02, 1), + fontsize=7, + ncol=1 if n_channels <= 16 else 2, + ) + + +PAIR_ORDER = ("R→E", "W→E", "W→R", "R→W", "E→W", "E→R") +PAIR_COLOR = { + "R→E": "#1f77b4", + "W→E": "#ff7f0e", + "W→R": "#2ca02c", + "R→W": "#5fa8d3", + "E→W": "#f4a261", + "E→R": "#6dce6d", +} + + +def _draw_table(ax, phases: list[Phase], n: int) -> None: + """Static per-PE @set_color_config: each cell is the union of rx→tx pairs.""" + n_channels = 2 * (n - 1) + routes: list[list[set[str]]] = [[set() for _ in range(n_channels)] for _ in range(n)] + for ph in phases: + for m in ph.matchings: + for lo, hi in m.pairs: + routes[lo][m.fwd].add("R→E") + routes[hi][m.fwd].add("W→R") + for mid in range(lo + 1, hi): + routes[mid][m.fwd].add("W→E") + routes[hi][m.bwd].add("R→W") + routes[lo][m.bwd].add("E→R") + for mid in range(lo + 1, hi): + routes[mid][m.bwd].add("E→W") + + ax.set_xlim(-0.5, n_channels - 0.5) + ax.set_ylim(n - 0.5, -0.5) + ax.set_xticks(range(n_channels)) + ax.set_yticks(range(n)) + ax.set_xlabel("Channel") + ax.set_ylabel("PE") + ax.set_title(f"Static rx→tx table, n={n} ({n_channels} colors); stacked pairs are a union") + ax.set_aspect("equal") + + for pe in range(n): + for ch in range(n_channels): + pairs = [p for p in PAIR_ORDER if p in routes[pe][ch]] + if not pairs: + ax.add_patch( + Rectangle( + (ch - 0.45, pe - 0.45), + 0.9, + 0.9, + facecolor="#f4f4f4", + edgecolor="0.85", + lw=0.4, + ) + ) + continue + band = 0.9 / len(pairs) + fontsize = 6 if len(pairs) == 1 else 5 + for i, pair in enumerate(pairs): + y0 = pe - 0.45 + i * band + ax.add_patch( + Rectangle( + (ch - 0.45, y0), + 0.9, + band, + facecolor=PAIR_COLOR[pair], + edgecolor="0.85", + lw=0.4, + ) + ) + ax.text(ch, y0 + band / 2, pair, ha="center", va="center", fontsize=fontsize, color="0.1") + + handles = [ + Line2D([0], [0], marker="s", color="w", markerfacecolor=PAIR_COLOR[p], markersize=10, label=p) + for p in PAIR_ORDER + ] + ax.legend(handles=handles, title="rx→tx", loc="upper left", bbox_to_anchor=(1.02, 1), fontsize=8) + + +def plot(n: int, view: str, outfile: str | None, show: bool) -> None: + phases = batcher_phases(n) + if view == "network": + n_slots = sum(max(len(ph.matchings), 1) for ph in phases) + fig, ax = plt.subplots(figsize=(max(8, 0.7 * n_slots + 0.8 * len(phases)), max(4, 0.45 * n))) + _draw_network(ax, phases, n) + elif view == "table": + fig, ax = plt.subplots(figsize=(max(8, 0.45 * 2 * (n - 1)), max(4, 0.45 * n))) + _draw_table(ax, phases, n) + else: + raise ValueError(f"unknown view {view}") + + fig.tight_layout() + if outfile: + fig.savefig(outfile, bbox_inches="tight") + png = outfile[:-4] + ".png" if outfile.endswith(".pdf") else outfile + ".png" + if outfile.endswith(".pdf"): + fig.savefig(png, dpi=160, bbox_inches="tight") + print(f"wrote {outfile}") + if show: + plt.show() + plt.close(fig) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--n", type=int, default=8, help="number of PEs, power of two (default 8)") + parser.add_argument( + "--view", + choices=("network", "table"), + default="network", + help="network: sorting-network comparators; table: static color config", + ) + parser.add_argument("--out", default=None, help="output PDF path") + parser.add_argument("--show", action="store_true", help="open an interactive window") + args = parser.parse_args() + outfile = args.out + if outfile is None and not args.show: + outfile = f"samples/spatial/sort/batcher_routing_n{args.n}_{args.view}.pdf" + plot(args.n, args.view, outfile, args.show) + + +if __name__ == "__main__": + main() diff --git a/tests/csl_runtime/test_batcher_oddeven_1d.sh b/tests/csl_runtime/test_batcher_oddeven_1d.sh new file mode 100644 index 00000000..ecdcc108 --- /dev/null +++ b/tests/csl_runtime/test_batcher_oddeven_1d.sh @@ -0,0 +1,49 @@ +#!/bin/sh +# E2E test: 1D Batcher odd-even mergesort (2^L PEs, one f32 key per PE). +# Kernel: batcher_oddeven_1D.sptl params: L +# Reference: OUT_out[:, 0, 0] == sort(inp[:, 0, 0]) +# Tested with L ∈ {1, 2, 3}. + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +. "$SCRIPT_DIR/_lib.sh" + +SORT_DIR="$(cd "$(dirname "$0")/../../samples/spatial/sort" && pwd)" +FOLDER="batcher_oddeven_1d_sptl" + +run_batcher() { + l=$1 + echo "--- batcher_oddeven_1d L=$l ---" + + sptlc "$SORT_DIR/batcher_oddeven_1D.sptl" "$FOLDER" -p L=$l + + python3 - < Date: Mon, 17 Aug 2026 14:57:21 +0200 Subject: [PATCH 11/68] Initial draft of bundles optimization --- irspec/docs/spatial/routing.md | 8 + irspec/docs/spatial/spatial.md | 10 +- .../sort/batcher_oddeven_bundled_1D.sptl | 141 ++++++ spada/lowering/spatial_ir_to_csl.py | 68 +++ spada/syntax/spatial_ir/canonicalization.py | 5 +- spada/syntax/spatial_ir/irnodes.py | 40 +- spada/syntax/spatial_ir/language.lark | 6 +- spada/syntax/spatial_ir/lark_to_ir.py | 29 +- spada/syntax/spatial_ir/shift_bundles.py | 383 ++++++++++++++++ tests/spatial_ir/test_shift_bundles.py | 415 ++++++++++++++++++ 10 files changed, 1099 insertions(+), 6 deletions(-) create mode 100644 samples/spatial/sort/batcher_oddeven_bundled_1D.sptl create mode 100644 spada/syntax/spatial_ir/shift_bundles.py create mode 100644 tests/spatial_ir/test_shift_bundles.py diff --git a/irspec/docs/spatial/routing.md b/irspec/docs/spatial/routing.md index 22f6130c..4152d5ad 100644 --- a/irspec/docs/spatial/routing.md +++ b/irspec/docs/spatial/routing.md @@ -156,6 +156,14 @@ Recall that sending onto the same stream [must be synchronized using completions to avoid data races](../spatial#streaming-data-with-send). Hence, sending through the same stream multiple times in the same phases is ok as long as the sends (and receives) are correctly synchronized. +When a stream declares an explicit compile-time `count = k`, the compiler may +serialize overlapping 1D interval shifts onto one channel per direction by +switching the router after a known number of fabric waves (forward-then-inject +on the source half, absorb-then-forward on the dest half). That program is +tied to the wave quotas of this phase, so a later phase with a different +shift distance uses a distinct color pair. `count = auto` (the default) does +not enable this: the stream is treated as unbounded. + Keep in mind that PEs transition between phases asynchronously, that is, a PE may advance to the next phase before another PE has completed the current phase. We exploit here implicitly that routers back-pressure when diff --git a/irspec/docs/spatial/spatial.md b/irspec/docs/spatial/spatial.md index db18b48e..e5cb7c1d 100644 --- a/irspec/docs/spatial/spatial.md +++ b/irspec/docs/spatial/spatial.md @@ -488,12 +488,20 @@ The routing configuration is set up as follows: stream stream_name = relative_stream(dx, dy) { // Optional routing declaration hops = [(dx_1, dy_1), (dx_2, dy_2), ... , (dx_n, dy_n)], - channel = channel_id + channel = channel_id, + count = k } ``` where `hops` is a list of relative hops that the data takes between the sender and receiver. Each hop is given by a pair of constant literals, the sum of their absolute value must be 1. The sum of all the hops must be equal to the relative position of the stream. +`count` is optional. If it is a compile-time integer `k`, each PE transfers exactly `k` +fabric words on this stream edge in the phase, which enables counted router switching +on overlapping 1D shifts. Switching uses two colors **per phase** (one per direction). +Wave quotas depend on the shift distance, so later phases with a different `d` receive +a fresh color pair rather than reloading the same two colors. If `count` is omitted or +`count = auto`, the stream is unbounded: the compiler does not infer a length and does +not apply counted switching. If two messages (elements of a `send`) are routed through a PE simultaneously, it must be ensured that they do not share a `channel`. diff --git a/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl b/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl new file mode 100644 index 00000000..0a0e71d6 --- /dev/null +++ b/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl @@ -0,0 +1,141 @@ +/** + * 1D Batcher odd-even mergesort over N = 2^L PEs, one f32 key per PE. + * Ascending: after the network, PE i holds the i-th smallest input. + * + * Merge stages l = 1 .. L, each of width 2^l. Within stage l: + * p = 1: every PE in each 2^l box compares at dist = 2^{l-1} + * p = 2 .. l: skip box endpoints, compare at dist = 2^{l-p} + * + * Each drawn comparator is two messages (low -> high, then high -> low). + * Lower index keeps min; higher index keeps max. + * + * Counted color switching (count = 1): overlapping matchings at distance d + * share one eastbound and one westbound color in that phase. Each PE + * forwards a known number of waves, then injects or absorbs. Distance-1 + * stages keep a static RAMP/EAST (WEST) pair. + * Each CAS phase uses 2 colors. Colors are not reused across phases: + * the wave quotas change with d, so a counted two-state program cannot + * be reloaded onto the same color. Total colors = 2 * (# of (l,p) phases). + * + * Example L=3 (n=8): + * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) + * (l=2,p=1) dist 2: (0,2)(1,3)(4,6)(5,7) + * (l=2,p=2) dist 1: (1,2)(5,6) + * (l=3,p=1) dist 4: (0,4)(1,5)(2,6)(3,7) + * (l=3,p=2) dist 2: (2,4)(3,5) + * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) + * + * Constraints: L >= 1 + **/ +kernel @batcher_oddeven_1d( + stream[1<[1< fwd = relative_stream(1<<(l-1), 0) { + hops = auto, + channel = auto, + count = 1 + } + stream bwd = relative_stream(-(1<<(l-1)), 0) { + hops = auto, + channel = auto, + count = 1 + } + } + dataflow i16 i, i16 j in [(r + (1<<(l-1))):1< fwd = relative_stream(1<<(l-1), 0) { + hops = auto, + channel = auto, + count = 1 + } + stream bwd = relative_stream(-(1<<(l-1)), 0) { + hops = auto, + channel = auto, + count = 1 + } + } + + compute i16 i, i16 j in [r:1< val else val + } + } + } + + // p = 2 .. l: skip first/last dist PEs of each 2^l box, dist = 1<<(l-p). + for i16 p in [2:l+1] { + phase { + for i16 r in [0:1<<(l-p)] { + for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { + hops = auto, + channel = auto, + count = 1 + } + stream bwd = relative_stream(-(1<<(l-p)), 0) { + hops = auto, + channel = auto, + count = 1 + } + } + dataflow i16 i, i16 j in [(b + ((2 * (1<<(l-p))) + r)):(b + (1< fwd = relative_stream(1<<(l-p), 0) { + hops = auto, + channel = auto, + count = 1 + } + stream bwd = relative_stream(-(1<<(l-p)), 0) { + hops = auto, + channel = auto, + count = 1 + } + } + + compute i16 i, i16 j in [(b + ((1<<(l-p)) + r)):((b + (1< val else val + } + } + } + } + } + } + + // Write sorted keys back to the host. + phase { + compute i16 i, i16 j in [0:1< spir.Kernel: kernel = canonicalization.reduce_streams(kernel) kernel = canonical_subgrids.canonicalize_subgrids(kernel) kernel = canonicalization.resolve_auto_hops(kernel) + from spada.syntax.spatial_ir.shift_bundles import coalesce_shift_bundles + coalesce_shift_bundles(kernel) kernel = canonicalization.inline_phases(kernel) return kernel @@ -156,6 +158,7 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, # Collect unique routes for all rectangles routes_per_rectangle = _collect_routes(rectangles, color_maps) + shift_schedule_code = _emit_shift_schedules(kernel, channel_to_color) if use_memcpy_mode: layout_code.write(f''' @@ -254,6 +257,9 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, for rinst in routing_instructions: layout_code.write(rinst + '\n') + if shift_schedule_code: + layout_code.write(shift_schedule_code) + # Emit symbol names for arguments and kernel layout_code.write('\n // Extern fields\n') # Gather extern fields from kernel arguments @@ -1236,6 +1242,10 @@ def _collect_routes(rectangles: list[Rectangle[PEBlock]], color_maps: list[dict[ if isinstance(stream.stream, spir.ExternStreamDeclaration): continue # Extern streams do not have on-chip routing + if stream.stream.routing is not None and stream.stream.routing.counted_switch: + # Counted interval shifts are emitted from kernel.shift_schedules. + continue + if isinstance(stream.stream, spir.MulticastRangeStreamDeclaration): if sent and received: raise ValueError( @@ -1390,6 +1400,64 @@ def _coord(k): # noqa: E731 return result +def _emit_shift_schedules(kernel: spir.Kernel, channel_to_color: dict[int, int]) -> str: + """ + Emit absolute @set_color_config plus spa_switch_after / spa_color_schedule + annotations for counted 1D shift bundles. + + Each color carries one two-state program (one phase). Wave quotas differ + across phases, so the same color is never reloaded with a new count. + """ + schedules = getattr(kernel, "shift_schedules", None) + if not schedules: + return "" + + from collections import defaultdict + from spada.syntax.spatial_ir.shift_bundles import ColorSchedule + + by_pe: dict[tuple[int, int, int], list[ColorSchedule]] = defaultdict(list) + for sched in schedules: + by_pe[(sched.x, sched.y, sched.channel)].append(sched) + + lines = ["\n // Counted shift-bundle color schedules\n"] + seen_config: set[str] = set() + for (x, y, channel), pe_scheds in sorted(by_pe.items()): + pe_scheds = sorted(pe_scheds, key=lambda s: s.phase_index) + if channel not in channel_to_color: + continue + if len(pe_scheds) > 1: + phases = [sched.phase_index for sched in pe_scheds] + raise ValueError( + f"Color {channel} on PE ({x},{y}) is used in phases {phases}; " + "counted switching assigns a distinct color pair per phase " + "because wave quotas differ." + ) + color_expr = f"@get_color({channel_to_color[channel]})" + sched = pe_scheds[0] + step_txt = " ; ".join( + f"{step.as_pair()} waves={step.waves}" for step in sched.steps + ) + lines.append( + f" // spa_color_schedule phase={sched.phase_index} pe={x},{y} " + f"ch={channel} : {step_txt}\n" + ) + first = sched.steps[0] + config = ( + f" @set_color_config({x}, {y}, {color_expr}, " + f".{{ .routes = .{{ .rx = .{{{first.rx}}}, .tx = .{{{first.tx}}} }} }});\n" + ) + if config not in seen_config: + lines.append(config) + seen_config.add(config) + if len(sched.steps) > 1: + nxt = sched.steps[1] + lines.append( + f" // spa_switch_after phase={sched.phase_index} pe={x},{y} " + f"ch={channel} waves={first.waves} rx={nxt.rx} tx={nxt.tx}\n" + ) + return "".join(lines) + + def _write_indented_block(current_code: StringIO, block: str, indent: str) -> None: block = textwrap.dedent(block).strip('\n') if not block: diff --git a/spada/syntax/spatial_ir/canonicalization.py b/spada/syntax/spatial_ir/canonicalization.py index 998a88b7..d8ccec89 100644 --- a/spada/syntax/spatial_ir/canonicalization.py +++ b/spada/syntax/spatial_ir/canonicalization.py @@ -311,11 +311,14 @@ def inline_phases(kernel: spir.Kernel) -> spir.Kernel: raise TypeError(f'Unexpected block type "{type(block).__name__}" in kernel. Was ``canonicalize_phases`` ' 'called?') - return spir.Kernel( + new_kernel = spir.Kernel( name=kernel.name, parameters=copy.deepcopy(kernel.parameters), arguments=copy.deepcopy(kernel.arguments), body=list(rect_place.values()) + list(rect_dataflow.values()) + list(rect_compute.values())) + if hasattr(kernel, "shift_schedules"): + new_kernel.shift_schedules = kernel.shift_schedules + return new_kernel @dataclass diff --git a/spada/syntax/spatial_ir/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index 28d029ef..89e7905e 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -566,7 +566,7 @@ def as_ir(self, indent: int = 0) -> str: @dataclass class RoutingDeclaration(SpatialNode): """ - A routing declaration for a stream, optionally specifying hops and channel. + A routing declaration for a stream, optionally specifying hops, channel, and count. The ``channel`` field may hold: * ``"auto"`` – the channel number is assigned automatically. @@ -576,9 +576,18 @@ class RoutingDeclaration(SpatialNode): meta-for loop variable such as ``stage``). It must evaluate to an integer by the time CSL lowering runs; use :attr:`resolved_channel` to obtain the concrete value. + + The ``count`` field is the number of fabric words on this stream edge per PE + per phase. ``"auto"`` (the default, also used when the field is omitted) means + the stream is unbounded: the compiler does not infer a length and does not + apply counted router switching. Counted switching requires an explicit + compile-time integer ``count``. """ hops: Union[list[RoutingHop], Literal["auto"]] = "auto" # list of hops or 'auto' channel: Union["Expression", int, Literal["auto"]] = "auto" + count: Union["Expression", int, Literal["auto"]] = "auto" + # Set by the shift-bundle pass; not part of the surface language. + counted_switch: bool = False def validate(self) -> None: if isinstance(self.hops, list): @@ -605,6 +614,23 @@ def resolved_channel(self) -> Union[int, Literal["auto"]]: ) return val + @property + def resolved_count(self) -> Union[int, Literal["auto"]]: + """ + Return the message count as a concrete integer, or ``"auto"`` if unbounded. + """ + if self.count == "auto": + return "auto" + if isinstance(self.count, int): + return self.count + val = self.count.eval() + if not isinstance(val, int): + raise ValueError( + f"Count expression '{self.count.as_ir()}' did not evaluate to an integer. " + "Ensure all parameters and loop variables are concretized before counted switching." + ) + return val + def as_ir(self, indent: int = 0) -> str: indent_str = ' ' * indent hops_str = "auto" if self.hops == "auto" else f"[{', '.join(hop.as_ir() for hop in self.hops)}]" @@ -614,7 +640,17 @@ def as_ir(self, indent: int = 0) -> str: channel_str = str(self.channel) else: channel_str = self.channel.as_ir() - return f"{indent_str}hops = {hops_str}, \n{indent_str}channel = {channel_str}" + lines = [ + f"{indent_str}hops = {hops_str}", + f"{indent_str}channel = {channel_str}", + ] + if self.count != "auto": + if isinstance(self.count, int): + count_str = str(self.count) + else: + count_str = self.count.as_ir() + lines.append(f"{indent_str}count = {count_str}") + return ", \n".join(lines) @dataclass diff --git a/spada/syntax/spatial_ir/language.lark b/spada/syntax/spatial_ir/language.lark index 4b4a981b..3df05056 100644 --- a/spada/syntax/spatial_ir/language.lark +++ b/spada/syntax/spatial_ir/language.lark @@ -116,7 +116,11 @@ subgrid_expression_2d : "[" range_expression "," range_expression "]" !direction : "in" | "out" hop : "(" posneg_integer_literal "," posneg_integer_literal ")" // 2D at the moment, might expand hops : "[" hop ("," hop)* "]" -routing : "hops" "=" (auto | hops) "," "channel" "=" (auto | value_expr) +routing_hops : "hops" "=" (auto | hops) +routing_channel : "channel" "=" (auto | value_expr) +routing_count : "count" "=" (auto | value_expr) +routing_field : routing_hops | routing_channel | routing_count +routing : routing_field ("," routing_field)* multicast_range : "[" range_expression "]" relative_stream_declaration : "relative_stream" "(" (value_expr | multicast_range) "," (value_expr | multicast_range) ")" ("{" routing "}")? extern_stream_declaration : "extern_stream" "(" direction ")" ("{" routing "}")? diff --git a/spada/syntax/spatial_ir/lark_to_ir.py b/spada/syntax/spatial_ir/lark_to_ir.py index da877497..b685d87e 100644 --- a/spada/syntax/spatial_ir/lark_to_ir.py +++ b/spada/syntax/spatial_ir/lark_to_ir.py @@ -177,7 +177,6 @@ def field_declaration(self, args): def hop(self, args): return irnodes.RoutingHop(tuple(args)) - routing = irnodes.RoutingDeclaration.from_lark subgrid_expression_2d = irnodes.SubgridExpression.from_lark def hop(self, args): @@ -315,6 +314,34 @@ def parameters(self, args): dataflow_body = list phase_body = list + def routing_hops(self, args): + return ('hops', args[0]) + + def routing_channel(self, args): + return ('channel', args[0]) + + def routing_count(self, args): + return ('count', args[0]) + + def routing_field(self, args): + return args[0] + + def routing(self, args): + kwargs = {'hops': 'auto', 'channel': 'auto', 'count': 'auto'} + seen: set[str] = set() + for key, value in args: + if key in seen: + raise ValueError(f'Duplicate routing field "{key}"') + seen.add(key) + kwargs[key] = value + if 'hops' not in seen or 'channel' not in seen: + raise ValueError('Routing declaration requires both hops and channel') + return irnodes.RoutingDeclaration( + hops=kwargs['hops'], + channel=kwargs['channel'], + count=kwargs['count'], + ) + def compute_body(self, args): if len(args) == 1 and isinstance(args[0], list): return args[0] diff --git a/spada/syntax/spatial_ir/shift_bundles.py b/spada/syntax/spatial_ir/shift_bundles.py new file mode 100644 index 00000000..9d71c60c --- /dev/null +++ b/spada/syntax/spatial_ir/shift_bundles.py @@ -0,0 +1,383 @@ +""" +Detect consecutive 1D interval shifts and schedule counted two-state color switching. + +A stream with an explicit ``count = k`` whose senders form a consecutive interval +``[L, L+m)`` at distance ``d`` (with ``1 < m <= d``) can share one color: each PE +forwards a known number of waves, then injects or absorbs. ``count = auto`` is +unbounded and is never rewritten. +""" +from __future__ import annotations + +from collections import defaultdict +from dataclasses import dataclass, replace +from typing import Literal + +from spada.syntax.spatial_ir import analysis +from spada.syntax.spatial_ir import irnodes as spir + + +@dataclass(frozen=True) +class ShiftBundle: + """A consecutive interval of sources shifted by an axis-aligned distance.""" + + phase_index: int + axis: Literal["x", "y"] + sign: Literal[1, -1] + start: int + length: int + dist: int + count: int + fixed: int + stream_names: tuple[spir.Identifier, ...] + # Physical channel, unique to this phase and direction. Detect leaves 0; + # apply_shift_bundles assigns a per-phase pair (fwd, bwd). + channel: int = 0 + + +@dataclass(frozen=True) +class ColorScheduleStep: + """One router config and how many fabric waves it stays active.""" + + rx: str + tx: str + waves: int + + def as_pair(self) -> str: + return f"{_short(self.rx)}->{_short(self.tx)}" + + +@dataclass(frozen=True) +class ColorSchedule: + """Per-PE counted switch program for one phase and logical channel.""" + + phase_index: int + x: int + y: int + channel: int + steps: tuple[ColorScheduleStep, ...] + + +def _short(port: str) -> str: + return {"RAMP": "R", "EAST": "E", "WEST": "W", "NORTH": "N", "SOUTH": "S"}.get(port, port) + + +def _resolved_count(routing: spir.RoutingDeclaration | None) -> int | None: + if routing is None: + return None + count = routing.resolved_count + if count == "auto": + return None + if count < 1: + raise ValueError(f"Routing count must be a positive integer, got {count}") + return count + + +def _axis_offset(stream: spir.RelativeStreamDeclaration) -> tuple[Literal["x", "y"], int] | None: + dx = stream.dx.eval() + dy = stream.dy.eval() + if not isinstance(dx, int) or not isinstance(dy, int): + return None + if dx != 0 and dy == 0: + return "x", dx + if dy != 0 and dx == 0: + return "y", dy + return None + + +def _hops_are_straight(stream: spir.RelativeStreamDeclaration, axis: str, delta: int) -> bool: + routing = stream.routing + if routing is None or routing.hops == "auto": + return True + if not isinstance(routing.hops, list): + return False + step = 1 if delta > 0 else -1 + expected = [(step, 0)] * abs(delta) if axis == "x" else [(0, step)] * abs(delta) + actual = [hop.offset for hop in routing.hops] + return actual == expected + + +def _block_points(block) -> list[tuple[int, int]]: + x0, x1, y0, y1 = block.get_grid_rect() + xs, ys = block.get_grid_stride() + return [(x, y) for x in range(x0, x1, xs) for y in range(y0, y1, ys)] + + +def _consecutive_runs(values: list[int]) -> list[tuple[int, int]]: + if not values: + return [] + ordered = sorted(set(values)) + runs = [] + start = prev = ordered[0] + for value in ordered[1:]: + if value == prev + 1: + prev = value + continue + runs.append((start, prev - start + 1)) + start = prev = value + runs.append((start, prev - start + 1)) + return runs + + +def detect_shift_bundles(kernel: spir.Kernel) -> list[ShiftBundle]: + """ + Find consecutive interval shifts with an explicit count in each phase. + + :param kernel: A kernel whose metaprogramming and auto-hops are already resolved, + and whose phases have not yet been inlined. + """ + bundles: list[ShiftBundle] = [] + phase_index = 0 + for block in kernel.body: + if not isinstance(block, spir.Phase): + continue + bundles.extend(_detect_in_phase(block, phase_index)) + phase_index += 1 + return bundles + + +def _detect_in_phase(phase: spir.Phase, phase_index: int) -> list[ShiftBundle]: + decls: dict[spir.Identifier, spir.RelativeStreamDeclaration] = {} + for dataflow in phase.dataflow: + for stmt in dataflow.statements: + stream = stmt.stream + if not isinstance(stream, spir.RelativeStreamDeclaration): + continue + existing = decls.get(stmt.stream_name) + if existing is None: + decls[stmt.stream_name] = stream + continue + if (existing.dx.eval(), existing.dy.eval()) != (stream.dx.eval(), stream.dy.eval()): + raise ValueError( + f'Stream "{stmt.stream_name.as_ir()}" is declared with conflicting offsets ' + f"in the same phase." + ) + + # (axis, delta, k, fixed_coord) -> (source coords along the axis, stream names) + groups: dict[tuple[str, int, int, int], tuple[set[int], set[spir.Identifier]]] = {} + for compute in phase.compute: + sent_recv = analysis.sends_and_receives(compute) + for name, stream in decls.items(): + sent, _received = sent_recv.get(name, (False, False)) + if not sent: + continue + k = _resolved_count(stream.routing) + if k is None: + continue + axis_delta = _axis_offset(stream) + if axis_delta is None: + continue + axis, delta = axis_delta + if not _hops_are_straight(stream, axis, delta): + continue + for x, y in _block_points(compute): + if axis == "x": + key = (axis, delta, k, y) + coord = x + else: + key = (axis, delta, k, x) + coord = y + bucket = groups.get(key) + if bucket is None: + bucket = (set(), set()) + groups[key] = bucket + bucket[0].add(coord) + bucket[1].add(name) + + bundles: list[ShiftBundle] = [] + for (axis, delta, k, fixed), (coords, names) in groups.items(): + dist = abs(delta) + sign: Literal[1, -1] = 1 if delta > 0 else -1 + for start, length in _consecutive_runs(list(coords)): + if length < 2 or length > dist: + continue + bundles.append( + ShiftBundle( + phase_index=phase_index, + axis=axis, + sign=sign, + start=start, + length=length, + dist=dist, + count=k, + fixed=fixed, + stream_names=tuple(names), + ) + ) + return bundles + + +def schedule_counted_switch(bundle: ShiftBundle) -> list[ColorSchedule]: + """ + Build the two-state per-PE schedule for one interval shift. + + East/south (+d): source half forwards then injects; dest half absorbs then forwards. + West/north (−d): the same lemma with the opposite ports, sources on the far side. + """ + if bundle.axis == "x": + forward_rx, forward_tx = ("WEST", "EAST") if bundle.sign > 0 else ("EAST", "WEST") + inject_tx = "EAST" if bundle.sign > 0 else "WEST" + absorb_rx = "WEST" if bundle.sign > 0 else "EAST" + + def pe(coord: int) -> tuple[int, int]: + return coord, bundle.fixed + else: + forward_rx, forward_tx = ("NORTH", "SOUTH") if bundle.sign > 0 else ("SOUTH", "NORTH") + inject_tx = "SOUTH" if bundle.sign > 0 else "NORTH" + absorb_rx = "NORTH" if bundle.sign > 0 else "SOUTH" + + def pe(coord: int) -> tuple[int, int]: + return bundle.fixed, coord + + L = bundle.start + m = bundle.length + d = bundle.dist + k = bundle.count + schedules: dict[tuple[int, int], list[ColorScheduleStep]] = {} + + def add(coord: int, steps: list[ColorScheduleStep]) -> None: + x, y = pe(coord) + kept = [step for step in steps if step.waves > 0] + if kept: + schedules[(x, y)] = kept + + if bundle.sign > 0: + for j in range(m): + src = L + j + add(src, [ + ColorScheduleStep(forward_rx, forward_tx, j * k), + ColorScheduleStep("RAMP", inject_tx, k), + ]) + for j in range(m): + dest = L + d + j + add(dest, [ + ColorScheduleStep(absorb_rx, "RAMP", k), + ColorScheduleStep(forward_rx, forward_tx, (m - 1 - j) * k), + ]) + for coord in range(L + m, L + d): + add(coord, [ColorScheduleStep(forward_rx, forward_tx, m * k)]) + else: + # Sources occupy [L, L+m); dests are at source - d. + for j in range(m): + src = L + j + add(src, [ + ColorScheduleStep(forward_rx, forward_tx, (m - 1 - j) * k), + ColorScheduleStep("RAMP", inject_tx, k), + ]) + for j in range(m): + dest = L - d + j + add(dest, [ + ColorScheduleStep(absorb_rx, "RAMP", k), + ColorScheduleStep(forward_rx, forward_tx, j * k), + ]) + for coord in range(L - d + m, L): + add(coord, [ColorScheduleStep(forward_rx, forward_tx, m * k)]) + + return [ + ColorSchedule(bundle.phase_index, x, y, bundle.channel, tuple(steps)) + for (x, y), steps in sorted(schedules.items()) + ] + + +def _phase_has_explicit_count_relative(phase: spir.Phase) -> bool: + for dataflow in phase.dataflow: + for stmt in dataflow.statements: + stream = stmt.stream + if not isinstance(stream, spir.RelativeStreamDeclaration): + continue + if _resolved_count(stream.routing) is not None: + return True + return False + + +def _phase_channel_bases(kernel: spir.Kernel, bundles: list[ShiftBundle]) -> dict[int, int]: + """ + Give each routed phase its own even/odd color pair. + + Wave quotas depend on the shift distance, so a counted two-state program + cannot be reused on the same color in a later phase. + """ + needed = {bundle.phase_index for bundle in bundles} + phase_index = 0 + for block in kernel.body: + if not isinstance(block, spir.Phase): + continue + if _phase_has_explicit_count_relative(block): + needed.add(phase_index) + phase_index += 1 + return {phase: 2 * i for i, phase in enumerate(sorted(needed))} + + +def apply_shift_bundles(kernel: spir.Kernel, bundles: list[ShiftBundle]) -> list[ColorSchedule]: + """ + Mark bundled streams for counted switching, assign two colors per phase, + and return the per-PE schedules. + """ + bases = _phase_channel_bases(kernel, bundles) + assigned = [ + replace(bundle, channel=bases[bundle.phase_index] + (0 if bundle.sign > 0 else 1)) + for bundle in bundles + ] + + names_by_phase: dict[int, dict[spir.Identifier, int]] = defaultdict(dict) + for bundle in assigned: + for name in bundle.stream_names: + names_by_phase[bundle.phase_index][name] = bundle.channel + + phase_index = 0 + for block in kernel.body: + if not isinstance(block, spir.Phase): + continue + rewrite = names_by_phase.get(phase_index, {}) + if rewrite: + for dataflow in block.dataflow: + for stmt in dataflow.statements: + channel = rewrite.get(stmt.stream_name) + if channel is None or stmt.stream.routing is None: + continue + stmt.stream.routing.channel = channel + stmt.stream.routing.counted_switch = True + phase_index += 1 + + schedules: list[ColorSchedule] = [] + for bundle in assigned: + schedules.extend(schedule_counted_switch(bundle)) + _assign_unit_hop_channels(kernel, bases) + return schedules + + +def _assign_unit_hop_channels(kernel: spir.Kernel, bases: dict[int, int]) -> None: + """Assign this phase's fwd/bwd pair to explicit-count distance-1 streams.""" + phase_index = 0 + for block in kernel.body: + if not isinstance(block, spir.Phase): + continue + base = bases.get(phase_index) + phase_index += 1 + if base is None: + continue + for dataflow in block.dataflow: + for stmt in dataflow.statements: + stream = stmt.stream + if not isinstance(stream, spir.RelativeStreamDeclaration) or stream.routing is None: + continue + if stream.routing.counted_switch: + continue + if _resolved_count(stream.routing) is None: + continue + axis_delta = _axis_offset(stream) + if axis_delta is None: + continue + _axis, delta = axis_delta + if abs(delta) != 1: + continue + stream.routing.channel = base + (0 if delta > 0 else 1) + + +def coalesce_shift_bundles(kernel: spir.Kernel) -> list[ColorSchedule]: + """ + Detect shift bundles, rewrite their channels, and attach schedules to ``kernel``. + """ + bundles = detect_shift_bundles(kernel) + schedules = apply_shift_bundles(kernel, bundles) + kernel.shift_schedules = schedules + return schedules diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py new file mode 100644 index 00000000..8ffb7602 --- /dev/null +++ b/tests/spatial_ir/test_shift_bundles.py @@ -0,0 +1,415 @@ +"""Tests for count= routing and counted 1D shift-bundle switching.""" +import os +import re + +import pytest + +from spada.lowering.spatial_ir_to_csl import lower_spatial_ir_to_csl +from spada.syntax.spatial_ir import canonical_subgrids, canonicalization, irnodes as spir, parser, passes +from spada.syntax.spatial_ir.shift_bundles import ( + detect_shift_bundles, + schedule_counted_switch, + coalesce_shift_bundles, +) + + +def _prepare(src: str, **params: int): + kernel = parser.parse_string(src) + kernel = passes.concretize_parameters(kernel, **params) + kernel = passes.constexpr_propagation(kernel) + kernel = canonicalization.inline_metaprogramming(kernel) + kernel = canonicalization.canonicalize_phases(kernel) + kernel = canonicalization.reduce_streams(kernel) + kernel = canonical_subgrids.canonicalize_subgrids(kernel) + kernel = canonicalization.resolve_auto_hops(kernel) + return kernel + + +_SHIFT = """ +kernel @shift_k( + stream[N, 1] readonly inp, + stream[N, 1] writeonly out +) { + place i16 i, i16 j in [0:N, 0] { + f32 val + } + phase { + compute i16 i, i16 j in [0:N, 0] { + await receive(val, inp[i, j]) + } + } + phase { + dataflow i16 i, i16 j in [0:4, 0] { + stream fwd = relative_stream(4, 0) { + hops = auto, + channel = auto, + count = 1 + } + } + compute i16 i, i16 j in [0:4, 0] { + await send(val, fwd) + } + compute i16 i, i16 j in [4:8, 0] { + await receive(val, fwd) + } + } + phase { + compute i16 i, i16 j in [0:N, 0] { + await send(val, out[i, j]) + } + } +} +""" + +_SHIFT_AUTO = _SHIFT.replace("count = 1", "count = auto") + +_SHIFT_OMITTED = _SHIFT.replace(",\n count = 1", "") + +_SHIFT_TOO_LONG = _SHIFT.replace("[0:4, 0]", "[0:6, 0]").replace("relative_stream(4, 0)", "relative_stream(2, 0)") + + +def test_parse_count_roundtrip(): + kernel = parser.parse_string(_SHIFT) + ir_1 = kernel.as_ir() + assert "count = 1" in ir_1 + ir_2 = parser.parse_string(ir_1).as_ir() + assert ir_1 == ir_2 + + +def test_omitted_count_has_no_count_line(): + kernel = parser.parse_string(_SHIFT_OMITTED) + assert "count =" not in kernel.as_ir() + + +def test_detect_interval_shift(): + kernel = _prepare(_SHIFT, N=8) + bundles = detect_shift_bundles(kernel) + assert len(bundles) == 1 + b = bundles[0] + assert (b.start, b.length, b.dist, b.count, b.sign, b.axis) == (0, 4, 4, 1, 1, "x") + + +def test_auto_count_is_not_rewritten(): + kernel = _prepare(_SHIFT_AUTO, N=8) + assert detect_shift_bundles(kernel) == [] + coalesce_shift_bundles(kernel) + assert kernel.shift_schedules == [] + + +def test_omitted_count_is_not_rewritten(): + kernel = _prepare(_SHIFT_OMITTED, N=8) + assert detect_shift_bundles(kernel) == [] + + +def test_m_greater_than_d_is_rejected(): + kernel = _prepare(_SHIFT_TOO_LONG, N=8) + assert detect_shift_bundles(kernel) == [] + + +def test_schedule_source_and_dest_halves(): + kernel = _prepare(_SHIFT, N=8) + bundle = detect_shift_bundles(kernel)[0] + by_pe = {(s.x, s.y): s.steps for s in schedule_counted_switch(bundle)} + assert [(st.rx, st.tx, st.waves) for st in by_pe[(0, 0)]] == [("RAMP", "EAST", 1)] + assert [(st.rx, st.tx, st.waves) for st in by_pe[(1, 0)]] == [("WEST", "EAST", 1), ("RAMP", "EAST", 1)] + assert [(st.rx, st.tx, st.waves) for st in by_pe[(3, 0)]] == [("WEST", "EAST", 3), ("RAMP", "EAST", 1)] + assert [(st.rx, st.tx, st.waves) for st in by_pe[(4, 0)]] == [("WEST", "RAMP", 1), ("WEST", "EAST", 3)] + assert [(st.rx, st.tx, st.waves) for st in by_pe[(7, 0)]] == [("WEST", "RAMP", 1)] + + +def test_westbound_schedule(): + west = """ +kernel @shift_w( + stream[N, 1] readonly inp, + stream[N, 1] writeonly out +) { + place i16 i, i16 j in [0:N, 0] { f32 val } + phase { + compute i16 i, i16 j in [0:N, 0] { await receive(val, inp[i, j]) } + } + phase { + dataflow i16 i, i16 j in [4:8, 0] { + stream bwd = relative_stream(-4, 0) { + hops = auto, + channel = auto, + count = 1 + } + } + compute i16 i, i16 j in [4:8, 0] { await send(val, bwd) } + compute i16 i, i16 j in [0:4, 0] { await receive(val, bwd) } + } + phase { + compute i16 i, i16 j in [0:N, 0] { await send(val, out[i, j]) } + } +} +""" + kernel = _prepare(west, N=8) + bundle = detect_shift_bundles(kernel)[0] + assert bundle.sign == -1 and bundle.start == 4 and bundle.length == 4 + by_pe = {(s.x, s.y): s.steps for s in schedule_counted_switch(bundle)} + assert [(st.rx, st.tx, st.waves) for st in by_pe[(7, 0)]] == [("RAMP", "WEST", 1)] + assert [(st.rx, st.tx, st.waves) for st in by_pe[(4, 0)]] == [("EAST", "WEST", 3), ("RAMP", "WEST", 1)] + assert [(st.rx, st.tx, st.waves) for st in by_pe[(0, 0)]] == [("EAST", "RAMP", 1)] + assert [(st.rx, st.tx, st.waves) for st in by_pe[(3, 0)]] == [("EAST", "RAMP", 1), ("EAST", "WEST", 3)] + + +def _batcher_prepared(n_log: int): + path = os.path.join( + os.path.dirname(__file__), "..", "..", "samples", "spatial", "sort", "batcher_oddeven_1D.sptl" + ) + kernel = parser.parse_file(path) + kernel = passes.concretize_parameters(kernel, L=n_log) + kernel = passes.constexpr_propagation(kernel) + kernel = canonicalization.inline_metaprogramming(kernel) + kernel = canonicalization.canonicalize_phases(kernel) + kernel = canonicalization.reduce_streams(kernel) + kernel = canonical_subgrids.canonicalize_subgrids(kernel) + kernel = canonicalization.resolve_auto_hops(kernel) + return kernel + + +def test_batcher_n16_p1_d8_bundle(): + kernel = _batcher_prepared(4) + bundles = detect_shift_bundles(kernel) + east = [b for b in bundles if b.sign > 0 and b.dist == 8] + west = [b for b in bundles if b.sign < 0 and b.dist == 8] + assert len(east) == 1 and east[0].start == 0 and east[0].length == 8 + assert len(west) == 1 and west[0].start == 8 and west[0].length == 8 + + +def test_batcher_n16_p2_d4_bundle(): + kernel = _batcher_prepared(4) + bundles = detect_shift_bundles(kernel) + east = [b for b in bundles if b.sign > 0 and b.dist == 4] + assert any(b.start == 4 and b.length == 4 for b in east) + + +def test_batcher_d1_has_no_counted_switch(): + kernel = _batcher_prepared(3) + bundles = detect_shift_bundles(kernel) + assert all(b.dist != 1 for b in bundles) + + +def _xy(bundle, coord): + return (coord, bundle.fixed) if bundle.axis == "x" else (bundle.fixed, coord) + + +def _consume(states, xy, role): + steps = states[xy] + assert steps, f"PE {xy} has no remaining config but needs {role}" + rx, tx, waves = steps[0] + if role == "inject": + assert rx == "RAMP" and tx != "RAMP", (xy, steps[0], role) + elif role == "absorb": + assert tx == "RAMP" and rx != "RAMP", (xy, steps[0], role) + elif role == "forward": + assert rx != "RAMP" and tx != "RAMP", (xy, steps[0], role) + else: + raise ValueError(role) + waves -= 1 + if waves == 0: + steps.pop(0) + else: + steps[0] = (rx, tx, waves) + + +def simulate_bundle_delivery(bundle, drop_second_step=False): + """ + West-first (eastbound) / east-first (westbound) serial delivery. + + Each fabric word consumes one wave at every PE on its path. After the + source-half forward quota the PE must have switched to inject; after the + dest-half absorb quota it must have switched to forward. Exhausting every + quota means the two-state program matches the matching. + """ + states = {} + for sched in schedule_counted_switch(bundle): + steps = [(st.rx, st.tx, st.waves) for st in sched.steps] + if drop_second_step: + steps = steps[:1] + states[(sched.x, sched.y)] = steps + order = range(bundle.length) if bundle.sign > 0 else range(bundle.length - 1, -1, -1) + for j in order: + src = bundle.start + j + dest = src + bundle.sign * bundle.dist + for _ in range(bundle.count): + _consume(states, _xy(bundle, src), "inject") + hop = src + bundle.sign + while hop != dest: + _consume(states, _xy(bundle, hop), "forward") + hop += bundle.sign + _consume(states, _xy(bundle, dest), "absorb") + leftover = {pe: steps for pe, steps in states.items() if steps} + assert leftover == {}, leftover + + +def test_two_state_switch_delivers_eastbound(): + kernel = _prepare(_SHIFT, N=8) + simulate_bundle_delivery(detect_shift_bundles(kernel)[0]) + + +def test_two_state_switch_delivers_westbound(): + west = """ +kernel @shift_w( + stream[N, 1] readonly inp, + stream[N, 1] writeonly out +) { + place i16 i, i16 j in [0:N, 0] { f32 val } + phase { + compute i16 i, i16 j in [0:N, 0] { await receive(val, inp[i, j]) } + } + phase { + dataflow i16 i, i16 j in [4:8, 0] { + stream bwd = relative_stream(-4, 0) { + hops = auto, + channel = auto, + count = 1 + } + } + compute i16 i, i16 j in [4:8, 0] { await send(val, bwd) } + compute i16 i, i16 j in [0:4, 0] { await receive(val, bwd) } + } + phase { + compute i16 i, i16 j in [0:N, 0] { await send(val, out[i, j]) } + } +} +""" + kernel = _prepare(west, N=8) + simulate_bundle_delivery(detect_shift_bundles(kernel)[0]) + + +def test_two_state_switch_delivers_count_2(): + kernel = _prepare(_SHIFT.replace("count = 1", "count = 2"), N=8) + bundle = detect_shift_bundles(kernel)[0] + assert bundle.count == 2 + simulate_bundle_delivery(bundle) + + +def test_without_the_switch_delivery_fails(): + kernel = _prepare(_SHIFT, N=8) + with pytest.raises(AssertionError): + simulate_bundle_delivery(detect_shift_bundles(kernel)[0], drop_second_step=True) + + +def test_batcher_d8_switch_delivers(): + kernel = _batcher_prepared(4) + bundles = detect_shift_bundles(kernel) + east = next(b for b in bundles if b.sign > 0 and b.dist == 8) + west = next(b for b in bundles if b.sign < 0 and b.dist == 8) + simulate_bundle_delivery(east) + simulate_bundle_delivery(west) + + +def test_two_state_lemma_never_needs_a_third_config(): + kernel = _prepare(_SHIFT, N=8) + bundle = detect_shift_bundles(kernel)[0] + for sched in schedule_counted_switch(bundle): + assert 1 <= len(sched.steps) <= 2 + if len(sched.steps) == 2: + first, second = sched.steps + forward_then_inject = first.rx != "RAMP" and first.tx != "RAMP" and second.rx == "RAMP" + absorb_then_forward = first.tx == "RAMP" and second.rx != "RAMP" and second.tx != "RAMP" + assert forward_then_inject or absorb_then_forward + + +def _onchip_channels_by_phase(kernel): + phases = [] + for block in kernel.body: + if not isinstance(block, spir.Phase): + continue + chans = set() + for dataflow in block.dataflow: + for stmt in dataflow.statements: + routing = getattr(stmt.stream, "routing", None) + if routing is None or routing.resolved_channel == "auto": + continue + chans.add(routing.resolved_channel) + phases.append(chans) + return phases + + +def test_batcher_two_colors_per_phase_not_globally(): + kernel = _batcher_prepared(3) + coalesce_shift_bundles(kernel) + routed = [chans for chans in _onchip_channels_by_phase(kernel) if chans] + assert len(routed) == 6 # L=3 has 6 (l,p) CAS phases + for chans in routed: + assert len(chans) == 2 + used = [c for chans in routed for c in chans] + assert len(used) == len(set(used)) + assert set(used) == set(range(12)) + + +def test_lowering_encodes_intra_phase_switch(): + kernel = parser.parse_string(_SHIFT) + kernel = passes.concretize_parameters(kernel, N=8) + kernel = passes.constexpr_propagation(kernel) + files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) + layout = next(f.code for f in files if "layout" in f.filename) + assert "spa_phase_reload" not in layout + # PE 1: forward 1 wave W→E, then switch to inject R→E. + assert re.search( + r"@set_color_config\(1, 0, @get_color\(0\), " + r"\.\{ \.routes = \.\{ \.rx = \.\{WEST\}, \.tx = \.\{EAST\} \} \}\)", + layout, + ) + assert re.search( + r"spa_switch_after phase=\d+ pe=1,0 ch=0 waves=1 rx=RAMP tx=EAST", + layout, + ) + # PE 4: absorb 1 wave W→R, then switch to forward W→E. + assert re.search( + r"@set_color_config\(4, 0, @get_color\(0\), " + r"\.\{ \.routes = \.\{ \.rx = \.\{WEST\}, \.tx = \.\{RAMP\} \} \}\)", + layout, + ) + assert re.search( + r"spa_switch_after phase=\d+ pe=4,0 ch=0 waves=1 rx=WEST tx=EAST", + layout, + ) + + +def test_batcher_lowering_two_colors_per_phase(): + path = os.path.join( + os.path.dirname(__file__), "..", "..", "samples", "spatial", "sort", "batcher_oddeven_1D.sptl" + ) + for n_log, n_cas, old_static in ((3, 6, 14), (4, 10, 30)): + kernel = parser.parse_file(path) + kernel = passes.concretize_parameters(kernel, L=n_log) + kernel = passes.constexpr_propagation(kernel) + files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) + layout = next(f.code for f in files if "layout" in f.filename) + assert "spa_color_schedule" in layout + assert "spa_switch_after" in layout + assert "spa_phase_reload" not in layout + colors = {int(c) for c in re.findall(r"@get_color\((\d+)\)", layout)} + assert colors == set(range(2 * n_cas)), ( + f"L={n_log} used colors {sorted(colors)}, expected 2 per phase " + f"({2 * n_cas}), not a kernel-wide pair and not {old_static} static" + ) + # Every two-step schedule has a matching counted switch onto the second pair. + schedules = re.findall( + r"spa_color_schedule phase=(\d+) pe=(\d+),(\d+) ch=(\d+) : ([^\n]+)", + layout, + ) + switches = { + (int(ph), int(x), int(y), int(ch)): (int(w), rx, tx) + for ph, x, y, ch, w, rx, tx in re.findall( + r"spa_switch_after phase=(\d+) pe=(\d+),(\d+) ch=(\d+) " + r"waves=(\d+) rx=(\w+) tx=(\w+)", + layout, + ) + } + short = {"R": "RAMP", "E": "EAST", "W": "WEST", "N": "NORTH", "S": "SOUTH"} + for ph, x, y, ch, step_txt in schedules: + parts = [p.strip() for p in step_txt.split(";")] + key = (int(ph), int(x), int(y), int(ch)) + if len(parts) < 2: + assert key not in switches + continue + first, second = parts[0], parts[1] + first_waves = int(re.search(r"waves=(\d+)", first).group(1)) + pair = re.search(r"(\w+)->(\w+)", second) + rx = short.get(pair.group(1), pair.group(1)) + tx = short.get(pair.group(2), pair.group(2)) + assert switches[key] == (first_waves, rx, tx) From 554c852edf40a199d89328d7851f5a968340eedc Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 17 Aug 2026 16:18:40 +0200 Subject: [PATCH 12/68] draft switch advances --- spada/lowering/spatial_ir_to_csl.py | 138 ++++++++++++++++++-- spada/syntax/csl/dsd_ops.py | 10 +- spada/syntax/csl/structures.py | 5 +- spada/syntax/spatial_ir/canonicalization.py | 55 +++++--- spada/syntax/spatial_ir/shift_bundles.py | 35 +++++ tests/spatial_ir/test_shift_bundles.py | 40 +++++- 6 files changed, 247 insertions(+), 36 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index ae37a398..f4de0bee 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -369,6 +369,8 @@ def generate_rectangle(kernel: spir.Kernel, raise ValueError(f"Error in {e.args[0].lineinfo}. Undefined identifier \"{e.args[0].as_ir()}\".") raise + stream_to_adv = _bind_switch_advances(rect.metadata, kernel, dsds, header) + # Fuse tasks as much as possible to reduce number of resources if task_fusion: orig_len = 0 @@ -441,6 +443,7 @@ def generate_rectangle(kernel: spir.Kernel, task_bindings, benchmark_code.kernel_postamble, indent=' ', + stream_to_adv=stream_to_adv, ) except KeyError as e: identifier = e.args[0] @@ -469,6 +472,7 @@ def generate_rectangle(kernel: spir.Kernel, task_bindings, benchmark_code.kernel_postamble, indent=' ', + stream_to_adv=stream_to_adv, ) except KeyError as e: identifier = e.args[0] @@ -1402,11 +1406,11 @@ def _coord(k): # noqa: E731 def _emit_shift_schedules(kernel: spir.Kernel, channel_to_color: dict[int, int]) -> str: """ - Emit absolute @set_color_config plus spa_switch_after / spa_color_schedule - annotations for counted 1D shift bundles. + Emit absolute @set_color_config with fabric switch pos1 for counted 1D + shift bundles, plus spa_switch_after / spa_color_schedule comments. - Each color carries one two-state program (one phase). Wave quotas differ - across phases, so the same color is never reloaded with a new count. + Each color carries one two-state program (one phase). Downstream routers + advance from pos0 to pos1 on a SWITCH_ADV control wavelet. """ schedules = getattr(kernel, "shift_schedules", None) if not schedules: @@ -1442,22 +1446,40 @@ def _emit_shift_schedules(kernel: spir.Kernel, channel_to_color: dict[int, int]) f"ch={channel} : {step_txt}\n" ) first = sched.steps[0] - config = ( - f" @set_color_config({x}, {y}, {color_expr}, " - f".{{ .routes = .{{ .rx = .{{{first.rx}}}, .tx = .{{{first.tx}}} }} }});\n" - ) - if config not in seen_config: - lines.append(config) - seen_config.add(config) + switch_fields = [".pop_mode = .{ .always_pop = true }"] if len(sched.steps) > 1: nxt = sched.steps[1] + pos1 = _pos1_field(first, nxt) + if pos1 is not None: + switch_fields.insert(0, f".pos1 = .{{ {pos1} }}") lines.append( f" // spa_switch_after phase={sched.phase_index} pe={x},{y} " f"ch={channel} waves={first.waves} rx={nxt.rx} tx={nxt.tx}\n" ) + config = ( + f" @set_color_config({x}, {y}, {color_expr}, " + f".{{ .routes = .{{ .rx = .{{{first.rx}}}, .tx = .{{{first.tx}}} }}, " + f".switches = .{{ {', '.join(switch_fields)} }} }});\n" + ) + if config not in seen_config: + lines.append(config) + seen_config.add(config) return "".join(lines) +def _pos1_field(first, nxt) -> str | None: + """CSL pos1 may set only rx or only tx; the two-state lemma changes one.""" + if first.rx != nxt.rx and first.tx == nxt.tx: + return f".rx = {nxt.rx}" + if first.tx != nxt.tx and first.rx == nxt.rx: + return f".tx = {nxt.tx}" + if first.rx == nxt.rx and first.tx == nxt.tx: + return None + raise ValueError( + f"Counted switch pos1 must change only rx or only tx, got {first} -> {nxt}" + ) + + def _write_indented_block(current_code: StringIO, block: str, indent: str) -> None: block = textwrap.dedent(block).strip('\n') if not block: @@ -1547,6 +1569,86 @@ def _generate_data_task( current_code.write(f"\n}}\n") +def _bind_switch_advances(pe_block: PEBlock, kernel: spir.Kernel, dsds: UniqueDSDDict, + header: StringIO) -> dict[str, tuple]: + """ + Map counted-switch send streams to their SWITCH_ADV program and emit + control-wavelet DSDs plus encoded payloads into ``header``. + """ + from spada.syntax.spatial_ir.shift_bundles import SwitchAdvance + + advances: list[SwitchAdvance] = getattr(kernel, "switch_advances", None) or [] + if not advances or pe_block.dataflow is None: + return {} + by_channel = {adv.channel: adv for adv in advances} + stream_to_adv: dict[str, tuple] = {} + for stmt in pe_block.dataflow.statements: + routing = getattr(stmt.stream, "routing", None) + if routing is None or not routing.counted_switch: + continue + adv = by_channel.get(routing.channel) + if adv is None: + continue + key = stmt.stream_name.as_ir() + out_name, out_dsd = _fabout_dsd(dsds, key) + if out_dsd is None: + continue + ctrl_name = f'{name_to_csl(stmt.stream_name)}_ctrl_out_dsd' + stream_to_adv[key] = (adv, ctrl_name, out_dsd) + + if not stream_to_adv: + return {} + + header.write('const ctrl = @import_module("");\n') + header.write('const tile_config = @import_module("");\n') + emitted_payloads: set[int] = set() + for adv, ctrl_name, out_dsd in stream_to_adv.values(): + ctrl_dsd = cslstruct.FabricDSD( + cslstruct.DSDType.fabout, out_dsd.color, 1, out_dsd.queue, control=True) + header.write(f'const {ctrl_name} = {ctrl_dsd.as_csl()};\n') + if adv.channel in emitted_payloads: + continue + emitted_payloads.add(adv.channel) + header.write(_encode_switch_payload(adv)) + header.write('\n') + return stream_to_adv + + +def _fabout_dsd(dsds: UniqueDSDDict, stream_key: str): + for dsd_name, dsd in dsds.get(stream_key, []): + if isinstance(dsd, cslstruct.FabricDSD) and dsd.dsd_type == cslstruct.DSDType.fabout and not dsd.control: + return dsd_name, dsd + return None, None + + +def _encode_switch_payload(adv) -> str: + n = len(adv.opcodes) + op_enum = {"SWITCH_ADV": "ctrl.opcode.SWITCH_ADV", "NOP": "ctrl.opcode.NOP"} + ops = ", ".join(op_enum[op] for op in adv.opcodes) + ignores = ", ".join("true" for _ in adv.opcodes) + prefix = f'switch_adv_ch{adv.channel}' + return ( + f'const {prefix}_cmds = [{n}]ctrl.opcode{{ {ops} }};\n' + f'const {prefix}_ignore = [{n}]bool{{ {ignores} }};\n' + f'const {prefix}_pld: u32 = ctrl.encode_payload(' + f'{n}, {prefix}_cmds, {prefix}_ignore, true, {{}});\n' + ) + + +def _emit_switch_advance_after_send(stmt: spir.SendStatement, stream_to_adv: dict) -> str: + stream_name = stmt.stream_name.array if isinstance(stmt.stream_name, spir.ArraySlice) else stmt.stream_name + entry = stream_to_adv.get(stream_name.as_ir()) + if entry is None: + return "" + adv, ctrl_name, _out_dsd = entry + coord = "tile_config.fabric_coord.X" if adv.axis == "x" else "tile_config.fabric_coord.Y" + return ( + f'if (tile_config.get_fabric_coord({coord}) != {adv.last_injector}) {{\n' + f' @mov32({ctrl_name}, switch_adv_ch{adv.channel}_pld);\n' + f'}}' + ) + + def _generate_task_code(rect: PEBlock, task_index: int, task: tdag.CSLTask, @@ -1559,7 +1661,8 @@ def _generate_task_code(rect: PEBlock, tasks: list[tdag.CSLTask], task_bindings: task_recycling.TaskBindingPlan, postamble: str, - indent: str = ' '): + indent: str = ' ', + stream_to_adv: dict | None = None): """ Generates a local task from a CSL task. This function converts statements to DSD operations or generates appropriate code. @@ -1620,6 +1723,17 @@ def _generate_task_code(rect: PEBlock, # Asynchronous DSD op. DSD line already contains activation or unblocking skip_activation = True + if stream_to_adv and isinstance(stmt, spir.SendStatement): + extra = _emit_switch_advance_after_send(stmt, stream_to_adv) + if extra: + extra_lines = extra.splitlines() + insert_at = next( + (i for i, ln in enumerate(lines) + if ln.lstrip().startswith(('@activate(', '@unblock('))), + len(lines), + ) + lines = lines[:insert_at] + extra_lines + lines[insert_at:] + for line in lines: current_code.write(f'{indent}{line}\n') diff --git a/spada/syntax/csl/dsd_ops.py b/spada/syntax/csl/dsd_ops.py index 0123754c..72c29263 100644 --- a/spada/syntax/csl/dsd_ops.py +++ b/spada/syntax/csl/dsd_ops.py @@ -23,11 +23,13 @@ def _append_async_suffix(self, base: str, dsd_objects: list[cslstruct.DataStruct async_target: Optional[AsyncTarget]) -> str: if async_target is None: return base - if any(isinstance(dsd, cslstruct.FabricDSD) for dsd in dsd_objects): + # CSL allows .async only when every operand is a DSD/DSR. A scalar + # source (CopyDSDOp.scalar_input) must complete synchronously, then + # activate/unblock the next task. + fabric_async = any(isinstance(dsd, cslstruct.FabricDSD) for dsd in dsd_objects) + if fabric_async and not getattr(self, 'scalar_input', False): return f'{base[:-2]}, .{{ .async = true, .{async_target.inter_task_edge} = {async_target.target_task} }});' - else: - # Pure Memory DSD operations are synchronous - return f'{base}\n@{async_target.inter_task_edge}({async_target.target_task});' + return f'{base}\n@{async_target.inter_task_edge}({async_target.target_task});' def as_csl(self, statement: spir.Statement, diff --git a/spada/syntax/csl/structures.py b/spada/syntax/csl/structures.py index 1855895c..2ae4f952 100644 --- a/spada/syntax/csl/structures.py +++ b/spada/syntax/csl/structures.py @@ -55,6 +55,7 @@ class FabricDSD(DataStructureDescriptor): color: str extent: int queue: int + control: bool = False def __post_init__(self): assert self.dsd_type in (DSDType.fabin, DSDType.fabout) @@ -63,7 +64,9 @@ def as_csl(self) -> str: direction = "in" if self.dsd_type == DSDType.fabin else "out" queue_type = "input_queue" if self.dsd_type == DSDType.fabin else "output_queue" fabric_color = f' .fabric_color = {self.color}_{direction},' if self.color else '' - return f'@get_dsd({self.dsd_type.name}_dsd, .{{ .extent = {self.extent},{fabric_color} .{queue_type} = @get_{queue_type}({self.queue}) }})' + control = ' .control = true,' if self.control else '' + return (f'@get_dsd({self.dsd_type.name}_dsd, .{{ .extent = {self.extent},{fabric_color}' + f'{control} .{queue_type} = @get_{queue_type}({self.queue}) }})') def __hash__(self): return hash(("FabricDSD", self.as_csl())) diff --git a/spada/syntax/spatial_ir/canonicalization.py b/spada/syntax/spatial_ir/canonicalization.py index d8ccec89..575c193a 100644 --- a/spada/syntax/spatial_ir/canonicalization.py +++ b/spada/syntax/spatial_ir/canonicalization.py @@ -318,6 +318,8 @@ def inline_phases(kernel: spir.Kernel) -> spir.Kernel: body=list(rect_place.values()) + list(rect_dataflow.values()) + list(rect_compute.values())) if hasattr(kernel, "shift_schedules"): new_kernel.shift_schedules = kernel.shift_schedules + if hasattr(kernel, "switch_advances"): + new_kernel.switch_advances = kernel.switch_advances return new_kernel @@ -502,28 +504,41 @@ def __init__(self, place: spir.PlaceBlock): def visit_ReceiveStatement(self, node: spir.ReceiveStatement): sz = node.get_size(self.identifier_sizes) - if len(sz) == 0: # Scalar receive - return self.generic_visit(node) + scalar_onchip = False + if len(sz) == 0: + # Memcpy/extern scalar receive keeps a direct assignment (`val = inp[0]`). + # On-chip scalar receive (`await receive(tmp, bwd)`) must become a + # one-wavelet foreach so CSL binds a data task to the fabric color. + if isinstance(node.stream_name, spir.ArraySlice): + return self.generic_visit(node) + sz = [1] + scalar_onchip = True + + if scalar_onchip: + body = [ + spir.AssignmentStatement( + copy.deepcopy(node.local_array), + spir.Expression(spir.Identifier('__x', 0))), + ] + else: + body = [ + spir.AssignmentStatement( + spir.ArraySlice( + copy.deepcopy(node.local_array), + [spir.Expression(spir.Identifier(f'__k{i}', 0)) for i in range(len(sz))]), + spir.Expression(spir.Identifier('__x', 0))), + ] - # Array receive, make a foreach node new_node = spir.ForeachStatement( [spir.TypedIdentifier(spir.ScalarType.u16, spir.Identifier(f'__k{i}', 0)) for i in range(len(sz))], [ - # ``0:size`` for every dimension spir.RangeExpression( spir.Expression(spir.ConstantLiteral(0, spir.ScalarType.u16)), spir.Expression(spir.ConstantLiteral(s, spir.ScalarType.u16))) for s in sz ], - spir.TypedIdentifier(self.identifier_dtypes[node.local_array], spir.Identifier(f'__x', 0)), + spir.TypedIdentifier(self.identifier_dtypes[node.local_array], spir.Identifier('__x', 0)), spir.ReceiveGenerator(node.stream_name), - [ - # ``arr[__k0, ...] = __x`` - spir.AssignmentStatement( - spir.ArraySlice( - copy.deepcopy(node.local_array), - [spir.Expression(spir.Identifier(f'__k{i}', 0)) for i in range(len(sz))]), - spir.Expression(spir.Identifier(f'__x', 0))), - ], + body, node.completion_name) new_node.lineinfo = node.lineinfo @@ -559,8 +574,8 @@ def visit_ReceiveStatement(self, node: spir.ReceiveStatement): def lower_bulk_communication(rectangles: list[Rectangle[PEBlock]]) -> None: """ - Lowers top-level array ``receive`` and ``send`` operations to foreach and for loops, respectively. - The array operations are shorthands for a row-major (C-order) loop over the communication operations. + Lowers top-level array ``receive`` operations to foreach loops, and on-chip + scalar ``receive`` from a named stream to a one-wavelet foreach. :param rectangles: A list of PE block rectangles to lower computations within. """ @@ -648,14 +663,20 @@ def visit_ForeachStatement(self, node: spir.ForeachStatement): if dsd_ops.get_dsd_op(self.dtypes, node) is not None: return self.generic_visit(node) - if isinstance(self.dtypes[node.receive_stream.stream_name], spir.StreamType): + sname = node.receive_stream.stream_name + if isinstance(sname, spir.ArraySlice): + sname = sname.array + stream_dtype = self.dtypes.get(sname) + # On-chip streams stay data tasks. Missing names are also on-chip + # (declared on a sender rectangle). Memcpy fields are ArrayType. + if stream_dtype is None or isinstance(stream_dtype, spir.StreamType): return self.generic_visit(node) body_statements = [self.visit(stmt) for stmt in node.body] loop_variables = [copy.deepcopy(var) for var in node.variables] loop_ranges = [copy.deepcopy(rng) for rng in node.parameter_range] stream_target = copy.deepcopy(node.receive_stream.stream_name) - if isinstance(self.dtypes[stream_target], spir.ArrayType) and loop_ranges: + if isinstance(stream_dtype, spir.ArrayType) and loop_ranges: index_exprs = [] for var in loop_variables: idx_identifier = copy.deepcopy(var.identifier) diff --git a/spada/syntax/spatial_ir/shift_bundles.py b/spada/syntax/spatial_ir/shift_bundles.py index 9d71c60c..e5519ede 100644 --- a/spada/syntax/spatial_ir/shift_bundles.py +++ b/spada/syntax/spatial_ir/shift_bundles.py @@ -57,6 +57,22 @@ class ColorSchedule: steps: tuple[ColorScheduleStep, ...] +@dataclass(frozen=True) +class SwitchAdvance: + """ + SWITCH_ADV control wavelet sent after a non-last injector's data waves. + + Each downstream router pops one opcode (always_pop): the next source and + the dest that just absorbed see SWITCH_ADV; hops in between see NOP. + A control wavelet holds at most 8 opcodes, so ``dist`` must be <= 8. + """ + + channel: int + axis: Literal["x", "y"] + last_injector: int + opcodes: tuple[str, ...] + + def _short(port: str) -> str: return {"RAMP": "R", "EAST": "E", "WEST": "W", "NORTH": "N", "SOUTH": "S"}.get(port, port) @@ -278,6 +294,22 @@ def add(coord: int, steps: list[ColorScheduleStep]) -> None: ] +_MAX_SWITCH_CMDS = 8 + + +def switch_advance_for_bundle(bundle: ShiftBundle) -> SwitchAdvance: + """Build the always_pop opcode chain for one interval shift of distance ``d``.""" + d = bundle.dist + if d > _MAX_SWITCH_CMDS: + raise ValueError( + f"Counted switching encodes one opcode per hop and a control wavelet " + f"holds at most {_MAX_SWITCH_CMDS} commands; got dist={d}." + ) + opcodes = ("SWITCH_ADV",) + ("NOP",) * max(d - 2, 0) + ("SWITCH_ADV",) + last_injector = bundle.start + bundle.length - 1 if bundle.sign > 0 else bundle.start + return SwitchAdvance(bundle.channel, bundle.axis, last_injector, opcodes) + + def _phase_has_explicit_count_relative(phase: spir.Phase) -> bool: for dataflow in phase.dataflow: for stmt in dataflow.statements: @@ -339,9 +371,12 @@ def apply_shift_bundles(kernel: spir.Kernel, bundles: list[ShiftBundle]) -> list phase_index += 1 schedules: list[ColorSchedule] = [] + advances: list[SwitchAdvance] = [] for bundle in assigned: schedules.extend(schedule_counted_switch(bundle)) + advances.append(switch_advance_for_bundle(bundle)) _assign_unit_hop_channels(kernel, bases) + kernel.switch_advances = advances return schedules diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py index 8ffb7602..38fa946a 100644 --- a/tests/spatial_ir/test_shift_bundles.py +++ b/tests/spatial_ir/test_shift_bundles.py @@ -10,6 +10,7 @@ detect_shift_bundles, schedule_counted_switch, coalesce_shift_bundles, + switch_advance_for_bundle, ) @@ -46,6 +47,13 @@ def _prepare(src: str, **params: int): count = 1 } } + dataflow i16 i, i16 j in [4:8, 0] { + stream fwd = relative_stream(4, 0) { + hops = auto, + channel = auto, + count = 1 + } + } compute i16 i, i16 j in [0:4, 0] { await send(val, fwd) } @@ -115,6 +123,9 @@ def test_schedule_source_and_dest_halves(): assert [(st.rx, st.tx, st.waves) for st in by_pe[(3, 0)]] == [("WEST", "EAST", 3), ("RAMP", "EAST", 1)] assert [(st.rx, st.tx, st.waves) for st in by_pe[(4, 0)]] == [("WEST", "RAMP", 1), ("WEST", "EAST", 3)] assert [(st.rx, st.tx, st.waves) for st in by_pe[(7, 0)]] == [("WEST", "RAMP", 1)] + adv = switch_advance_for_bundle(bundle) + assert adv.last_injector == 3 + assert adv.opcodes == ("SWITCH_ADV", "NOP", "NOP", "SWITCH_ADV") def test_westbound_schedule(): @@ -340,6 +351,24 @@ def test_batcher_two_colors_per_phase_not_globally(): assert set(used) == set(range(12)) +def test_batcher_scalar_receive_lowers_to_data_task(): + path = os.path.join( + os.path.dirname(__file__), "..", "..", "samples", "spatial", "sort", "batcher_oddeven_1D.sptl" + ) + kernel = parser.parse_file(path) + kernel = passes.concretize_parameters(kernel, L=1) + kernel = passes.constexpr_propagation(kernel) + files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) + pe_codes = [f.code for f in files if "code_" in f.filename] + assert pe_codes + for code in pe_codes: + assert "tmp = bwd" not in code + assert ".async = true" not in code + pe0 = next(f.code for f in files if "code_0_0" in f.filename) + assert "task dtask_" in pe0 + assert "tmp = __x" in pe0 + + def test_lowering_encodes_intra_phase_switch(): kernel = parser.parse_string(_SHIFT) kernel = passes.concretize_parameters(kernel, N=8) @@ -350,23 +379,30 @@ def test_lowering_encodes_intra_phase_switch(): # PE 1: forward 1 wave W→E, then switch to inject R→E. assert re.search( r"@set_color_config\(1, 0, @get_color\(0\), " - r"\.\{ \.routes = \.\{ \.rx = \.\{WEST\}, \.tx = \.\{EAST\} \} \}\)", + r"\.\{ \.routes = \.\{ \.rx = \.\{WEST\}, \.tx = \.\{EAST\} \}", layout, ) assert re.search( r"spa_switch_after phase=\d+ pe=1,0 ch=0 waves=1 rx=RAMP tx=EAST", layout, ) + assert ".pos1 = .{ .rx = RAMP }" in layout + assert ".pop_mode = .{ .always_pop = true }" in layout # PE 4: absorb 1 wave W→R, then switch to forward W→E. assert re.search( r"@set_color_config\(4, 0, @get_color\(0\), " - r"\.\{ \.routes = \.\{ \.rx = \.\{WEST\}, \.tx = \.\{RAMP\} \} \}\)", + r"\.\{ \.routes = \.\{ \.rx = \.\{WEST\}, \.tx = \.\{RAMP\} \}", layout, ) assert re.search( r"spa_switch_after phase=\d+ pe=4,0 ch=0 waves=1 rx=WEST tx=EAST", layout, ) + assert ".pos1 = .{ .tx = EAST }" in layout + pe0 = next(f.code for f in files if "code_0_0" in f.filename) + assert "ctrl.opcode.SWITCH_ADV" in pe0 + assert "encode_payload" in pe0 + assert "get_fabric_coord" in pe0 def test_batcher_lowering_two_colors_per_phase(): From 04c05faf09b6926114abb0375136aa0aa28d5f99 Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Mon, 17 Aug 2026 09:31:34 -0700 Subject: [PATCH 13/68] Factor out DSD detection so that streams in loops can be detected --- spada/lowering/spatial_ir_to_csl.py | 122 ++++++++++++++++++++-------- tests/spatial_ir/test_dsd_ops.py | 78 ++++++++++++++++++ 2 files changed, 164 insertions(+), 36 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 6e543259..56248c0f 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -1007,6 +1007,50 @@ def allocate_output_queue(stream: spir.Identifier) -> int: output_queue_of[key] = csl.OUTPUT_QUEUE_IDS[output_queue_id_ctr % len(csl.OUTPUT_QUEUE_IDS)] output_queue_id_ctr += 1 return output_queue_of[key] + + def _visit_foreach(stmt: spir.ForeachStatement) -> None: + """ + Registers the fabric input DSD for a ``foreach`` that draws from a stream. + + An array ``receive`` is canonicalized into one of these (see ``_BulkCommunicationLowerer``), + so this is the path every bulk receive takes -- including one nested inside a sequential + ``for``, which is why this is a function rather than inline in the statement walk below. + """ + # If the foreach statement has a stream generator, it is a DSD + # unless only the receive generator is given (streaming, no range provided). + stream_name = ( + stmt.receive_stream.stream_name.array + if isinstance(stmt.receive_stream.stream_name, spir.ArraySlice) else stmt.receive_stream.stream_name) + if not stmt.parameter_range: + if stream_name not in stream_args: + raise SyntaxError(f'Foreach generator "{stream_name.as_ir()}" without a defined ' + f'range must only be used with a kernel argument or extern_stream.' + f'\n In line {stmt.lineinfo}') + # A data task will be created instead (handled in _generate_data_task) + return + if stream_name.as_ir() not in stream_candidates: + return + if memcpy_mode and stream_name in stream_args: + # If memcpy mode is enabled, the stream contents will have already been copied to the PE + dsd = _dsd_from_stream(stream_candidates, stream_name) + dsds[stream_name.as_ir()].append((f"{name_to_csl(stream_name)}_dsd", dsd)) + return + dsd_name = f'{name_to_csl(stream_name)}_in_dsd' + extents = stream_candidates[stream_name.as_ir()][1] + if extents is not None: # Use buffer size + extents = extents if isinstance(extents, int) else extents.eval() + else: # Infer from foreach range + if len(stmt.parameter_range) != 1: + raise SyntaxError( + f'Expected one-dimensional foreach range for stream "{stream_name.as_ir()}", ' + f'got {stmt.parameter_range}.\n In line {stmt.lineinfo}') + start, end, step = (stmt.parameter_range[0].start, stmt.parameter_range[0].stop, + stmt.parameter_range[0].step) + extents = (end.eval() - start.eval()) // (step.eval() if step is not None else 1) + fabric_color = f'{name_to_csl(stream_name)}_color' + dsd = cslstruct.FabricDSD(cslstruct.DSDType.fabin, fabric_color, extents, allocate_input_queue(stream_name)) + dsds[stream_name.as_ir()].append((dsd_name, dsd)) + for stmt in rect.compute.statements: # Find out if compute block uses this stream for receive/send if isinstance(stmt, (spir.ReceiveStatement, spir.SendStatement)): @@ -1078,40 +1122,7 @@ def allocate_output_queue(stream: spir.Identifier) -> int: dsds[arr.as_ir()].append((f"{name_to_csl(arr)}_dsd", dsd)) elif isinstance(stmt, spir.ForeachStatement): - # If the foreach statement has a stream generator, it is a DSD - # unless only the receive generator is given (streaming, no range provided). - stream_name = ( - stmt.receive_stream.stream_name.array - if isinstance(stmt.receive_stream.stream_name, spir.ArraySlice) else stmt.receive_stream.stream_name) - if not stmt.parameter_range: - if stream_name not in stream_args: - raise SyntaxError(f'Foreach generator "{stream_name.as_ir()}" without a defined ' - f'range must only be used with a kernel argument or extern_stream.' - f'\n In line {stmt.lineinfo}') - # A data task will be created instead (handled in _generate_data_task) - else: - if stream_name.as_ir() in stream_candidates: - if memcpy_mode and stream_name in stream_args: - # If memcpy mode is enabled, the stream contents will have already been copied to the PE - dsd = _dsd_from_stream(stream_candidates, stream_name) - dsds[stream_name.as_ir()].append((f"{name_to_csl(stream_name)}_dsd", dsd)) - else: - dsd_name = f'{name_to_csl(stream_name)}_in_dsd' - extents = stream_candidates[stream_name.as_ir()][1] - if extents is not None: # Use buffer size - extents = extents if isinstance(extents, int) else extents.eval() - else: # Infer from foreach range - if len(stmt.parameter_range) != 1: - raise SyntaxError( - f'Expected one-dimensional foreach range for stream "{stream_name.as_ir()}", got {stmt.parameter_range}.\n In line {stmt.lineinfo}' - ) - start, end, step = stmt.parameter_range[0].start, stmt.parameter_range[ - 0].stop, stmt.parameter_range[0].step - extents = (end.eval() - start.eval()) // (step.eval() if step is not None else 1) - fabric_color = f'{name_to_csl(stream_name)}_color' - dsd = cslstruct.FabricDSD(cslstruct.DSDType.fabin, fabric_color, extents, - allocate_input_queue(stream_name)) - dsds[stream_name.as_ir()].append((dsd_name, dsd)) + _visit_foreach(stmt) def _visit_nested_send(substmt: spir.SendStatement): if substmt.stream_name.as_ir() not in stream_candidates: @@ -1162,13 +1173,36 @@ def _visit_nested_receive(substmt: spir.ReceiveStatement): dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, allocate_input_queue(stream_name)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) + def _visit_local_array(operand): + """ + Registers the memory DSD for the local side of a nested transfer. + + A top-level send gets this from the explicit walk over ``stmt.local_array`` above; a + nested one is only reached through the visitor, which stops at the statement and never + descends to the operand. + """ + name = operand.identifier if isinstance(operand, spir.TypedIdentifier) else operand + if not isinstance(name, spir.Identifier): + return + if name.as_ir() not in array_candidates or name.as_ir() in dsds: + return + dsds[name.as_ir()].append((f"{name_to_csl(name)}_dsd", _dsd_from_array(array_candidates, name))) + def _visit_dsd(substmt, in_scope, in_assignment): + if isinstance(substmt, spir.ForeachStatement): + # Only reached for a foreach nested inside a sequential ``for``; a top-level one is + # registered by the statement walk above. + if in_scope: + _visit_foreach(substmt) + return if isinstance(substmt, spir.SendStatement) and in_scope: _visit_nested_send(substmt) + _visit_local_array(substmt.local_array) return if isinstance(substmt, spir.ReceiveStatement): if in_scope: _visit_nested_receive(substmt) + _visit_local_array(substmt.local_array) return if (isinstance(substmt, spir.Identifier) and substmt.as_ir() in array_candidates and substmt.as_ir() not in dsds): @@ -1230,6 +1264,10 @@ def __init__(self, callback, toplevel: bool): super().__init__() def visit_ForeachStatement(self, node: spir.ForeachStatement): + # Nested inside a sequential ``for``, this foreach is not visited by the statement walk that + # registers stream DSDs, so report it here. ``in_for`` is false at the top level, where that + # walk has already handled it. + self.callback(node, self.in_for, self.in_assignment) old_scope = self.in_foreach self.in_foreach = True self.generic_visit(node) @@ -1247,16 +1285,28 @@ def visit_ForStatement(self, node: spir.ForStatement): self.generic_visit(node) self.in_for = old_scope + @property + def in_transfer_scope(self) -> bool: + """ + Whether a send or receive here transfers a whole array and so needs a fabric DSD. + + A sequential ``for`` counts: it lowers to a real CSL loop, and each iteration moves the + whole local array, exactly as one outside the loop would. Element accesses do *not* count + (see :meth:`visit_Identifier`) -- ``a[k]`` inside a sequential loop is one element per + iteration, which is a scalar access and not a DSD. + """ + return self.in_foreach or self.in_map or self.in_for + def visit_Identifier(self, node: spir.Identifier): self.callback(node, self.in_foreach or self.in_map, self.in_assignment) return def visit_SendStatement(self, node: spir.SendStatement): - self.callback(node, self.in_foreach or self.in_map, self.in_assignment) + self.callback(node, self.in_transfer_scope, self.in_assignment) return def visit_ReceiveStatement(self, node: spir.ReceiveStatement): - self.callback(node, self.in_foreach or self.in_map, self.in_assignment) + self.callback(node, self.in_transfer_scope, self.in_assignment) return def visit_ArraySlice(self, node: spir.ArraySlice): diff --git a/tests/spatial_ir/test_dsd_ops.py b/tests/spatial_ir/test_dsd_ops.py index 1e7c2a97..61cf3737 100644 --- a/tests/spatial_ir/test_dsd_ops.py +++ b/tests/spatial_ir/test_dsd_ops.py @@ -131,7 +131,85 @@ def test_dsd_op_detection_receive_op_send(): assert dsd_stmt.destination.as_ir() == 'blue' +def test_transfers_inside_a_sequential_for_get_fabric_dsds(): + """ + A sequential ``for`` lowers to a real CSL loop, so a send or receive inside it still moves a + whole array per iteration and still needs a fabric DSD. + + The DSD walk only inspects top-level compute statements, so nested transfers reach it through + ``DSDVisitor``. It used to report them as out of scope -- it tracked ``in_for`` but never + consulted it -- and codegen then failed looking up a DSD that was never created. + """ + kernel = parser.parse_string(code=""" +kernel @looped (stream[2, 1] readonly src, + stream[2, 1] writeonly dst) { + place i16 i, i16 j in [0:2, 0] { + f32[K] val + f32[K] other + } + dataflow i16 i, i16 j in [0:2, 0] { + stream east = relative_stream(1, 0) { + hops = [(1, 0)], + channel = 0 + } + } + compute i16 i, i16 j in [0:2, 0] { + await receive(val, src[i, j]) + } + compute i16 i, i16 j in [0:1, 0] { + for i32 t in [0:M] { + await send(val, east) + } + await send(val, dst[i, j]) + } + compute i16 i, i16 j in [1:2, 0] { + for i32 t in [0:M] { + await receive(other, east) + } + await send(other, dst[i, j]) + } +}""") + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, K=4, M=3)) + files = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel)} + + sender, receiver = files['code_0_0.csl'], files['code_1_0.csl'] + assert 'fabout_dsd' in sender, sender + assert 'fabin_dsd' in receiver, receiver + # The loop stays a loop: the body is emitted once with the trip count in the range, not M times. + assert 'for (@range(i32, 0, 3, 1))' in sender, sender + assert 'for (@range(i32, 0, 3, 1))' in receiver, receiver + + +def test_sequential_for_does_not_make_element_accesses_into_dsds(): + """ + The counterpart to the test above: ``a[k]`` inside a sequential ``for`` is one element per + iteration, a scalar access, and must not be promoted to a DSD operation. + """ + kernel = parser.parse_string(code=""" +kernel @scan (stream[1, 1] readonly src, + stream[1, 1] writeonly dst) { + place i16 i, i16 j in [0, 0] { + f32[K] a + } + compute i16 i, i16 j in [0, 0] { + await receive(a, src[i, j]) + for i32 k in [1:K] { + a[k] = a[k-1] + a[k] + } + await send(a, dst[i, j]) + } +}""") + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, K=8)) + code = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel)}['code_0_0.csl'] + + assert 'for (@range' in code, code + # A scalar recurrence, not a vector add over a DSD. + assert '@fadds(a_dsd' not in code, code + + if __name__ == '__main__': test_dsd_op_detection() test_dsd_op_detection_constant_folding() test_dsd_op_detection_receive_op_send() + test_transfers_inside_a_sequential_for_get_fabric_dsds() + test_sequential_for_does_not_make_element_accesses_into_dsds() From a7e7c69f86fa9342aa7291cc3e2ab63f4b9582ef Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 17 Aug 2026 19:47:48 +0200 Subject: [PATCH 14/68] fix spmv example --- samples/spatial/blas/spmv.sptl | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/samples/spatial/blas/spmv.sptl b/samples/spatial/blas/spmv.sptl index fc7994ba..db28f519 100644 --- a/samples/spatial/blas/spmv.sptl +++ b/samples/spatial/blas/spmv.sptl @@ -35,6 +35,8 @@ kernel @spmv( f32[NZ] A_val // COO nonzero values (zero-padded) i16[NZ] A_row // COO row indices i16[NZ] A_col // COO column indices + i16 tmp + i16 tmp2 f32[K] x // x chunk for this PE column (populated by multicast) f32[K] z // local partial result; accumulated during reduction f32[K] y_block // y chunk; read from host, used only at i=0 @@ -82,11 +84,13 @@ kernel @spmv( // Phase 4: Local COO SpMV: z[A_row[p]] += A_val[p] * x[A_col[p]]. phase { compute i16 i, i16 j in [0:PX, 0:PY] { - for i16 k in [0:K] { + await map i16 k in [0:K] { z[k] = 0.0 } for i16 p in [0:NZ] { - z[A_row[p]] = z[A_row[p]] + A_val[p] * x[A_col[p]] + tmp = A_row[p] + tmp2 = A_col[p] + z[tmp] = z[tmp] + A_val[p] * x[tmp2] } } } @@ -131,7 +135,7 @@ kernel @spmv( await foreach i16 k, f32 v in [0:K], receive(green) { z[k] = z[k] + v } - for i16 k in [0:K] { + await map i16 k in [0:K] { z[k] = alpha * z[k] + beta * y_block[k] } await send(z, out[i, j]) From 3f29a518964442480036e8b81fc1a6147d65818f Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 17 Aug 2026 19:50:14 +0200 Subject: [PATCH 15/68] add shift bundle test, fix shift bundling --- samples/spatial/simple/shift_bundle_1D.sptl | 56 +++++++++++++++++ spada/lowering/spatial_ir_to_csl.py | 28 ++++++--- spada/syntax/spatial_ir/shift_bundles.py | 5 +- tests/csl_runtime/test_shift_bundle_1d.sh | 50 +++++++++++++++ tests/spatial_ir/test_shift_bundles.py | 67 ++++++++++++++++++++- 5 files changed, 194 insertions(+), 12 deletions(-) create mode 100644 samples/spatial/simple/shift_bundle_1D.sptl create mode 100755 tests/csl_runtime/test_shift_bundle_1d.sh diff --git a/samples/spatial/simple/shift_bundle_1D.sptl b/samples/spatial/simple/shift_bundle_1D.sptl new file mode 100644 index 00000000..d09128e0 --- /dev/null +++ b/samples/spatial/simple/shift_bundle_1D.sptl @@ -0,0 +1,56 @@ +/** + * Counted 1D interval shift, and nothing else. + * + * 2M PEs on a line. Sources [0:M) each send one f32 to dest M steps east. + * Dest i+M overwrites its value with the word from i. Sources keep theirs. + * + * Overlapping paths share one eastbound color. A WSE color holds a single + * (rx, tx) pair at a time; the compiler sequences two pairs with a counted + * switch (forward-then-inject on the source half, absorb-then-forward on + * the dest half). There is no westbound stream, no CAS, and no d=1 stage. + * + * Constraints: M >= 2 (M=1 is a static hop; no overlap). + **/ +kernel @shift_bundle_1d( + stream[2 * M, 1] readonly inp, + stream[2 * M, 1] writeonly out +) { + place i16 i, i16 j in [0:2 * M, 0] { + f32 val + } + + phase { + compute i16 i, i16 j in [0:2 * M, 0] { + await receive(val, inp[i, j]) + } + } + + phase { + dataflow i16 i, i16 j in [0:M, 0] { + stream fwd = relative_stream(M, 0) { + hops = auto, + channel = auto, + count = 1 + } + } + dataflow i16 i, i16 j in [M:2 * M, 0] { + stream fwd = relative_stream(M, 0) { + hops = auto, + channel = auto, + count = 1 + } + } + compute i16 i, i16 j in [0:M, 0] { + await send(val, fwd) + } + compute i16 i, i16 j in [M:2 * M, 0] { + await receive(val, fwd) + } + } + + phase { + compute i16 i, i16 j in [0:2 * M, 0] { + await send(val, out[i, j]) + } + } +} diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index f4de0bee..c6f3d228 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -1446,16 +1446,24 @@ def _emit_shift_schedules(kernel: spir.Kernel, channel_to_color: dict[int, int]) f"ch={channel} : {step_txt}\n" ) first = sched.steps[0] - switch_fields = [".pop_mode = .{ .always_pop = true }"] - if len(sched.steps) > 1: - nxt = sched.steps[1] - pos1 = _pos1_field(first, nxt) - if pos1 is not None: - switch_fields.insert(0, f".pos1 = .{{ {pos1} }}") - lines.append( - f" // spa_switch_after phase={sched.phase_index} pe={x},{y} " - f"ch={channel} waves={first.waves} rx={nxt.rx} tx={nxt.tx}\n" - ) + inject_only = first.rx == "RAMP" and len(sched.steps) == 1 + if inject_only: + # The sender's router sees the control wavelet on RAMP. always_pop + # would consume the first opcode before it reaches downstream PEs. + switch_fields = [".pop_mode = .{ .no_pop = true }"] + else: + # Pop ADV on a real switch and NOP on pass-through hops; do not pop + # SWITCH_ADV after this PE has already reached pos1 (later inject). + switch_fields = [".pop_mode = .{ .pop_on_advance_nop = true }"] + if len(sched.steps) > 1: + nxt = sched.steps[1] + pos1 = _pos1_field(first, nxt) + if pos1 is not None: + switch_fields.insert(0, f".pos1 = .{{ {pos1} }}") + lines.append( + f" // spa_switch_after phase={sched.phase_index} pe={x},{y} " + f"ch={channel} waves={first.waves} rx={nxt.rx} tx={nxt.tx}\n" + ) config = ( f" @set_color_config({x}, {y}, {color_expr}, " f".{{ .routes = .{{ .rx = .{{{first.rx}}}, .tx = .{{{first.tx}}} }}, " diff --git a/spada/syntax/spatial_ir/shift_bundles.py b/spada/syntax/spatial_ir/shift_bundles.py index e5519ede..0c9c9dcc 100644 --- a/spada/syntax/spatial_ir/shift_bundles.py +++ b/spada/syntax/spatial_ir/shift_bundles.py @@ -298,7 +298,10 @@ def add(coord: int, steps: list[ColorScheduleStep]) -> None: def switch_advance_for_bundle(bundle: ShiftBundle) -> SwitchAdvance: - """Build the always_pop opcode chain for one interval shift of distance ``d``.""" + """Build the opcode chain popped by each downstream hop (distance ``d``). + + The injecting PE uses ``no_pop``, so the first opcode is for its neighbor. + """ d = bundle.dist if d > _MAX_SWITCH_CMDS: raise ValueError( diff --git a/tests/csl_runtime/test_shift_bundle_1d.sh b/tests/csl_runtime/test_shift_bundle_1d.sh new file mode 100755 index 00000000..ec2976d4 --- /dev/null +++ b/tests/csl_runtime/test_shift_bundle_1d.sh @@ -0,0 +1,50 @@ +#!/bin/sh +# E2E: eastbound counted interval shift only (no CAS, no westbound). +# Kernel: shift_bundle_1D.sptl params: M +# After the shift, OUT_out[M:2M] == inp[0:M] and OUT_out[0:M] == inp[0:M]. + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +. "$SCRIPT_DIR/_lib.sh" + +SAMPLE="$(cd "$(dirname "$0")/../../samples/spatial/simple" && pwd)/shift_bundle_1D.sptl" +FOLDER="shift_bundle_1d_sptl" + +run_shift() { + m=$1 + echo "--- shift_bundle_1d M=$m ---" + + sptlc "$SAMPLE" "$FOLDER" -p M=$m + + python3 - < Date: Mon, 17 Aug 2026 20:54:02 +0200 Subject: [PATCH 16/68] Prototype the counted switch in hand-written CSL Establishes the hardware contract the shift-bundle lowering will rely on, after the per-hop SWITCH_ADV payload turned out to be inert in slots 1-7: senders hand their router over in descending send order, triggered by a wavelet they emit themselves, and statically routed receivers use a counter filter to pick out the block addressed to them. Also pins the filter arithmetic, which the manual describes ambiguously: a wavelet reaches the compute element iff counter <= max_counter, counting modulo limit1 + 1 from init_counter. --- .../handwritten/shift_bundle/layout.csl | 122 ++++++++++++++++++ .../handwritten/shift_bundle/receiver.csl | 52 ++++++++ .../handwritten/shift_bundle/run.py | 76 +++++++++++ .../handwritten/shift_bundle/sender.csl | 67 ++++++++++ .../csl_runtime/test_shift_bundle_filters.sh | 46 +++++++ 5 files changed, 363 insertions(+) create mode 100644 tests/csl_runtime/handwritten/shift_bundle/layout.csl create mode 100644 tests/csl_runtime/handwritten/shift_bundle/receiver.csl create mode 100644 tests/csl_runtime/handwritten/shift_bundle/run.py create mode 100644 tests/csl_runtime/handwritten/shift_bundle/sender.csl create mode 100644 tests/csl_runtime/test_shift_bundle_filters.sh diff --git a/tests/csl_runtime/handwritten/shift_bundle/layout.csl b/tests/csl_runtime/handwritten/shift_bundle/layout.csl new file mode 100644 index 00000000..6e027c5d --- /dev/null +++ b/tests/csl_runtime/handwritten/shift_bundle/layout.csl @@ -0,0 +1,122 @@ +// Counted switching by send order, with counter filters at the receivers. +// +// M senders at x in [0, M) each ship K f32 words M steps east, to the receiver at x + M. A +// single switchable color carries all of it, even though the paths overlap pairwise. +// +// The senders go east to west -- the one closest to the receivers first. That order is what +// makes this cheap: a sender injects its own words and only then has to become a relay for +// the senders west of it, so the hand-over is triggered by an event it knows locally and can +// signal itself with one SWITCH_ADV. No router has to be switched from a distance, which is +// the thing a control wavelet cannot do selectively. +// +// The receivers never switch. Each one transmits to its ramp *and* onward east, so every +// receiver's router sees the whole stream and the outermost one terminates it. Which words a +// receiver keeps is decided by a counter filter, when FILTER != 0. +// +// Counter filter semantics, as measured on the simulator (the manual's "(exclusive)" aside +// notwithstanding, and matching its "reject all wavelets whose active counter is greater +// than max_counter"): +// +// * the counter starts at init_counter and increments on every data wavelet, +// * it wraps to zero after limit1, so it cycles through limit1 + 1 values, +// * a wavelet reaches the compute element iff counter <= max_counter. +// +// So a window of K words out of a stream of M*K is limit1 = M*K - 1, max_counter = K - 1, +// and init_counter placing the wanted word at counter zero. + +param M: i16; +param K: i16; + +// 0: no filter; every receiver takes all M*K words. Establishes the arrival order. +// 1: counter filters; every receiver takes only the K words addressed to it. +param FILTER: i16; + +const memcpy = @import_module("", .{ + .width = 2 * M, + .height = 1, +}); + +const CHANNEL: i16 = 1; + +const STREAM: i16 = M * K; + +// Sends run east to west, so the words from sender q arrive as block M - 1 - q of the +// stream. Receiver M + q wants that block, so its counter must read zero when the block +// starts: shift the start of the cycle back by as many words as precede the block. +fn window_start(q: i16) i16 { + return @as(i16, ((q + 1) * K) % STREAM); +} + +layout { + @set_rectangle(2 * M, 1); + + for (@range(i16, 0, M, 1)) |x| { + @set_tile_code(x, 0, "sender.csl", .{ + .memcpy_params = memcpy.get_params(x), + .stream = STREAM, + .words = K, + // The westmost sender is the last to send and has nothing to relay for. + .hands_over = x > 0, + }); + } + for (@range(i16, 0, M, 1)) |q| { + @set_tile_code(M + q, 0, "receiver.csl", .{ + .memcpy_params = memcpy.get_params(M + q), + .stream = STREAM, + .words = K, + .takes = if (FILTER == 0) STREAM else K, + }); + } + + // Senders: inject from the ramp, then relay from the west. + @set_color_config(0, 0, @get_color(CHANNEL), .{ .routes = .{ .rx = .{RAMP}, .tx = .{EAST} } }); + for (@range(i16, 1, M, 1)) |x| { + @set_color_config(x, 0, @get_color(CHANNEL), .{ + .routes = .{ .rx = .{RAMP}, .tx = .{EAST} }, + .switches = .{ .pos1 = .{ .rx = WEST } }, + }); + } + + // Receivers: statically routed, duplicating to the ramp and onward east. The outermost + // one terminates the stream instead of running off the east edge, which also makes its + // router the thing that removes the words nobody kept. + for (@range(i16, 0, M, 1)) |q| { + if (FILTER == 0) { + if (q == M - 1) { + @set_color_config(M + q, 0, @get_color(CHANNEL), + .{ .routes = .{ .rx = .{WEST}, .tx = .{RAMP} } }); + } else { + @set_color_config(M + q, 0, @get_color(CHANNEL), + .{ .routes = .{ .rx = .{WEST}, .tx = .{RAMP, EAST} } }); + } + } else { + if (q == M - 1) { + @set_color_config(M + q, 0, @get_color(CHANNEL), .{ + .routes = .{ .rx = .{WEST}, .tx = .{RAMP} }, + .filter = .{ + .kind = .{ .counter = true }, + .count_data = true, + .init_counter = window_start(q), + .limit1 = STREAM - 1, + .max_counter = K - 1, + }, + }); + } else { + @set_color_config(M + q, 0, @get_color(CHANNEL), .{ + .routes = .{ .rx = .{WEST}, .tx = .{RAMP, EAST} }, + .filter = .{ + .kind = .{ .counter = true }, + .count_data = true, + .init_counter = window_start(q), + .limit1 = STREAM - 1, + .max_counter = K - 1, + }, + }); + } + } + } + + @export_name("val", [*]f32, true); + @export_name("got", [*]f32, true); + @export_name("main", fn()void); +} diff --git a/tests/csl_runtime/handwritten/shift_bundle/receiver.csl b/tests/csl_runtime/handwritten/shift_bundle/receiver.csl new file mode 100644 index 00000000..7bb2e07d --- /dev/null +++ b/tests/csl_runtime/handwritten/shift_bundle/receiver.csl @@ -0,0 +1,52 @@ +// One receiver of the bundle. Its router never switches; a counter filter in the layout +// decides which of the words passing through are handed to this compute element. +// +// ``takes`` is how many that is: the whole stream when no filter is configured (used to +// establish the arrival order), otherwise the ``words`` addressed to this receiver. + +param memcpy_params: comptime_struct; +param stream: i16; +param words: i16; +param takes: i16; + +const sys_mod = @import_module("", memcpy_params); + +const channel: color = @get_color(1); + +var val: [words]f32; +var got: [stream]f32; +var __val_ptr: [*]f32 = &val; +var __got_ptr: [*]f32 = &got; + +const got_dsd = @get_dsd(mem1d_dsd, .{ .tensor_access = |i|{takes} -> got[i] }); +const in_dsd = @get_dsd(fabin_dsd, .{ + .extent = takes, + .fabric_color = channel, + .input_queue = @get_input_queue(0), +}); + +const recv_id = @get_local_task_id(8); +const done_id = @get_local_task_id(9); + +task recv_task() void { + @fmovs(got_dsd, in_dsd, .{ .async = true, .activate = done_id }); +} + +task done_task() void { + sys_mod.unblock_cmd_stream(); +} + +fn main() void { + for (@range(i16, 0, stream, 1)) |i| { + got[i] = -1.0; + } + @activate(recv_id); +} + +comptime { + @export_symbol(__val_ptr, "val"); + @export_symbol(__got_ptr, "got"); + @bind_local_task(recv_task, recv_id); + @bind_local_task(done_task, done_id); + @export_symbol(main, "main"); +} diff --git a/tests/csl_runtime/handwritten/shift_bundle/run.py b/tests/csl_runtime/handwritten/shift_bundle/run.py new file mode 100644 index 00000000..a8f99c75 --- /dev/null +++ b/tests/csl_runtime/handwritten/shift_bundle/run.py @@ -0,0 +1,76 @@ +#!/usr/bin/env cs_python +""" +Host side of the counted-switch prototype. + +Sender x holds K words 10*(x+1) + w and ships them M steps east. With --filter 0 every +receiver takes the whole stream, which reports the order the words arrive in; with +--filter 1 the counter filters are active and every receiver should end up with exactly +the block addressed to it. + +Exits non-zero on a mismatch, after printing the full picture either way. +""" +import argparse +import sys + +import numpy as np +from cerebras.sdk.runtime import sdkruntimepybind as crt + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument('outdir') + parser.add_argument('--M', type=int, required=True) + parser.add_argument('--K', type=int, default=1) + parser.add_argument('--filter', type=int, default=0) + args = parser.parse_args() + + m, k = args.M, args.K + width = 2 * m + stream = m * k + + runner = crt.SdkRuntime(args.outdir, suppress_simfab_trace=True) + val_id = runner.get_id('val') + got_id = runner.get_id('got') + + vals = np.array([[10 * (x + 1) + w for w in range(k)] for x in range(width)], + dtype=np.float32) + + runner.load() + runner.run() + runner.memcpy_h2d(val_id, vals.ravel(), 0, 0, width, 1, k, streaming=False, + data_type=crt.MemcpyDataType.MEMCPY_32BIT, + order=crt.MemcpyOrder.ROW_MAJOR, nonblock=False) + runner.launch('main', nonblock=False) + got = np.zeros(width * stream, dtype=np.float32) + runner.memcpy_d2h(got, got_id, 0, 0, width, 1, stream, streaming=False, + data_type=crt.MemcpyDataType.MEMCPY_32BIT, + order=crt.MemcpyOrder.ROW_MAJOR, nonblock=False) + runner.stop() + + got = got.reshape(width, stream) + for x in range(width): + role = 'sender ' if x < m else 'receiver' + print(f' PE {x} ({role}) val={vals[x]} got={got[x]}') + + if args.filter == 0: + # Sends run east to west, so sender m-1's block arrives first. + expected = np.concatenate([vals[x] for x in range(m - 1, -1, -1)]) + bad = [x for x in range(m, width) if not np.allclose(got[x], expected, atol=1e-6)] + if bad: + print(f'FAILED: receivers {bad} did not see {expected}') + return 1 + print(f'Passed M={m} K={k}: every receiver saw the whole stream as {expected}.') + return 0 + + # Receiver m + q is the destination of sender q. + bad = [q for q in range(m) if not np.allclose(got[m + q][:k], vals[q], atol=1e-6)] + if bad: + for q in bad: + print(f'FAILED: receiver {m + q} wanted {vals[q]}, kept {got[m + q][:k]}') + return 1 + print(f'Passed M={m} K={k}: every receiver kept exactly the block addressed to it.') + return 0 + + +if __name__ == '__main__': + sys.exit(main()) diff --git a/tests/csl_runtime/handwritten/shift_bundle/sender.csl b/tests/csl_runtime/handwritten/shift_bundle/sender.csl new file mode 100644 index 00000000..5d8b4e57 --- /dev/null +++ b/tests/csl_runtime/handwritten/shift_bundle/sender.csl @@ -0,0 +1,67 @@ +// One sender of the bundle: inject this PE's words, then hand the router over to relay mode. +// +// All senders are launched at once and the order sorts itself out in the fabric: a sender +// west of us cannot push its words through our router while we still receive from the ramp, +// so it waits on the link until our SWITCH_ADV has moved us to relay mode. That is the whole +// serialisation mechanism -- no barrier, no counting. + +param memcpy_params: comptime_struct; +param stream: i16; +param words: i16; +param hands_over: bool; + +const sys_mod = @import_module("", memcpy_params); +const ctrl = @import_module(""); + +const channel: color = @get_color(1); + +var val: [words]f32; +var got: [stream]f32; +var __val_ptr: [*]f32 = &val; +var __got_ptr: [*]f32 = &got; + +const val_dsd = @get_dsd(mem1d_dsd, .{ .tensor_access = |i|{words} -> val[i] }); +const out_dsd = @get_dsd(fabout_dsd, .{ + .extent = words, + .fabric_color = channel, + .output_queue = @get_output_queue(2), +}); +// The control wavelet has to leave through the same queue as the data, so that it stays +// behind our own words and only advances our router once they are out. +const switch_dsd = @get_dsd(fabout_dsd, .{ + .extent = 1, + .fabric_color = channel, + .control = true, + .output_queue = @get_output_queue(2), +}); + +const send_id = @get_local_task_id(8); +const done_id = @get_local_task_id(9); + +task send_task() void { + @fmovs(out_dsd, val_dsd, .{ .async = true, .activate = done_id }); +} + +task done_task() void { + if (hands_over) { + @mov32(switch_dsd, ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0)); + } + sys_mod.unblock_cmd_stream(); +} + +fn main() void { + // Senders keep nothing, so the host sees the sentinel and cannot mistake a stray + // delivery here for a real one. + for (@range(i16, 0, stream, 1)) |i| { + got[i] = -1.0; + } + @activate(send_id); +} + +comptime { + @export_symbol(__val_ptr, "val"); + @export_symbol(__got_ptr, "got"); + @bind_local_task(send_task, send_id); + @bind_local_task(done_task, done_id); + @export_symbol(main, "main"); +} diff --git a/tests/csl_runtime/test_shift_bundle_filters.sh b/tests/csl_runtime/test_shift_bundle_filters.sh new file mode 100644 index 00000000..a298cd4a --- /dev/null +++ b/tests/csl_runtime/test_shift_bundle_filters.sh @@ -0,0 +1,46 @@ +#!/bin/sh +# Hand-written prototype of the counted switch used by 1D shift bundles. +# +# Stage 1 runs without filters and checks the order words arrive in, which is what proves the +# send-order hand-over: every receiver must see the senders in descending order, including the +# case where senders relay for each other. +# +# Stage 2 turns the counter filters on and checks that each receiver keeps only its own block. +# Together they pin down the hardware contract the compiler relies on: +# +# * a sender's own SWITCH_ADV advances its own router, and over-advancing a router that is +# already on its last position is harmless, +# * a receiver transmitting to RAMP and EAST duplicates rather than consumes, +# * a filter withholds a wavelet from the compute element without removing it from the +# network, and a wavelet nobody keeps is dropped by the terminating router, +# * the counter delivers iff counter <= max_counter, counting modulo limit1 + 1 from +# init_counter. + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +SRC="$SCRIPT_DIR/handwritten/shift_bundle" +OUT="shift_bundle_hw" + +compile_and_run() { + m=$1 + k=$2 + filter=$3 + width=$((2 * m)) + + rm -rf "$OUT" + cslc --arch=wse2 "$SRC/layout.csl" -o "$OUT" \ + --fabric-dims=$((7 + width)),3 --fabric-offsets=4,1 --memcpy --channels=1 \ + --params=M:$m,K:$k,FILTER:$filter + timeout -s 9 240 cs_python "$SRC/run.py" "$OUT" --M "$m" --K "$k" --filter "$filter" +} + +for case in "3 1" "4 1" "3 2"; do + set -- $case + echo "--- stage 1: arrival order, M=$1 K=$2 (no filters) ---" + compile_and_run "$1" "$2" 0 + echo "--- stage 2: filtered delivery, M=$1 K=$2 ---" + compile_and_run "$1" "$2" 1 +done + +rm -rf "$OUT" +echo "Passed: counted switch by send order with filtered delivery." From 127183e6b560241c0c17d73e1332c8fb5cd18eb7 Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 17 Aug 2026 22:37:59 +0200 Subject: [PATCH 17/68] new bundling draft --- irspec/docs/spatial/routing.md | 110 ++- irspec/docs/spatial/spatial.md | 13 +- .../spatial/simple/exchange_bundle_1D.sptl | 61 ++ samples/spatial/simple/shift_bundle_1D.sptl | 48 +- .../sort/batcher_oddeven_bundled_1D.sptl | 96 +-- spada/lowering/spatial_ir_to_csl.py | 40 +- spada/runtime/runtime.py | 6 +- spada/syntax/csl/constants.py | 8 + spada/syntax/csl/routing.py | 303 +++++++- spada/syntax/spatial_ir/canonicalization.py | 4 - spada/syntax/spatial_ir/irnodes.py | 35 +- spada/syntax/spatial_ir/language.lark | 3 +- spada/syntax/spatial_ir/lark_to_ir.py | 11 +- spada/syntax/spatial_ir/shift_bundles.py | 534 +++++--------- spada/syntax/spatial_ir/stream_lifetime.py | 115 ++- .../test_batcher_oddeven_bundled_1d.sh | 50 ++ tests/csl_runtime/test_exchange_bundle_1d.sh | 56 ++ tests/csl_runtime/test_shift_bundle_1d.sh | 31 +- tests/spatial_ir/test_shift_bundles.py | 677 ++++++------------ 19 files changed, 1167 insertions(+), 1034 deletions(-) create mode 100644 samples/spatial/simple/exchange_bundle_1D.sptl create mode 100755 tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh create mode 100755 tests/csl_runtime/test_exchange_bundle_1d.sh diff --git a/irspec/docs/spatial/routing.md b/irspec/docs/spatial/routing.md index 6c9b5caa..9ef4cead 100644 --- a/irspec/docs/spatial/routing.md +++ b/irspec/docs/spatial/routing.md @@ -251,22 +251,116 @@ control message* on the channel, one per position to be traversed. It follows th using the configuration that is being retired, and advances the router of each PE it traverses, after all data of the epoch. -!!! warning "WSE: Advances Are Not Selective" +!!! warning "WSE: A Control Message's Payload Selects Nothing" A CSL control wavelet nominally carries up to eight per-router switching commands (``'s `MAX_CMDS`), which would let one message advance some routers on a path and leave - others alone. **On the WSE hardware it does not work that way.** Measured on the simulator, only - command slot 0 is ever executed, and **every** switch-configured router the wavelet reaches - applies it; slots 1–7 had no effect in any topology tested — the sender's own router, one hop, - two hops through a plain relay, and two switch-configured routers in sequence. The compiler - therefore emits `encode_single_payload`, which writes slot 0 only. + others alone. **On the WSE hardware the payload does not work that way.** Measured on the + simulator, only command slot 0 is ever executed, and **every** switch-configured router the + wavelet reaches applies it; slots 1–7 had no effect in any topology tested — the sender's own + router, one hop, two hops through a plain relay, and two switch-configured routers in sequence. + A generated kernel built on the opposite assumption, advancing the fourth router of a path with + an `[ADV, NOP, NOP, ADV]` chain, stalled in the fabric. The compiler therefore emits + `encode_single_payload`, which writes slot 0 only. The consequence is that a message cannot advance one router while leaving another on the same path where it is. *If the routers along one path would have to advance by different amounts, a - compile error is raised.* A router that is already on its last configuration is exempt: it never - routes anything again, so a message passing through may over-advance it harmlessly. + compile error is raised.* A router that is already on its last position is exempt: outside + `ring_mode` an advance past the last position is a no-op, so a message passing through may + over-advance it harmlessly. (Such a router is not necessarily finished — it may keep relaying the + same configuration for the rest of the kernel, which is exactly what the bundle below relies + on.) + + This is a property of the *payload*, not of switching: a router can still be switched at a time + only it knows, and delivery to a compute element can still be made selective, by the two + mechanisms the next section combines. Because the control message travels the path of the retired configuration in order behind the data, a receiving PE needs to emit nothing to advance its own router: the ordering required by the [lemma above](#undefined-behavior) is provided by the fabric. A receiver's `close` therefore has no runtime effect; it exists so that the lifetime of the stream — and hence the number of elements it carries — is stated by every participant and can be checked. + +## Overlapping Interval Shifts + +The correctness conditions above rule out one shape that occurs constantly: a run of consecutive PEs +all shifting the same distance $d$ along an axis, as in + +``` +dataflow i16 i, i16 j in [0:D + M, 0] { + stream fwd = relative_stream(D, 0) { hops = auto, channel = 0 } +} +``` + +with the PEs in `[0:M)` sending and those in `[D:D+M)` receiving. Source $p$'s word passes through +the routers of sources $p+1, \dotsc, M-1$, so the paths share PEs within one epoch. Written as one +stream per source, that is a channel each, $M$ colors for a shift; the alternative is a chain of +single-hop stores and forwards, which serializes the whole run behind $d$ hops of copying. + +Neither is necessary. The compiler recognizes this pattern — `detect_shift_bundles` — and lowers the +whole run onto **one channel**, giving the routers configurations that are switched only by events +the PE owning them knows locally: + +``` +PE: 0 1 2 3 4 5 (M = 3, D = 3) +role: src0 src1 src2 dst0 dst1 dst2 +sends: 3rd 2nd 1st -- -- -- +routes: R->E R->E R->E W->{R,E} W->{R,E} W->R +pos1: -- W->E W->E -- -- -- +filter: -- -- -- win 2 win 1 win 0 +``` + +**Sources inject, then relay.** A source starts at `rx = RAMP, tx = {EAST}` with `pos1` taking +`rx = WEST`, sends its own words, and its `close` advances its own router into relay mode. The +trigger is local — *"my own send is done"* — which is what makes it expressible at all, given that +the payload of the resulting message [selects nothing](#lowering-to-switches). + +**The order is descending, and enforces itself.** The source nearest the destinations goes first. No +schedule or barrier is needed: a source further away cannot push a word through its neighbour's +router while that neighbour is still injecting from its ramp, so it waits on the link. Backpressure +serializes the run in exactly the order the switches expect. + +**Destinations do not switch; a filter picks their words.** Each destination is statically routed to +`tx = {RAMP, EAST}`, which *duplicates* rather than consumes: every destination's router sees the +entire stream, in one order, and the one the stream reaches last uses `tx = {RAMP}` to take it out of +the network. Which words a destination hands to its compute element is decided by a counter filter +on that color, one linear function of the PE coordinate, so a single `@set_color_config` covers the +whole run. A control message passes such a router without being counted (`count_data = true`) and +without being filtered. + +!!! note "Note: Counter Filter Arithmetic" + As measured on the simulator, a counter filter starts at `init_counter`, increments on every data + wavelet, wraps to zero after `limit1`, and hands a wavelet to the compute element iff the counter + is at most `max_counter`. A window of `words` out of a stream of `length * words` is therefore + `limit1 = length * words - 1`, `max_counter = words - 1`, and an `init_counter` chosen so that + the counter reads zero as the wanted block arrives. + `tests/csl_runtime/test_shift_bundle_filters.sh` is the hand-written layout this was measured + with. + +!!! danger "Error: Too Many Wavelet Filters" + WSE-2 has four filters per PE, of which the `memcpy` module reserves one, so **three** are usable + (`FILTERS_PER_PE`). A PE needs one per color it filters. *If a PE would need more, a compile error + is raised*; the fix is to give some of the streams their own channels, which trades filters for + colors. `batcher_oddeven_bundled_1D.sptl` hits this at $2^4$ PEs, where a PE receives a bundle in + six phases. + + Reconfiguring a filter while wavelets are still in flight on its color is a data race, and the + destination that terminates the stream cannot be reconfigured until the stream has drained, + because its router is what removes the wavelets from the network. Filters are consequently set up + once, at layout time, and never reused between phases. + +Bundling applies only when every run it decomposes into has at least two sources and is no longer +than the shift distance, so that no PE is both a source and a destination; a shift of one PE is left +alone, since a chain at distance one is already sequenced by ordinary switch positions. Anything +else falls back to the per-hop lowering, and to the errors above if that conflicts. + +This arrangement is the one used in Schnyder's *Distributed Sorting on the Cerebras Wafer-Scale +Engine* (fig. 7.6) for the 2D reduce-scatter, and is known to run on WSE-2. + +!!! note "Note: Multiple Rounds on One Color" + Two mechanisms are deliberately left unused, and are what to reach for if the four switch + positions or three filters run out. `SWITCH_RST` restores the initial configuration of every + router a message passes, which retires a whole path with one wavelet. Teardown-based + reconfiguration reprograms the routers between rounds outright, which is the only known way to + put an arbitrary *sequence* of sends and receives on one color: for a long enough sequence there + is a PE for which no fixed cycle of switch positions exists. Until then, splitting the rounds + across channels — as `bitonic_sort_1D.sptl` does — remains the per-kernel fallback. diff --git a/irspec/docs/spatial/spatial.md b/irspec/docs/spatial/spatial.md index 42605014..79749a8f 100644 --- a/irspec/docs/spatial/spatial.md +++ b/irspec/docs/spatial/spatial.md @@ -493,20 +493,15 @@ The routing configuration is set up as follows: stream stream_name = relative_stream(dx, dy) { // Optional routing declaration hops = [(dx_1, dy_1), (dx_2, dy_2), ... , (dx_n, dy_n)], - channel = channel_id, - count = k + channel = channel_id } ``` where `hops` is a list of relative hops that the data takes between the sender and receiver. Each hop is given by a pair of constant literals, the sum of their absolute value must be 1. The sum of all the hops must be equal to the relative position of the stream. -`count` is optional. If it is a compile-time integer `k`, each PE transfers exactly `k` -fabric words on this stream edge in the phase, which enables counted router switching -on overlapping 1D shifts. Switching uses two colors **per phase** (one per direction). -Wave quotas depend on the shift distance, so later phases with a different `d` receive -a fresh color pair rather than reloading the same two colors. If `count` is omitted or -`count = auto`, the stream is unbounded: the compiler does not infer a length and does -not apply counted switching. +How many words the stream carries is not stated here but by its type: a bounded +`stream` closes after `BOUND` elements, which is what frees its channel +(see [Streams](#streams) and [closing streams](#closing-streams-with-close)). If two messages (elements of a `send`) are routed through a PE simultaneously, it must be ensured that they do not share a `channel`. diff --git a/samples/spatial/simple/exchange_bundle_1D.sptl b/samples/spatial/simple/exchange_bundle_1D.sptl new file mode 100644 index 00000000..41abc8d4 --- /dev/null +++ b/samples/spatial/simple/exchange_bundle_1D.sptl @@ -0,0 +1,61 @@ +/** + * Two overlapping 1D interval shifts, in opposite directions, repeated R times. + * + * D + M PEs on a line. The PEs in [0:M) and those in [D:D+M) swap values pairwise: PE i + * trades with PE i + D. Each direction is a shift bundle of its own on a color of its own, + * so one exchange costs two colors however many pairs there are. M <= D, so the two halves + * are disjoint. + * + * This is the skeleton of a sorting network's phase (see the bundled Batcher), reduced to + * one operation: every PE is a source on one color and a filtered destination on the other. + * R repeats it in R phases, none of which shares a color with another, so every PE ends up + * with R wavelet filters and R is what the per-PE filter budget limits. + * + * Constraints: 2 <= M <= D, R >= 1. + **/ +kernel @exchange_bundle_1d( + stream[D + M, 1] readonly inp, + stream[D + M, 1] writeonly out +) { + place i16 i, i16 j in [0:D + M, 0] { + f32 val + f32 tmp + } + + phase { + compute i16 i, i16 j in [0:D + M, 0] { + await receive(val, inp[i, j]) + } + } + + for i16 r in [0:R] { + phase { + dataflow i16 i, i16 j in [0:D + M, 0] { + stream fwd = relative_stream(D, 0) { + hops = auto, + channel = auto + } + stream bwd = relative_stream(-D, 0) { + hops = auto, + channel = auto + } + } + compute i16 i, i16 j in [0:M, 0] { + await send(val, fwd) + await receive(tmp, bwd) + val = tmp + } + compute i16 i, i16 j in [D:D + M, 0] { + await receive(tmp, fwd) + await send(val, bwd) + val = tmp + } + } + } + + phase { + compute i16 i, i16 j in [0:D + M, 0] { + await send(val, out[i, j]) + } + } +} diff --git a/samples/spatial/simple/shift_bundle_1D.sptl b/samples/spatial/simple/shift_bundle_1D.sptl index d09128e0..1db7dee7 100644 --- a/samples/spatial/simple/shift_bundle_1D.sptl +++ b/samples/spatial/simple/shift_bundle_1D.sptl @@ -1,55 +1,51 @@ /** - * Counted 1D interval shift, and nothing else. + * An overlapping 1D interval shift, and nothing else. * - * 2M PEs on a line. Sources [0:M) each send one f32 to dest M steps east. - * Dest i+M overwrites its value with the word from i. Sources keep theirs. + * D + M PEs on a line. Sources [0:M) each send one f32 to the PE D steps east, which + * overwrites its own value with it. Sources keep theirs. M <= D, so no PE is both a + * source and a destination. * - * Overlapping paths share one eastbound color. A WSE color holds a single - * (rx, tx) pair at a time; the compiler sequences two pairs with a counted - * switch (forward-then-inject on the source half, absorb-then-forward on - * the dest half). There is no westbound stream, no CAS, and no d=1 stage. + * The paths overlap: source i's word passes through the routers of sources i+1 .. M-1. + * A color holds one (rx, tx) pair at a time, so the sources take turns -- nearest the + * destinations first, each handing its router over to relay mode once its own word is + * out. The destinations are statically routed and pick their word out of the stream + * with a counter filter. All of it on one color. * - * Constraints: M >= 2 (M=1 is a static hop; no overlap). + * Constraints: 2 <= M <= D. **/ -kernel @shift_bundle_1d( - stream[2 * M, 1] readonly inp, - stream[2 * M, 1] writeonly out +kernel @shift_bundle_1d( + stream[D + M, 1] readonly inp, + stream[D + M, 1] writeonly out ) { - place i16 i, i16 j in [0:2 * M, 0] { + place i16 i, i16 j in [0:D + M, 0] { f32 val } phase { - compute i16 i, i16 j in [0:2 * M, 0] { + compute i16 i, i16 j in [0:D + M, 0] { await receive(val, inp[i, j]) } } phase { - dataflow i16 i, i16 j in [0:M, 0] { - stream fwd = relative_stream(M, 0) { + // One declaration for the whole line: both halves have to agree on the channel, and the + // bundle is what lets them share it despite the overlap. + dataflow i16 i, i16 j in [0:D + M, 0] { + stream fwd = relative_stream(D, 0) { hops = auto, - channel = auto, - count = 1 - } - } - dataflow i16 i, i16 j in [M:2 * M, 0] { - stream fwd = relative_stream(M, 0) { - hops = auto, - channel = auto, - count = 1 + channel = 0 } } compute i16 i, i16 j in [0:M, 0] { await send(val, fwd) } - compute i16 i, i16 j in [M:2 * M, 0] { + compute i16 i, i16 j in [D:D + M, 0] { await receive(val, fwd) } } phase { - compute i16 i, i16 j in [0:2 * M, 0] { + compute i16 i, i16 j in [0:D + M, 0] { await send(val, out[i, j]) } } diff --git a/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl b/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl index 0a0e71d6..793ce8a9 100644 --- a/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl +++ b/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl @@ -9,13 +9,23 @@ * Each drawn comparator is two messages (low -> high, then high -> low). * Lower index keeps min; higher index keeps max. * - * Counted color switching (count = 1): overlapping matchings at distance d - * share one eastbound and one westbound color in that phase. Each PE - * forwards a known number of waves, then injects or absorbs. Distance-1 - * stages keep a static RAMP/EAST (WEST) pair. - * Each CAS phase uses 2 colors. Colors are not reused across phases: - * the wave quotas change with d, so a counted two-state program cannot - * be reloaded onto the same color. Total colors = 2 * (# of (l,p) phases). + * This is the bundled variant of batcher_oddeven_1D: a whole phase uses two colors, one per + * direction, where the static version needs one per matching. The comparators of a phase at + * distance d partition the line into alternating blocks of d PEs, and each block ships its + * keys d steps into the next -- an overlapping interval shift, which is what a bundle is. The + * low PEs of a block take turns nearest-the-partner-first, handing their routers over to relay + * mode as they finish; the high PEs are statically routed and pick their key out of the stream + * with a counter filter. See irspec/docs/spatial/routing.md. + * + * Both streams of a phase are declared once for the whole line, which is what puts every + * matching of that phase on one channel. Colors are not reused across phases: the router + * configurations depend on d, and a color holds one set of them. + * + * What caps L here is the three filters a PE can use, not the 21 colors. A PE needs one per + * phase it receives a bundle in, and only the d = 1 phases are unbundled (a chain of shifts + * of one is sequenced by ordinary switch positions), so the count is the number of phases + * with d >= 2: three at L = 3, six at L = 4. Beyond L = 3 the phases have to go back to a + * color per matching, as batcher_oddeven_1D does. * * Example L=3 (n=8): * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) @@ -25,7 +35,7 @@ * (l=3,p=2) dist 2: (2,4)(3,5) * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) * - * Constraints: L >= 1 + * Constraints: 1 <= L <= 3 **/ kernel @batcher_oddeven_1d( stream[1<( for i16 l in [1:L+1] { // p = 1: all PEs participate, dist = 1<<(l-1). - // Offset r is one disjoint matching; counted switching shares two colors. phase { - for i16 r in [0:1<<(l-1)] { - dataflow i16 i, i16 j in [r:1< fwd = relative_stream(1<<(l-1), 0) { - hops = auto, - channel = auto, - count = 1 - } - stream bwd = relative_stream(-(1<<(l-1)), 0) { - hops = auto, - channel = auto, - count = 1 - } + dataflow i16 i, i16 j in [0:1< fwd = relative_stream(1<<(l-1), 0) { + hops = auto, + channel = auto } - dataflow i16 i, i16 j in [(r + (1<<(l-1))):1< fwd = relative_stream(1<<(l-1), 0) { - hops = auto, - channel = auto, - count = 1 - } - stream bwd = relative_stream(-(1<<(l-1)), 0) { - hops = auto, - channel = auto, - count = 1 - } + stream bwd = relative_stream(-(1<<(l-1)), 0) { + hops = auto, + channel = auto } + } + // Offset r is one disjoint matching; together they are the phase's shift. + for i16 r in [0:1<<(l-1)] { compute i16 i, i16 j in [r:1<( // p = 2 .. l: skip first/last dist PEs of each 2^l box, dist = 1<<(l-p). for i16 p in [2:l+1] { phase { + dataflow i16 i, i16 j in [0:1< fwd = relative_stream(1<<(l-p), 0) { + hops = auto, + channel = auto + } + stream bwd = relative_stream(-(1<<(l-p)), 0) { + hops = auto, + channel = auto + } + } + for i16 r in [0:1<<(l-p)] { for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { - hops = auto, - channel = auto, - count = 1 - } - stream bwd = relative_stream(-(1<<(l-p)), 0) { - hops = auto, - channel = auto, - count = 1 - } - } - dataflow i16 i, i16 j in [(b + ((2 * (1<<(l-p))) + r)):(b + (1< fwd = relative_stream(1<<(l-p), 0) { - hops = auto, - channel = auto, - count = 1 - } - stream bwd = relative_stream(-(1<<(l-p)), 0) { - hops = auto, - channel = auto, - count = 1 - } - } - compute i16 i, i16 j in [(b + ((1<<(l-p)) + r)):((b + (1< bool: return self.rx != previous.rx and self.tx != previous.tx +@dataclass(frozen=True) +class FilterConfig: + """ + A counter filter: which of the wavelets passing a router are handed to its compute element. + + The router keeps a counter per filter. It starts at ``init_counter``, advances on every wavelet + the filter counts, and wraps to zero after ``limit1``, so it cycles through ``limit1 + 1`` + values. A wavelet is delivered iff the counter is at most ``max_counter``, and *withheld* + otherwise -- withheld is not the same as consumed: the wavelet carries on along the router's + ``tx`` directions, so PEs further along still see it. Only a router that transmits to the ramp + alone drops what it withholds, which is what takes a wavelet out of the network. + (Measured; see ``tests/csl_runtime/test_shift_bundle_filters.sh``. The manual describes + ``max_counter`` as exclusive, but a wavelet arriving at ``counter == max_counter`` is delivered.) + + The fields are expression strings rather than integers because a filter's window generally + depends on where the PE sits: a shift bundle's receivers share one ``@set_color_config`` whose + ``init_counter`` is a function of ``pe_x``. + + A PE can hold only ``constants.FILTERS_PER_PE`` of these across all of its colors. + """ + init_counter: str + limit1: str + max_counter: str + + def as_csl(self) -> str: + return ('.{ .kind = .{ .counter = true }, .count_data = true, .init_counter = %s, ' + '.limit1 = %s, .max_counter = %s }' + % (self.init_counter, self.limit1, self.max_counter)) + + def logical_positions(configs: list[RouteConfig]) -> tuple[list[RouteConfig], bool]: """ Reduces a router's configuration sequence to one period, reporting whether it repeats. @@ -136,9 +166,13 @@ class ColorSwitchPlan: ``positions[0]`` is the base configuration, and every further entry becomes a switch position. Consecutive identical configurations are collapsed by :meth:`add`, so a router that keeps the same configuration across an epoch boundary consumes no switch position and needs no advance. + + A :class:`FilterConfig` applies to the color as a whole rather than to a position: filters cannot + be switched, and rewriting one while wavelets are in flight is unsafe. """ positions: list[RouteConfig] = field(default_factory=list) ring_mode: bool = False + filter: 'FilterConfig | None' = None def add(self, config: RouteConfig) -> None: if self.positions and self.positions[-1] == config: @@ -166,18 +200,19 @@ def uses_switches(self) -> bool: return len(self.cycle[0]) > 1 def as_csl(self) -> str: - base = self.positions[0].as_csl() - if not self.uses_switches: - return '.{ .routes = %s }' % base - - hardware = self.hardware_positions - switches = [ - '.pos%d = %s' % (index, config.as_switch_position(hardware[index - 1])) - for index, config in enumerate(hardware[1:], start=1) - ] - if self.cycle[1]: - switches.append('.ring_mode = true') - return '.{ .routes = %s, .switches = .{ %s } }' % (base, ', '.join(switches)) + fields = ['.routes = %s' % self.positions[0].as_csl()] + if self.uses_switches: + hardware = self.hardware_positions + switches = [ + '.pos%d = %s' % (index, config.as_switch_position(hardware[index - 1])) + for index, config in enumerate(hardware[1:], start=1) + ] + if self.cycle[1]: + switches.append('.ring_mode = true') + fields.append('.switches = .{ %s }' % ', '.join(switches)) + if self.filter is not None: + fields.append('.filter = %s' % self.filter.as_csl()) + return '.{ %s }' % ', '.join(fields) def validate(self, color: int, location: str) -> None: """ @@ -311,10 +346,18 @@ class _RouteSite: """ One ``@set_color_config`` target: the rectangle of PEs (already shifted by the relay offset) that receive a route configuration for one color. + + A site is normally a compute rectangle shifted by a relay offset, and is configured from that + rectangle's layout loop as ``pe_x + offset``. ``absolute`` marks the exception: a site whose PEs + are not any rectangle shifted as a whole -- the pure relays between the two halves of a shift + bundle, say, whose count has nothing to do with either half's width -- and which therefore gets a + layout loop of its own. It is excluded from equality so that :func:`_find_site` can still look a + site up by its PEs alone. """ color: int x_range: tuple[int, int, int] y_range: tuple[int, int, int] + absolute: bool = field(default=False, compare=False) def as_rectangle(self) -> Rectangle: return Rectangle(self.x_range, self.y_range, None) @@ -341,6 +384,112 @@ class _RouteEntry: stream: spir.Identifier #: Routing identity of the stream; stable across the per-rectangle renaming of ``inline_phases`` group: str = '' + #: Which of the wavelets reaching these routers are handed to their compute element. Belongs to + #: the color rather than to this one configuration, so all entries of a site must agree. + filter: FilterConfig | None = None + + +def _bundle_ports(bundle: shift_bundles.ShiftBundle) -> tuple[str, str]: + """Returns the ``(incoming, outgoing)`` router ports along a bundle's direction of travel.""" + if bundle.axis == 'x': + return ('WEST', 'EAST') if bundle.sign > 0 else ('EAST', 'WEST') + return ('NORTH', 'SOUTH') if bundle.sign > 0 else ('SOUTH', 'NORTH') + + +def _window_start(variable: str, first: int, step: int, words: int) -> str: + """ + Returns the counter value a destination's filter starts at, as an expression in the loop variable. + + Every destination sees the whole stream, in one order, so which words a destination keeps is + decided by where it sits: the ``p``-th destination along the direction of travel keeps the block + the ``p``-th-from-last source sent, which begins ``(p + 1) * words`` short of the end of the + cycle. Starting the counter there brings it to zero just as that block arrives. + + :param first: The coordinate of the destination the stream reaches first, where ``p`` is zero. + :param step: ``+1`` or ``-1``, the direction the coordinate grows in as ``p`` grows. + """ + offset = 1 - first if step > 0 else first + 1 + if step > 0: + inner = variable if offset == 0 else f'{variable} + {offset}' if offset > 0 else f'{variable} - {-offset}' + else: + inner = f'{offset} - {variable}' + if words == 1: + return inner + return f'({inner}) * {words}' + + +def _bundle_route_entries(bundle: shift_bundles.ShiftBundle, color: int, order: tuple[int, int, int], + rect_index: int, stream_name: spir.Identifier) -> list[tuple['_RouteSite', '_RouteEntry']]: + """ + Returns the route configurations of one shift bundle, replacing the per-hop ones. + + Four kinds of router take part, and none of them is a compute rectangle shifted as a whole -- the + relays in between are as many as the shift distance minus the run length, which is neither half's + width -- so every site here is a standalone one. + + The sources all get the same pair of configurations, injecting and then relaying, including the + one furthest from the destinations which has nothing to relay for. Giving it a switch position it + never uses costs nothing and keeps one ``@set_color_config`` for the whole run; its own + switch-advance wavelet moves it into a configuration that never carries anything, and the + wavelets of the sources behind it pass routers already sitting on their last position, which is a + no-op. + + :param order: The switch-position order key of the send that this bundle carries. + """ + incoming, outgoing = _bundle_ports(bundle) + variable = 'pe_x' if bundle.axis == 'x' else 'pe_y' + first, step = bundle.destination_order() + limit1 = str(bundle.length * bundle.words - 1) + max_counter = str(bundle.words - 1) + + collected: list[tuple[_RouteSite, _RouteEntry]] = [] + + def add(span: tuple[int, int], configs: list[RouteConfig], + wavelet_filter: FilterConfig | None = None) -> None: + start, stop = span + if start >= stop: + return + along = (start, stop, 1) + site = _RouteSite(color=color, + x_range=along if bundle.axis == 'x' else bundle.cross, + y_range=bundle.cross if bundle.axis == 'x' else along, + absolute=True) + for config in configs: + collected.append((site, _RouteEntry(config, order, rect_index, (0, 0), stream_name, + bundle.group, wavelet_filter))) + + add(bundle.sources(), [RouteConfig(('RAMP', ), (outgoing, )), RouteConfig((incoming, ), (outgoing, ))]) + add(bundle.relays(), [RouteConfig((incoming, ), (outgoing, ))]) + + # The destinations the stream still has to travel past hand a copy to their ramp and pass it on; + # the one it reaches last has nowhere to pass it and so is what removes it from the network. + low, high = bundle.destinations() + terminal = high - 1 if step > 0 else low + passing = (low, high - 1) if step > 0 else (low + 1, high) + add(passing, [RouteConfig((incoming, ), ('RAMP', outgoing))], + FilterConfig(_window_start(variable, first, step, bundle.words), limit1, max_counter)) + add((terminal, terminal + 1), [RouteConfig((incoming, ), ('RAMP', ))], + FilterConfig('0', limit1, max_counter)) + return collected + + +def _bundle_owners(rectangles: list[Rectangle[PEBlock]], + bundles: dict[str, list[shift_bundles.ShiftBundle]]) -> dict[str, int]: + """ + Picks the rectangle that contributes each bundle's routing. + + A bundle's sources are often declared by several compute blocks, each of which sends the same + stream; the configurations only have to be contributed once. + """ + owners: dict[str, int] = {} + for rect_index, rect in enumerate(rectangles): + sends_recvs = analysis.sends_and_receives(rect.metadata.compute) + for declaration in rect.metadata.dataflow.statements: + sent, _received = sends_recvs.get(declaration.stream_name, (False, False)) + group = stream_lifetime.stream_group_key(declaration) + if sent and group in bundles: + owners.setdefault(group, rect_index) + return owners def _stream_use_order(compute: spir.ComputeBlock) -> dict[spir.Identifier, dict[str, tuple[int, int]]]: @@ -372,7 +521,8 @@ def _offset_expression(axis: str, offset: int) -> str: def collect_routes(rectangles: list[Rectangle[PEBlock]], color_maps: list[dict[str, int]], - disable_switching: bool = False) -> dict[tuple[int, int], str]: + disable_switching: bool = False, + grid_offset: tuple[int, int] = (0, 0)) -> tuple[dict[tuple[int, int], str], list[str]]: """ Creates a parametric version of the Routing Graph (see the Spatial IR specification for more information) and returns a dictionary of code segements to add to the layout CSL file based on the streams. @@ -386,20 +536,23 @@ def collect_routes(rectangles: list[Rectangle[PEBlock]], :param color_maps: Per-rectangle mapping of stream names to colors. :param disable_switching: If True, emit each configuration as its own ``@set_color_config`` instead of merging them into switch positions. - :return: A dictionary mapping the starting point of each rectangle to a string representing the layout instructions. + :param grid_offset: Where the PE grid sits in the fabric rectangle, applied to the loop bounds of + standalone sites. Sites belonging to a rectangle inherit it from that + rectangle's loop instead. + :return: The layout instructions to place inside each rectangle's loop, keyed by the starting + point of the rectangle, together with the standalone loops of the sites that belong to + no rectangle. """ INDENT = 12 * ' ' - entries: dict[_RouteSite, list[_RouteEntry]] = {} - for rect_index, (rect, color_map) in enumerate(zip(rectangles, color_maps)): - for site, entry in _rectangle_route_entries(rect_index, rect, color_map): - entries.setdefault(site, []).append(entry) - + entries = _route_sites(rectangles, color_maps) _check_site_overlap(entries) + _check_filter_budget(entries) result = {(rect.x_range[0], rect.y_range[0]): '' for rect in rectangles} + standalone: list[str] = [] for site, site_entries in entries.items(): - site_entries.sort(key=lambda entry: entry.order) + color = f'@get_color({site.color})' # The site is configured from the loop of one rectangle: the one that owns these PEs if # there is one, otherwise the first relay that reaches them. @@ -408,9 +561,8 @@ def collect_routes(rectangles: list[Rectangle[PEBlock]], key = (owner_rect.x_range[0], owner_rect.y_range[0]) x = _offset_expression('pe_x', owner.origin_offset[0]) y = _offset_expression('pe_y', owner.origin_offset[1]) - color = f'@get_color({site.color})' - if disable_switching: + if disable_switching and not site.absolute: for entry in site_entries: plan = ColorSwitchPlan([entry.config]) text = set_color_config(x, y, color, plan, INDENT) @@ -418,23 +570,93 @@ def collect_routes(rectangles: list[Rectangle[PEBlock]], result[key] += text continue - plan = ColorSwitchPlan() + plan = ColorSwitchPlan(filter=_site_filter(site, site_entries)) for entry in site_entries: plan.add(entry.config) plan.validate(site.color, site.describe()) - result[key] += set_color_config(x, y, color, plan, INDENT) + if site.absolute: + standalone.append(_standalone_site(site, color, plan, grid_offset)) + else: + result[key] += set_color_config(x, y, color, plan, INDENT) - return result + return result, standalone + + +def _site_filter(site: '_RouteSite', site_entries: list['_RouteEntry']) -> FilterConfig | None: + """ + Returns the filter of a site, checking that every configuration contributed to it agrees. + + A filter is a property of the color at a router, not of one route configuration: it cannot be + switched along with them. + """ + filters = {entry.filter for entry in site_entries} + if len(filters) > 1: + raise SyntaxError( + f'Color {site.color} at {site.describe()} is given more than one wavelet filter, but a ' + 'router holds one filter per color and it cannot be switched.\n' + ' note: give the streams that disagree separate channels, at the cost of an ' + 'additional color') + return filters.pop() if filters else None + + +def _standalone_site(site: '_RouteSite', color: str, plan: ColorSwitchPlan, + grid_offset: tuple[int, int]) -> str: + """ + Renders a site that belongs to no rectangle as a layout loop of its own. + + :param grid_offset: Where the PE grid sits in the fabric rectangle. + """ + xb, xe, xs = site.x_range + yb, ye, ys = site.y_range + body = set_color_config('pe_x', 'pe_y', color, plan, 12 * ' ') + return (f' for (@range(i16, {xb + grid_offset[0]}, {xe + grid_offset[0]}, {xs})) |pe_x| {{\n' + f' for (@range(i16, {yb + grid_offset[1]}, {ye + grid_offset[1]}, {ys})) |pe_y| {{\n' + f'{body}' + f' }}\n' + f' }}\n') + + +def _check_filter_budget(entries: dict['_RouteSite', list['_RouteEntry']]) -> None: + """ + Raises a ``SyntaxError`` if some PE would need more wavelet filters than a router has. + + Filters are counted per PE across all colors, which is why sites of *different* colors are + compared here -- unlike switch positions, which are a per-color resource. A PE needs one filter + per *color* it filters, however many sites of that color it belongs to: the sites of one bundle + that carry a filter are disjoint, and two bundles on one color are as well. + """ + colors_per_pe: dict[tuple[int, int], set[int]] = {} + for site, site_entries in entries.items(): + if all(entry.filter is None for entry in site_entries): + continue + for x in range(*site.x_range): + for y in range(*site.y_range): + colors_per_pe.setdefault((x, y), set()).add(site.color) + for (x, y), colors in colors_per_pe.items(): + if len(colors) > constants.FILTERS_PER_PE: + raise SyntaxError( + f'PE ({x}, {y}) would need {len(colors)} wavelet filters (colors ' + f'{sorted(colors)}), but a PE can use at most {constants.FILTERS_PER_PE} on ' + f'{constants.ARCH}.\n' + ' note: one of the four hardware filters is reserved by the memcpy module\n' + ' note: filtered delivery is what lets several streams share a color; using fewer ' + 'channels here needs more filters, and using more channels needs more colors') def _route_sites(rectangles: list[Rectangle[PEBlock]], color_maps: list[dict[str, int]]) -> dict['_RouteSite', list['_RouteEntry']]: """ Collects the route configurations of every site, sorted into switch-position order. + + Sorting is stable, which is what keeps configurations contributed by one statement -- a shift + bundle's inject-then-relay pair, whose order is geometric rather than a matter of statement + order -- in the sequence they were added in. """ + bundles = shift_bundles.bundles_by_group(shift_bundles.detect_shift_bundles(rectangles)) + owners = _bundle_owners(rectangles, bundles) entries: dict[_RouteSite, list[_RouteEntry]] = {} for rect_index, (rect, color_map) in enumerate(zip(rectangles, color_maps)): - for site, entry in _rectangle_route_entries(rect_index, rect, color_map): + for site, entry in _rectangle_route_entries(rect_index, rect, color_map, bundles, owners): entries.setdefault(site, []).append(entry) for site_entries in entries.values(): site_entries.sort(key=lambda entry: entry.order) @@ -547,8 +769,10 @@ def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: advances[site.describe()] = indices[position + 1] - indices[position] elif ring: advances[site.describe()] = total - indices[position] - # Otherwise this router is on its last configuration for this color and never routes - # anything again, so wavelets passing through may over-advance it harmlessly. + # Otherwise this router is on its last position for this color, and outside ring mode + # an advance past it is a no-op, so wavelets passing through over-advance it + # harmlessly. It is not necessarily finished with the color: a shift bundle's source + # keeps relaying its last configuration long after reaching it. distinct = set(advances.values()) if not distinct: @@ -626,10 +850,20 @@ def _check_site_overlap(entries: dict['_RouteSite', list['_RouteEntry']]) -> Non def _rectangle_route_entries(rect_index: int, rect: Rectangle[PEBlock], - color_map: dict[str, int]) -> list[tuple['_RouteSite', '_RouteEntry']]: + color_map: dict[str, int], + bundles: dict[str, list[shift_bundles.ShiftBundle]] | None = None, + bundle_owners: dict[str, int] | None = None + ) -> list[tuple['_RouteSite', '_RouteEntry']]: """ Collects every route configuration a single rectangle contributes, as ``(site, entry)`` pairs. + + :param bundles: The shift bundles of the kernel, keyed by stream group. A stream that is bundled + is routed by :func:`_bundle_route_entries` instead of hop by hop. + :param bundle_owners: Which rectangle contributes each bundle, so that a bundle declared by + several sending blocks is only contributed once. """ + bundles = bundles or {} + bundle_owners = bundle_owners or {} # Test whether a receive/send statement are called for creating inbound/outbound routes sends_recvs = analysis.sends_and_receives(rect.metadata.compute) use_order = _stream_use_order(rect.metadata.compute) @@ -665,6 +899,15 @@ def add(offset: tuple[int, int], color: int, rx: tuple[str, ...], tx: tuple[str, if isinstance(stream.stream, spir.ExternStreamDeclaration): continue # Extern streams do not have on-chip routing + if group in bundles: + # The whole bundle -- both halves and the relays between them -- is contributed at once, + # by one of its sending rectangles, so the receiving side adds nothing here. + if sent and bundle_owners.get(group) == rect_index: + for bundle in bundles[group]: + collected.extend(_bundle_route_entries(bundle, color_outbound, send_order, + rect_index, stream.stream_name)) + continue + if isinstance(stream.stream, spir.MulticastRangeStreamDeclaration): if sent and received: raise ValueError( diff --git a/spada/syntax/spatial_ir/canonicalization.py b/spada/syntax/spatial_ir/canonicalization.py index 4e95b4ad..a89b3443 100644 --- a/spada/syntax/spatial_ir/canonicalization.py +++ b/spada/syntax/spatial_ir/canonicalization.py @@ -426,10 +426,6 @@ def inline_phases(kernel: spir.Kernel) -> spir.Kernel: parameters=copy.deepcopy(kernel.parameters), arguments=copy.deepcopy(kernel.arguments), body=list(rect_place.values()) + list(rect_dataflow.values()) + list(rect_compute.values())) - if hasattr(kernel, "shift_schedules"): - new_kernel.shift_schedules = kernel.shift_schedules - if hasattr(kernel, "switch_advances"): - new_kernel.switch_advances = kernel.switch_advances return new_kernel diff --git a/spada/syntax/spatial_ir/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index cfd70e05..87943a5c 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -569,7 +569,7 @@ def as_ir(self, indent: int = 0) -> str: @dataclass class RoutingDeclaration(SpatialNode): """ - A routing declaration for a stream, optionally specifying hops, channel, and count. + A routing declaration for a stream, optionally specifying hops and channel. The ``channel`` field may hold: * ``"auto"`` – the channel number is assigned automatically. @@ -580,17 +580,11 @@ class RoutingDeclaration(SpatialNode): an integer by the time CSL lowering runs; use :attr:`resolved_channel` to obtain the concrete value. - The ``count`` field is the number of fabric words on this stream edge per PE - per phase. ``"auto"`` (the default, also used when the field is omitted) means - the stream is unbounded: the compiler does not infer a length and does not - apply counted router switching. Counted switching requires an explicit - compile-time integer ``count``. + How many words a stream carries is stated by its type (``stream``), not + here; see :class:`StreamType` and ``stream_lifetime``. """ hops: Union[list[RoutingHop], Literal["auto"]] = "auto" # list of hops or 'auto' channel: Union["Expression", int, Literal["auto"]] = "auto" - count: Union["Expression", int, Literal["auto"]] = "auto" - # Set by the shift-bundle pass; not part of the surface language. - counted_switch: bool = False def validate(self) -> None: if isinstance(self.hops, list): @@ -617,23 +611,6 @@ def resolved_channel(self) -> Union[int, Literal["auto"]]: ) return val - @property - def resolved_count(self) -> Union[int, Literal["auto"]]: - """ - Return the message count as a concrete integer, or ``"auto"`` if unbounded. - """ - if self.count == "auto": - return "auto" - if isinstance(self.count, int): - return self.count - val = self.count.eval() - if not isinstance(val, int): - raise ValueError( - f"Count expression '{self.count.as_ir()}' did not evaluate to an integer. " - "Ensure all parameters and loop variables are concretized before counted switching." - ) - return val - def as_ir(self, indent: int = 0) -> str: indent_str = ' ' * indent hops_str = "auto" if self.hops == "auto" else f"[{', '.join(hop.as_ir() for hop in self.hops)}]" @@ -647,12 +624,6 @@ def as_ir(self, indent: int = 0) -> str: f"{indent_str}hops = {hops_str}", f"{indent_str}channel = {channel_str}", ] - if self.count != "auto": - if isinstance(self.count, int): - count_str = str(self.count) - else: - count_str = self.count.as_ir() - lines.append(f"{indent_str}count = {count_str}") return ", \n".join(lines) diff --git a/spada/syntax/spatial_ir/language.lark b/spada/syntax/spatial_ir/language.lark index 38b49675..3eb7724b 100644 --- a/spada/syntax/spatial_ir/language.lark +++ b/spada/syntax/spatial_ir/language.lark @@ -118,8 +118,7 @@ hop : "(" posneg_integer_literal "," posneg_integer_literal ")" // 2D at the mo hops : "[" hop ("," hop)* "]" routing_hops : "hops" "=" (auto | hops) routing_channel : "channel" "=" (auto | value_expr) -routing_count : "count" "=" (auto | value_expr) -routing_field : routing_hops | routing_channel | routing_count +routing_field : routing_hops | routing_channel routing : routing_field ("," routing_field)* multicast_range : "[" range_expression "]" relative_stream_declaration : "relative_stream" "(" (value_expr | multicast_range) "," (value_expr | multicast_range) ")" ("{" routing "}")? diff --git a/spada/syntax/spatial_ir/lark_to_ir.py b/spada/syntax/spatial_ir/lark_to_ir.py index 3e4e4be7..c8cdefed 100644 --- a/spada/syntax/spatial_ir/lark_to_ir.py +++ b/spada/syntax/spatial_ir/lark_to_ir.py @@ -331,14 +331,11 @@ def routing_hops(self, args): def routing_channel(self, args): return ('channel', args[0]) - def routing_count(self, args): - return ('count', args[0]) - def routing_field(self, args): return args[0] def routing(self, args): - kwargs = {'hops': 'auto', 'channel': 'auto', 'count': 'auto'} + kwargs = {'hops': 'auto', 'channel': 'auto'} seen: set[str] = set() for key, value in args: if key in seen: @@ -347,11 +344,7 @@ def routing(self, args): kwargs[key] = value if 'hops' not in seen or 'channel' not in seen: raise ValueError('Routing declaration requires both hops and channel') - return irnodes.RoutingDeclaration( - hops=kwargs['hops'], - channel=kwargs['channel'], - count=kwargs['count'], - ) + return irnodes.RoutingDeclaration(hops=kwargs['hops'], channel=kwargs['channel']) def compute_body(self, args): if len(args) == 1 and isinstance(args[0], list): diff --git a/spada/syntax/spatial_ir/shift_bundles.py b/spada/syntax/spatial_ir/shift_bundles.py index 0c9c9dcc..fc6ad9ae 100644 --- a/spada/syntax/spatial_ir/shift_bundles.py +++ b/spada/syntax/spatial_ir/shift_bundles.py @@ -1,421 +1,213 @@ """ -Detect consecutive 1D interval shifts and schedule counted two-state color switching. - -A stream with an explicit ``count = k`` whose senders form a consecutive interval -``[L, L+m)`` at distance ``d`` (with ``1 < m <= d``) can share one color: each PE -forwards a known number of waves, then injects or absorbs. ``count = auto`` is -unbounded and is never rewritten. +Overlapping 1D interval shifts on one color. + +A run of consecutive PEs each shifting the same distance ``d`` along an axis has overlapping paths: +the router of source ``p + 1`` carries source ``p``'s words. One color per direction still suffices, +because the routers can be time-multiplexed -- but only if every switch is triggered by something +the PE that owns it knows locally, since a control wavelet advances *every* switch-configured router +it reaches (see ``irspec/docs/spatial/routing.md``). + +Send order is what provides that. The sources go nearest-the-destinations first, so a source's +router changes from injecting to relaying exactly when that source has finished its own send, which +it can signal itself with one switch-advance wavelet. Nothing needs to be told from a distance, and +the order enforces itself: a source further from the destinations cannot push a word through its +neighbour's router while that neighbour is still injecting, so it waits on the link. + +The destinations do not switch at all. Each transmits to its ramp *and* onward, so all of them see +the whole stream and a counter filter decides which words each one keeps; the last one transmits to +its ramp alone and thereby takes the stream out of the network. + +This is the arrangement Schnyder's 2D reduce-scatter uses ("Distributed Sorting on the Cerebras +Wafer-Scale Engine", fig. 7.6), and ``tests/csl_runtime/test_shift_bundle_filters.sh`` is a +hand-written version of it that pins down the hardware behaviour relied on here. """ from __future__ import annotations -from collections import defaultdict -from dataclasses import dataclass, replace -from typing import Literal +from dataclasses import dataclass +from typing import Literal, Optional -from spada.syntax.spatial_ir import analysis +from spada.syntax.spatial_ir import analysis, stream_lifetime from spada.syntax.spatial_ir import irnodes as spir +from spada.syntax.spatial_ir.canonicalization import PEBlock +from spada.syntax.spatial_ir.grid_geometry import Rectangle + +#: A shift of one PE needs no bundling: consecutive sources at distance 1 form a chain, whose +#: routers are sequenced by the ordinary receive-then-send switch positions. +MIN_BUNDLE_DISTANCE = 2 @dataclass(frozen=True) class ShiftBundle: - """A consecutive interval of sources shifted by an axis-aligned distance.""" + """ + One run of consecutive sources shifted onto an equally long run of destinations. - phase_index: int - axis: Literal["x", "y"] + The sources occupy ``[start, start + length)`` along ``axis``; source ``c`` sends to + ``c + sign * dist``. ``length <= dist`` keeps the two runs apart, so no PE is both a source and a + destination of the same bundle. + """ + axis: Literal['x', 'y'] sign: Literal[1, -1] start: int length: int dist: int - count: int - fixed: int - stream_names: tuple[spir.Identifier, ...] - # Physical channel, unique to this phase and direction. Detect leaves 0; - # apply_shift_bundles assigns a per-phase pair (fwd, bwd). - channel: int = 0 - - -@dataclass(frozen=True) -class ColorScheduleStep: - """One router config and how many fabric waves it stays active.""" - - rx: str - tx: str - waves: int - - def as_pair(self) -> str: - return f"{_short(self.rx)}->{_short(self.tx)}" - - -@dataclass(frozen=True) -class ColorSchedule: - """Per-PE counted switch program for one phase and logical channel.""" - - phase_index: int - x: int - y: int + #: Words each source sends, from the stream's bound. The filter windows are this wide. + words: int + #: The range of the *other* axis, as ``(start, stop, stride)``. Every PE in it runs an + #: independent copy of the bundle with the same router configurations. + cross: tuple[int, int, int] channel: int - steps: tuple[ColorScheduleStep, ...] - - -@dataclass(frozen=True) -class SwitchAdvance: + #: Routing identity of the stream, as :func:`stream_lifetime.stream_group_key` defines it. + group: str + + def sources(self) -> tuple[int, int]: + """The source run, as a half-open interval in ascending coordinates.""" + return self.start, self.start + self.length + + def destinations(self) -> tuple[int, int]: + """The destination run, as a half-open interval in ascending coordinates.""" + first = self.start + self.sign * self.dist + return first, first + self.length + + def relays(self) -> tuple[int, int]: + """ + The PEs between the two runs that only pass the stream through, as a half-open interval. + + Empty when ``length == dist``, which is the densest a bundle gets. + """ + if self.sign > 0: + return self.sources()[1], self.destinations()[0] + return self.destinations()[1], self.sources()[0] + + def destination_order(self) -> tuple[int, int]: + """ + Returns ``(first, step)``: the destination the stream reaches first, and the step from one + destination to the next along the direction of travel. + + The stream passes the destination run from the side it arrives on, and each destination sees + the whole stream, so this is what maps a destination onto the words it should keep. + """ + low, high = self.destinations() + return (low, 1) if self.sign > 0 else (high - 1, -1) + + def describe(self) -> str: + low, high = self.sources() + direction = {('x', 1): 'east', ('x', -1): 'west', + ('y', 1): 'south', ('y', -1): 'north'}[(self.axis, self.sign)] + return f'{self.axis} in [{low}:{high}] shifted {self.dist} {direction}' + + +def _straight_shift(declaration: spir.StreamDeclaration) -> Optional[tuple[Literal['x', 'y'], int]]: """ - SWITCH_ADV control wavelet sent after a non-last injector's data waves. + Returns the axis and signed distance of a stream that runs straight along one axis. - Each downstream router pops one opcode (always_pop): the next source and - the dest that just absorbed see SWITCH_ADV; hops in between see NOP. - A control wavelet holds at most 8 opcodes, so ``dist`` must be <= 8. + ``None`` for anything else: a stream that is not a relative one, that moves diagonally, or whose + hop list does not walk the axis one PE at a time. """ - - channel: int - axis: Literal["x", "y"] - last_injector: int - opcodes: tuple[str, ...] - - -def _short(port: str) -> str: - return {"RAMP": "R", "EAST": "E", "WEST": "W", "NORTH": "N", "SOUTH": "S"}.get(port, port) - - -def _resolved_count(routing: spir.RoutingDeclaration | None) -> int | None: - if routing is None: + stream = declaration.stream + if not isinstance(stream, spir.RelativeStreamDeclaration) or stream.routing is None: return None - count = routing.resolved_count - if count == "auto": + try: + dx, dy = int(stream.dx.eval()), int(stream.dy.eval()) + except Exception: # pragma: no cover - defensive: a non-constant offset return None - if count < 1: - raise ValueError(f"Routing count must be a positive integer, got {count}") - return count - - -def _axis_offset(stream: spir.RelativeStreamDeclaration) -> tuple[Literal["x", "y"], int] | None: - dx = stream.dx.eval() - dy = stream.dy.eval() - if not isinstance(dx, int) or not isinstance(dy, int): + if (dx == 0) == (dy == 0): return None - if dx != 0 and dy == 0: - return "x", dx - if dy != 0 and dx == 0: - return "y", dy - return None - + axis: Literal['x', 'y'] = 'x' if dy == 0 else 'y' + delta = dx if axis == 'x' else dy -def _hops_are_straight(stream: spir.RelativeStreamDeclaration, axis: str, delta: int) -> bool: - routing = stream.routing - if routing is None or routing.hops == "auto": - return True - if not isinstance(routing.hops, list): - return False - step = 1 if delta > 0 else -1 - expected = [(step, 0)] * abs(delta) if axis == "x" else [(0, step)] * abs(delta) - actual = [hop.offset for hop in routing.hops] - return actual == expected + hops = stream.routing.hops + if isinstance(hops, list): + step = (1 if delta > 0 else -1) + expected = [(step, 0)] * abs(delta) if axis == 'x' else [(0, step)] * abs(delta) + if [hop.offset for hop in hops] != expected: + return None + return axis, delta -def _block_points(block) -> list[tuple[int, int]]: - x0, x1, y0, y1 = block.get_grid_rect() - xs, ys = block.get_grid_stride() - return [(x, y) for x in range(x0, x1, xs) for y in range(y0, y1, ys)] +def _bound(declaration: spir.StreamDeclaration) -> Optional[int]: + if declaration.dtype.bound is None: + return None + try: + value = declaration.dtype.bound.eval() + except Exception: # pragma: no cover - defensive: a non-constant bound + return None + return value if isinstance(value, int) and value > 0 else None -def _consecutive_runs(values: list[int]) -> list[tuple[int, int]]: +def _consecutive_runs(values: set[int]) -> list[tuple[int, int]]: + """Splits a set of coordinates into ``(start, length)`` runs of consecutive values.""" if not values: return [] - ordered = sorted(set(values)) - runs = [] - start = prev = ordered[0] + runs: list[tuple[int, int]] = [] + ordered = sorted(values) + start = previous = ordered[0] for value in ordered[1:]: - if value == prev + 1: - prev = value + if value == previous + 1: + previous = value continue - runs.append((start, prev - start + 1)) - start = prev = value - runs.append((start, prev - start + 1)) + runs.append((start, previous - start + 1)) + start = previous = value + runs.append((start, previous - start + 1)) return runs -def detect_shift_bundles(kernel: spir.Kernel) -> list[ShiftBundle]: +def detect_shift_bundles(rectangles: list[Rectangle[PEBlock]]) -> list[ShiftBundle]: """ - Find consecutive interval shifts with an explicit count in each phase. + Finds the interval shifts in a kernel whose paths overlap, and which therefore need bundling. - :param kernel: A kernel whose metaprogramming and auto-hops are already resolved, - and whose phases have not yet been inlined. - """ - bundles: list[ShiftBundle] = [] - phase_index = 0 - for block in kernel.body: - if not isinstance(block, spir.Phase): - continue - bundles.extend(_detect_in_phase(block, phase_index)) - phase_index += 1 - return bundles + Sources are collected per channel and shift, across rectangles: one logical shift is often + declared by several compute blocks -- a sorting network's matchings at successive offsets, for + instance -- and only their union shows which PEs form a consecutive run. + A shift is bundled only if *every* run it decomposes into can be: at least two sources, and no + longer than the shift distance, so that sources and destinations stay disjoint. A shift that + fails this is left to the ordinary per-hop lowering, which reports the conflict if there is one. -def _detect_in_phase(phase: spir.Phase, phase_index: int) -> list[ShiftBundle]: - decls: dict[spir.Identifier, spir.RelativeStreamDeclaration] = {} - for dataflow in phase.dataflow: - for stmt in dataflow.statements: - stream = stmt.stream - if not isinstance(stream, spir.RelativeStreamDeclaration): - continue - existing = decls.get(stmt.stream_name) - if existing is None: - decls[stmt.stream_name] = stream - continue - if (existing.dx.eval(), existing.dy.eval()) != (stream.dx.eval(), stream.dy.eval()): - raise ValueError( - f'Stream "{stmt.stream_name.as_ir()}" is declared with conflicting offsets ' - f"in the same phase." - ) - - # (axis, delta, k, fixed_coord) -> (source coords along the axis, stream names) - groups: dict[tuple[str, int, int, int], tuple[set[int], set[spir.Identifier]]] = {} - for compute in phase.compute: - sent_recv = analysis.sends_and_receives(compute) - for name, stream in decls.items(): - sent, _received = sent_recv.get(name, (False, False)) + :param rectangles: The consolidated PE rectangles of the kernel, with channels already resolved. + :return: The bundles, in a deterministic order. + """ + # (channel, axis, signed distance, cross-axis range) -> (source coordinates, words, group) + groups: dict[tuple[int, str, int, tuple[int, int, int]], tuple[set[int], int, str]] = {} + for rect in rectangles: + sends_recvs = analysis.sends_and_receives(rect.metadata.compute) + for declaration in rect.metadata.dataflow.statements: + sent, _received = sends_recvs.get(declaration.stream_name, (False, False)) if not sent: continue - k = _resolved_count(stream.routing) - if k is None: - continue - axis_delta = _axis_offset(stream) - if axis_delta is None: + shift = _straight_shift(declaration) + words = _bound(declaration) + if shift is None or words is None: continue - axis, delta = axis_delta - if not _hops_are_straight(stream, axis, delta): + axis, delta = shift + if abs(delta) < MIN_BUNDLE_DISTANCE: continue - for x, y in _block_points(compute): - if axis == "x": - key = (axis, delta, k, y) - coord = x - else: - key = (axis, delta, k, x) - coord = y - bucket = groups.get(key) - if bucket is None: - bucket = (set(), set()) - groups[key] = bucket - bucket[0].add(coord) - bucket[1].add(name) - - bundles: list[ShiftBundle] = [] - for (axis, delta, k, fixed), (coords, names) in groups.items(): - dist = abs(delta) - sign: Literal[1, -1] = 1 if delta > 0 else -1 - for start, length in _consecutive_runs(list(coords)): - if length < 2 or length > dist: + channel = declaration.stream.routing.resolved_channel + if channel == 'auto': continue - bundles.append( - ShiftBundle( - phase_index=phase_index, - axis=axis, - sign=sign, - start=start, - length=length, - dist=dist, - count=k, - fixed=fixed, - stream_names=tuple(names), - ) - ) - return bundles + along = rect.x_range if axis == 'x' else rect.y_range + cross = rect.y_range if axis == 'x' else rect.x_range + key = (channel, axis, delta, cross) + coords, _, _ = groups.setdefault(key, (set(), words, stream_lifetime.stream_group_key(declaration))) + coords.update(range(along[0], along[1], along[2])) -def schedule_counted_switch(bundle: ShiftBundle) -> list[ColorSchedule]: - """ - Build the two-state per-PE schedule for one interval shift. - - East/south (+d): source half forwards then injects; dest half absorbs then forwards. - West/north (−d): the same lemma with the opposite ports, sources on the far side. - """ - if bundle.axis == "x": - forward_rx, forward_tx = ("WEST", "EAST") if bundle.sign > 0 else ("EAST", "WEST") - inject_tx = "EAST" if bundle.sign > 0 else "WEST" - absorb_rx = "WEST" if bundle.sign > 0 else "EAST" - - def pe(coord: int) -> tuple[int, int]: - return coord, bundle.fixed - else: - forward_rx, forward_tx = ("NORTH", "SOUTH") if bundle.sign > 0 else ("SOUTH", "NORTH") - inject_tx = "SOUTH" if bundle.sign > 0 else "NORTH" - absorb_rx = "NORTH" if bundle.sign > 0 else "SOUTH" - - def pe(coord: int) -> tuple[int, int]: - return bundle.fixed, coord - - L = bundle.start - m = bundle.length - d = bundle.dist - k = bundle.count - schedules: dict[tuple[int, int], list[ColorScheduleStep]] = {} - - def add(coord: int, steps: list[ColorScheduleStep]) -> None: - x, y = pe(coord) - kept = [step for step in steps if step.waves > 0] - if kept: - schedules[(x, y)] = kept - - if bundle.sign > 0: - for j in range(m): - src = L + j - add(src, [ - ColorScheduleStep(forward_rx, forward_tx, j * k), - ColorScheduleStep("RAMP", inject_tx, k), - ]) - for j in range(m): - dest = L + d + j - add(dest, [ - ColorScheduleStep(absorb_rx, "RAMP", k), - ColorScheduleStep(forward_rx, forward_tx, (m - 1 - j) * k), - ]) - for coord in range(L + m, L + d): - add(coord, [ColorScheduleStep(forward_rx, forward_tx, m * k)]) - else: - # Sources occupy [L, L+m); dests are at source - d. - for j in range(m): - src = L + j - add(src, [ - ColorScheduleStep(forward_rx, forward_tx, (m - 1 - j) * k), - ColorScheduleStep("RAMP", inject_tx, k), - ]) - for j in range(m): - dest = L - d + j - add(dest, [ - ColorScheduleStep(absorb_rx, "RAMP", k), - ColorScheduleStep(forward_rx, forward_tx, j * k), - ]) - for coord in range(L - d + m, L): - add(coord, [ColorScheduleStep(forward_rx, forward_tx, m * k)]) - - return [ - ColorSchedule(bundle.phase_index, x, y, bundle.channel, tuple(steps)) - for (x, y), steps in sorted(schedules.items()) - ] - - -_MAX_SWITCH_CMDS = 8 - - -def switch_advance_for_bundle(bundle: ShiftBundle) -> SwitchAdvance: - """Build the opcode chain popped by each downstream hop (distance ``d``). - - The injecting PE uses ``no_pop``, so the first opcode is for its neighbor. - """ - d = bundle.dist - if d > _MAX_SWITCH_CMDS: - raise ValueError( - f"Counted switching encodes one opcode per hop and a control wavelet " - f"holds at most {_MAX_SWITCH_CMDS} commands; got dist={d}." - ) - opcodes = ("SWITCH_ADV",) + ("NOP",) * max(d - 2, 0) + ("SWITCH_ADV",) - last_injector = bundle.start + bundle.length - 1 if bundle.sign > 0 else bundle.start - return SwitchAdvance(bundle.channel, bundle.axis, last_injector, opcodes) - - -def _phase_has_explicit_count_relative(phase: spir.Phase) -> bool: - for dataflow in phase.dataflow: - for stmt in dataflow.statements: - stream = stmt.stream - if not isinstance(stream, spir.RelativeStreamDeclaration): - continue - if _resolved_count(stream.routing) is not None: - return True - return False - - -def _phase_channel_bases(kernel: spir.Kernel, bundles: list[ShiftBundle]) -> dict[int, int]: - """ - Give each routed phase its own even/odd color pair. - - Wave quotas depend on the shift distance, so a counted two-state program - cannot be reused on the same color in a later phase. - """ - needed = {bundle.phase_index for bundle in bundles} - phase_index = 0 - for block in kernel.body: - if not isinstance(block, spir.Phase): + bundles: list[ShiftBundle] = [] + for (channel, axis, delta, cross), (coords, words, group) in sorted(groups.items()): + runs = _consecutive_runs(coords) + if not all(2 <= length <= abs(delta) for _start, length in runs): continue - if _phase_has_explicit_count_relative(block): - needed.add(phase_index) - phase_index += 1 - return {phase: 2 * i for i, phase in enumerate(sorted(needed))} + for start, length in runs: + bundles.append(ShiftBundle(axis=axis, sign=1 if delta > 0 else -1, start=start, + length=length, dist=abs(delta), words=words, cross=cross, + channel=channel, group=group)) + return bundles -def apply_shift_bundles(kernel: spir.Kernel, bundles: list[ShiftBundle]) -> list[ColorSchedule]: - """ - Mark bundled streams for counted switching, assign two colors per phase, - and return the per-PE schedules. - """ - bases = _phase_channel_bases(kernel, bundles) - assigned = [ - replace(bundle, channel=bases[bundle.phase_index] + (0 if bundle.sign > 0 else 1)) - for bundle in bundles - ] - - names_by_phase: dict[int, dict[spir.Identifier, int]] = defaultdict(dict) - for bundle in assigned: - for name in bundle.stream_names: - names_by_phase[bundle.phase_index][name] = bundle.channel - - phase_index = 0 - for block in kernel.body: - if not isinstance(block, spir.Phase): - continue - rewrite = names_by_phase.get(phase_index, {}) - if rewrite: - for dataflow in block.dataflow: - for stmt in dataflow.statements: - channel = rewrite.get(stmt.stream_name) - if channel is None or stmt.stream.routing is None: - continue - stmt.stream.routing.channel = channel - stmt.stream.routing.counted_switch = True - phase_index += 1 - - schedules: list[ColorSchedule] = [] - advances: list[SwitchAdvance] = [] - for bundle in assigned: - schedules.extend(schedule_counted_switch(bundle)) - advances.append(switch_advance_for_bundle(bundle)) - _assign_unit_hop_channels(kernel, bases) - kernel.switch_advances = advances - return schedules - - -def _assign_unit_hop_channels(kernel: spir.Kernel, bases: dict[int, int]) -> None: - """Assign this phase's fwd/bwd pair to explicit-count distance-1 streams.""" - phase_index = 0 - for block in kernel.body: - if not isinstance(block, spir.Phase): - continue - base = bases.get(phase_index) - phase_index += 1 - if base is None: - continue - for dataflow in block.dataflow: - for stmt in dataflow.statements: - stream = stmt.stream - if not isinstance(stream, spir.RelativeStreamDeclaration) or stream.routing is None: - continue - if stream.routing.counted_switch: - continue - if _resolved_count(stream.routing) is None: - continue - axis_delta = _axis_offset(stream) - if axis_delta is None: - continue - _axis, delta = axis_delta - if abs(delta) != 1: - continue - stream.routing.channel = base + (0 if delta > 0 else 1) - - -def coalesce_shift_bundles(kernel: spir.Kernel) -> list[ColorSchedule]: +def bundles_by_group(bundles: list[ShiftBundle]) -> dict[str, list[ShiftBundle]]: """ - Detect shift bundles, rewrite their channels, and attach schedules to ``kernel``. + Indexes bundles by the routing identity of their stream, for the route collector to look up. """ - bundles = detect_shift_bundles(kernel) - schedules = apply_shift_bundles(kernel, bundles) - kernel.shift_schedules = schedules - return schedules + grouped: dict[str, list[ShiftBundle]] = {} + for bundle in bundles: + grouped.setdefault(bundle.group, []).append(bundle) + return grouped diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py index 0f75f843..d8081e99 100644 --- a/spada/syntax/spatial_ir/stream_lifetime.py +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -126,14 +126,96 @@ def _declared_stream_names(kernel: spir.Kernel) -> set[spir.Identifier]: } +def _dataflow_declarations(kernel: spir.Kernel) -> dict[spir.Identifier, spir.StreamDeclaration]: + return { + statement.stream_name: statement + for node in kernel.walk() + if isinstance(node, spir.DataflowBlock) + for statement in node.statements + } + + +def _kernel_identifier_sizes(kernel: spir.Kernel) -> dict[spir.Identifier, list[int]]: + """ + Returns the shape of every field placed in the kernel, merged across its ``place`` blocks. + + A field that two place blocks give different shapes is dropped rather than guessed at, which + makes the counting that uses this give up instead of counting the wrong array. + """ + sizes: dict[spir.Identifier, list[int]] = {} + conflicting: set[spir.Identifier] = set() + for node in kernel.walk(): + if not isinstance(node, spir.PlaceBlock): + continue + for name, shape in _identifier_sizes(node).items(): + if name in sizes and sizes[name] != shape: + conflicting.add(name) + sizes[name] = shape + for name in conflicting: + del sizes[name] + return sizes + + +def _is_synchronous(statement: spir.Statement) -> bool: + """ + Returns whether a statement, and everything nested in it, completes before the next one starts. + + A statement that names a completion may still be in flight afterwards, so nothing may be + concluded from having executed it. + """ + return all(getattr(node, 'completion_name', None) is None for node in statement.walk()) + + +def _bound_exhausted_at(compute: spir.ComputeBlock, use: StreamUse, + declaration: Optional[spir.StreamDeclaration], + sizes: dict[spir.Identifier, list[int]]) -> Optional[int]: + """ + Returns the index of the statement that transfers the last element of a bounded stream, or + ``None`` if that statement cannot be identified. + + This is where a bounded stream closes itself, which is earlier than the end of the phase and + sometimes has to be: a PE that hands its router over to the next sender of a shift bundle at its + close cannot wait for the rest of the phase, since the rest of the phase may be waiting on the + traffic that the hand-over lets through. + + ``None`` is returned unless every use up to that statement is synchronous and no use follows it, + so that the close is only placed where the stream is demonstrably finished. + """ + if declaration is None or declaration.dtype.bound is None: + return None + try: + bound = declaration.dtype.bound.eval() + except Exception: # pragma: no cover - defensive: a non-constant bound + return None + if not isinstance(bound, int): + return None + + transferred = {'send': 0, 'receive': 0} + directions = [kind for kind in ('send', 'receive') if (use.sent if kind == 'send' else use.received)] + for index in use.uses: + statement = compute.statements[index] + if not _is_synchronous(statement): + return None + for kind in directions: + count = _transferred_elements(statement, use.name, kind, sizes) + if count is None: + return None + transferred[kind] += count + if all(transferred[kind] >= bound for kind in directions): + # Anything after this exceeds the bound; ``verify_stream_bounds`` reports it. + return index if index == use.uses[-1] else None + return None + + def insert_implicit_closes(kernel: spir.Kernel) -> spir.Kernel: """ Materializes the implicit close of every stream at the end of its scope. - For each compute block, an ``awaitall`` followed by ``await .close()`` is appended for - every stream the block uses and does not already close, in the phase in which that block last - uses it. The closes are emitted *after* the barrier because the phase's implicit awaits may be - waiting on operations that are still using those very streams. + A bounded stream closes itself where its bound is exhausted, so its close goes directly after + the statement that transfers its last element. Everything else is closed at the end of the phase + in which the block last uses it, as an ``awaitall`` followed by ``await .close()``. Those + closes are emitted *after* the barrier because the phase's implicit awaits may be waiting on + operations that are still using those very streams. Must run after ``canonicalize_phases`` (so that ``kernel.body`` contains only phases and place blocks) and ``uniquify_stream_names`` (so that a stream name means one stream), and before @@ -145,6 +227,8 @@ def insert_implicit_closes(kernel: spir.Kernel) -> spir.Kernel: declared = _declared_stream_names(kernel) if not declared: return kernel + declarations = _dataflow_declarations(kernel) + sizes = _kernel_identifier_sizes(kernel) phases = [block for block in kernel.body if isinstance(block, spir.Phase)] @@ -175,11 +259,28 @@ def insert_implicit_closes(kernel: spir.Kernel) -> spir.Kernel: if not to_close: continue - compute.statements.append(spir.AwaitAllStatement()) - for use in to_close: + def make_close(use: StreamUse) -> spir.CloseStatement: close = spir.CloseStatement(copy.deepcopy(use.expression)) close.lineinfo = getattr(use.expression, 'lineinfo', None) - compute.statements.append(close) + return close + + self_closing: dict[int, list[StreamUse]] = {} + at_end_of_phase: list[StreamUse] = [] + for use in to_close: + index = _bound_exhausted_at(compute, use, declarations.get(use.name), sizes) + if index is None: + at_end_of_phase.append(use) + else: + self_closing.setdefault(index, []).append(use) + + # Back to front, so that the indices of the insertions still to come stay valid. + for index in sorted(self_closing, reverse=True): + closes = [make_close(use) for use in self_closing[index]] + compute.statements[index + 1:index + 1] = closes + + if at_end_of_phase: + compute.statements.append(spir.AwaitAllStatement()) + compute.statements.extend(make_close(use) for use in at_end_of_phase) return kernel diff --git a/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh b/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh new file mode 100755 index 00000000..35954b62 --- /dev/null +++ b/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh @@ -0,0 +1,50 @@ +#!/bin/sh +# E2E test: the bundled 1D Batcher odd-even mergesort (2^L PEs, one f32 key per PE). +# Kernel: batcher_oddeven_bundled_1D.sptl params: L +# Same result as batcher_oddeven_1D, but each phase runs on two colors instead of one per +# comparator: OUT_out[:, 0, 0] == sort(inp[:, 0, 0]). +# L <= 3, which is where the three filters a PE can use run out. + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +. "$SCRIPT_DIR/_lib.sh" + +SORT_DIR="$(cd "$(dirname "$0")/../../samples/spatial/sort" && pwd)" +FOLDER="batcher_oddeven_bundled_1d_sptl" + +run_batcher() { + l=$1 + echo "--- batcher_oddeven_bundled_1d L=$l ---" + + sptlc "$SORT_DIR/batcher_oddeven_bundled_1D.sptl" "$FOLDER" -p L=$l + + python3 - < M is the case with pure relays between the two halves. set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" @@ -12,15 +13,14 @@ FOLDER="shift_bundle_1d_sptl" run_shift() { m=$1 - echo "--- shift_bundle_1d M=$m ---" + d=$2 + echo "--- shift_bundle_1d M=$m D=$d ---" - sptlc "$SAMPLE" "$FOLDER" -p M=$m + sptlc "$SAMPLE" "$FOLDER" -p M=$m -p D=$d python3 - <( - stream[N, 1] readonly inp, - stream[N, 1] writeonly out +kernel @shift( + stream[D + M, 1] readonly inp, + stream[D + M, 1] writeonly out ) { - place i16 i, i16 j in [0:N, 0] { - f32 val + place i16 i, i16 j in [0:D + M, 0] { + f32[K] val } phase { - compute i16 i, i16 j in [0:N, 0] { + compute i16 i, i16 j in [0:D + M, 0] { await receive(val, inp[i, j]) } } phase { - dataflow i16 i, i16 j in [0:4, 0] { - stream fwd = relative_stream(4, 0) { - hops = auto, - channel = auto, - count = 1 - } - } - dataflow i16 i, i16 j in [4:8, 0] { - stream fwd = relative_stream(4, 0) { + dataflow i16 i, i16 j in [0:D + M, 0] { + stream fwd = relative_stream(D, 0) { hops = auto, - channel = auto, - count = 1 + channel = 0 } } - compute i16 i, i16 j in [0:4, 0] { + compute i16 i, i16 j in [0:M, 0] { await send(val, fwd) } - compute i16 i, i16 j in [4:8, 0] { + compute i16 i, i16 j in [D:D + M, 0] { await receive(val, fwd) } } phase { - compute i16 i, i16 j in [0:N, 0] { + compute i16 i, i16 j in [0:D + M, 0] { await send(val, out[i, j]) } } } """ -_SHIFT_AUTO = _SHIFT.replace("count = 1", "count = auto") +_WESTBOUND = _SHIFT.replace('relative_stream(D, 0)', 'relative_stream(-D, 0)') \ + .replace('compute i16 i, i16 j in [0:M, 0] {\n await send(val, fwd)', + 'compute i16 i, i16 j in [D:D + M, 0] {\n await send(val, fwd)') \ + .replace('compute i16 i, i16 j in [D:D + M, 0] {\n await receive(val, fwd)', + 'compute i16 i, i16 j in [0:M, 0] {\n await receive(val, fwd)') + +_UNBOUNDED = _SHIFT.replace('stream fwd', 'stream fwd') + -_SHIFT_OMITTED = _SHIFT.replace(",\n count = 1", "") +def _rectangles(source: str, **params: int): + kernel = parser.parse_string(source) + kernel = passes.concretize_parameters(kernel, **params) + kernel = passes.constexpr_propagation(kernel) + kernel = canonicalize_kernel(kernel) + return canonicalization.consolidate_rectangles_to_equivalence_classes(kernel) + + +def _bundles(source: str, **params: int): + return detect_shift_bundles(_rectangles(source, **params)) -_SHIFT_TOO_LONG = _SHIFT.replace("[0:4, 0]", "[0:6, 0]").replace("relative_stream(4, 0)", "relative_stream(2, 0)") +def _layout(source: str, **params: int) -> str: + kernel = parser.parse_string(source) + kernel = passes.concretize_parameters(kernel, **params) + kernel = passes.constexpr_propagation(kernel) + files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) + return next(f.code for f in files if 'layout' in f.filename) -def test_parse_count_roundtrip(): - kernel = parser.parse_string(_SHIFT) - ir_1 = kernel.as_ir() - assert "count = 1" in ir_1 - ir_2 = parser.parse_string(ir_1).as_ir() - assert ir_1 == ir_2 +def _codes(source: str, **params: int) -> dict[str, str]: + kernel = parser.parse_string(source) + kernel = passes.concretize_parameters(kernel, **params) + kernel = passes.constexpr_propagation(kernel) + files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) + return {f.filename: f.code for f in files if f.filename.startswith('code_')} + + +def _configs(layout: str) -> dict[tuple[int, int], str]: + """ + Maps each ``@set_color_config`` loop in a layout to its configuration, keyed by the PE range. -def test_omitted_count_has_no_count_line(): - kernel = parser.parse_string(_SHIFT_OMITTED) - assert "count =" not in kernel.as_ir() + Only the standalone loops a shift bundle produces are keyed this way; they hold one call each. + """ + found = {} + for start, stop, body in re.findall( + r'for \(@range\(i16, (\d+), (\d+), 1\)\) \|pe_x\| \{\s*' + r'for \(@range\(i16, \d+, \d+, 1\)\) \|pe_y\| \{\s*' + r'(@set_color_config\([^\n]*\);)', layout): + found[(int(start), int(stop))] = body + return found -def test_detect_interval_shift(): - kernel = _prepare(_SHIFT, N=8) - bundles = detect_shift_bundles(kernel) +def test_detects_the_overlapping_shift(): + bundles = _bundles(_SHIFT, M=3, D=3, K=1) assert len(bundles) == 1 - b = bundles[0] - assert (b.start, b.length, b.dist, b.count, b.sign, b.axis) == (0, 4, 4, 1, 1, "x") + bundle = bundles[0] + assert (bundle.axis, bundle.sign, bundle.start, bundle.length, bundle.dist, bundle.words) \ + == ('x', 1, 0, 3, 3, 1) + assert (bundle.sources(), bundle.destinations(), bundle.relays()) == ((0, 3), (3, 6), (3, 3)) -def test_auto_count_is_not_rewritten(): - kernel = _prepare(_SHIFT_AUTO, N=8) - assert detect_shift_bundles(kernel) == [] - coalesce_shift_bundles(kernel) - assert kernel.shift_schedules == [] +def test_a_gap_between_the_halves_becomes_relays(): + bundle = _bundles(_SHIFT, M=3, D=5, K=1)[0] + assert (bundle.sources(), bundle.relays(), bundle.destinations()) == ((0, 3), (3, 5), (5, 8)) -def test_omitted_count_is_not_rewritten(): - kernel = _prepare(_SHIFT_OMITTED, N=8) - assert detect_shift_bundles(kernel) == [] +def test_westbound_shift_is_detected_mirrored(): + bundle = _bundles(_WESTBOUND, M=3, D=5, K=1)[0] + assert (bundle.sign, bundle.sources(), bundle.relays(), bundle.destinations()) \ + == (-1, (5, 8), (3, 5), (0, 3)) + # The stream arrives on the destinations' east side, so it is the westmost one it reaches last. + assert bundle.destination_order() == (2, -1) -def test_m_greater_than_d_is_rejected(): - kernel = _prepare(_SHIFT_TOO_LONG, N=8) - assert detect_shift_bundles(kernel) == [] +def test_an_unbounded_stream_is_not_bundled(): + # Without a bound there is no self-close, so nothing would advance the sources' switches. + assert _bundles(_UNBOUNDED, M=3, D=3, K=1) == [] -def test_schedule_source_and_dest_halves(): - kernel = _prepare(_SHIFT, N=8) - bundle = detect_shift_bundles(kernel)[0] - by_pe = {(s.x, s.y): s.steps for s in schedule_counted_switch(bundle)} - assert [(st.rx, st.tx, st.waves) for st in by_pe[(0, 0)]] == [("RAMP", "EAST", 1)] - assert [(st.rx, st.tx, st.waves) for st in by_pe[(1, 0)]] == [("WEST", "EAST", 1), ("RAMP", "EAST", 1)] - assert [(st.rx, st.tx, st.waves) for st in by_pe[(3, 0)]] == [("WEST", "EAST", 3), ("RAMP", "EAST", 1)] - assert [(st.rx, st.tx, st.waves) for st in by_pe[(4, 0)]] == [("WEST", "RAMP", 1), ("WEST", "EAST", 3)] - assert [(st.rx, st.tx, st.waves) for st in by_pe[(7, 0)]] == [("WEST", "RAMP", 1)] - adv = switch_advance_for_bundle(bundle) - assert adv.last_injector == 3 - assert adv.opcodes == ("SWITCH_ADV", "NOP", "NOP", "SWITCH_ADV") +def test_a_single_source_is_not_bundled(): + assert _bundles(_SHIFT, M=1, D=3, K=1) == [] -def test_westbound_schedule(): - west = """ -kernel @shift_w( - stream[N, 1] readonly inp, - stream[N, 1] writeonly out -) { - place i16 i, i16 j in [0:N, 0] { f32 val } - phase { - compute i16 i, i16 j in [0:N, 0] { await receive(val, inp[i, j]) } - } - phase { - dataflow i16 i, i16 j in [4:8, 0] { - stream bwd = relative_stream(-4, 0) { - hops = auto, - channel = auto, - count = 1 - } - } - compute i16 i, i16 j in [4:8, 0] { await send(val, bwd) } - compute i16 i, i16 j in [0:4, 0] { await receive(val, bwd) } - } - phase { - compute i16 i, i16 j in [0:N, 0] { await send(val, out[i, j]) } - } -} -""" - kernel = _prepare(west, N=8) - bundle = detect_shift_bundles(kernel)[0] - assert bundle.sign == -1 and bundle.start == 4 and bundle.length == 4 - by_pe = {(s.x, s.y): s.steps for s in schedule_counted_switch(bundle)} - assert [(st.rx, st.tx, st.waves) for st in by_pe[(7, 0)]] == [("RAMP", "WEST", 1)] - assert [(st.rx, st.tx, st.waves) for st in by_pe[(4, 0)]] == [("EAST", "WEST", 3), ("RAMP", "WEST", 1)] - assert [(st.rx, st.tx, st.waves) for st in by_pe[(0, 0)]] == [("EAST", "RAMP", 1)] - assert [(st.rx, st.tx, st.waves) for st in by_pe[(3, 0)]] == [("EAST", "RAMP", 1), ("EAST", "WEST", 3)] +def test_a_shift_of_one_is_not_bundled(): + # Consecutive sources at distance one form a chain, which the ordinary receive-then-send switch + # positions already sequence. + assert _bundles(_SHIFT, M=1, D=1, K=1) == [] -def _batcher_prepared(n_log: int): - path = os.path.join( - os.path.dirname(__file__), "..", "..", "samples", "spatial", "sort", "batcher_oddeven_1D.sptl" - ) +def test_sources_longer_than_the_shift_are_not_bundled(): + # Sources would be destinations of the same bundle, which this arrangement cannot express. + assert _bundles(_SHIFT, M=4, D=2, K=1) == [] + + +def test_sources_inject_then_relay(): + configs = _configs(_layout(_SHIFT, M=3, D=3, K=1)) + sources = configs[(0, 3)] + assert '.routes = .{ .rx = .{RAMP}, .tx = .{EAST} }' in sources + assert '.switches = .{ .pos1 = .{ .rx = WEST } }' in sources + assert 'ring_mode' not in sources + # One call covers the whole run, the westmost source included: a switch position it never uses + # is cheaper than a second configuration. + assert '.filter' not in sources + + +def test_destinations_are_static_and_filtered(): + configs = _configs(_layout(_SHIFT, M=3, D=3, K=1)) + passing = configs[(3, 5)] + assert '.routes = .{ .rx = .{WEST}, .tx = .{RAMP, EAST} }' in passing + assert '.switches' not in passing + # Destination 3 keeps the last of the three words, destination 4 the second: the counter has to + # start one and two words short of the end of the cycle respectively. + assert '.init_counter = pe_x - 2' in passing + assert '.limit1 = 2, .max_counter = 0' in passing + + +def test_the_last_destination_terminates_the_stream(): + configs = _configs(_layout(_SHIFT, M=3, D=3, K=1)) + terminal = configs[(5, 6)] + assert '.routes = .{ .rx = .{WEST}, .tx = .{RAMP} }' in terminal + assert '.init_counter = 0' in terminal + + +def test_relays_pass_the_stream_through_unchanged(): + configs = _configs(_layout(_SHIFT, M=3, D=5, K=1)) + relays = configs[(3, 5)] + assert '.routes = .{ .rx = .{WEST}, .tx = .{EAST} }' in relays + assert '.switches' not in relays and '.filter' not in relays + + +def test_westbound_layout_mirrors_the_eastbound_one(): + configs = _configs(_layout(_WESTBOUND, M=3, D=3, K=1)) + assert '.routes = .{ .rx = .{RAMP}, .tx = .{WEST} }' in configs[(3, 6)] + assert '.switches = .{ .pos1 = .{ .rx = EAST } }' in configs[(3, 6)] + assert '.routes = .{ .rx = .{EAST}, .tx = .{RAMP, WEST} }' in configs[(1, 3)] + assert '.init_counter = 3 - pe_x' in configs[(1, 3)] + assert '.routes = .{ .rx = .{EAST}, .tx = .{RAMP} }' in configs[(0, 1)] + assert '.init_counter = 0' in configs[(0, 1)] + + +def test_windows_are_as_wide_as_the_stream_bound(): + configs = _configs(_layout(_SHIFT, M=3, D=3, K=2)) + # Six words in the cycle, two of which each destination keeps. + assert '.limit1 = 5, .max_counter = 1' in configs[(3, 5)] + assert '.init_counter = (pe_x - 2) * 2' in configs[(3, 5)] + assert '.limit1 = 5, .max_counter = 1' in configs[(5, 6)] + + +def test_each_source_advances_its_own_switch_once(): + codes = _codes(_SHIFT, M=3, D=3, K=1) + sending = codes['code_0_0.csl'] + assert sending.count('ctrl.opcode.SWITCH_ADV') == 1 + assert '.control = true' in sending + # The destinations do not switch, so nothing is emitted there. + assert 'SWITCH_ADV' not in codes['code_3_0.csl'] + + +def test_one_color_carries_the_whole_bundle(): + layout = _layout(_SHIFT, M=3, D=3, K=1) + routes = layout[layout.index('// Routes'):] + assert {int(color) for color in re.findall(r'@get_color\((\d+)\)', routes)} == {0} + + +def test_no_port_is_a_union_on_the_receiving_side(): + # A switch position accepts a single rx direction, and only a destination unions its tx. + layout = _layout(_SHIFT, M=3, D=3, K=1) + assert re.search(r'\.rx = \.\{[A-Z]+, [A-Z]+\}', layout) is None + assert re.search(r'\.pos\d = \.\{ \.rx = [A-Z]+, ', layout) is None + + +def test_filter_renders_as_a_color_config_field(): + plan = cslrouting.ColorSwitchPlan( + [cslrouting.RouteConfig(('WEST', ), ('RAMP', 'EAST'))], + filter=cslrouting.FilterConfig('pe_x - 2', '2', '0')) + assert plan.as_csl() == ('.{ .routes = .{ .rx = .{WEST}, .tx = .{RAMP, EAST} }, ' + '.filter = .{ .kind = .{ .counter = true }, .count_data = true, ' + '.init_counter = pe_x - 2, .limit1 = 2, .max_counter = 0 } }') + + +def test_a_pe_cannot_use_more_filters_than_the_hardware_has(): + from spada.syntax.csl import constants + + def site(color: int): + return cslrouting._RouteSite(color=color, x_range=(0, 4, 1), y_range=(0, 1, 1), absolute=True) + + def entry(): + return cslrouting._RouteEntry(cslrouting.RouteConfig(('WEST', ), ('RAMP', )), (0, 0, 0), 0, + (0, 0), None, '', cslrouting.FilterConfig('0', '1', '0')) + + entries = {site(color): [entry()] for color in range(constants.FILTERS_PER_PE + 1)} + with pytest.raises(SyntaxError, match='wavelet filters'): + cslrouting._check_filter_budget(entries) + + +def _bundled_batcher(l: int): + path = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', + 'batcher_oddeven_bundled_1D.sptl') kernel = parser.parse_file(path) - kernel = passes.concretize_parameters(kernel, L=n_log) + kernel = passes.concretize_parameters(kernel, L=l) kernel = passes.constexpr_propagation(kernel) - kernel = canonicalization.inline_metaprogramming(kernel) - kernel = canonicalization.canonicalize_phases(kernel) - kernel = canonicalization.reduce_streams(kernel) - kernel = canonical_subgrids.canonicalize_subgrids(kernel) - kernel = canonicalization.resolve_auto_hops(kernel) - return kernel - - -def test_batcher_n16_p1_d8_bundle(): - kernel = _batcher_prepared(4) - bundles = detect_shift_bundles(kernel) - east = [b for b in bundles if b.sign > 0 and b.dist == 8] - west = [b for b in bundles if b.sign < 0 and b.dist == 8] - assert len(east) == 1 and east[0].start == 0 and east[0].length == 8 - assert len(west) == 1 and west[0].start == 8 and west[0].length == 8 - - -def test_batcher_n16_p2_d4_bundle(): - kernel = _batcher_prepared(4) - bundles = detect_shift_bundles(kernel) - east = [b for b in bundles if b.sign > 0 and b.dist == 4] - assert any(b.start == 4 and b.length == 4 for b in east) - - -def test_batcher_d1_has_no_counted_switch(): - kernel = _batcher_prepared(3) - bundles = detect_shift_bundles(kernel) - assert all(b.dist != 1 for b in bundles) - - -def _xy(bundle, coord): - return (coord, bundle.fixed) if bundle.axis == "x" else (bundle.fixed, coord) - - -def _consume(states, xy, role): - steps = states[xy] - assert steps, f"PE {xy} has no remaining config but needs {role}" - rx, tx, waves = steps[0] - if role == "inject": - assert rx == "RAMP" and tx != "RAMP", (xy, steps[0], role) - elif role == "absorb": - assert tx == "RAMP" and rx != "RAMP", (xy, steps[0], role) - elif role == "forward": - assert rx != "RAMP" and tx != "RAMP", (xy, steps[0], role) - else: - raise ValueError(role) - waves -= 1 - if waves == 0: - steps.pop(0) - else: - steps[0] = (rx, tx, waves) - - -def simulate_bundle_delivery(bundle, drop_second_step=False): - """ - West-first (eastbound) / east-first (westbound) serial delivery. + return lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) - Each fabric word consumes one wave at every PE on its path. After the - source-half forward quota the PE must have switched to inject; after the - dest-half absorb quota it must have switched to forward. Exhausting every - quota means the two-state program matches the matching. - """ - states = {} - for sched in schedule_counted_switch(bundle): - steps = [(st.rx, st.tx, st.waves) for st in sched.steps] - if drop_second_step: - steps = steps[:1] - states[(sched.x, sched.y)] = steps - order = range(bundle.length) if bundle.sign > 0 else range(bundle.length - 1, -1, -1) - for j in order: - src = bundle.start + j - dest = src + bundle.sign * bundle.dist - for _ in range(bundle.count): - _consume(states, _xy(bundle, src), "inject") - hop = src + bundle.sign - while hop != dest: - _consume(states, _xy(bundle, hop), "forward") - hop += bundle.sign - _consume(states, _xy(bundle, dest), "absorb") - leftover = {pe: steps for pe, steps in states.items() if steps} - assert leftover == {}, leftover - - -def test_two_state_switch_delivers_eastbound(): - kernel = _prepare(_SHIFT, N=8) - simulate_bundle_delivery(detect_shift_bundles(kernel)[0]) - - -def test_two_state_switch_delivers_westbound(): - west = """ -kernel @shift_w( - stream[N, 1] readonly inp, - stream[N, 1] writeonly out -) { - place i16 i, i16 j in [0:N, 0] { f32 val } - phase { - compute i16 i, i16 j in [0:N, 0] { await receive(val, inp[i, j]) } - } - phase { - dataflow i16 i, i16 j in [4:8, 0] { - stream bwd = relative_stream(-4, 0) { - hops = auto, - channel = auto, - count = 1 - } - } - compute i16 i, i16 j in [4:8, 0] { await send(val, bwd) } - compute i16 i, i16 j in [0:4, 0] { await receive(val, bwd) } - } - phase { - compute i16 i, i16 j in [0:N, 0] { await send(val, out[i, j]) } - } -} -""" - kernel = _prepare(west, N=8) - simulate_bundle_delivery(detect_shift_bundles(kernel)[0]) - - -def test_two_state_switch_delivers_count_2(): - kernel = _prepare(_SHIFT.replace("count = 1", "count = 2"), N=8) - bundle = detect_shift_bundles(kernel)[0] - assert bundle.count == 2 - simulate_bundle_delivery(bundle) - - -def test_without_the_switch_delivery_fails(): - kernel = _prepare(_SHIFT, N=8) - with pytest.raises(AssertionError): - simulate_bundle_delivery(detect_shift_bundles(kernel)[0], drop_second_step=True) - - -def test_batcher_d8_switch_delivers(): - kernel = _batcher_prepared(4) - bundles = detect_shift_bundles(kernel) - east = next(b for b in bundles if b.sign > 0 and b.dist == 8) - west = next(b for b in bundles if b.sign < 0 and b.dist == 8) - simulate_bundle_delivery(east) - simulate_bundle_delivery(west) - - -def test_two_state_lemma_never_needs_a_third_config(): - kernel = _prepare(_SHIFT, N=8) - bundle = detect_shift_bundles(kernel)[0] - for sched in schedule_counted_switch(bundle): - assert 1 <= len(sched.steps) <= 2 - if len(sched.steps) == 2: - first, second = sched.steps - forward_then_inject = first.rx != "RAMP" and first.tx != "RAMP" and second.rx == "RAMP" - absorb_then_forward = first.tx == "RAMP" and second.rx != "RAMP" and second.tx != "RAMP" - assert forward_then_inject or absorb_then_forward - - -def _onchip_channels_by_phase(kernel): - phases = [] - for block in kernel.body: - if not isinstance(block, spir.Phase): - continue - chans = set() - for dataflow in block.dataflow: - for stmt in dataflow.statements: - routing = getattr(stmt.stream, "routing", None) - if routing is None or routing.resolved_channel == "auto": - continue - chans.add(routing.resolved_channel) - phases.append(chans) - return phases - - -def test_batcher_two_colors_per_phase_not_globally(): - kernel = _batcher_prepared(3) - coalesce_shift_bundles(kernel) - routed = [chans for chans in _onchip_channels_by_phase(kernel) if chans] - assert len(routed) == 6 # L=3 has 6 (l,p) CAS phases - for chans in routed: - assert len(chans) == 2 - used = [c for chans in routed for c in chans] - assert len(used) == len(set(used)) - assert set(used) == set(range(12)) + +@pytest.mark.parametrize('l, phases', [(2, 3), (3, 6)]) +def test_a_batcher_phase_costs_two_colors(l: int, phases: int): + # The matchings of a phase are declared as one stream over the whole line and share its channel, + # so the count is per phase and per direction rather than per comparator. + layout = next(f.code for f in _bundled_batcher(l) if 'layout' in f.filename) + routes = layout[layout.index('// Routes'):] + assert len({int(color) for color in re.findall(r'@get_color\((\d+)\)', routes)}) == 2 * phases + + +def test_the_batcher_runs_out_of_filters_before_it_runs_out_of_colors(): + # A PE receives a bundle in every phase whose distance is at least two, and cannot filter more + # than three colors; L = 4 has six such phases. + with pytest.raises(SyntaxError, match='wavelet filters'): + _bundled_batcher(4) def test_batcher_scalar_receive_lowers_to_data_task(): path = os.path.join( - os.path.dirname(__file__), "..", "..", "samples", "spatial", "sort", "batcher_oddeven_1D.sptl" + os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_1D.sptl' ) kernel = parser.parse_file(path) kernel = passes.concretize_parameters(kernel, L=1) kernel = passes.constexpr_propagation(kernel) files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) - pe_codes = [f.code for f in files if "code_" in f.filename] + pe_codes = [f.code for f in files if 'code_' in f.filename] assert pe_codes for code in pe_codes: - assert "tmp = bwd" not in code - assert ".async = true" not in code - pe0 = next(f.code for f in files if "code_0_0" in f.filename) - assert "task dtask_" in pe0 - assert "tmp = __x" in pe0 - - -def test_lowering_encodes_intra_phase_switch(): - kernel = parser.parse_string(_SHIFT) - kernel = passes.concretize_parameters(kernel, N=8) - kernel = passes.constexpr_propagation(kernel) - files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) - layout = next(f.code for f in files if "layout" in f.filename) - assert "spa_phase_reload" not in layout - # PE 1: forward 1 wave W→E, then switch to inject R→E. - assert re.search( - r"@set_color_config\(1, 0, @get_color\(0\), " - r"\.\{ \.routes = \.\{ \.rx = \.\{WEST\}, \.tx = \.\{EAST\} \}", - layout, - ) - assert re.search( - r"spa_switch_after phase=\d+ pe=1,0 ch=0 waves=1 rx=RAMP tx=EAST", - layout, - ) - assert ".pos1 = .{ .rx = RAMP }" in layout - assert ".pop_mode = .{ .pop_on_advance_nop = true }" in layout - assert ".pop_mode = .{ .no_pop = true }" in layout - # PE 4: absorb 1 wave W→R, then switch to forward W→E. - assert re.search( - r"@set_color_config\(4, 0, @get_color\(0\), " - r"\.\{ \.routes = \.\{ \.rx = \.\{WEST\}, \.tx = \.\{RAMP\} \}", - layout, - ) - assert re.search( - r"spa_switch_after phase=\d+ pe=4,0 ch=0 waves=1 rx=WEST tx=EAST", - layout, - ) - assert ".pos1 = .{ .tx = EAST }" in layout - pe0 = next(f.code for f in files if "code_0_0" in f.filename) - assert "ctrl.opcode.SWITCH_ADV" in pe0 - assert "encode_payload" in pe0 - assert "get_fabric_coord" in pe0 - - -def test_batcher_lowering_two_colors_per_phase(): - path = os.path.join( - os.path.dirname(__file__), "..", "..", "samples", "spatial", "sort", "batcher_oddeven_1D.sptl" - ) - for n_log, n_cas, old_static in ((3, 6, 14), (4, 10, 30)): - kernel = parser.parse_file(path) - kernel = passes.concretize_parameters(kernel, L=n_log) - kernel = passes.constexpr_propagation(kernel) - files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) - layout = next(f.code for f in files if "layout" in f.filename) - assert "spa_color_schedule" in layout - assert "spa_switch_after" in layout - assert "spa_phase_reload" not in layout - colors = {int(c) for c in re.findall(r"@get_color\((\d+)\)", layout)} - assert colors == set(range(2 * n_cas)), ( - f"L={n_log} used colors {sorted(colors)}, expected 2 per phase " - f"({2 * n_cas}), not a kernel-wide pair and not {old_static} static" - ) - # Every two-step schedule has a matching counted switch onto the second pair. - schedules = re.findall( - r"spa_color_schedule phase=(\d+) pe=(\d+),(\d+) ch=(\d+) : ([^\n]+)", - layout, - ) - switches = { - (int(ph), int(x), int(y), int(ch)): (int(w), rx, tx) - for ph, x, y, ch, w, rx, tx in re.findall( - r"spa_switch_after phase=(\d+) pe=(\d+),(\d+) ch=(\d+) " - r"waves=(\d+) rx=(\w+) tx=(\w+)", - layout, - ) - } - short = {"R": "RAMP", "E": "EAST", "W": "WEST", "N": "NORTH", "S": "SOUTH"} - for ph, x, y, ch, step_txt in schedules: - parts = [p.strip() for p in step_txt.split(";")] - key = (int(ph), int(x), int(y), int(ch)) - if len(parts) < 2: - assert key not in switches - continue - first, second = parts[0], parts[1] - first_waves = int(re.search(r"waves=(\d+)", first).group(1)) - pair = re.search(r"(\w+)->(\w+)", second) - rx = short.get(pair.group(1), pair.group(1)) - tx = short.get(pair.group(2), pair.group(2)) - assert switches[key] == (first_waves, rx, tx) - - -_SAMPLE_SHIFT = os.path.join( - os.path.dirname(__file__), "..", "..", "samples", "spatial", "simple", "shift_bundle_1D.sptl" -) - -_DUAL_PORT = re.compile( - r"\.(?:rx|tx) = \.\{[A-Z]+, [A-Z]+\}" -) -_ABS_COLOR_CONFIG = re.compile( - r"@set_color_config\((\d+), (\d+), @get_color\((\d+)\)," -) - - -def _sample_shift(m: int = 4): - kernel = parser.parse_file(_SAMPLE_SHIFT) - kernel = passes.concretize_parameters(kernel, M=m) - kernel = passes.constexpr_propagation(kernel) - return kernel - - -def test_sample_shift_bundle_is_one_eastbound_color(): - kernel = _prepare(open(_SAMPLE_SHIFT).read(), M=4) - bundles = detect_shift_bundles(kernel) - assert len(bundles) == 1 - b = bundles[0] - assert (b.start, b.length, b.dist, b.sign, b.axis, b.count) == (0, 4, 4, 1, "x", 1) - coalesce_shift_bundles(kernel) - routed = [chans for chans in _onchip_channels_by_phase(kernel) if chans] - assert routed == [{0}] - - -def test_sample_shift_bundle_one_rx_tx_pair_at_a_time(): - """A color installs one (rx, tx) pair; pos1 changes rx or tx, never both.""" - kernel = _prepare(open(_SAMPLE_SHIFT).read(), M=4) - bundle = detect_shift_bundles(kernel)[0] - for sched in schedule_counted_switch(bundle): - assert 1 <= len(sched.steps) <= 2 - for step in sched.steps: - assert step.rx in {"RAMP", "WEST", "EAST"} - assert step.tx in {"RAMP", "WEST", "EAST"} - assert step.rx != step.tx - if len(sched.steps) == 2: - first, second = sched.steps - changed_rx = first.rx != second.rx - changed_tx = first.tx != second.tx - assert changed_rx ^ changed_tx, ( - f"PE ({sched.x},{sched.y}) changes both rx and tx: " - f"{first.as_pair()} then {second.as_pair()}" - ) - - -def test_sample_shift_bundle_layout_does_not_union_ports(): - kernel = _sample_shift(4) - files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) - layout = next(f.code for f in files if "layout" in f.filename) - assert _DUAL_PORT.search(layout) is None - keys = _ABS_COLOR_CONFIG.findall(layout) - assert keys - assert len(keys) == len(set(keys)), f"duplicate (PE, color) configs: {keys}" - assert "spa_switch_after" in layout - assert ".pos1 = .{ .rx = RAMP }" in layout - assert ".pos1 = .{ .tx = EAST }" in layout - simulate_bundle_delivery(detect_shift_bundles(_prepare(open(_SAMPLE_SHIFT).read(), M=4))[0]) + # A scalar receive becomes a data task, not an undeclared assignment, and a scalar source + # cannot be moved asynchronously. + assert 'tmp = bwd' not in code + assert '.async = true' not in code + pe0 = next(f.code for f in files if 'code_0_0' in f.filename) + assert 'task dtask_' in pe0 + assert 'tmp = __x' in pe0 From 590b44a9c91427d0bb190ea95788de173fe469d9 Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 17 Aug 2026 23:26:17 +0200 Subject: [PATCH 18/68] fix data task recycling --- spada/lowering/spatial_ir_to_csl.py | 13 ++++++- .../spatial_ir/test_task_recycling_codegen.py | 37 +++++++++++++++++++ 2 files changed, 48 insertions(+), 2 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 83edfc3a..75c7d12f 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -514,7 +514,8 @@ def generate_rectangle(kernel: spir.Kernel, for i, task in enumerate(tasks): if task.task_type != 'data': continue - _generate_data_task(rect.metadata, i, task, current_code, header, footer, dsds, dtypes, color_map, tasks) + _generate_data_task(rect.metadata, i, task, current_code, header, footer, dsds, dtypes, color_map, tasks, + task_bindings) footer.write(f' @bind_data_task(dtask_{i}, dtask_{i}_id);\n') if task.blocked: footer.write(f' @block(dtask_{i}_id);\n') @@ -1360,6 +1361,7 @@ def _generate_data_task( dtypes: dict[spir.Identifier, spir.IRType], color_map: dict[str, int], tasks: list[tdag.CSLTask], + task_bindings: task_recycling.TaskBindingPlan, ): """ Generates a data task from a foreach loop. @@ -1373,6 +1375,8 @@ def _generate_data_task( :param dtypes: A dictionary mapping identifiers to their defined types. :param color_map: Dictionary mapping each stream to its respective color id ({name}_color also works). :param tasks: A list of all tasks in the kernel. + :param task_bindings: The local-task binding plan, needed to install the state of a recycled + successor slot before handing control to it. """ # * If index is requested: before unblocking task, set k; inc at end of task # * Wavelet-triggered task as fallback @@ -1396,7 +1400,12 @@ def _generate_data_task( next_task_code = f'@{itedge_code}(exit_task_id);' else: prefix = "d" if next_task_type == 'data' else "" - next_task_code = f'@{itedge_code}({prefix}task_{next_task}_id);' + lines = [] + if next_task_type == 'local': + lines.extend(task_bindings.emit_local_transition_preamble( + next_task, tasks[next_task].blocked, indent='').splitlines()) + lines.append(f'@{itedge_code}({prefix}task_{next_task}_id);') + next_task_code = '\n '.join(lines) var_dtype_csl = dtype_as_csl(stmt.variables[0].dtype) param_range = stmt.parameter_range[0] diff --git a/tests/spatial_ir/test_task_recycling_codegen.py b/tests/spatial_ir/test_task_recycling_codegen.py index dc023cf0..1bb94da3 100644 --- a/tests/spatial_ir/test_task_recycling_codegen.py +++ b/tests/spatial_ir/test_task_recycling_codegen.py @@ -55,6 +55,43 @@ def test_csl_runtime_task_recycling_sample_lowers(filename: str): assert '__task_slot_' in combined, 'expected task-ID recycling in generated CSL' +def test_data_tasks_install_the_state_of_a_recycled_successor(): + """A data task handing control to a recycled slot must install that slot's state first. + + Without the assignment the dispatcher runs whichever branch was installed last -- in the + bundled Batcher at L=3 that meant a PE silently skipped its comparator and the fabric + deadlocked behind the send it never made. + """ + sample = os.path.join( + os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_bundled_1D.sptl') + kernel = parser.parse_file(sample) + kernel = passes.concretize_parameters(kernel, L=3) + kernel = passes.constexpr_propagation(kernel) + + csl_files = lower_spatial_ir_to_csl(kernel) + + checked = 0 + for file in csl_files: + code = file.code + hardware_ids: dict[str, list[str]] = {} + for task_index, hardware_id in re.findall(r'const task_(\d+)_id = @get_local_task_id\((\d+)\)', code): + hardware_ids.setdefault(hardware_id, []).append(task_index) + recycled = {task for tasks in hardware_ids.values() if len(tasks) > 1 for task in tasks} + + for body in re.findall(r'task dtask_\d+\([^)]*\) void \{(.*?)\n\}', code, re.S): + for match in re.finditer(r'@(?:activate|unblock)\(task_(\d+)_id\);', body): + if match.group(1) not in recycled: + continue + written = [line.strip() for line in body[:match.start()].splitlines() if line.strip()] + preceding = written[-1] if written else '' + assert re.fullmatch(r'__task_slot_\d+_state = \d+;', preceding), ( + f'{file.filename}: @activate(task_{match.group(1)}_id) is not preceded by its ' + f'slot state assignment, but by "{preceding}"') + checked += 1 + + assert checked, 'sample no longer exercises a data task triggering a recycled local task' + + def test_codegen_avoids_local_task_id_color_overlap(): path = os.path.join(_CSL_RUNTIME_TASK_RECYCLING_SAMPLES, 'task_color_overlap_many_channels.sptl') kernel = parser.parse_file(path) From 67de3151786c5c1355c141212f48f33d92906ebe Mon Sep 17 00:00:00 2001 From: glukas Date: Tue, 18 Aug 2026 08:21:04 +0200 Subject: [PATCH 19/68] implement 16 bit memcpy in runtime --- spada/runtime/runtime.py | 71 ++++++++++++++++++++++++++++++++++++---- 1 file changed, 64 insertions(+), 7 deletions(-) diff --git a/spada/runtime/runtime.py b/spada/runtime/runtime.py index 2fc2cdf6..ce9615c4 100644 --- a/spada/runtime/runtime.py +++ b/spada/runtime/runtime.py @@ -94,6 +94,63 @@ def from_json(cls, json_data: Union[str, Dict[str, Any]]) -> "ProgramMetadata": ######################################################## +def memcpy_data_type(dtype: np.dtype) -> "crt.MemcpyDataType": + """ + Pick the transfer width for a kernel argument of ``dtype``. + + :param dtype: The dtype the kernel declared for the argument. + :return: The ``MemcpyDataType`` to pass alongside the buffer. + """ + if dtype.itemsize == 4: + return crt.MemcpyDataType.MEMCPY_32BIT + if dtype.itemsize == 2: + return crt.MemcpyDataType.MEMCPY_16BIT + raise ValueError(f"Cannot transfer {dtype} arrays: the SDK moves either 16 or 32 bits per " + f"element, so a kernel argument must be 2 or 4 bytes wide.") + + +def memcpy_word_dtype(dtype: np.dtype) -> np.dtype: + """ + Give the host-buffer dtype for a kernel argument of ``dtype``, one element per 32-bit word. + + ``memcpy_h2d`` and ``memcpy_d2h`` reject a buffer whose elements are not 32 bits ("Internal + data type of any memcpy_d2h() or memcpy_h2d() operation should be 32 bit") even when the + device-side array is 16-bit: ``MEMCPY_16BIT`` means only the low half of each word travels. + + :param dtype: The dtype the kernel declared for the argument. + :return: ``dtype`` itself when it is already 32 bits wide, else a 32-bit word dtype. + """ + return dtype if dtype.itemsize == 4 else np.dtype(np.uint32) + + +def as_memcpy_words(data: np.ndarray) -> np.ndarray: + """ + Widen a 16-bit array into the 32-bit words ``memcpy_h2d`` expects. + + The widening is bit-for-bit rather than by value, so that a negative ``i16`` and an ``f16`` + both arrive on the device unchanged. + + :param data: A contiguous array in the dtype the kernel declared. + :return: ``data`` itself when it is already 32 bits wide, else a widened copy. + """ + if data.dtype.itemsize == 4: + return data + return data.view(np.uint16).astype(np.uint32) + + +def from_memcpy_words(words: np.ndarray, dtype: np.dtype) -> np.ndarray: + """ + Undo :func:`as_memcpy_words` for data copied back from the device. + + :param words: The buffer ``memcpy_d2h`` filled. + :param dtype: The dtype the kernel declared for the output. + :return: ``words`` reinterpreted in ``dtype``, keeping the shape. + """ + if dtype.itemsize == 4: + return words + return words.astype(np.uint16).view(dtype) + + def flatten_copy( name: str, data: np.ndarray, shape: List[int], runtime: crt.SdkRuntime, metadata: ProgramMetadata, benchmark: bool ): @@ -117,14 +174,14 @@ def flatten_copy( runtime.memcpy_h2d( buffer_id, - src.ravel(), + as_memcpy_words(src).ravel(), metadata.inputs[name].rect_offset_used[0], # PE offset in x direction metadata.inputs[name].rect_offset_used[1], # PE offset in y direction shape[0], # Width (number of PEs in x) shape[1], # Height (number of PEs in y) shape[2], streaming=not metadata.memcpy_mode, # Use streaming if not in memcpy mode - data_type=crt.MemcpyDataType.MEMCPY_32BIT if data.dtype == np.float32 else crt.MemcpyDataType.MEMCPY_16BIT, + data_type=memcpy_data_type(data.dtype), order=crt.MemcpyOrder.ROW_MAJOR if not metadata.inputs[name].column_major else crt.MemcpyOrder.COL_MAJOR, nonblock=not benchmark, # Non-blocking copy if not benchmarking ) @@ -145,11 +202,11 @@ def copy_unflatten(name: str, data: np.ndarray, shape: List[int], runtime: crt.S if buffer_id is None: raise ValueError(f"Buffer ID for '{name}' not found in program.") - # The SDK returns A[h][w][elem_per_pe]; allocate a buffer in that layout. - sdk_buf = np.empty((shape[1], shape[0], shape[2]), dtype=data.dtype) + # The SDK returns A[h][w][elem_per_pe]; allocate a buffer in that layout, one element per word. + words = np.empty((shape[1], shape[0], shape[2]), dtype=memcpy_word_dtype(data.dtype)) runtime.memcpy_d2h( - sdk_buf.ravel(), + words.ravel(), buffer_id, metadata.outputs[name].rect_offset_used[0], # PE offset in x direction metadata.outputs[name].rect_offset_used[1], # PE offset in y direction @@ -157,13 +214,13 @@ def copy_unflatten(name: str, data: np.ndarray, shape: List[int], runtime: crt.S shape[1], # Height (number of PEs in y) shape[2], streaming=not metadata.memcpy_mode, # Use streaming if not in memcpy mode - data_type=crt.MemcpyDataType.MEMCPY_32BIT if data.dtype == np.float32 else crt.MemcpyDataType.MEMCPY_16BIT, + data_type=memcpy_data_type(data.dtype), order=crt.MemcpyOrder.ROW_MAJOR if not metadata.outputs[name].column_major else crt.MemcpyOrder.COL_MAJOR, nonblock=False, # Blocking copy to ensure data is ready after copy ) # Transpose back from (h, w, elem) to (w, h, elem) to match our convention. - np.copyto(data, sdk_buf.transpose(1, 0, 2)) + np.copyto(data, from_memcpy_words(words, data.dtype).transpose(1, 0, 2)) def convert_timestamp(hw_timestamp: npt.NDArray[np.uint32]) -> npt.NDArray[np.uint64]: From 1c810c9dd4b76478d41b063cff9f37d1cb1b11b0 Mon Sep 17 00:00:00 2001 From: glukas Date: Tue, 18 Aug 2026 09:51:03 +0200 Subject: [PATCH 20/68] update plotting script --- samples/spatial/sort/plot_batcher_routing.py | 365 +++++++++++++++---- 1 file changed, 296 insertions(+), 69 deletions(-) diff --git a/samples/spatial/sort/plot_batcher_routing.py b/samples/spatial/sort/plot_batcher_routing.py index f56beab1..d26a0449 100644 --- a/samples/spatial/sort/plot_batcher_routing.py +++ b/samples/spatial/sort/plot_batcher_routing.py @@ -1,34 +1,44 @@ #!/usr/bin/env python3 """ -Visualize the static channel assignment used by batcher_oddeven_1D.sptl. +Visualize the channel assignment of the 1D Batcher samples. -At distance d, offset r in [0, d) uses: - fwd (east, +d) : channel 2*((d - 1) + r) - bwd (west, -d) : channel 2*((d - 1) + r) + 1 - -Stages that share the same d reuse those colors. Total colors = 2*(n - 1). +Versions +-------- + static batcher_oddeven_1D.sptl. At distance d, offset r uses + fwd (east, +d) : channel 2*((d - 1) + r) + bwd (west, -d) : channel 2*((d - 1) + r) + 1 + Stages that share d reuse those colors, so a PE that plays a + different role at the same d occupies several switch positions. + bundled batcher_oddeven_bundled_1D.sptl. Each phase uses two colors, one + per direction. Phases with d >= 2 whose comparators form a run of + at least two sources are a shift bundle: sources inject then + relay (pos0 / pos1), destinations stay put and pick their word + with a counter filter. Colors are not reused across phases. Views ----- network Knuth-style sorting network. Concurrent matchings (same phase) occupy adjacent sub-columns so overlapping spans stay visible. Each comparator is two arrows: down = fwd (east, +d), up = bwd - (west, -d), each colored by its own channel. - table Per-PE color table of the static @set_color_config. Each cell is - the union of rx→tx pairs on that (PE, channel); a reused color - may install both RAMP→EAST and WEST→RAMP on the same PE. + (west, -d), each colored by its channel. + table Per-PE @set_color_config, with switch positions resolved the way + WSE-2 stores them (a both-sides change is split through a relay + intermediate). Stacked bands are pos0, pos1, … in that order. + A destination filter is the small ``fN`` in the cell, N being + the filter's init_counter. Examples -------- python samples/spatial/sort/plot_batcher_routing.py --n 8 - python samples/spatial/sort/plot_batcher_routing.py --n 8 --view table --show + python samples/spatial/sort/plot_batcher_routing.py --n 8 --version bundled --view table + python samples/spatial/sort/plot_batcher_routing.py --n 8 --version static bundled --view table """ from __future__ import annotations import argparse import math -from dataclasses import dataclass +from dataclasses import dataclass, replace import matplotlib @@ -39,6 +49,14 @@ from matplotlib.patches import Rectangle +VERSIONS = ("static", "bundled") +MIN_BUNDLE_DISTANCE = 2 +MIN_BUNDLE_LENGTH = 2 + +TX_ORDER = ("RAMP", "EAST", "WEST") +DIR_LETTER = {"RAMP": "R", "EAST": "E", "WEST": "W"} + + @dataclass(frozen=True) class Matching: l: int @@ -59,12 +77,32 @@ class Phase: matchings: tuple[Matching, ...] -def fwd_channel(dist: int, offset: int) -> int: +@dataclass(frozen=True) +class Route: + """One router configuration: receive from ``rx``, transmit to ``tx``.""" + + rx: str + tx: tuple[str, ...] + + def label(self) -> str: + tx = "".join(DIR_LETTER[d] for d in self.tx) + return f"{DIR_LETTER[self.rx]}→{tx}" + + +@dataclass(frozen=True) +class Cell: + """Resolved hardware switch positions of one (PE, channel), plus its filter.""" + + positions: tuple[Route, ...] + filter_init: int | None = None + + +def fwd_channel_static(dist: int, offset: int) -> int: return 2 * ((dist - 1) + offset) -def bwd_channel(dist: int, offset: int) -> int: - return fwd_channel(dist, offset) + 1 +def bwd_channel_static(dist: int, offset: int) -> int: + return fwd_channel_static(dist, offset) + 1 def batcher_phases(n: int) -> list[Phase]: @@ -79,7 +117,7 @@ def batcher_phases(n: int) -> list[Phase]: for r in range(dist): pairs = tuple((i, i + dist) for i in range(r, n, 1 << l)) matchings.append( - Matching(l, 1, dist, r, pairs, fwd_channel(dist, r), bwd_channel(dist, r)) + Matching(l, 1, dist, r, pairs, fwd_channel_static(dist, r), bwd_channel_static(dist, r)) ) phases.append(Phase(index, l, 1, dist, tuple(matchings))) index += 1 @@ -94,13 +132,30 @@ def batcher_phases(n: int) -> list[Phase]: for lo in range(start, stop, 2 * dist): pairs.append((lo, lo + dist)) matchings.append( - Matching(l, p, dist, r, tuple(pairs), fwd_channel(dist, r), bwd_channel(dist, r)) + Matching(l, p, dist, r, tuple(pairs), fwd_channel_static(dist, r), bwd_channel_static(dist, r)) ) phases.append(Phase(index, l, p, dist, tuple(matchings))) index += 1 return phases +def assign_channels(phases: list[Phase], version: str) -> list[Phase]: + """Rewrite matching channels to match the sample of ``version``.""" + if version == "static": + return phases + assigned = [] + for ph in phases: + fwd, bwd = 2 * ph.index, 2 * ph.index + 1 + matchings = tuple(replace(m, fwd=fwd, bwd=bwd) for m in ph.matchings) + assigned.append(replace(ph, matchings=matchings)) + return assigned + + +def channel_count(phases: list[Phase]) -> int: + used = [ch for ph in phases for m in ph.matchings for ch in (m.fwd, m.bwd)] + return (max(used) + 1) if used else 0 + + def channel_color(channel: int, n_channels: int, cmap_name: str = "tab20"): cmap = plt.get_cmap(cmap_name) if n_channels <= 20: @@ -118,9 +173,9 @@ def _dir_arrow(ax, x: float, y_from: float, y_to: float, color) -> None: ) -def _draw_network(ax, phases: list[Phase], n: int) -> None: +def _draw_network(ax, phases: list[Phase], n: int, version: str) -> None: """Draw concurrent matchings as sub-columns; fwd and bwd as offset arrows.""" - n_channels = 2 * (n - 1) + n_channels = channel_count(phases) slot = 0.38 gap = 0.55 dx = 0.06 @@ -137,7 +192,7 @@ def _draw_network(ax, phases: list[Phase], n: int) -> None: ax.set_yticklabels([str(i) for i in range(n)]) ax.set_ylabel("PE index") ax.set_xlabel("Phase (left arrow ↓ fwd, right arrow ↑ bwd)") - ax.set_title(f"Batcher network, n={n}: downward = fwd (east), upward = bwd (west)") + ax.set_title(f"Batcher network ({version}), n={n}: downward = fwd (east), upward = bwd (west)") tick_pos = [] tick_lab = [] @@ -172,16 +227,17 @@ def _draw_network(ax, phases: list[Phase], n: int) -> None: handles = [] for ch in range(n_channels): - arrow = "↓" if ch % 2 == 0 else "↑" - kind = "fwd" if ch % 2 == 0 else "bwd" + if version == "static": + arrow = "↓" if ch % 2 == 0 else "↑" + kind = "fwd" if ch % 2 == 0 else "bwd" + label = f"{arrow} ch {ch} ({kind})" + else: + phase = ch // 2 + arrow = "↓" if ch % 2 == 0 else "↑" + kind = "fwd" if ch % 2 == 0 else "bwd" + label = f"{arrow} ch {ch} (p{phase} {kind})" handles.append( - Line2D( - [0], - [0], - color=channel_color(ch, n_channels), - lw=2.0, - label=f"{arrow} ch {ch} ({kind})", - ) + Line2D([0], [0], color=channel_color(ch, n_channels), lw=2.0, label=label) ) ax.legend( handles=handles, @@ -193,46 +249,168 @@ def _draw_network(ax, phases: list[Phase], n: int) -> None: ) -PAIR_ORDER = ("R→E", "W→E", "W→R", "R→W", "E→W", "E→R") -PAIR_COLOR = { +ROUTE_COLOR = { "R→E": "#1f77b4", "W→E": "#ff7f0e", "W→R": "#2ca02c", + "W→RE": "#9467bd", "R→W": "#5fa8d3", "E→W": "#f4a261", "E→R": "#6dce6d", + "E→RW": "#c77dff", } +ROUTE_ORDER = ("R→E", "W→E", "W→RE", "W→R", "R→W", "E→W", "E→RW", "E→R") + + +def _tx(*dirs: str) -> tuple[str, ...]: + wanted = set(dirs) + return tuple(d for d in TX_ORDER if d in wanted) + + +def _add(seq: list[Route], config: Route) -> None: + if not seq or seq[-1] != config: + seq.append(config) + + +def _resolve_hardware(configs: list[Route]) -> tuple[Route, ...]: + """WSE-2: a position names one side; a both-sides change is split in two.""" + if not configs: + return () + positions = [configs[0]] + for config in configs[1:]: + previous = positions[-1] + if previous.rx != config.rx and previous.tx != config.tx: + positions.append(Route(previous.rx, config.tx)) + positions.append(config) + return tuple(positions) + + +def _shift_runs(pairs: tuple[tuple[int, int], ...], dist: int) -> list[tuple[int, int]]: + """Group ``(lo, lo+dist)`` pairs into maximal consecutive source runs ``(start, length)``.""" + sources = sorted(lo for lo, hi in pairs if hi - lo == dist) + runs: list[tuple[int, int]] = [] + i = 0 + while i < len(sources): + start = sources[i] + length = 1 + while i + length < len(sources) and sources[i + length] == start + length: + length += 1 + runs.append((start, length)) + i += length + return runs + + +def _install_hop(configs: list[list[list[Route]]], pe: int, ch: int, route: Route) -> None: + _add(configs[pe][ch], route) + + +def _install_ordinary_pair(configs: list[list[list[Route]]], lo: int, hi: int, fwd: int, bwd: int) -> None: + _install_hop(configs, lo, fwd, Route("RAMP", _tx("EAST"))) + _install_hop(configs, hi, fwd, Route("WEST", _tx("RAMP"))) + for mid in range(lo + 1, hi): + _install_hop(configs, mid, fwd, Route("WEST", _tx("EAST"))) + _install_hop(configs, hi, bwd, Route("RAMP", _tx("WEST"))) + _install_hop(configs, lo, bwd, Route("EAST", _tx("RAMP"))) + for mid in range(lo + 1, hi): + _install_hop(configs, mid, bwd, Route("EAST", _tx("WEST"))) + + +def _window_init(pe: int, first: int, step: int) -> int: + """``init_counter`` of a one-word destination filter, as ``_window_start`` emits it.""" + offset = 1 - first if step > 0 else first + 1 + return pe + offset if step > 0 else offset - pe + + +def _install_bundle( + configs: list[list[list[Route]]], + filters: list[list[int | None]], + start: int, + length: int, + dist: int, + fwd: int, + bwd: int, +) -> None: + """Eastbound inject-then-relay plus westbound mirror, with destination filters.""" + src_lo, src_hi = start, start + length + dst_lo, dst_hi = start + dist, start + dist + length + + for pe in range(src_lo, src_hi): + _install_hop(configs, pe, fwd, Route("RAMP", _tx("EAST"))) + _install_hop(configs, pe, fwd, Route("WEST", _tx("EAST"))) + for pe in range(src_hi, dst_lo): + _install_hop(configs, pe, fwd, Route("WEST", _tx("EAST"))) + # Stream travels east: last destination terminates, the others copy-and-forward. + for pe in range(dst_lo, dst_hi - 1): + _install_hop(configs, pe, fwd, Route("WEST", _tx("RAMP", "EAST"))) + filters[pe][fwd] = _window_init(pe, dst_lo, 1) + _install_hop(configs, dst_hi - 1, fwd, Route("WEST", _tx("RAMP"))) + filters[dst_hi - 1][fwd] = 0 + + for pe in range(dst_lo, dst_hi): + _install_hop(configs, pe, bwd, Route("RAMP", _tx("WEST"))) + _install_hop(configs, pe, bwd, Route("EAST", _tx("WEST"))) + for pe in range(src_hi, dst_lo): + _install_hop(configs, pe, bwd, Route("EAST", _tx("WEST"))) + # Stream travels west: lowest destination terminates. + for pe in range(src_lo + 1, src_hi): + _install_hop(configs, pe, bwd, Route("EAST", _tx("RAMP", "WEST"))) + filters[pe][bwd] = _window_init(pe, src_hi - 1, -1) + _install_hop(configs, src_lo, bwd, Route("EAST", _tx("RAMP"))) + filters[src_lo][bwd] = 0 + + +def pe_table(phases: list[Phase], n: int, version: str) -> list[list[Cell]]: + """Build the resolved per-PE switch table of ``version``.""" + n_channels = channel_count(phases) + configs: list[list[list[Route]]] = [[[] for _ in range(n_channels)] for _ in range(n)] + filters: list[list[int | None]] = [[None] * n_channels for _ in range(n)] - -def _draw_table(ax, phases: list[Phase], n: int) -> None: - """Static per-PE @set_color_config: each cell is the union of rx→tx pairs.""" - n_channels = 2 * (n - 1) - routes: list[list[set[str]]] = [[set() for _ in range(n_channels)] for _ in range(n)] for ph in phases: - for m in ph.matchings: - for lo, hi in m.pairs: - routes[lo][m.fwd].add("R→E") - routes[hi][m.fwd].add("W→R") - for mid in range(lo + 1, hi): - routes[mid][m.fwd].add("W→E") - routes[hi][m.bwd].add("R→W") - routes[lo][m.bwd].add("E→R") - for mid in range(lo + 1, hi): - routes[mid][m.bwd].add("E→W") + pairs = tuple(pair for m in ph.matchings for pair in m.pairs) + if not pairs: + continue + if version == "bundled": + fwd, bwd = ph.matchings[0].fwd, ph.matchings[0].bwd + if ph.dist >= MIN_BUNDLE_DISTANCE: + for start, length in _shift_runs(pairs, ph.dist): + if length >= MIN_BUNDLE_LENGTH: + _install_bundle(configs, filters, start, length, ph.dist, fwd, bwd) + else: + _install_ordinary_pair(configs, start, start + ph.dist, fwd, bwd) + else: + for lo, hi in pairs: + _install_ordinary_pair(configs, lo, hi, fwd, bwd) + else: + for m in ph.matchings: + for lo, hi in m.pairs: + _install_ordinary_pair(configs, lo, hi, m.fwd, m.bwd) + + return [ + [Cell(_resolve_hardware(configs[pe][ch]), filters[pe][ch]) for ch in range(n_channels)] + for pe in range(n) + ] + +def _draw_table(ax, phases: list[Phase], n: int, version: str) -> None: + """Resolved per-PE switch positions; destination filters as a small ``fN``.""" + table = pe_table(phases, n, version) + n_channels = channel_count(phases) ax.set_xlim(-0.5, n_channels - 0.5) ax.set_ylim(n - 0.5, -0.5) ax.set_xticks(range(n_channels)) ax.set_yticks(range(n)) ax.set_xlabel("Channel") ax.set_ylabel("PE") - ax.set_title(f"Static rx→tx table, n={n} ({n_channels} colors); stacked pairs are a union") + ax.set_title( + f"Resolved switch positions ({version}, WSE-2), n={n} ({n_channels} colors); " + "stacked bands are pos0, pos1, ...; fN is the filter init_counter" + ) ax.set_aspect("equal") for pe in range(n): for ch in range(n_channels): - pairs = [p for p in PAIR_ORDER if p in routes[pe][ch]] - if not pairs: + cell = table[pe][ch] + if not cell.positions: ax.add_patch( Rectangle( (ch - 0.45, pe - 0.45), @@ -244,69 +422,118 @@ def _draw_table(ax, phases: list[Phase], n: int) -> None: ) ) continue - band = 0.9 / len(pairs) - fontsize = 6 if len(pairs) == 1 else 5 - for i, pair in enumerate(pairs): + n_pos = len(cell.positions) + # Leave a thin strip at the bottom when a filter is present. + usable = 0.78 if cell.filter_init is not None else 0.9 + band = usable / n_pos + fontsize = 6 if n_pos == 1 else 5 + for i, route in enumerate(cell.positions): y0 = pe - 0.45 + i * band + label = route.label() ax.add_patch( Rectangle( (ch - 0.45, y0), 0.9, band, - facecolor=PAIR_COLOR[pair], + facecolor=ROUTE_COLOR.get(label, "#bbbbbb"), edgecolor="0.85", lw=0.4, ) ) - ax.text(ch, y0 + band / 2, pair, ha="center", va="center", fontsize=fontsize, color="0.1") + text = f"{i}:{label}" if n_pos > 1 else label + ax.text(ch, y0 + band / 2, text, ha="center", va="center", fontsize=fontsize, color="0.1") + if cell.filter_init is not None: + ax.text( + ch, + pe + 0.38, + f"f{cell.filter_init}", + ha="center", + va="center", + fontsize=5, + color="0.15", + fontweight="bold", + ) handles = [ - Line2D([0], [0], marker="s", color="w", markerfacecolor=PAIR_COLOR[p], markersize=10, label=p) - for p in PAIR_ORDER + Line2D([0], [0], marker="s", color="w", markerfacecolor=ROUTE_COLOR[p], markersize=10, label=p) + for p in ROUTE_ORDER + if any(route.label() == p for row in table for cell in row for route in cell.positions) ] + if any(cell.filter_init is not None for row in table for cell in row): + handles.append( + Line2D([0], [0], marker="$f$", color="0.15", markerfacecolor="w", markersize=10, label="fN filter init") + ) ax.legend(handles=handles, title="rx→tx", loc="upper left", bbox_to_anchor=(1.02, 1), fontsize=8) -def plot(n: int, view: str, outfile: str | None, show: bool) -> None: - phases = batcher_phases(n) +def plot(n: int, view: str, version: str, outfile: str | None, show: bool) -> None: + phases = assign_channels(batcher_phases(n), version) if view == "network": n_slots = sum(max(len(ph.matchings), 1) for ph in phases) fig, ax = plt.subplots(figsize=(max(8, 0.7 * n_slots + 0.8 * len(phases)), max(4, 0.45 * n))) - _draw_network(ax, phases, n) + _draw_network(ax, phases, n, version) elif view == "table": - fig, ax = plt.subplots(figsize=(max(8, 0.45 * 2 * (n - 1)), max(4, 0.45 * n))) - _draw_table(ax, phases, n) + n_channels = channel_count(phases) + fig, ax = plt.subplots(figsize=(max(8, 0.45 * n_channels), max(4, 0.5 * n))) + _draw_table(ax, phases, n, version) else: raise ValueError(f"unknown view {view}") fig.tight_layout() if outfile: fig.savefig(outfile, bbox_inches="tight") - png = outfile[:-4] + ".png" if outfile.endswith(".pdf") else outfile + ".png" if outfile.endswith(".pdf"): - fig.savefig(png, dpi=160, bbox_inches="tight") + fig.savefig(outfile[:-4] + ".png", dpi=160, bbox_inches="tight") print(f"wrote {outfile}") if show: plt.show() plt.close(fig) +def _expand_versions(requested: list[str]) -> list[str]: + if "all" in requested: + return list(VERSIONS) + # Preserve order, drop duplicates. + seen: set[str] = set() + versions = [] + for version in requested: + if version not in seen: + seen.add(version) + versions.append(version) + return versions + + def main() -> None: parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) parser.add_argument("--n", type=int, default=8, help="number of PEs, power of two (default 8)") + parser.add_argument( + "--version", + nargs="+", + choices=VERSIONS + ("all",), + default=["static"], + help="static (one color per matching) and/or bundled (two colors per phase). " + "'all' is both. Repeatable.", + ) parser.add_argument( "--view", choices=("network", "table"), default="network", - help="network: sorting-network comparators; table: static color config", + help="network: sorting-network comparators; table: resolved switch positions and filters", ) - parser.add_argument("--out", default=None, help="output PDF path") + parser.add_argument("--out", default=None, help="output PDF path (version is inserted before the extension if several)") parser.add_argument("--show", action="store_true", help="open an interactive window") args = parser.parse_args() - outfile = args.out - if outfile is None and not args.show: - outfile = f"samples/spatial/sort/batcher_routing_n{args.n}_{args.view}.pdf" - plot(args.n, args.view, outfile, args.show) + versions = _expand_versions(args.version) + for version in versions: + outfile = args.out + if outfile is None and not args.show: + outfile = f"samples/spatial/sort/batcher_routing_{version}_n{args.n}_{args.view}.pdf" + elif outfile is not None and len(versions) > 1: + if outfile.endswith(".pdf"): + outfile = f"{outfile[:-4]}_{version}.pdf" + else: + outfile = f"{outfile}_{version}" + plot(args.n, args.view, version, outfile, args.show) if __name__ == "__main__": From 7fbaa1e085e8a6dea57c4fbd86e750af1500ea70 Mon Sep 17 00:00:00 2001 From: glukas Date: Tue, 18 Aug 2026 11:46:39 +0200 Subject: [PATCH 21/68] implement data task recycling logic --- spada/lowering/spatial_ir_to_csl.py | 251 ++++++++++++------ spada/syntax/csl/task_recycling.py | 176 +++++++++++- .../spatial_ir/test_task_recycling_codegen.py | 67 +++++ 3 files changed, 402 insertions(+), 92 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 75c7d12f..029893c7 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -403,14 +403,20 @@ def generate_rectangle(kernel: spir.Kernel, if len(tasks) != len_for_reporting: print(f'P{rect.x_range[0]},{rect.y_range[0]}: Reduced from {len_for_reporting} to {len(tasks)} tasks.') - task_bindings = task_recycling.plan_task_bindings(tasks, task_creation_behavior, set(color_map.values())) + data_task_colors = { + i: _data_task_color(rect.metadata, i, task, color_map) + for i, task in enumerate(tasks) if task.task_type == 'data' + } + task_bindings = task_recycling.plan_task_bindings(tasks, task_creation_behavior, set(color_map.values()), + data_task_colors) place_block_bytes = _place_block_storage_bytes(rect.metadata.place) print(f'Stats P{rect.x_range[0]},{rect.y_range[0]}: {place_block_bytes} bytes/PE, ' f'{sum(1 if t.task_type == "local" else 0 for t in tasks)} local tasks across ' f'{len(task_bindings.local_slots)} local task IDs, ' - f'{sum(1 if t.task_type == "data" else 0 for t in tasks)} data tasks, ' + f'{sum(1 if t.task_type == "data" else 0 for t in tasks)} data tasks across ' + f'{len(task_bindings.data_slots)} colors, ' f'{len(set(color_map.values()))} colors') # Declare each logical local task ID alias. @@ -422,24 +428,15 @@ def generate_rectangle(kernel: spir.Kernel, current_code.write(f'var {task_bindings.state_var(representative)}: u16 = ' f'{task_bindings.invalid_state_literal(representative)};\n') - # Declare each data task ID. - for i, task in enumerate(tasks): - if task.task_type != "data": - continue - - stmt = rect.metadata.compute.statements[task.statements[0]] - assert isinstance(stmt, spir.ForeachStatement) - sname = stmt.receive_stream.stream_name - if isinstance(sname, spir.ArraySlice): - sname = sname.array - if name_to_csl(sname) + "_H2D" in color_map: - color = color_map[name_to_csl(sname) + "_H2D"] - elif name_to_csl(sname) + "_IN" in color_map: - color = color_map[name_to_csl(sname) + "_IN"] - else: - print(color_map) - raise ValueError(f'Cannot find color for stream "{name_to_csl(sname)}" in data task {i}') - current_code.write(f'const dtask_{i}_id = @get_data_task_id(@get_color({color}));\n') + # Declare each data task ID. Data tasks that share a color are aliases of one hardware ID, and + # a state variable selects which of them the shared task runs as. + for slot in task_bindings.data_slots: + for task_index in slot.task_indices: + current_code.write(f'const dtask_{task_index}_id = @get_data_task_id(@get_color({slot.color}));\n') + if slot.recycled: + representative = slot.representative_task_index + current_code.write(f'var {task_bindings.data_state_var(representative)}: u16 = ' + f'{task_bindings.data_state(representative)};\n') # Generate each local slot as one hardware task. for slot in task_bindings.local_slots: @@ -510,15 +507,14 @@ def generate_rectangle(kernel: spir.Kernel, if not slot.recycled and tasks[slot.representative_task_index].blocked: footer.write(f' @block(task_{slot.representative_task_index}_id);\n') - # Generate each data task. - for i, task in enumerate(tasks): - if task.task_type != 'data': - continue - _generate_data_task(rect.metadata, i, task, current_code, header, footer, dsds, dtypes, color_map, tasks, - task_bindings) - footer.write(f' @bind_data_task(dtask_{i}, dtask_{i}_id);\n') - if task.blocked: - footer.write(f' @block(dtask_{i}_id);\n') + # Generate each color's data task, dispatching between the receives that share it. + for slot in task_bindings.data_slots: + _generate_data_task_slot(rect.metadata, slot, current_code, header, dsds, dtypes, tasks, task_bindings) + representative = slot.representative_task_index + footer.write(f' @bind_data_task({task_bindings.data_function_name(slot)}, ' + f'dtask_{representative}_id);\n') + if tasks[representative].blocked: + footer.write(f' @block(dtask_{representative}_id);\n') max_task_id = max((slot.hardware_task_id for slot in task_bindings.local_slots), default=csl.LOCAL_TASK_IDS[0] - 1) @@ -559,6 +555,14 @@ def generate_rectangle(kernel: spir.Kernel, current_code.write(f' {task_bindings.state_var(representative)} = ' f'{task_bindings.invalid_state_literal(representative)};\n') + # Reset recycled data-slot state to the receive that runs first on each color. + for slot in task_bindings.data_slots: + if not slot.recycled: + continue + representative = slot.representative_task_index + current_code.write(f' {task_bindings.data_state_var(representative)} = ' + f'{task_bindings.data_state(representative)};\n') + # Reset data task counters and re-block dedicated tasks. for i, task in enumerate(tasks): if task.task_type == "data": @@ -569,7 +573,10 @@ def generate_rectangle(kernel: spir.Kernel, if stmt.parameter_range: param_range = stmt.parameter_range[0] current_code.write(f' __num_dtask_{i} = {param_range.start.as_ir()};\n') - if task.task_type == 'data' and task.blocked: + # Blocking a shared color is the first receive's business: the later ones are installed and + # unblocked by their predecessors, and blocking here would hold up the first one. + if (task.task_type == 'data' and task.blocked + and i == task_bindings.data_slot(i).representative_task_index): current_code.write(f' @block(dtask_{i}_id);\n') if task.task_type == 'local' and task.blocked and not task_bindings.is_recycled_local_task(i): current_code.write(f' @block(task_{i}_id);\n') @@ -1350,93 +1357,169 @@ def _write_indented_block(current_code: StringIO, block: str, indent: str) -> No current_code.write(f'{indent}{line}\n') -def _generate_data_task( +def _data_task_color(rect: PEBlock, task_index: int, task: tdag.CSLTask, color_map: dict[str, int]) -> int: + """ + Returns the color a data task listens on, which is also its hardware task ID. + + :param task_index: Only used to name the task in the error message. + """ + stmt = rect.compute.statements[task.statements[0]] + assert isinstance(stmt, spir.ForeachStatement) + sname = stmt.receive_stream.stream_name + if isinstance(sname, spir.ArraySlice): + sname = sname.array + if name_to_csl(sname) + '_H2D' in color_map: + return color_map[name_to_csl(sname) + '_H2D'] + if name_to_csl(sname) + '_IN' in color_map: + return color_map[name_to_csl(sname) + '_IN'] + raise ValueError(f'Cannot find color for stream "{name_to_csl(sname)}" in data task {task_index}') + + +def _generate_data_task_slot( rect: PEBlock, - task_index: int, - task: tdag.CSLTask, + slot: task_recycling.DataTaskSlot, current_code: StringIO, header: StringIO, - footer: StringIO, dsds: UniqueDSDDict, dtypes: dict[spir.Identifier, spir.IRType], - color_map: dict[str, int], tasks: list[tdag.CSLTask], task_bindings: task_recycling.TaskBindingPlan, ): """ - Generates a data task from a foreach loop. + Generates the one CSL data task that a color binds. + + A color that carries a single receive becomes that receive's task. A color reused by several + receives becomes a dispatcher over them, in the order they run, selected by the slot's state + variable -- the same shape :mod:`spada.syntax.csl.task_recycling` gives an overrun local task, + except that here sharing is forced rather than chosen. :param rect: The rectangle PE block to generate. - :param task: The data task to generate. + :param slot: The color and the data tasks bound to it. :param current_code: The caret to the code generator at the current position (global). :param header: A code generator stream for a file's header (where the declarations are). - :param footer: A code generator stream for a file's footer (the comptime block where the array would be exported). :param dsds: A dictionary mapping names to unique data structure descriptor objects. :param dtypes: A dictionary mapping identifiers to their defined types. - :param color_map: Dictionary mapping each stream to its respective color id ({name}_color also works). :param tasks: A list of all tasks in the kernel. - :param task_bindings: The local-task binding plan, needed to install the state of a recycled - successor slot before handing control to it. + :param task_bindings: The binding plan, which supplies the state variable and the states. + """ + generated = [index for index in slot.task_indices + if _declare_data_task_counter(rect, index, tasks[index], current_code)] + if not generated: + return + + representative = rect.compute.statements[tasks[generated[0]].statements[0]] + argtype_csl = dtype_as_csl(representative.stream_variable.dtype) + argname = name_to_csl(representative.stream_variable.identifier) + current_code.write(f'task {task_bindings.data_function_name(slot)}({argname}: {argtype_csl}) void {{\n') + + if len(generated) == 1: + _generate_data_task_body(rect, generated[0], tasks[generated[0]], current_code, header, dsds, dtypes, + tasks, task_bindings, argname, indent=' ', self_block=False) + else: + state_var = task_bindings.data_state_var(generated[0]) + for branch, task_index in enumerate(generated): + keyword = 'if' if branch == 0 else 'else if' + current_code.write(f' {keyword} ({state_var} == {task_bindings.data_state(task_index)}) {{\n') + _generate_data_task_body(rect, task_index, tasks[task_index], current_code, header, dsds, dtypes, + tasks, task_bindings, argname, indent=' ', self_block=True) + current_code.write(' }\n') + current_code.write('}\n') + + +def _declare_data_task_counter(rect: PEBlock, task_index: int, task: tdag.CSLTask, + current_code: StringIO) -> bool: + """ + Declares the counter that tells a data task when it has received its last wavelet. + + :return: Whether the task has a body to generate at all. """ - # * If index is requested: before unblocking task, set k; inc at end of task - # * Wavelet-triggered task as fallback assert task.task_type == 'data' assert len(task.statements) == 1 stmt_id = task.statements[0] - if isinstance(stmt_id, int) and stmt_id >= 0: - stmt = rect.compute.statements[stmt_id] - assert isinstance(stmt, spir.ForeachStatement) - else: - return + if not isinstance(stmt_id, int) or stmt_id < 0: + return False + stmt = rect.compute.statements[stmt_id] + assert isinstance(stmt, spir.ForeachStatement) + if stmt.parameter_range: + assert len(stmt.parameter_range) == 1, 'Only one-dimensional foreach loops are supported in data tasks' + var_dtype_csl = dtype_as_csl(stmt.variables[0].dtype) + current_code.write(f'var __num_dtask_{task_index}: {var_dtype_csl} = ' + f'{stmt.parameter_range[0].start.as_ir()};\n') + return True + + +def _generate_data_task_body( + rect: PEBlock, + task_index: int, + task: tdag.CSLTask, + current_code: StringIO, + header: StringIO, + dsds: UniqueDSDDict, + dtypes: dict[spir.Identifier, spir.IRType], + tasks: list[tdag.CSLTask], + task_bindings: task_recycling.TaskBindingPlan, + argname: str, + indent: str, + self_block: bool, +): + """ + Generates what one data task does with a wavelet, without the surrounding task frame. + + :param argname: The wavelet parameter of the generated task, which the receives sharing a color + have in common; a receive that names it differently gets an alias. + :param self_block: Whether the task blocks its color once its last wavelet has arrived. Set for a + shared color, where leaving it live would let the next epoch's wavelets be + taken by this branch. + """ + # * If index is requested: before unblocking task, set k; inc at end of task + # * Wavelet-triggered task as fallback + stmt = rect.compute.statements[task.statements[0]] + assert isinstance(stmt, spir.ForeachStatement) next_task, itedge = task.outgoing[0] next_task_type = tasks[next_task].task_type if next_task != -1 else 'local' itedge_code = 'unblock' if itedge == tdag.InterTaskEdge.UNBLOCK else 'activate' - # If a range was specified, write counter and add code to execute next task if stmt.parameter_range: - assert len(stmt.parameter_range) == 1, 'Only one-dimensional foreach loops are supported in data tasks' + lines = [] + if self_block: + lines.append(f'@block(dtask_{task_index}_id);') if next_task == -1: - next_task_code = f'@{itedge_code}(exit_task_id);' + lines.append(f'@{itedge_code}(exit_task_id);') else: - prefix = "d" if next_task_type == 'data' else "" - lines = [] + prefix = 'd' if next_task_type == 'data' else '' if next_task_type == 'local': lines.extend(task_bindings.emit_local_transition_preamble( next_task, tasks[next_task].blocked, indent='').splitlines()) + else: + lines.extend(task_bindings.emit_data_transition_preamble(next_task, indent='').splitlines()) lines.append(f'@{itedge_code}({prefix}task_{next_task}_id);') - next_task_code = '\n '.join(lines) - var_dtype_csl = dtype_as_csl(stmt.variables[0].dtype) param_range = stmt.parameter_range[0] - current_code.write(f"var __num_dtask_{task_index}: {var_dtype_csl} = {param_range.start.as_ir()};\n") - - next_task_code = f""" - __num_dtask_{task_index} += {1 if param_range.step is None else param_range.step.as_ir()}; - if (__num_dtask_{task_index} == {param_range.stop.as_ir()}) {{ - {next_task_code} - }}""" + step = 1 if param_range.step is None else param_range.step.as_ir() + body = f'\n{indent} '.join(lines) + next_task_code = (f'{indent}__num_dtask_{task_index} += {step};\n' + f'{indent}if (__num_dtask_{task_index} == {param_range.stop.as_ir()}) {{\n' + f'{indent} {body}\n' + f'{indent}}}\n') else: - next_task_code = "" + next_task_code = '' - # Write frame for data task - argtype_csl = dtype_as_csl(stmt.stream_variable.dtype) - argname = name_to_csl(stmt.stream_variable.identifier) - current_code.write(f"task dtask_{task_index}({argname}: {argtype_csl}) void {{\n") if stmt.variables: - current_code.write( - f' var {name_to_csl(stmt.variables[0].identifier)}: {var_dtype_csl} = __num_dtask_{task_index};\n') + var_dtype_csl = dtype_as_csl(stmt.variables[0].dtype) + current_code.write(f'{indent}var {name_to_csl(stmt.variables[0].identifier)}: {var_dtype_csl} = ' + f'__num_dtask_{task_index};\n') + own_argname = name_to_csl(stmt.stream_variable.identifier) + if own_argname != argname: + current_code.write(f'{indent}const {own_argname} = {argname};\n') # Write op contents for substmt in stmt.body: code = cslstmt.generate_csl_statement(substmt, dsds, dtypes, None, header, in_foreach_or_map=True) - for line in code.splitlines(): - current_code.write(f' {line}\n') + current_code.write(f'{indent}{line}\n') - # Write footer current_code.write(next_task_code) - current_code.write(f"\n}}\n") def _generate_task_code(rect: PEBlock, @@ -1495,15 +1578,14 @@ def _generate_task_code(rect: PEBlock, next_task, tasks[next_task].blocked, indent=indent) task_id = f'task_{next_task}_id' else: + transition_preamble = task_bindings.emit_data_transition_preamble( + next_task, indent=indent) task_id = f'dtask_{next_task}_id' async_target = dsd_ops.AsyncTarget(task_id, itedge.name.lower()) else: async_target = None - if transition_preamble: - current_code.write(transition_preamble) - code = cslstmt.generate_csl_statement(stmt, dsds, dtypes, async_target, header) lines = code.splitlines() @@ -1514,6 +1596,13 @@ def _generate_task_code(rect: PEBlock, if async_target is not None and async_target.target_task in code: skip_activation = True + # The preamble installs the successor's state, so it belongs immediately before whatever + # hands control over. When the operation does that itself the preamble has to precede it; + # otherwise it waits for the explicit activate/unblock written further down. + if transition_preamble and skip_activation: + current_code.write(transition_preamble) + transition_preamble = '' + for line in lines: current_code.write(f'{indent}{line}\n') @@ -1533,12 +1622,14 @@ def _generate_task_code(rect: PEBlock, task_id = 'exit_task_id' else: if tasks[next_task].task_type == 'local': - if not transition_preamble: - current_code.write( - task_bindings.emit_local_transition_preamble( - next_task, tasks[next_task].blocked, indent=indent)) + current_code.write( + transition_preamble or task_bindings.emit_local_transition_preamble( + next_task, tasks[next_task].blocked, indent=indent)) task_id = f'task_{next_task}_id' else: + current_code.write( + transition_preamble or task_bindings.emit_data_transition_preamble( + next_task, indent=indent)) task_id = f'dtask_{next_task}_id' if itedge == tdag.InterTaskEdge.ACTIVATE: current_code.write(f'{indent}@activate({task_id});\n') diff --git a/spada/syntax/csl/task_recycling.py b/spada/syntax/csl/task_recycling.py index 96dc6d78..e57742c6 100644 --- a/spada/syntax/csl/task_recycling.py +++ b/spada/syntax/csl/task_recycling.py @@ -1,7 +1,18 @@ """ -This module plans how logical CSL local tasks can share a smaller set of -hardware local-task IDs when the program contains more local tasks than the -target architecture exposes in :mod:`spada.syntax.csl.constants`. +This module plans how logical CSL tasks share hardware task IDs. + +For local tasks that is an optimization: it is needed only when the program +contains more of them than the target architecture exposes in +:mod:`spada.syntax.csl.constants`. + +For data tasks it is mandatory. A data task's hardware ID *is* the color it +listens on, so two receives that a PE performs on one channel have no choice but +to share, and ``cslc`` rejects the alternative outright ("task ID '0' bound to +more than one task"). Channel reuse across epochs is what makes a channel a +reusable resource in the first place (see ``irspec/docs/spatial/routing.md``), so +the two receives are a shape the backend has to be able to express. Data-task +slots are therefore planned here alongside the local ones; see +:func:`plan_data_task_slots` for what they additionally require of codegen. Terminology ----------- @@ -33,14 +44,16 @@ 1. Collect local tasks - Only ``task.task_type == 'local'`` participates in recycling. Data tasks - have their own binding scheme and are not handled here. + Only ``task.task_type == 'local'`` participates in slot *assignment*: a + local task may go to any free hardware ID, whereas a data task's ID is + dictated by its color. 2. Decide whether recycling is needed If the requested task-creation behavior forbids recycling, or if the number of local tasks already fits in the available hardware IDs, the planner emits - a trivial one-task-per-slot mapping. + a trivial one-task-per-slot mapping. Data tasks are grouped by color + regardless, since that grouping is not a choice. 3. Assign overflow tasks to slots @@ -126,6 +139,10 @@ * Before activating or unblocking a recycled local task, lowering emits the transition preamble returned by this module. +Recycled *data* slots follow the same shape, with one addition: a branch blocks +its own color once it has received the last wavelet it expects. See +:func:`plan_data_task_slots`. + Determinism ----------- @@ -140,7 +157,7 @@ """ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field, replace import heapq from typing import Iterable @@ -178,6 +195,30 @@ def recycled(self) -> bool: return len(self.task_indices) > 1 +@dataclass(frozen=True) +class DataTaskSlot: + """ + The data tasks a PE binds to one color. + + Unlike :class:`LocalTaskSlot` this is not an allocation decision: the color a + receive listens on determines the hardware ID, so every logical data task on + that color lands here. ``task_indices`` is in the order the tasks run. + """ + + color: int + task_indices: tuple[int, ...] + + @property + def representative_task_index(self) -> int: + """Return the first logical data task bound to this color.""" + return self.task_indices[0] + + @property + def recycled(self) -> bool: + """Whether the color carries more than one logical data task.""" + return len(self.task_indices) > 1 + + @dataclass(frozen=True) class TaskBindingPlan: """ @@ -194,6 +235,9 @@ class TaskBindingPlan: local_slots: tuple[LocalTaskSlot, ...] task_to_local_slot: dict[int, int] task_to_local_state: dict[int, int] + data_slots: tuple[DataTaskSlot, ...] = () + task_to_data_slot: dict[int, int] = field(default_factory=dict) + task_to_data_state: dict[int, int] = field(default_factory=dict) @property def uses_recycling(self) -> bool: @@ -263,17 +307,55 @@ def emit_local_transition_preamble( lines.append(f'{indent}{state_var} = {state_value};') return '\n'.join(lines) + '\n' + def data_slot(self, task_index: int) -> DataTaskSlot: + """Return the slot holding the color that ``task_index`` receives on.""" + return self.data_slots[self.task_to_data_slot[task_index]] + + def data_state(self, task_index: int) -> int: + """Return the per-color state number assigned to ``task_index``.""" + return self.task_to_data_state[task_index] + + def is_recycled_data_task(self, task_index: int) -> bool: + """Return whether ``task_index`` shares its color with another data task.""" + return self.data_slot(task_index).recycled + + def data_state_var(self, task_index: int) -> str: + """Return the generated CSL state variable name for ``task_index``'s color.""" + return f'__dtask_color_{self.data_slot(task_index).color}_state' + + def data_function_name(self, slot: DataTaskSlot) -> str: + """Return the generated task name for ``slot``. + + A color with one receive keeps the plain ``dtask_`` name, so the + overwhelmingly common case reads as it did before recycling existed. + """ + if not slot.recycled: + return f'dtask_{slot.representative_task_index}' + return f'dtask_color_{slot.color}' + + def emit_data_transition_preamble(self, task_index: int, indent: str = ' ') -> str: + """Emit the state assignment required before unblocking a recycled data task. + + Unlike a local slot this never re-blocks: the caller unblocks the color + immediately afterwards, and what keeps the color inert in the meantime is + the branch that blocked it when its own last wavelet arrived. + """ + if not self.is_recycled_data_task(task_index): + return '' + return f'{indent}{self.data_state_var(task_index)} = {self.data_state(task_index)};\n' + def plan_task_bindings( tasks: list[tdag.CSLTask], task_creation_behavior: tdag.TaskCreationBehavior, disallowed_task_ids: Optional[set[int]] = None, + data_task_colors: dict[int, int] | None = None, ) -> TaskBindingPlan: - """Compute a local-task binding plan for the generated CSL. + """Compute the task binding plan for the generated CSL. - Returns either a trivial one-task-per-slot mapping when recycling is not - required or not allowed, or a state-machine-compatible sharing plan when - local-task overrun occurs. + For local tasks, returns either a trivial one-task-per-slot mapping when + recycling is not required or not allowed, or a state-machine-compatible + sharing plan when local-task overrun occurs. ``STATE_MACHINE_ON_OVERRUN`` is the only mode that attempts recycling. Other modes either keep a unique mapping or raise when the local task count @@ -282,11 +364,81 @@ def plan_task_bindings( When recycling is needed, all tasks are colored together using load-balanced greedy coloring in degeneracy order, distributing tasks evenly across hardware slots to minimise dispatcher state machine size. + + :param data_task_colors: The color each data task listens on, keyed by task + index. Data tasks are grouped by it unconditionally; + omitting the mapping leaves ``data_slots`` empty. """ + data_slots, task_to_data_slot, task_to_data_state = plan_data_task_slots(tasks, data_task_colors or {}) + plan = _plan_local_bindings(tasks, task_creation_behavior, disallowed_task_ids or set()) + return replace(plan, + data_slots=data_slots, + task_to_data_slot=task_to_data_slot, + task_to_data_state=task_to_data_state) + + +def plan_data_task_slots( + tasks: list[tdag.CSLTask], + data_task_colors: dict[int, int], +) -> tuple[tuple[DataTaskSlot, ...], dict[int, int], dict[int, int]]: + """Group the data tasks by the color they listen on. + + Sharing a color is sound only if the receives take it in turns, which is the + same criterion local slots use: every trigger source of the later task must + be reachable from the earlier one. That much orders the *installation* of the + later branch, but not the arrival of its wavelets, which the fabric may + deliver while the earlier branch is still installed. Codegen closes that gap + by having each branch of a recycled slot ``@block`` its own color once its + last wavelet has arrived, so wavelets of the next epoch wait in the queue + until their branch is installed and unblocked. + + :param data_task_colors: The color each data task listens on, keyed by task index. + :return: ``(slots, task_to_slot, task_to_state)``, the last two mapping a task + index to its slot number and to its state within that slot. + :raises SyntaxError: If two data tasks share a color without being ordered. + """ + by_color: dict[int, list[int]] = {} + for task_index, task in enumerate(tasks): + if task.task_type != 'data': + continue + if task_index not in data_task_colors: + raise ValueError(f'No color given for data task {task_index}') + by_color.setdefault(data_task_colors[task_index], []).append(task_index) + + reachable = _compute_reachability(tasks) + trigger_sources = _trigger_sources(tasks) + + slots: list[DataTaskSlot] = [] + task_to_slot: dict[int, int] = {} + task_to_state: dict[int, int] = {} + for slot_index, color in enumerate(sorted(by_color)): + # Task indices follow the topological order of the completion DAG, so this is the order the + # receives run in; the check below is what makes sure of it. + task_indices = tuple(sorted(by_color[color])) + for earlier, later in zip(task_indices, task_indices[1:]): + if not _precedes_all_trigger_sources(earlier, later, trigger_sources, reachable): + raise SyntaxError( + f'Two receives on channel {color} at one PE are not ordered, so they cannot ' + f'share the data task the channel binds (tasks {earlier} and {later}).\n' + ' note: close the earlier stream before the later one is used, or assign the ' + 'later one a different channel') + slots.append(DataTaskSlot(color, task_indices)) + for state, task_index in enumerate(task_indices): + task_to_slot[task_index] = slot_index + task_to_state[task_index] = state + + return tuple(slots), task_to_slot, task_to_state + + +def _plan_local_bindings( + tasks: list[tdag.CSLTask], + task_creation_behavior: tdag.TaskCreationBehavior, + disallowed_task_ids: set[int], +) -> TaskBindingPlan: + """Assign local tasks to hardware slots, recycling them when they overrun.""" local_task_indices = [i for i, task in enumerate(tasks) if task.task_type == 'local'] if not local_task_indices: return TaskBindingPlan((), {}, {}) - disallowed_task_ids = disallowed_task_ids or set() allowed_local_task_ids = [t for t in constants.LOCAL_TASK_IDS if t not in disallowed_task_ids] diff --git a/tests/spatial_ir/test_task_recycling_codegen.py b/tests/spatial_ir/test_task_recycling_codegen.py index 1bb94da3..c8545987 100644 --- a/tests/spatial_ir/test_task_recycling_codegen.py +++ b/tests/spatial_ir/test_task_recycling_codegen.py @@ -4,6 +4,8 @@ import pytest from spada.lowering.spatial_ir_to_csl import lower_spatial_ir_to_csl +from spada.syntax.csl import task_recycling +from spada.syntax.csl import tasks as tdag from spada.syntax.spatial_ir import parser, passes _CSL_RUNTIME_TASK_RECYCLING_SAMPLES = os.path.join( @@ -92,6 +94,71 @@ def test_data_tasks_install_the_state_of_a_recycled_successor(): assert checked, 'sample no longer exercises a data task triggering a recycled local task' +def test_a_reused_channel_binds_one_data_task_that_dispatches_on_its_epoch(): + """Two receives on one channel at one PE share the data task the channel binds. + + A data task's hardware ID is the color, so binding two of them is not merely wasteful + but rejected by cslc ("task ID '0' bound to more than one task"). The static Batcher + reaches that shape at L=3, where a PE compares against its neighbour in two phases that + the sample gives the same channel. + """ + sample = os.path.join( + os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_1D.sptl') + kernel = parser.parse_file(sample) + kernel = passes.concretize_parameters(kernel, L=3) + kernel = passes.constexpr_propagation(kernel) + + csl_files = lower_spatial_ir_to_csl(kernel) + + shared = 0 + for file in csl_files: + code = file.code + colors: dict[str, list[str]] = {} + for task_index, color in re.findall(r'const dtask_(\d+)_id = @get_data_task_id\(@get_color\((\d+)\)\)', code): + colors.setdefault(color, []).append(task_index) + + bound = re.findall(r'@bind_data_task\(\w+, dtask_(\d+)_id\);', code) + assert len(bound) == len(colors), ( + f'{file.filename}: binds {len(bound)} data tasks for {len(colors)} colors') + + for color, task_indices in colors.items(): + if len(task_indices) == 1: + continue + shared += 1 + dispatcher = re.search(rf'task dtask_color_{color}\([^)]*\) void \{{(.*?)\n\}}', code, re.S) + assert dispatcher, f'{file.filename}: color {color} is reused but has no dispatcher' + body = dispatcher.group(1) + states = re.findall(r'(?:else )?if \(__dtask_color_' + color + r'_state == (\d+)\)', body) + assert states == [str(state) for state in range(len(task_indices))], ( + f'{file.filename}: color {color} dispatches on {states} for {len(task_indices)} receives') + # A branch that keeps its color live would take the next epoch's wavelets as its own. + for task_index in task_indices: + assert f'@block(dtask_{task_index}_id);' in body, ( + f'{file.filename}: dtask_{task_index} does not block color {color} when it is done') + + assert shared, 'sample no longer reuses a channel for two receives at one PE' + + +def test_data_tasks_sharing_a_channel_must_take_turns(): + """Receives that could run concurrently cannot share a channel's data task.""" + def task(index: int, task_type: str, successor: int) -> tdag.CSLTask: + edge = tdag.InterTaskEdge.SEQUENCE if successor == -1 else tdag.InterTaskEdge.UNBLOCK + return tdag.CSLTask(index, task_type, [index], [(successor, edge)], blocked=task_type == 'data') + + # 0 -> 1 (receive) -> 2 -> 3 (receive), with 1 and 3 on the same channel. + ordered = [task(0, 'local', 1), task(1, 'data', 2), task(2, 'local', 3), task(3, 'data', -1)] + + slots, task_to_slot, task_to_state = task_recycling.plan_data_task_slots(ordered, {1: 5, 3: 5}) + assert [slot.task_indices for slot in slots] == [(1, 3)] + assert task_to_slot == {1: 0, 3: 0} + assert task_to_state == {1: 0, 3: 1} + + # Dropping the edge from the first receive to the second one's trigger leaves both live at once. + concurrent = [task(0, 'local', 1), task(1, 'data', -1), task(2, 'local', 3), task(3, 'data', -1)] + with pytest.raises(SyntaxError, match='not ordered'): + task_recycling.plan_data_task_slots(concurrent, {1: 5, 3: 5}) + + def test_codegen_avoids_local_task_id_color_overlap(): path = os.path.join(_CSL_RUNTIME_TASK_RECYCLING_SAMPLES, 'task_color_overlap_many_channels.sptl') kernel = parser.parse_file(path) From 8410a1f31b5fd7b2da4159a1808552f4b9c5fcc9 Mon Sep 17 00:00:00 2001 From: glukas Date: Tue, 18 Aug 2026 12:55:05 +0200 Subject: [PATCH 22/68] fix batcher sample --- samples/spatial/sort/batcher_oddeven_1D.sptl | 26 +++++++++++++------- 1 file changed, 17 insertions(+), 9 deletions(-) diff --git a/samples/spatial/sort/batcher_oddeven_1D.sptl b/samples/spatial/sort/batcher_oddeven_1D.sptl index 7d05268c..1156c18c 100644 --- a/samples/spatial/sort/batcher_oddeven_1D.sptl +++ b/samples/spatial/sort/batcher_oddeven_1D.sptl @@ -11,11 +11,19 @@ * * Static channel assignment (no router reconfiguration): * At distance d, the d interleaved matchings (offset r = 0 .. d-1) overlap - * on the 1D mesh, so each (d, r) pair gets its own colors: - * fwd (east, +d) : channel 2*(((d - 1) + r)) - * bwd (west, -d) : channel 2*(((d - 1) + r)) + 1 - * Stages that share the same d reuse those colors (identical hop table). - * Total colors = 2*(N - 1). Intended for small N (e.g. L <= 3). + * on the 1D mesh, so each (d, r) pair gets its own colors. A channel may only + * be reused by a later stage if every PE keeps the same role on it: a PE that + * sends on a channel and later receives on it would have to swap both the + * input and the output of its router, and on WSE-2 that costs two switch + * advances, of which a sender can only ever emit one (the first advance takes + * RAMP off its input, so the second never leaves the PE). + * + * The p = 1 stages each get their own block, since a PE's role there depends + * on the stage. The p >= 2 stages of any distance d agree on roles -- sender + * iff index = d + r (mod 2d) -- and so share one pair of channels: + * p = 1 : fwd 2*((2^(l-1) - 1) + r), bwd fwd + 1 + * p >= 2 : fwd 2*(N - 1) + 2*((d - 1) + r), bwd fwd + 1 + * Total colors = 3*N - 4. Intended for small N (e.g. L <= 3). * * Example L=3 (n=8): * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) @@ -90,21 +98,21 @@ kernel @batcher_oddeven_1d( dataflow i16 i, i16 j in [(b + ((1<<(l-p)) + r)):((b + (1< fwd = relative_stream(1<<(l-p), 0) { hops = auto, - channel = 2 * (((1<<(l-p)) - 1) + r) + channel = (2 * ((1< bwd = relative_stream(-(1<<(l-p)), 0) { hops = auto, - channel = (2 * (((1<<(l-p)) - 1) + r)) + 1 + channel = ((2 * ((1< fwd = relative_stream(1<<(l-p), 0) { hops = auto, - channel = 2 * (((1<<(l-p)) - 1) + r) + channel = (2 * ((1< bwd = relative_stream(-(1<<(l-p)), 0) { hops = auto, - channel = (2 * (((1<<(l-p)) - 1) + r)) + 1 + channel = ((2 * ((1< Date: Tue, 18 Aug 2026 12:56:14 +0200 Subject: [PATCH 23/68] additional csl unit test --- .../samples/data_task_two_epochs.sptl | 68 +++++++++++++++++++ .../csl_runtime/test_data_task_two_epochs.sh | 44 ++++++++++++ 2 files changed, 112 insertions(+) create mode 100644 tests/csl_runtime/samples/data_task_two_epochs.sptl create mode 100755 tests/csl_runtime/test_data_task_two_epochs.sh diff --git a/tests/csl_runtime/samples/data_task_two_epochs.sptl b/tests/csl_runtime/samples/data_task_two_epochs.sptl new file mode 100644 index 00000000..c923d289 --- /dev/null +++ b/tests/csl_runtime/samples/data_task_two_epochs.sptl @@ -0,0 +1,68 @@ +/** + * One channel, two epochs, and nothing else. + * + * PE0 sends a word to PE1 in one phase and another word in the next, both on channel 0. + * The channel is a reusable resource, so this is legal -- but PE1 receives twice on it, and + * the data task a channel binds is the color itself. The two receives therefore have to + * share one hardware task and take turns, which is what this sample exercises. + * + * Constraints: R >= 1 (repeats the pair of epochs R times). + **/ +kernel @data_task_two_epochs( + stream[2, 1] readonly inp, + stream[2, 1] writeonly out +) { + place i16 i, i16 j in [0:2, 0] { + f32[2] val + f32 a + f32 b + } + + phase { + compute i16 i, i16 j in [0:2, 0] { + await receive(val, inp[i, j]) + a = val[0] + b = val[1] + } + } + + for i16 r in [0:R] { + phase { + dataflow i16 i, i16 j in [0:2, 0] { + stream first = relative_stream(1, 0) { + hops = auto, + channel = 0 + } + } + compute i16 i, i16 j in [0:1, 0] { + await send(a, first) + } + compute i16 i, i16 j in [1:2, 0] { + await receive(a, first) + } + } + + phase { + dataflow i16 i, i16 j in [0:2, 0] { + stream second = relative_stream(1, 0) { + hops = auto, + channel = 0 + } + } + compute i16 i, i16 j in [0:1, 0] { + await send(b, second) + } + compute i16 i, i16 j in [1:2, 0] { + await receive(b, second) + } + } + } + + phase { + compute i16 i, i16 j in [0:2, 0] { + val[0] = a + val[1] = b + await send(val, out[i, j]) + } + } +} diff --git a/tests/csl_runtime/test_data_task_two_epochs.sh b/tests/csl_runtime/test_data_task_two_epochs.sh new file mode 100755 index 00000000..49f4f33c --- /dev/null +++ b/tests/csl_runtime/test_data_task_two_epochs.sh @@ -0,0 +1,44 @@ +#!/bin/sh +# E2E: one channel carrying two epochs between the same pair of PEs, and nothing else. +# Kernel: samples/data_task_two_epochs.sptl params: R (repeats of the pair of epochs). +# PE1 receives twice on channel 0, so both receives share the data task the channel binds. +# After the run both PEs hold PE0's two keys. + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +. "$SCRIPT_DIR/_lib.sh" + +SAMPLE="$SCRIPT_DIR/samples/data_task_two_epochs.sptl" +FOLDER="data_task_two_epochs_sptl" + +run_epochs() { + r=$1 + echo "--- data_task_two_epochs R=$r ---" + + sptlc "$SAMPLE" "$FOLDER" -p R=$r + + python3 - < Date: Tue, 18 Aug 2026 14:01:44 +0200 Subject: [PATCH 24/68] update routing script --- samples/spatial/sort/plot_batcher_routing.py | 39 ++++++++++++++------ 1 file changed, 28 insertions(+), 11 deletions(-) diff --git a/samples/spatial/sort/plot_batcher_routing.py b/samples/spatial/sort/plot_batcher_routing.py index d26a0449..55ca892c 100644 --- a/samples/spatial/sort/plot_batcher_routing.py +++ b/samples/spatial/sort/plot_batcher_routing.py @@ -4,11 +4,13 @@ Versions -------- - static batcher_oddeven_1D.sptl. At distance d, offset r uses - fwd (east, +d) : channel 2*((d - 1) + r) - bwd (west, -d) : channel 2*((d - 1) + r) + 1 - Stages that share d reuse those colors, so a PE that plays a - different role at the same d occupies several switch positions. + static batcher_oddeven_1D.sptl. No router reconfiguration at all: a + channel is only reused where every PE keeps its role on it, so each + cell of the table holds a single switch position. The p = 1 stages + get their own block, since a PE's role there depends on the stage; + the p >= 2 stages of one distance d agree on roles and share: + p = 1 : fwd 2*((d - 1) + r), bwd fwd + 1 + p >= 2 : fwd 2*(n - 1) + 2*((d - 1) + r), bwd fwd + 1 bundled batcher_oddeven_bundled_1D.sptl. Each phase uses two colors, one per direction. Phases with d >= 2 whose comparators form a run of at least two sources are a shift bundle: sources inject then @@ -97,12 +99,19 @@ class Cell: filter_init: int | None = None -def fwd_channel_static(dist: int, offset: int) -> int: - return 2 * ((dist - 1) + offset) +def fwd_channel_static(dist: int, offset: int, p: int, n: int) -> int: + """Eastbound channel of matching ``(dist, offset)`` of a stage with the given ``p``. + The p = 1 stages live in their own block of ``2*(n - 1)`` channels because a + PE's role on such a channel depends on the stage; the p >= 2 stages of one + distance agree on roles and so share a channel above that block. + """ + block = 0 if p == 1 else 2 * (n - 1) + return block + 2 * ((dist - 1) + offset) -def bwd_channel_static(dist: int, offset: int) -> int: - return fwd_channel_static(dist, offset) + 1 + +def bwd_channel_static(dist: int, offset: int, p: int, n: int) -> int: + return fwd_channel_static(dist, offset, p, n) + 1 def batcher_phases(n: int) -> list[Phase]: @@ -117,7 +126,11 @@ def batcher_phases(n: int) -> list[Phase]: for r in range(dist): pairs = tuple((i, i + dist) for i in range(r, n, 1 << l)) matchings.append( - Matching(l, 1, dist, r, pairs, fwd_channel_static(dist, r), bwd_channel_static(dist, r)) + Matching( + l, 1, dist, r, pairs, + fwd_channel_static(dist, r, 1, n), + bwd_channel_static(dist, r, 1, n), + ) ) phases.append(Phase(index, l, 1, dist, tuple(matchings))) index += 1 @@ -132,7 +145,11 @@ def batcher_phases(n: int) -> list[Phase]: for lo in range(start, stop, 2 * dist): pairs.append((lo, lo + dist)) matchings.append( - Matching(l, p, dist, r, tuple(pairs), fwd_channel_static(dist, r), bwd_channel_static(dist, r)) + Matching( + l, p, dist, r, tuple(pairs), + fwd_channel_static(dist, r, p, n), + bwd_channel_static(dist, r, p, n), + ) ) phases.append(Phase(index, l, p, dist, tuple(matchings))) index += 1 From 1d13afcc43f5db6c32e23d43178205ba406a1637 Mon Sep 17 00:00:00 2001 From: glukas Date: Tue, 18 Aug 2026 16:58:37 +0200 Subject: [PATCH 25/68] Improve batcher sort, extend support to 16 PEs --- irspec/docs/spatial/routing.md | 17 ++- .../sort/batcher_oddeven_bundled_1D.sptl | 106 ++++++++++++------ samples/spatial/sort/plot_batcher_routing.py | 99 ++++++++++------ .../test_batcher_oddeven_bundled_1d.sh | 8 +- tests/spatial_ir/test_shift_bundles.py | 36 ++++-- .../spatial_ir/test_task_recycling_codegen.py | 8 +- 6 files changed, 185 insertions(+), 89 deletions(-) diff --git a/irspec/docs/spatial/routing.md b/irspec/docs/spatial/routing.md index 9ef4cead..734992a8 100644 --- a/irspec/docs/spatial/routing.md +++ b/irspec/docs/spatial/routing.md @@ -340,8 +340,9 @@ without being filtered. WSE-2 has four filters per PE, of which the `memcpy` module reserves one, so **three** are usable (`FILTERS_PER_PE`). A PE needs one per color it filters. *If a PE would need more, a compile error is raised*; the fix is to give some of the streams their own channels, which trades filters for - colors. `batcher_oddeven_bundled_1D.sptl` hits this at $2^4$ PEs, where a PE receives a bundle in - six phases. + colors. Three filters therefore means at most three bundled phases per PE, whatever the kernel: + `batcher_oddeven_bundled_1D.sptl` would want ten at $2^4$ PEs and bundles only its three widest + phases, which is where most of the colors are saved anyway. Reconfiguring a filter while wavelets are still in flight on its color is a data race, and the destination that terminates the stream cannot be reconfigured until the stream has drained, @@ -353,6 +354,18 @@ than the shift distance, so that no PE is both a source and a destination; a shi alone, since a chain at distance one is already sequenced by ordinary switch positions. Anything else falls back to the per-hop lowering, and to the errors above if that conflicts. +Which shifts are bundled is decided by the channel assignment rather than by an attribute: a bundle +is what several overlapping matchings on *one* channel become, so giving each matching a channel of +its own is how a kernel declines the trade. What it then costs is colors, and those can be won back +by reusing a channel across phases. Two unbundled shifts may share one safely when they agree on +axis and signed distance and their sources agree modulo twice that distance, because a PE's role — +source, relay or destination — is then a function of its position modulo twice the distance alone, +so one static configuration serves every phase in the pool. Sharing on any other basis risks a PE +that sends on the color in one phase and receives on it in another, which needs a two-sided switch +change that a sender cannot drive (see *Lowering to Switches*), and nothing in the compiler +currently rejects it. `batcher_oddeven_bundled_1D.sptl` pools on exactly this rule, and +`batcher_oddeven_1D.sptl` is the same rule written out as arithmetic. + This arrangement is the one used in Schnyder's *Distributed Sorting on the Cerebras Wafer-Scale Engine* (fig. 7.6) for the 2D reduce-scatter, and is known to run on WSE-2. diff --git a/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl b/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl index 793ce8a9..93d678b4 100644 --- a/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl +++ b/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl @@ -9,23 +9,34 @@ * Each drawn comparator is two messages (low -> high, then high -> low). * Lower index keeps min; higher index keeps max. * - * This is the bundled variant of batcher_oddeven_1D: a whole phase uses two colors, one per - * direction, where the static version needs one per matching. The comparators of a phase at - * distance d partition the line into alternating blocks of d PEs, and each block ships its - * keys d steps into the next -- an overlapping interval shift, which is what a bundle is. The - * low PEs of a block take turns nearest-the-partner-first, handing their routers over to relay - * mode as they finish; the high PEs are statically routed and pick their key out of the stream - * with a counter filter. See irspec/docs/spatial/routing.md. + * This is the bundled variant of batcher_oddeven_1D. The comparators of a phase at distance d + * partition the line into alternating blocks of d PEs, and each block ships its keys d steps + * into the next -- an overlapping interval shift, which is what a bundle is. The low PEs of a + * block take turns nearest-the-partner-first, handing their routers over to relay mode as they + * finish; the high PEs are statically routed and pick their key out of the stream with a + * counter filter. See irspec/docs/spatial/routing.md. * - * Both streams of a phase are declared once for the whole line, which is what puts every - * matching of that phase on one channel. Colors are not reused across phases: the router - * configurations depend on d, and a color holds one set of them. + * Bundling is chosen by the channel assignment, since a bundle is what several overlapping + * matchings on one channel become. It costs one filter per participating PE and saves the + * colors the matchings would otherwise need, and a PE has only three filters, so only the three + * widest phases are bundled -- the rule 4*d >= N picks exactly those for any L >= 3, and they + * are where the saving is largest. A bundled color cannot be reused by a later phase: its + * sources advance to pos1 and nothing resets them. * - * What caps L here is the three filters a PE can use, not the 21 colors. A PE needs one per - * phase it receives a bundle in, and only the d = 1 phases are unbundled (a chain of shifts - * of one is sequenced by ordinary switch positions), so the count is the number of phases - * with d >= 2: three at L = 3, six at L = 4. Beyond L = 3 the phases have to go back to a - * color per matching, as batcher_oddeven_1D does. + * The phases left unbundled are routed per matching, as in batcher_oddeven_1D, but their colors + * are pooled across phases. Two of them may share when they agree on direction, distance, and + * the source residue mod 2d, because a PE's role is then decided by its position mod 2d alone + * and one static configuration serves every phase in the pool. Sharing on any looser rule would + * let a PE send on a color in one phase and receive on it in another, which needs two switch + * advances from a sender that can only emit one. + * + * bundled (4*d >= N) : fwd N + 2*(l*(L+1) + p), bwd fwd + 1 + * pooled : fwd 2*((2*d - 2) + c), bwd fwd + 1 + * + * with c = r for p = 1 and c = d + r for p >= 2; the pooled blocks for successive d are disjoint + * because the block for d runs from 2*d-2 to 4*d-3, and unbundled phases have 4*d < N, so every + * pooled channel stays below N. Colors: 10 at L = 3 and 18 at L = 4, of the 21 available. L = 5 + * would need 34, which is what caps this kernel. * * Example L=3 (n=8): * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) @@ -35,7 +46,7 @@ * (l=3,p=2) dist 2: (2,4)(3,5) * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) * - * Constraints: 1 <= L <= 3 + * Constraints: 1 <= L <= 4 **/ kernel @batcher_oddeven_1d( stream[1<( for i16 l in [1:L+1] { // p = 1: all PEs participate, dist = 1<<(l-1). phase { - dataflow i16 i, i16 j in [0:1< fwd = relative_stream(1<<(l-1), 0) { - hops = auto, - channel = auto + // Offset r is one disjoint matching. The matchings share a channel where the phase is + // bundled, and take one each where it is not, which is what turns bundling off. + for i16 r in [0:1<<(l-1)] { + dataflow i16 i, i16 j in [r:1< fwd = relative_stream(1<<(l-1), 0) { + hops = auto, + channel = ((1<= (1< bwd = relative_stream(-(1<<(l-1)), 0) { + hops = auto, + channel = (((1<= (1< bwd = relative_stream(-(1<<(l-1)), 0) { - hops = auto, - channel = auto + dataflow i16 i, i16 j in [(r + (1<<(l-1))):1< fwd = relative_stream(1<<(l-1), 0) { + hops = auto, + channel = ((1<= (1< bwd = relative_stream(-(1<<(l-1)), 0) { + hops = auto, + channel = (((1<= (1<( // p = 2 .. l: skip first/last dist PEs of each 2^l box, dist = 1<<(l-p). for i16 p in [2:l+1] { phase { - dataflow i16 i, i16 j in [0:1< fwd = relative_stream(1<<(l-p), 0) { - hops = auto, - channel = auto - } - stream bwd = relative_stream(-(1<<(l-p)), 0) { - hops = auto, - channel = auto - } - } - for i16 r in [0:1<<(l-p)] { for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { + hops = auto, + channel = ((1<= (1< bwd = relative_stream(-(1<<(l-p)), 0) { + hops = auto, + channel = (((1<= (1< fwd = relative_stream(1<<(l-p), 0) { + hops = auto, + channel = ((1<= (1< bwd = relative_stream(-(1<<(l-p)), 0) { + hops = auto, + channel = (((1<= (1<= 2 stages of one distance d agree on roles and share: p = 1 : fwd 2*((d - 1) + r), bwd fwd + 1 p >= 2 : fwd 2*(n - 1) + 2*((d - 1) + r), bwd fwd + 1 - bundled batcher_oddeven_bundled_1D.sptl. Each phase uses two colors, one - per direction. Phases with d >= 2 whose comparators form a run of - at least two sources are a shift bundle: sources inject then - relay (pos0 / pos1), destinations stay put and pick their word - with a counter filter. Colors are not reused across phases. + bundled Each phase uses two colors, one per direction. Phases with d >= 2 + whose comparators form a run of at least two sources are a shift + bundle: sources inject then relay (pos0 / pos1), destinations stay + put and pick their word with a counter filter. Colors are not + reused across phases. This is what the sample did before the filter + budget capped it at L = 3; kept here for comparison. + hybrid batcher_oddeven_bundled_1D.sptl as it stands. Only the three widest + phases bundle, which is exactly the three filters a PE has: the + rule is 4*d >= n. The rest are routed per matching, on colors + pooled across phases by the source residue mod 2d, so that a PE's + role on a pooled color is the same in every phase that uses it. Views ----- @@ -51,7 +57,7 @@ from matplotlib.patches import Rectangle -VERSIONS = ("static", "bundled") +VERSIONS = ("static", "bundled", "hybrid") MIN_BUNDLE_DISTANCE = 2 MIN_BUNDLE_LENGTH = 2 @@ -156,16 +162,51 @@ def batcher_phases(n: int) -> list[Phase]: return phases -def assign_channels(phases: list[Phase], version: str) -> list[Phase]: +def is_bundled(phase: Phase, n: int, version: str) -> bool: + """Whether ``version`` serializes a phase's matchings onto one colour pair.""" + if version == "static" or phase.dist < MIN_BUNDLE_DISTANCE: + return False + if version == "bundled": + return True + return 4 * phase.dist >= n # hybrid: the three widest phases, one per filter + + +def assign_channels(phases: list[Phase], version: str, n: int) -> list[Phase]: """Rewrite matching channels to match the sample of ``version``.""" if version == "static": return phases + log_n = int(math.log2(n)) assigned = [] for ph in phases: - fwd, bwd = 2 * ph.index, 2 * ph.index + 1 - matchings = tuple(replace(m, fwd=fwd, bwd=bwd) for m in ph.matchings) - assigned.append(replace(ph, matchings=matchings)) - return assigned + matchings = [] + for m in ph.matchings: + if version == "bundled": + fwd = 2 * ph.index + elif is_bundled(ph, n, version): + fwd = n + 2 * (ph.l * (log_n + 1) + ph.p) + else: + # Pooled: the source residue mod 2d decides the colour, so phases that agree on + # it agree on every router configuration and can share. + residue = m.offset if ph.p == 1 else ph.dist + m.offset + fwd = 2 * ((2 * ph.dist - 2) + residue) + matchings.append(replace(m, fwd=fwd, bwd=fwd + 1)) + assigned.append(replace(ph, matchings=tuple(matchings))) + return _compact_channels(assigned) + + +def _compact_channels(phases: list[Phase]) -> list[Phase]: + """ + Renumbers the channels densely, as the compiler's colour allocation does. + + A hand-written channel formula generally leaves gaps, and only the channels a kernel actually + uses are given a colour, in ascending order. Renumbering here is what makes the drawn channel + axis the colour axis of the emitted layout. + """ + used = sorted({ch for ph in phases for m in ph.matchings for ch in (m.fwd, m.bwd)}) + color_of = {channel: index for index, channel in enumerate(used)} + return [replace(ph, matchings=tuple(replace(m, fwd=color_of[m.fwd], bwd=color_of[m.bwd]) + for m in ph.matchings)) + for ph in phases] def channel_count(phases: list[Phase]) -> int: @@ -244,15 +285,12 @@ def _draw_network(ax, phases: list[Phase], n: int, version: str) -> None: handles = [] for ch in range(n_channels): - if version == "static": - arrow = "↓" if ch % 2 == 0 else "↑" - kind = "fwd" if ch % 2 == 0 else "bwd" - label = f"{arrow} ch {ch} ({kind})" + arrow = "↓" if ch % 2 == 0 else "↑" + kind = "fwd" if ch % 2 == 0 else "bwd" + if version == "bundled": + label = f"{arrow} ch {ch} (p{ch // 2} {kind})" else: - phase = ch // 2 - arrow = "↓" if ch % 2 == 0 else "↑" - kind = "fwd" if ch % 2 == 0 else "bwd" - label = f"{arrow} ch {ch} (p{phase} {kind})" + label = f"{arrow} ch {ch} ({kind})" handles.append( Line2D([0], [0], color=channel_color(ch, n_channels), lw=2.0, label=label) ) @@ -386,17 +424,13 @@ def pe_table(phases: list[Phase], n: int, version: str) -> list[list[Cell]]: pairs = tuple(pair for m in ph.matchings for pair in m.pairs) if not pairs: continue - if version == "bundled": + if is_bundled(ph, n, version): fwd, bwd = ph.matchings[0].fwd, ph.matchings[0].bwd - if ph.dist >= MIN_BUNDLE_DISTANCE: - for start, length in _shift_runs(pairs, ph.dist): - if length >= MIN_BUNDLE_LENGTH: - _install_bundle(configs, filters, start, length, ph.dist, fwd, bwd) - else: - _install_ordinary_pair(configs, start, start + ph.dist, fwd, bwd) - else: - for lo, hi in pairs: - _install_ordinary_pair(configs, lo, hi, fwd, bwd) + for start, length in _shift_runs(pairs, ph.dist): + if length >= MIN_BUNDLE_LENGTH: + _install_bundle(configs, filters, start, length, ph.dist, fwd, bwd) + else: + _install_ordinary_pair(configs, start, start + ph.dist, fwd, bwd) else: for m in ph.matchings: for lo, hi in m.pairs: @@ -484,7 +518,7 @@ def _draw_table(ax, phases: list[Phase], n: int, version: str) -> None: def plot(n: int, view: str, version: str, outfile: str | None, show: bool) -> None: - phases = assign_channels(batcher_phases(n), version) + phases = assign_channels(batcher_phases(n), version, n) if view == "network": n_slots = sum(max(len(ph.matchings), 1) for ph in phases) fig, ax = plt.subplots(figsize=(max(8, 0.7 * n_slots + 0.8 * len(phases)), max(4, 0.45 * n))) @@ -528,8 +562,9 @@ def main() -> None: nargs="+", choices=VERSIONS + ("all",), default=["static"], - help="static (one color per matching) and/or bundled (two colors per phase). " - "'all' is both. Repeatable.", + help="static (one color per matching), bundled (two per phase) and/or hybrid (the " + "sample: the three widest phases bundled, the rest pooled). 'all' is every one. " + "Repeatable.", ) parser.add_argument( "--view", diff --git a/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh b/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh index 35954b62..09274fe8 100755 --- a/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh +++ b/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh @@ -1,9 +1,10 @@ #!/bin/sh # E2E test: the bundled 1D Batcher odd-even mergesort (2^L PEs, one f32 key per PE). # Kernel: batcher_oddeven_bundled_1D.sptl params: L -# Same result as batcher_oddeven_1D, but each phase runs on two colors instead of one per -# comparator: OUT_out[:, 0, 0] == sort(inp[:, 0, 0]). -# L <= 3, which is where the three filters a PE can use run out. +# Same result as batcher_oddeven_1D, but the three widest phases run on two colors each instead +# of one per comparator: OUT_out[:, 0, 0] == sort(inp[:, 0, 0]). +# L <= 4: three phases is what the filter budget allows, and at L = 5 the phases left unbundled +# need more than the 21 colors. set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" @@ -48,3 +49,4 @@ PYEOF run_batcher 1 run_batcher 2 run_batcher 3 +run_batcher 4 diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py index 2f1b3463..6b83ec34 100644 --- a/tests/spatial_ir/test_shift_bundles.py +++ b/tests/spatial_ir/test_shift_bundles.py @@ -253,20 +253,34 @@ def _bundled_batcher(l: int): return lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) -@pytest.mark.parametrize('l, phases', [(2, 3), (3, 6)]) -def test_a_batcher_phase_costs_two_colors(l: int, phases: int): - # The matchings of a phase are declared as one stream over the whole line and share its channel, - # so the count is per phase and per direction rather than per comparator. - layout = next(f.code for f in _bundled_batcher(l) if 'layout' in f.filename) +def _colors_of(layout: str, pattern: str = '') -> set[int]: routes = layout[layout.index('// Routes'):] - assert len({int(color) for color in re.findall(r'@get_color\((\d+)\)', routes)}) == 2 * phases + return {int(color) for color in re.findall(r'@get_color\((\d+)\)[^;]*' + pattern, routes)} -def test_the_batcher_runs_out_of_filters_before_it_runs_out_of_colors(): - # A PE receives a bundle in every phase whose distance is at least two, and cannot filter more - # than three colors; L = 4 has six such phases. - with pytest.raises(SyntaxError, match='wavelet filters'): - _bundled_batcher(4) +@pytest.mark.parametrize('l, colors', [(2, 6), (3, 10), (4, 18)]) +def test_the_batcher_fits_the_colors_it_has(l: int, colors: int): + # A bundled phase puts all of its matchings on one color pair; the phases left unbundled take a + # pair per matching, but share those across phases wherever their sources agree mod 2d. + from spada.syntax.csl import constants + + used = _colors_of(next(f.code for f in _bundled_batcher(l) if 'layout' in f.filename)) + assert len(used) == colors + assert len(used) <= len(constants.COLORS) + + +def test_only_the_widest_batcher_phases_are_bundled(): + # Bundling costs one filter at every participating PE and a PE has three, so the sample bundles + # the three widest phases only. Lowering at all is the check that no PE needs a fourth, since + # ``_check_filter_budget`` would refuse. + from spada.syntax.csl import constants + + layout = next(f.code for f in _bundled_batcher(4) if 'layout' in f.filename) + used, filtered, switched = _colors_of(layout), _colors_of(layout, r'\.filter'), _colors_of(layout, r'\.switches') + + assert len(filtered) == 2 * constants.FILTERS_PER_PE # three phases, two directions each + assert filtered == switched # a bundled color is one whose sources hand over to relay mode + assert not (used - filtered) & switched # the pooled ones hold a single static configuration def test_batcher_scalar_receive_lowers_to_data_task(): diff --git a/tests/spatial_ir/test_task_recycling_codegen.py b/tests/spatial_ir/test_task_recycling_codegen.py index c8545987..43ea2460 100644 --- a/tests/spatial_ir/test_task_recycling_codegen.py +++ b/tests/spatial_ir/test_task_recycling_codegen.py @@ -60,14 +60,14 @@ def test_csl_runtime_task_recycling_sample_lowers(filename: str): def test_data_tasks_install_the_state_of_a_recycled_successor(): """A data task handing control to a recycled slot must install that slot's state first. - Without the assignment the dispatcher runs whichever branch was installed last -- in the - bundled Batcher at L=3 that meant a PE silently skipped its comparator and the fabric - deadlocked behind the send it never made. + Without the assignment the dispatcher runs whichever branch was installed last, which means a + PE silently skips its comparator and the fabric deadlocks behind the send it never made. The + bundled Batcher reaches that shape once it has enough phases sharing a colour, at L=4. """ sample = os.path.join( os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_bundled_1D.sptl') kernel = parser.parse_file(sample) - kernel = passes.concretize_parameters(kernel, L=3) + kernel = passes.concretize_parameters(kernel, L=4) kernel = passes.constexpr_propagation(kernel) csl_files = lower_spatial_ir_to_csl(kernel) From 25955f9bad12b23ab613903d49db9aa912078513 Mon Sep 17 00:00:00 2001 From: glukas Date: Tue, 18 Aug 2026 17:51:12 +0200 Subject: [PATCH 26/68] support wse3 e2e tests --- tests/csl_runtime/Makefile | 9 ++++++++- tests/csl_runtime/run-in-lima.sh | 16 +++++++++++++++- 2 files changed, 23 insertions(+), 2 deletions(-) diff --git a/tests/csl_runtime/Makefile b/tests/csl_runtime/Makefile index 778841a4..60948164 100644 --- a/tests/csl_runtime/Makefile +++ b/tests/csl_runtime/Makefile @@ -12,9 +12,14 @@ SDK_DIR := $(THIS_DIR)cerebras-sdk CSL_SDK_DIR ?= $(SDK_DIR) SDK_PATH_PREFIX := $(abspath $(CSL_SDK_DIR)): +# Cerebras generation to compile and simulate for. The compiler reads it from the environment and +# passes it to cslc as --arch; test scripts read it to skip cases the generation cannot express. +WSE_ARCH ?= wse2 + TEST_ENV = PATH="$(SDK_PATH_PREFIX)$$PATH" \ PYTHONPATH="$(REPO_ROOT)$${PYTHONPATH:+:$$PYTHONPATH}" \ - SINGULARITY_BIND="$(REPO_ROOT)$${SINGULARITY_BIND:+,$$SINGULARITY_BIND}" + SINGULARITY_BIND="$(REPO_ROOT)$${SINGULARITY_BIND:+,$$SINGULARITY_BIND}" \ + WSE_ARCH="$(WSE_ARCH)" .DEFAULT_GOAL := help @@ -29,6 +34,7 @@ help: @echo " make -C tests/csl_runtime check-sdk [CSL_SDK_DIR=/path/to/sdk]" @echo " make -C tests/csl_runtime test [CSL_SDK_DIR=/path/to/sdk]" @echo " make -C tests/csl_runtime test-one TEST=test_add.sh [CSL_SDK_DIR=/path/to/sdk]" + @echo " make -C tests/csl_runtime test WSE_ARCH=wse3" @echo " make -C tests/csl_runtime shell [CSL_SDK_DIR=/path/to/sdk]" @echo " make -C tests/csl_runtime smoke-sdk SDK_EXAMPLES_DIR=/path/to/csl-extras-*" @echo "" @@ -40,6 +46,7 @@ help: @echo " - On Apple Silicon macOS, use run-in-lima.sh instead:" @echo " tests/csl_runtime/run-in-lima.sh --sdk-url " @echo " - CSL_SDK_DIR defaults to tests/csl_runtime/cerebras-sdk/ (populated by setup-sdk)." + @echo " - WSE_ARCH selects the Cerebras generation (wse2 or wse3); it defaults to wse2." # ── SDK download and extraction ─────────────────────────────────────────────── # Download the SDK tarball. Requires CSL_SDK_URL to be set: diff --git a/tests/csl_runtime/run-in-lima.sh b/tests/csl_runtime/run-in-lima.sh index 2eaeeedb..d315846c 100755 --- a/tests/csl_runtime/run-in-lima.sh +++ b/tests/csl_runtime/run-in-lima.sh @@ -12,6 +12,7 @@ # tests/csl_runtime/run-in-lima.sh --sdk /path/to/cs_sdk --smoke /path/to/csl-extras-* # tests/csl_runtime/run-in-lima.sh --sdk /path/to/cs_sdk --shell # tests/csl_runtime/run-in-lima.sh --sdk /path/to/cs_sdk --check +# tests/csl_runtime/run-in-lima.sh --sdk /path/to/cs_sdk --arch wse3 # # The SDK directory and the repository must both be under your Mac home # directory ($HOME), which Lima mounts automatically. @@ -32,6 +33,7 @@ SDK_URL="" TEST_NAME="" SMOKE_DIR="" MODE="test" # test | test-one | smoke | shell | check +WSE_ARCH="wse2" usage() { cat <<'EOF' @@ -59,6 +61,11 @@ Usage (run from the repo root): tests/csl_runtime/run-in-lima.sh --sdk /path/to/cs_sdk (same --test / --smoke / --shell / --check flags work with --sdk too) +Architecture: + --arch Cerebras generation to compile and simulate for (default wse2). + The compiler reads it from WSE_ARCH and passes it to cslc; tests + that a generation cannot express skip themselves with a note. + The repository must reside under $HOME, which Lima mounts automatically. Prerequisites (install once): @@ -75,6 +82,7 @@ while [[ $# -gt 0 ]]; do --smoke) MODE="smoke"; SMOKE_DIR="$(cd "$2" && pwd)"; shift 2 ;; --shell) MODE="shell"; shift ;; --check) MODE="check"; shift ;; + --arch) WSE_ARCH="$2"; shift 2 ;; -h|--help) usage ;; *) echo "Unknown argument: $1"; usage ;; esac @@ -90,6 +98,10 @@ if [[ -n "$SDK_DIR" && -n "$SDK_URL" ]]; then echo "" usage fi +if [[ "$WSE_ARCH" != "wse2" && "$WSE_ARCH" != "wse3" ]]; then + echo "ERROR: --arch must be wse2 or wse3, got '$WSE_ARCH'." + exit 1 +fi # ── Validate paths are under $HOME ──────────────────────────────────────────── check_under_home() { @@ -190,7 +202,8 @@ vm "if ! python3 -m pip --version >/dev/null 2>&1; then \ python3 -m pip install --no-deps --quiet -e '$REPO_ROOT'" # ── Delegate to the Makefile ────────────────────────────────────────────────── -MAKE_ARGS="CSL_SDK_DIR=$SDK_DIR" +MAKE_ARGS="CSL_SDK_DIR=$SDK_DIR WSE_ARCH=$WSE_ARCH" +echo "==> Target architecture: $WSE_ARCH" case "$MODE" in check) @@ -214,6 +227,7 @@ case "$MODE" in limactl shell "$VM_NAME" -- bash -lc \ "export PATH='$SDK_DIR:\$PATH'; \ export PYTHONPATH='$REPO_ROOT\${PYTHONPATH:+:\$PYTHONPATH}'; \ + export WSE_ARCH='$WSE_ARCH'; \ cd '$REPO_ROOT'; exec bash" ;; esac From c78fe8bf4924a0c9fdb56d176c32912892ab932b Mon Sep 17 00:00:00 2001 From: glukas Date: Tue, 18 Aug 2026 21:57:56 +0200 Subject: [PATCH 27/68] Generalize sorts to K elems per PE, fix queue assignments --- README.md | 2 +- samples/spatial/sort/batcher_oddeven_1D.sptl | 128 +++++++++++++-- .../sort/batcher_oddeven_bundled_1D.sptl | 154 +++++++++++++++--- samples/spatial/sort/plot_batcher_routing.py | 43 +++-- spada/lowering/spatial_ir_to_csl.py | 126 +++++++++++--- spada/syntax/spatial_ir/stream_lifetime.py | 56 +++++++ tests/csl_runtime/test_batcher_oddeven_1d.sh | 33 ++-- .../test_batcher_oddeven_bundled_1d.sh | 42 +++-- .../csl_runtime/test_shift_bundle_filters.sh | 5 +- tests/spatial_ir/test_dsd_ops.py | 59 +++++++ tests/spatial_ir/test_shift_bundles.py | 47 +++++- tests/spatial_ir/test_stream_lifetime.py | 38 +++++ .../spatial_ir/test_task_recycling_codegen.py | 94 ++++++++--- 13 files changed, 681 insertions(+), 146 deletions(-) diff --git a/README.md b/README.md index f9ba38d8..78051e35 100644 --- a/README.md +++ b/README.md @@ -104,7 +104,7 @@ Sample SPADA programs are in `samples/`: | `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy` | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Sorting networks: `batcher_oddeven_1D` | +| `samples/spatial/sort/` | Sorting networks over `2^L` PEs holding `K` keys each: `batcher_oddeven_1D` (a color per matching) and `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | diff --git a/samples/spatial/sort/batcher_oddeven_1D.sptl b/samples/spatial/sort/batcher_oddeven_1D.sptl index 1156c18c..560c0e22 100644 --- a/samples/spatial/sort/batcher_oddeven_1D.sptl +++ b/samples/spatial/sort/batcher_oddeven_1D.sptl @@ -1,13 +1,35 @@ /** - * 1D Batcher odd-even mergesort over N = 2^L PEs, one f32 key per PE. - * Ascending: after the network, PE i holds the i-th smallest input. + * 1D Batcher odd-even mergesort over N = 2^L PEs, each holding K f32 keys: N*K keys in all. + * Ascending: after the network, PE i holds keys i*K .. i*K + (K-1) of the sorted sequence. * * Merge stages l = 1 .. L, each of width 2^l. Within stage l: * p = 1: every PE in each 2^l box compares at dist = 2^{l-1} * p = 2 .. l: skip box endpoints, compare at dist = 2^{l-p} * - * Each drawn comparator is two messages (low -> high, then high -> low). - * Lower index keeps min; higher index keeps max. + * Each drawn comparator is two messages (low -> high, then high -> low), a block of K keys each. + * Lower index keeps the small keys; higher index keeps the large ones. + * + * A block per PE + * -------------- + * Every PE keeps its block sorted. The load phase establishes that and every comparator preserves + * it, because a comparator becomes a *compare-split*, which is the standard block form of one: the + * partners trade blocks, the lower index keeps the K smallest of the 2K keys and the higher index + * the K largest, and both halves come out ascending since each is a merge of two ascending runs. + * A network that sorts N keys sorts N sorted blocks this way, so concatenating the blocks left to + * right gives the sorted sequence. K = 1 is the one-key-per-PE network again, and K need not be a + * power of two. + * + * Two ascending runs are merged by a K-step walk with one index into each. No bounds guard is + * needed: the walk takes K steps and each step advances exactly one index, so the low PE reads + * val[pv] and tmp[pt] with pv + pt = m <= K-1, and the high PE, walking down from the two block + * ends, keeps (K-1-pv) + (K-1-pt) = m <= K-1. The walk cannot run in place -- it writes the m-th + * result while still reading val at an index it has already passed -- so it writes res and copies + * back. + * + * The load phase sorts the block with an insertion network. Its inner trip count is K-1 rather + * than the m of a plain insertion sort, because a compute-level loop bound has to be a + * compile-time expression; the surplus trips clamp onto the (0, 1) pair, and a compare-exchange of + * an ordered pair is a no-op. * * Static channel assignment (no router reconfiguration): * At distance d, the d interleaved matchings (offset r = 0 .. d-1) overlap @@ -23,7 +45,8 @@ * iff index = d + r (mod 2d) -- and so share one pair of channels: * p = 1 : fwd 2*((2^(l-1) - 1) + r), bwd fwd + 1 * p >= 2 : fwd 2*(N - 1) + 2*((d - 1) + r), bwd fwd + 1 - * Total colors = 3*N - 4. Intended for small N (e.g. L <= 3). + * Total colors = 3*N - 4, whatever K is: a wider block is more wavelets on a + * channel, not more channels. Intended for small N (e.g. L <= 3). * * Example L=3 (n=8): * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) @@ -33,21 +56,38 @@ * (l=3,p=2) dist 2: (2,4)(3,5) * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) * - * Constraints: L >= 1 + * Constraints: L >= 1, K >= 1 **/ -kernel @batcher_oddeven_1d( - stream[1<[1<( + stream[1<[1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } } } @@ -80,12 +120,39 @@ kernel @batcher_oddeven_1d( compute i16 i, i16 j in [r:1< val else val + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = tmp[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } } } } @@ -119,12 +186,39 @@ kernel @batcher_oddeven_1d( compute i16 i, i16 j in [(b + ((1<<(l-p)) + r)):((b + (1< val else val + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = tmp[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } } } } @@ -132,7 +226,7 @@ kernel @batcher_oddeven_1d( } } - // Write sorted keys back to the host. + // Write the sorted blocks back to the host. phase { compute i16 i, i16 j in [0:1< high, then high -> low). - * Lower index keeps min; higher index keeps max. + * Each drawn comparator is two messages (low -> high, then high -> low), a block of K keys each. + * Lower index keeps the small keys; higher index keeps the large ones. * + * A block per PE + * -------------- + * Every PE keeps its block sorted. The load phase establishes that and every comparator preserves + * it, because a comparator becomes a *compare-split*, which is the standard block form of one: the + * partners trade blocks, the lower index keeps the K smallest of the 2K keys and the higher index + * the K largest, and both halves come out ascending since each is a merge of two ascending runs. + * A network that sorts N keys sorts N sorted blocks this way, so concatenating the blocks left to + * right gives the sorted sequence. K = 1 is the one-key-per-PE network again, and K need not be a + * power of two. + * + * Two ascending runs are merged by a K-step walk with one index into each. No bounds guard is + * needed: the walk takes K steps and each step advances exactly one index, so the low PE reads + * val[pv] and tmp[pt] with pv + pt = m <= K-1, and the high PE, walking down from the two block + * ends, keeps (K-1-pv) + (K-1-pt) = m <= K-1. The walk cannot run in place -- it writes the m-th + * result while still reading val at an index it has already passed -- so it writes res and copies + * back. + * + * The load phase sorts the block with an insertion network. Its inner trip count is K-1 rather + * than the m of a plain insertion sort, because a compute-level loop bound has to be a + * compile-time expression; the surplus trips clamp onto the (0, 1) pair, and a compare-exchange of + * an ordered pair is a no-op. + * + * Bundling + * -------- * This is the bundled variant of batcher_oddeven_1D. The comparators of a phase at distance d * partition the line into alternating blocks of d PEs, and each block ships its keys d steps * into the next -- an overlapping interval shift, which is what a bundle is. The low PEs of a * block take turns nearest-the-partner-first, handing their routers over to relay mode as they - * finish; the high PEs are statically routed and pick their key out of the stream with a + * finish; the high PEs are statically routed and pick their keys out of the stream with a * counter filter. See irspec/docs/spatial/routing.md. * * Bundling is chosen by the channel assignment, since a bundle is what several overlapping @@ -23,6 +47,12 @@ * are where the saving is largest. A bundled color cannot be reused by a later phase: its * sources advance to pos1 and nothing resets them. * + * K only widens the wavelet windows, never the color or filter count: a bundle of M sources + * carries M*K wavelets per epoch instead of M, and each destination keeps the K of them that its + * counter filter windows out (limit1 = M*K - 1, max_counter = K - 1). Streams are bounded here, + * unlike in batcher_oddeven_1D, because a bundle's source hands its router on when its stream + * closes, and a bound is what closes it as soon as the K-th wavelet is out. + * * The phases left unbundled are routed per matching, as in batcher_oddeven_1D, but their colors * are pooled across phases. Two of them may share when they agree on direction, distance, and * the source residue mod 2d, because a PE's role is then decided by its position mod 2d alone @@ -36,7 +66,8 @@ * with c = r for p = 1 and c = d + r for p >= 2; the pooled blocks for successive d are disjoint * because the block for d runs from 2*d-2 to 4*d-3, and unbundled phases have 4*d < N, so every * pooled channel stays below N. Colors: 10 at L = 3 and 18 at L = 4, of the 21 available. L = 5 - * would need 34, which is what caps this kernel. + * would need 34, which is what caps this kernel. On wse2 an interior PE at L = 4 has three inbound + * colors live at once, and a PE has two input queues, so L <= 3 there; wse3 has six queues. * * Example L=3 (n=8): * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) @@ -46,21 +77,38 @@ * (l=3,p=2) dist 2: (2,4)(3,5) * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) * - * Constraints: 1 <= L <= 4 + * Constraints: 1 <= L <= 4 (L <= 3 on wse2), K >= 1 **/ -kernel @batcher_oddeven_1d( - stream[1<[1<( + stream[1<[1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } } } @@ -71,21 +119,21 @@ kernel @batcher_oddeven_1d( // bundled, and take one each where it is not, which is what turns bundling off. for i16 r in [0:1<<(l-1)] { dataflow i16 i, i16 j in [r:1< fwd = relative_stream(1<<(l-1), 0) { + stream fwd = relative_stream(1<<(l-1), 0) { hops = auto, channel = ((1<= (1< bwd = relative_stream(-(1<<(l-1)), 0) { + stream bwd = relative_stream(-(1<<(l-1)), 0) { hops = auto, channel = (((1<= (1< fwd = relative_stream(1<<(l-1), 0) { + stream fwd = relative_stream(1<<(l-1), 0) { hops = auto, channel = ((1<= (1< bwd = relative_stream(-(1<<(l-1)), 0) { + stream bwd = relative_stream(-(1<<(l-1)), 0) { hops = auto, channel = (((1<= (1<( compute i16 i, i16 j in [r:1< val else val + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = tmp[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } } } } @@ -110,21 +185,21 @@ kernel @batcher_oddeven_1d( for i16 r in [0:1<<(l-p)] { for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { + stream fwd = relative_stream(1<<(l-p), 0) { hops = auto, channel = ((1<= (1< bwd = relative_stream(-(1<<(l-p)), 0) { + stream bwd = relative_stream(-(1<<(l-p)), 0) { hops = auto, channel = (((1<= (1< fwd = relative_stream(1<<(l-p), 0) { + stream fwd = relative_stream(1<<(l-p), 0) { hops = auto, channel = ((1<= (1< bwd = relative_stream(-(1<<(l-p)), 0) { + stream bwd = relative_stream(-(1<<(l-p)), 0) { hops = auto, channel = (((1<= (1<( compute i16 i, i16 j in [(b + ((1<<(l-p)) + r)):((b + (1< val else val + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = tmp[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } } } } @@ -146,7 +248,7 @@ kernel @batcher_oddeven_1d( } } - // Write sorted keys back to the host. + // Write the sorted blocks back to the host. phase { compute i16 i, i16 j in [0:1< int: - """``init_counter`` of a one-word destination filter, as ``_window_start`` emits it.""" +def _window_init(pe: int, first: int, step: int, words: int) -> int: + """``init_counter`` of a destination filter keeping ``words`` of them, as ``_window_start`` emits it.""" offset = 1 - first if step > 0 else first + 1 - return pe + offset if step > 0 else offset - pe + return (pe + offset if step > 0 else offset - pe) * words def _install_bundle( @@ -384,6 +386,7 @@ def _install_bundle( dist: int, fwd: int, bwd: int, + words: int, ) -> None: """Eastbound inject-then-relay plus westbound mirror, with destination filters.""" src_lo, src_hi = start, start + length @@ -397,7 +400,7 @@ def _install_bundle( # Stream travels east: last destination terminates, the others copy-and-forward. for pe in range(dst_lo, dst_hi - 1): _install_hop(configs, pe, fwd, Route("WEST", _tx("RAMP", "EAST"))) - filters[pe][fwd] = _window_init(pe, dst_lo, 1) + filters[pe][fwd] = _window_init(pe, dst_lo, 1, words) _install_hop(configs, dst_hi - 1, fwd, Route("WEST", _tx("RAMP"))) filters[dst_hi - 1][fwd] = 0 @@ -409,13 +412,13 @@ def _install_bundle( # Stream travels west: lowest destination terminates. for pe in range(src_lo + 1, src_hi): _install_hop(configs, pe, bwd, Route("EAST", _tx("RAMP", "WEST"))) - filters[pe][bwd] = _window_init(pe, src_hi - 1, -1) + filters[pe][bwd] = _window_init(pe, src_hi - 1, -1, words) _install_hop(configs, src_lo, bwd, Route("EAST", _tx("RAMP"))) filters[src_lo][bwd] = 0 -def pe_table(phases: list[Phase], n: int, version: str) -> list[list[Cell]]: - """Build the resolved per-PE switch table of ``version``.""" +def pe_table(phases: list[Phase], n: int, version: str, words: int = 1) -> list[list[Cell]]: + """Build the resolved per-PE switch table of ``version``, for ``words`` keys per PE.""" n_channels = channel_count(phases) configs: list[list[list[Route]]] = [[[] for _ in range(n_channels)] for _ in range(n)] filters: list[list[int | None]] = [[None] * n_channels for _ in range(n)] @@ -428,7 +431,7 @@ def pe_table(phases: list[Phase], n: int, version: str) -> list[list[Cell]]: fwd, bwd = ph.matchings[0].fwd, ph.matchings[0].bwd for start, length in _shift_runs(pairs, ph.dist): if length >= MIN_BUNDLE_LENGTH: - _install_bundle(configs, filters, start, length, ph.dist, fwd, bwd) + _install_bundle(configs, filters, start, length, ph.dist, fwd, bwd, words) else: _install_ordinary_pair(configs, start, start + ph.dist, fwd, bwd) else: @@ -442,9 +445,9 @@ def pe_table(phases: list[Phase], n: int, version: str) -> list[list[Cell]]: ] -def _draw_table(ax, phases: list[Phase], n: int, version: str) -> None: +def _draw_table(ax, phases: list[Phase], n: int, version: str, words: int) -> None: """Resolved per-PE switch positions; destination filters as a small ``fN``.""" - table = pe_table(phases, n, version) + table = pe_table(phases, n, version, words) n_channels = channel_count(phases) ax.set_xlim(-0.5, n_channels - 0.5) ax.set_ylim(n - 0.5, -0.5) @@ -453,7 +456,7 @@ def _draw_table(ax, phases: list[Phase], n: int, version: str) -> None: ax.set_xlabel("Channel") ax.set_ylabel("PE") ax.set_title( - f"Resolved switch positions ({version}, WSE-2), n={n} ({n_channels} colors); " + f"Resolved switch positions ({version}, WSE-2), n={n}, K={words} ({n_channels} colors); " "stacked bands are pos0, pos1, ...; fN is the filter init_counter" ) ax.set_aspect("equal") @@ -517,7 +520,7 @@ def _draw_table(ax, phases: list[Phase], n: int, version: str) -> None: ax.legend(handles=handles, title="rx→tx", loc="upper left", bbox_to_anchor=(1.02, 1), fontsize=8) -def plot(n: int, view: str, version: str, outfile: str | None, show: bool) -> None: +def plot(n: int, view: str, version: str, words: int, outfile: str | None, show: bool) -> None: phases = assign_channels(batcher_phases(n), version, n) if view == "network": n_slots = sum(max(len(ph.matchings), 1) for ph in phases) @@ -526,7 +529,7 @@ def plot(n: int, view: str, version: str, outfile: str | None, show: bool) -> No elif view == "table": n_channels = channel_count(phases) fig, ax = plt.subplots(figsize=(max(8, 0.45 * n_channels), max(4, 0.5 * n))) - _draw_table(ax, phases, n, version) + _draw_table(ax, phases, n, version, words) else: raise ValueError(f"unknown view {view}") @@ -557,6 +560,13 @@ def _expand_versions(requested: list[str]) -> list[str]: def main() -> None: parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) parser.add_argument("--n", type=int, default=8, help="number of PEs, power of two (default 8)") + parser.add_argument( + "--k", + type=int, + default=1, + help="keys per PE (default 1). A comparator trades K of them, so a bundle's cycle is " + "K times longer and every destination filter starts K times further back.", + ) parser.add_argument( "--version", nargs="+", @@ -579,13 +589,14 @@ def main() -> None: for version in versions: outfile = args.out if outfile is None and not args.show: - outfile = f"samples/spatial/sort/batcher_routing_{version}_n{args.n}_{args.view}.pdf" + suffix = "" if args.k == 1 else f"_k{args.k}" + outfile = f"samples/spatial/sort/batcher_routing_{version}_n{args.n}{suffix}_{args.view}.pdf" elif outfile is not None and len(versions) > 1: if outfile.endswith(".pdf"): outfile = f"{outfile[:-4]}_{version}.pdf" else: outfile = f"{outfile}_{version}" - plot(args.n, args.view, version, outfile, args.show) + plot(args.n, args.view, version, args.k, outfile, args.show) if __name__ == "__main__": diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 029893c7..3fa78d22 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -7,6 +7,7 @@ import functools from io import StringIO import textwrap +from typing import Optional from spada.syntax.common.types import BIT_WIDTH from spada.syntax.spatial_ir import irnodes as spir, canonicalization, analysis, passes from spada.syntax.spatial_ir import copy_elimination @@ -383,7 +384,9 @@ def generate_rectangle(kernel: spir.Kernel, dtypes = _collect_identifier_types(rect.metadata, kernel.arguments) try: - dsds = _collect_unique_dsds(tasks, rect.metadata, header, dtypes, kernel, use_memcpy_mode) + dsds = _collect_unique_dsds( + tasks, rect.metadata, header, dtypes, kernel, use_memcpy_mode, + location=f'PEs [{rect.x_range[0]}:{rect.x_range[1]}, {rect.y_range[0]}:{rect.y_range[1]}]') except KeyError as e: if e.args and isinstance(e.args[0], spir.Identifier): raise ValueError(f"Error in {e.args[0].lineinfo}. Undefined identifier \"{e.args[0].as_ir()}\".") @@ -948,6 +951,83 @@ def _declare_queue_initialization(dsds: UniqueDSDDict, rect: PEBlock, footer: St footer.write(f' @initialize_queue(@get_{kind}({queue}), .{{ .color = {color} }});\n') +def _queue_spans(compute: spir.ComputeBlock, names: set[spir.Identifier], queue_key, + inbound: bool) -> dict[str, tuple[int, int]]: + """ + Occupancy of each queue key along the linearized send/receive order of this PE. + + Nested transfers in a loop body are ordered by the walk, so sequential halo exchanges in one + ``for`` do not look concurrent. Uses of the same channel still collapse to one span, so a + colour that comes back after a gap keeps its queue for the whole of that span. + + :param compute: The compute block being lowered. + :param names: Streams that actually bind a fabric queue in this direction. + :param queue_key: Maps a stream identifier to its grouping key (channel, or the name itself). + :param inbound: True to walk receives, False to walk sends. + :return: Mapping of grouping key to ``(first_use, last_use)`` in linearized order. + """ + points: list[spir.Identifier] = [] + for statement in compute.statements: + for node in statement.walk(): + stream = _fabric_transfer_stream(node, inbound) + if stream is None or stream not in names: + continue + points.append(stream) + + spans: dict[str, tuple[int, int]] = {} + for index, stream in enumerate(points): + key = queue_key(stream) + if key in spans: + start, _ = spans[key] + spans[key] = (start, index) + else: + spans[key] = (index, index) + return spans + + +def _fabric_transfer_stream(node: spir.SpatialNode, inbound: bool) -> Optional[spir.Identifier]: + """ + The stream a node transfers in the requested direction, or ``None``. + """ + if inbound: + if isinstance(node, spir.ReceiveStatement): + return stream_lifetime.underlying_stream(node.stream_name) + if (isinstance(node, spir.ForeachStatement) and node.parameter_range + and node.receive_stream is not None): + return stream_lifetime.underlying_stream(node.receive_stream.stream_name) + return None + if isinstance(node, spir.SendStatement): + return stream_lifetime.underlying_stream(node.stream_name) + return None + + +def _streams_with_fabric_dsds(compute: spir.ComputeBlock, memcpy_mode: bool, + stream_args: set[spir.Identifier], inbound: bool) -> set[spir.Identifier]: + """ + Streams that lower to a fabric DSD in one direction, so they need a hardware queue. + + A data-task receive (``foreach`` with no range) binds the color itself and does not take a + queue. Memcpy arguments are already in local memory, so they do not either. + + :param compute: The compute block being lowered. + :param memcpy_mode: Whether memcpy mode is used. + :param stream_args: Kernel-argument streams, which memcpy has already copied. + :param inbound: True for receives, False for sends. + :return: The stream identifiers that need a queue in that direction. + """ + result: set[spir.Identifier] = set() + argument_names = {name.as_ir() for name in stream_args} + for statement in compute.statements: + for node in statement.walk(): + name = _fabric_transfer_stream(node, inbound) + if name is None: + continue + if memcpy_mode and name.as_ir() in argument_names: + continue + result.add(name) + return result + + def _collect_unique_dsds( tasks: list[tdag.CSLTask], rect: PEBlock, @@ -955,6 +1035,7 @@ def _collect_unique_dsds( dtypes: dict[spir.Identifier, spir.IRType], kernel: spir.Kernel, memcpy_mode: bool, + location: str = 'PEs', ) -> UniqueDSDDict: """ Returns a list of DSDs and generates them in the header. @@ -982,27 +1063,22 @@ def _collect_unique_dsds( for place_statement in rect.place.statements: if isinstance(place_statement, spir.FieldDeclaration): if isinstance(place_statement.dtype, spir.ArrayType): - try: - eval_shape = [s if isinstance(s, int) else s.eval() for s in place_statement.dtype.shape] - # If the product of the shape is 1, it is a scalar - if not eval_shape or all(s == 1 for s in eval_shape): - # Scalar, no DSD - continue - except ValueError: - # Dynamic shape, must create a DSD - pass + # An array declared without extents is a scalar and gets no DSD. One of a single + # element still gets one: CSL takes it as an array wherever a DSD is called for and + # rejects the bare name as an operand -- "only DSD/DSR operands are allowed for + # async operations" for a transfer, and a type error for a move between two memory + # locations. + if not place_statement.dtype.shape: + continue array_candidates[place_statement.field_name.as_ir()] = (place_statement, place_statement.dtype.shape) # Find used DSDs in compute block - # TODO: Infer input/output queue ID based on concurrency - input_queue_id_ctr = 0 - output_queue_id_ctr = 0 - # Streams that share a channel share a color, and a color binds to exactly one fabric queue per # PE -- the hardware rejects "two master input queues for the same color". Queues are therefore # handed out per channel; streams on ``auto`` channels get a color to themselves, so they key on - # their own name. + # their own name. Sequential channels may share a queue only when their occupancy spans on this + # PE do not overlap; see ``stream_lifetime.assign_fabric_queues``. channel_of_stream = { declaration.stream_name.as_ir(): declaration.stream.routing.resolved_channel for declaration in rect.dataflow.statements @@ -1013,23 +1089,27 @@ def queue_key(stream: spir.Identifier) -> str: channel = channel_of_stream.get(stream.as_ir(), 'auto') return stream.as_ir() if channel == 'auto' else f'channel {channel}' - input_queue_of: dict[str, int] = {} - output_queue_of: dict[str, int] = {} + input_names = _streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=True) + output_names = _streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=False) + input_queue_of = stream_lifetime.assign_fabric_queues( + _queue_spans(rect.compute, input_names, queue_key, inbound=True), csl.INPUT_QUEUE_IDS, + kind='input', architecture=csl.ARCH, location=location) + output_queue_of = stream_lifetime.assign_fabric_queues( + _queue_spans(rect.compute, output_names, queue_key, inbound=False), csl.OUTPUT_QUEUE_IDS, + kind='output', architecture=csl.ARCH, location=location) def allocate_input_queue(stream: spir.Identifier) -> int: - nonlocal input_queue_id_ctr key = queue_key(stream) if key not in input_queue_of: - input_queue_of[key] = csl.INPUT_QUEUE_IDS[input_queue_id_ctr % len(csl.INPUT_QUEUE_IDS)] - input_queue_id_ctr += 1 + raise SyntaxError( + f'{location}: no input queue was reserved for {key} (stream "{stream.as_ir()}").') return input_queue_of[key] def allocate_output_queue(stream: spir.Identifier) -> int: - nonlocal output_queue_id_ctr key = queue_key(stream) if key not in output_queue_of: - output_queue_of[key] = csl.OUTPUT_QUEUE_IDS[output_queue_id_ctr % len(csl.OUTPUT_QUEUE_IDS)] - output_queue_id_ctr += 1 + raise SyntaxError( + f'{location}: no output queue was reserved for {key} (stream "{stream.as_ir()}").') return output_queue_of[key] def _visit_foreach(stmt: spir.ForeachStatement) -> None: diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py index d8081e99..97ab447e 100644 --- a/spada/syntax/spatial_ir/stream_lifetime.py +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -13,6 +13,7 @@ * :func:`verify_stream_bounds` -- checks ``stream`` against the transferred element count. * :func:`check_use_after_close` -- rejects any use of a stream past its close. * :func:`check_channel_conflicts` -- rejects concurrent use of a channel. +* :func:`assign_fabric_queues` -- colors live channel spans onto hardware fabric queues. * :func:`elide_redundant_closes` -- drops closes whose channel is never reused. """ from collections import defaultdict @@ -739,6 +740,61 @@ def _never_concurrent(first: str, second: str, uses_per_rect: list[dict[spir.Ide return used_anywhere +def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int], *, kind: str, + architecture: str, location: str) -> dict[str, int]: + """ + Assigns hardware fabric queues to stream groups from their occupancy spans on one PE. + + A group is one channel (or one ``auto`` stream). It occupies a single interval from its first + use on this PE to its last, including gaps between epochs: wavelets of that colour can still + arrive in a gap, and remapping the queue onto another color while they sit there is what the + hardware rejects. Two groups may share a queue only when those spans do not overlap. + + The spans form an interval graph, so colouring them in start-time order is optimal. + + :param spans: Mapping of grouping key to an inclusive ``(first_use, last_use)`` statement index + pair on this PE. + :param queue_ids: The hardware queue identifiers this direction may use, in the order they + should be handed out. + :param kind: ``'input'`` or ``'output'``, for the diagnostic. + :param architecture: The target name, for the diagnostic. + :param location: The PE rectangle, for the diagnostic. + :return: Mapping of grouping key to a queue identifier from ``queue_ids``. + """ + if not spans: + return {} + if not queue_ids: + raise SyntaxError( + f'{location} needs {kind} queues, but {architecture} has none that a program may use.') + + assigned: dict[str, int] = {} + for key in sorted(spans, key=lambda name: (spans[name][0], spans[name][1], name)): + start, end = spans[key] + used = { + assigned[other] + for other in assigned + if start <= spans[other][1] and spans[other][0] <= end + } + for queue in queue_ids: + if queue not in used: + assigned[key] = queue + break + else: + overlapping = sorted( + other for other, (other_start, other_end) in spans.items() + if other != key and start <= other_end and other_start <= end + ) + raise SyntaxError( + f'{location} would need {len(used) + 1} concurrent {kind} queues ' + f'(live groups {[key] + overlapping}), but a PE can use at most ' + f'{len(queue_ids)} on {architecture}.\n' + f' note: a fabric queue is remapped when a new color uses it, and the hardware ' + f'rejects that while wavelets remain\n' + f' note: a channel keeps one queue for the whole of its lifetime on the PE, ' + f'including gaps between epochs') + return assigned + + ### # Optimization passes ### diff --git a/tests/csl_runtime/test_batcher_oddeven_1d.sh b/tests/csl_runtime/test_batcher_oddeven_1d.sh index ecdcc108..091d5dbb 100644 --- a/tests/csl_runtime/test_batcher_oddeven_1d.sh +++ b/tests/csl_runtime/test_batcher_oddeven_1d.sh @@ -1,8 +1,11 @@ #!/bin/sh -# E2E test: 1D Batcher odd-even mergesort (2^L PEs, one f32 key per PE). -# Kernel: batcher_oddeven_1D.sptl params: L -# Reference: OUT_out[:, 0, 0] == sort(inp[:, 0, 0]) -# Tested with L ∈ {1, 2, 3}. +# E2E test: 1D Batcher odd-even mergesort (2^L PEs, a block of K f32 keys per PE). +# Kernel: batcher_oddeven_1D.sptl params: L, K +# Every comparator is a compare-split, so the network sorts all 2^L * K keys and PE i ends up with +# keys i*K .. i*K + K-1 of the sorted sequence. +# Reference: OUT_out.reshape(n*k) == sort(inp.reshape(n*k)) +# Tested with (L, K) ∈ {(1,1), (2,2), (3,1), (3,4)}: K = 1 is the one-key-per-PE network, and K = 4 +# is not the block size of any phase, which is what the K-element merge has to be independent of. set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" @@ -13,37 +16,39 @@ FOLDER="batcher_oddeven_1d_sptl" run_batcher() { l=$1 - echo "--- batcher_oddeven_1d L=$l ---" + k=$2 + echo "--- batcher_oddeven_1d L=$l K=$k ---" - sptlc "$SORT_DIR/batcher_oddeven_1D.sptl" "$FOLDER" -p L=$l + sptlc "$SORT_DIR/batcher_oddeven_1D.sptl" "$FOLDER" -p L=$l -p K=$k python3 - < 1 is what exercises that +# window on hardware. set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" @@ -15,38 +19,44 @@ FOLDER="batcher_oddeven_bundled_1d_sptl" run_batcher() { l=$1 - echo "--- batcher_oddeven_bundled_1d L=$l ---" + k=$2 + echo "--- batcher_oddeven_bundled_1d L=$l K=$k ---" - sptlc "$SORT_DIR/batcher_oddeven_bundled_1D.sptl" "$FOLDER" -p L=$l + sptlc "$SORT_DIR/batcher_oddeven_bundled_1D.sptl" "$FOLDER" -p L=$l -p K=$k python3 - < (stream[1, 1] readonly src, + stream[1, 1] writeonly dst) { + place i16 i, i16 j in [0, 0] { + f32[K] val + f32[K] res + } + compute i16 i, i16 j in [0, 0] { + await receive(res, src[i, j]) + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, dst[i, j]) + } +}""") + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, K=k)) + code = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel)}['code_0_0.csl'] + + # Every operand of every move is a DSD, whatever the arrays happen to be called. + moves = re.findall(r'@fmovs\((\w+), (\w+)[,)]', code) + assert moves, code + for destination, source in moves: + assert destination.endswith('_dsd') and source.endswith('_dsd'), code + + +def test_a_reused_color_keeps_its_input_queue_across_a_gap(): + """ + On the interior Batcher PE, one inbound color is used on both sides of another. Sharing the + queue across that gap is what the simulator rejects: remapping it onto the middle color while + wavelets of the outer color remain. The outer color must keep the queue for its whole span. + """ + path = os.path.join( + os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_1D.sptl' + ) + kernel = parser.parse_file(path) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, L=3, K=1)) + files = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_benchmarking=True)} + code = files['code_2_0.csl'] + + colors = dict(re.findall(r'const (\w+)_color_in: color = @get_color\((\d+)\);', code)) + queues = dict(re.findall( + r'const (\w+)_in_dsd = @get_dsd\(fabin_dsd, .*?input_queue = @get_input_queue\((\d+)\)', + code)) + # fwd__17 and fwd__37 share a color and straddle bwd__30. + assert colors['fwd__17'] == colors['fwd__37'] + assert colors['fwd__17'] != colors['bwd__30'] + assert queues['fwd__17'] == queues['fwd__37'] + assert queues['fwd__17'] != queues['bwd__30'] + + if __name__ == '__main__': test_dsd_op_detection() test_dsd_op_detection_constant_folding() diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py index 6b83ec34..f681e551 100644 --- a/tests/spatial_ir/test_shift_bundles.py +++ b/tests/spatial_ir/test_shift_bundles.py @@ -244,11 +244,11 @@ def entry(): cslrouting._check_filter_budget(entries) -def _bundled_batcher(l: int): +def _bundled_batcher(l: int, k: int = 1): path = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_bundled_1D.sptl') kernel = parser.parse_file(path) - kernel = passes.concretize_parameters(kernel, L=l) + kernel = passes.concretize_parameters(kernel, L=l, K=k) kernel = passes.constexpr_propagation(kernel) return lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) @@ -258,7 +258,7 @@ def _colors_of(layout: str, pattern: str = '') -> set[int]: return {int(color) for color in re.findall(r'@get_color\((\d+)\)[^;]*' + pattern, routes)} -@pytest.mark.parametrize('l, colors', [(2, 6), (3, 10), (4, 18)]) +@pytest.mark.parametrize('l, colors', [(2, 6), (3, 10)]) def test_the_batcher_fits_the_colors_it_has(l: int, colors: int): # A bundled phase puts all of its matchings on one color pair; the phases left unbundled take a # pair per matching, but share those across phases wherever their sources agree mod 2d. @@ -269,13 +269,25 @@ def test_the_batcher_fits_the_colors_it_has(l: int, colors: int): assert len(used) <= len(constants.COLORS) +def test_sixteen_keys_need_three_overlapping_input_queues(): + """ + At L = 4 a reused inbound color stays live across a gap that already holds two other colors. + WSE-2 has two input queues, so lowering must refuse rather than remap a busy queue. + """ + from spada.syntax.csl import constants + + if len(constants.INPUT_QUEUE_IDS) >= 3: + pytest.skip(f'{constants.ARCH} has {len(constants.INPUT_QUEUE_IDS)} input queues, enough for L=4') + with pytest.raises(SyntaxError, match='concurrent input queues'): + _bundled_batcher(4) + + def test_only_the_widest_batcher_phases_are_bundled(): # Bundling costs one filter at every participating PE and a PE has three, so the sample bundles - # the three widest phases only. Lowering at all is the check that no PE needs a fourth, since - # ``_check_filter_budget`` would refuse. + # every phase that satisfies 4d >= N. At L = 3 that is already three phases (d = 4, 2, 2). from spada.syntax.csl import constants - layout = next(f.code for f in _bundled_batcher(4) if 'layout' in f.filename) + layout = next(f.code for f in _bundled_batcher(3) if 'layout' in f.filename) used, filtered, switched = _colors_of(layout), _colors_of(layout, r'\.filter'), _colors_of(layout, r'\.switches') assert len(filtered) == 2 * constants.FILTERS_PER_PE # three phases, two directions each @@ -283,12 +295,29 @@ def test_only_the_widest_batcher_phases_are_bundled(): assert not (used - filtered) & switched # the pooled ones hold a single static configuration -def test_batcher_scalar_receive_lowers_to_data_task(): +def test_a_wider_block_costs_wavelets_not_colors(): + # The Batcher trades K keys per comparator instead of one. A bundle of M sources then carries + # M*K wavelets per epoch, of which each destination keeps the K its filter windows out -- the + # colors and the filters stay as they are, only the counters grow. + narrow = next(f.code for f in _bundled_batcher(3, 1) if 'layout' in f.filename) + wide = next(f.code for f in _bundled_batcher(3, 4) if 'layout' in f.filename) + + assert _colors_of(wide) == _colors_of(narrow) + assert _colors_of(wide, r'\.filter') == _colors_of(narrow, r'\.filter') + # At L=3 the bundled phases have two and four sources, so cycles of 8 and 16 words. + assert '.limit1 = 7, .max_counter = 3' in wide + assert '.limit1 = 15, .max_counter = 3' in wide + assert '.init_counter = (pe_x - 1) * 4' in wide + + +def test_a_scalar_receive_lowers_to_a_data_task(): + # A bundle destination that keeps a single key holds it in a scalar, which arrives as the + # argument of a data task rather than as a move out of the fabric. path = os.path.join( - os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_1D.sptl' + os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'simple', 'exchange_bundle_1D.sptl' ) kernel = parser.parse_file(path) - kernel = passes.concretize_parameters(kernel, L=1) + kernel = passes.concretize_parameters(kernel, M=2, D=2, R=1) kernel = passes.constexpr_propagation(kernel) files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) pe_codes = [f.code for f in files if 'code_' in f.filename] diff --git a/tests/spatial_ir/test_stream_lifetime.py b/tests/spatial_ir/test_stream_lifetime.py index 0b363c6e..6d41639f 100644 --- a/tests/spatial_ir/test_stream_lifetime.py +++ b/tests/spatial_ir/test_stream_lifetime.py @@ -458,5 +458,43 @@ def test_auto_channel_kernels_emit_no_closes(): assert all(not _closed_streams(rect.metadata.compute) for rect in rects) +### +# assign_fabric_queues +### + + +def test_sequential_spans_share_a_queue(): + """A queue is remapped only when the previous color's span on this PE has ended.""" + assigned = stream_lifetime.assign_fabric_queues( + {'channel 0': (0, 2), 'channel 1': (3, 5)}, [0, 1], + kind='input', architecture='wse2', location='PE (0, 0)') + assert assigned == {'channel 0': 0, 'channel 1': 0} + + +def test_overlapping_spans_take_distinct_queues(): + assigned = stream_lifetime.assign_fabric_queues( + {'channel 0': (0, 4), 'channel 1': (2, 6)}, [0, 1], + kind='input', architecture='wse2', location='PE (0, 0)') + assert assigned['channel 0'] != assigned['channel 1'] + + +def test_a_channel_keeps_its_queue_across_a_gap(): + """ + Wavelets of a reused color can still arrive between its epochs, so a different color that + sits in the gap cannot steal the queue. + """ + assigned = stream_lifetime.assign_fabric_queues( + {'channel 0': (0, 10), 'channel 1': (3, 5)}, [0, 1], + kind='input', architecture='wse2', location='PE (0, 0)') + assert assigned['channel 0'] != assigned['channel 1'] + + +def test_three_overlapping_spans_exhaust_two_queues(): + with pytest.raises(SyntaxError, match='concurrent input queues'): + stream_lifetime.assign_fabric_queues( + {'channel 0': (0, 10), 'channel 1': (2, 8), 'channel 2': (4, 6)}, [0, 1], + kind='input', architecture='wse2', location='PE (0, 0)') + + if __name__ == '__main__': pytest.main([__file__]) diff --git a/tests/spatial_ir/test_task_recycling_codegen.py b/tests/spatial_ir/test_task_recycling_codegen.py index 43ea2460..b0f149e4 100644 --- a/tests/spatial_ir/test_task_recycling_codegen.py +++ b/tests/spatial_ir/test_task_recycling_codegen.py @@ -11,6 +11,67 @@ _CSL_RUNTIME_TASK_RECYCLING_SAMPLES = os.path.join( os.path.dirname(__file__), '..', 'csl_runtime', 'samples') +# Two PEs trading a scalar back and forth for R phases, both directions pinned to a channel of +# their own so that every phase reuses them. Each receive is a data task on that channel, and each +# send a local task, so R controls how many of both a PE ends up with: past the local task IDs the +# hardware has, the slots start being recycled, and the receives of one channel have to share the +# one data task its color binds. Those are the two shapes the tests below check. +_SCALAR_EXCHANGE_CHAIN = """ +kernel @scalar_exchange_chain( + stream[2, 1] readonly inp, + stream[2, 1] writeonly out +) { + place i16 i, i16 j in [0:2, 0] { + f32 val + f32 tmp + } + phase { + compute i16 i, i16 j in [0:2, 0] { + await receive(val, inp[i, j]) + } + } + for i16 r in [0:R] { + phase { + dataflow i16 i, i16 j in [0:2, 0] { + stream fwd = relative_stream(1, 0) { + hops = auto, + channel = 0 + } + stream bwd = relative_stream(-1, 0) { + hops = auto, + channel = 1 + } + } + compute i16 i, i16 j in [0:1, 0] { + await send(val, fwd) + await receive(tmp, bwd) + val = tmp if tmp < val else val + } + compute i16 i, i16 j in [1:2, 0] { + await receive(tmp, fwd) + await send(val, bwd) + val = tmp if tmp > val else val + } + } + } + phase { + compute i16 i, i16 j in [0:2, 0] { + await send(val, out[i, j]) + } + } +} +""" + +# Twelve phases outrun the local task IDs of either generation, so the slots are recycled there. +_CHAIN_PHASES = 12 + + +def _scalar_exchange_chain(phases: int = _CHAIN_PHASES): + kernel = parser.parse_string(_SCALAR_EXCHANGE_CHAIN) + kernel = passes.concretize_parameters(kernel, R=phases) + kernel = passes.constexpr_propagation(kernel) + return lower_spatial_ir_to_csl(kernel) + def test_task_recycling_codegen_uses_else_if_dispatch_for_recycled_slots(): sample = os.path.join( @@ -61,16 +122,9 @@ def test_data_tasks_install_the_state_of_a_recycled_successor(): """A data task handing control to a recycled slot must install that slot's state first. Without the assignment the dispatcher runs whichever branch was installed last, which means a - PE silently skips its comparator and the fabric deadlocks behind the send it never made. The - bundled Batcher reaches that shape once it has enough phases sharing a colour, at L=4. + PE silently skips a phase of its own and the fabric deadlocks behind the send it never made. """ - sample = os.path.join( - os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_bundled_1D.sptl') - kernel = parser.parse_file(sample) - kernel = passes.concretize_parameters(kernel, L=4) - kernel = passes.constexpr_propagation(kernel) - - csl_files = lower_spatial_ir_to_csl(kernel) + csl_files = _scalar_exchange_chain() checked = 0 for file in csl_files: @@ -80,7 +134,7 @@ def test_data_tasks_install_the_state_of_a_recycled_successor(): hardware_ids.setdefault(hardware_id, []).append(task_index) recycled = {task for tasks in hardware_ids.values() if len(tasks) > 1 for task in tasks} - for body in re.findall(r'task dtask_\d+\([^)]*\) void \{(.*?)\n\}', code, re.S): + for body in re.findall(r'task dtask_(?:color_)?\d+\([^)]*\) void \{(.*?)\n\}', code, re.S): for match in re.finditer(r'@(?:activate|unblock)\(task_(\d+)_id\);', body): if match.group(1) not in recycled: continue @@ -91,24 +145,18 @@ def test_data_tasks_install_the_state_of_a_recycled_successor(): f'slot state assignment, but by "{preceding}"') checked += 1 - assert checked, 'sample no longer exercises a data task triggering a recycled local task' + assert checked, 'the chain no longer exercises a data task triggering a recycled local task' def test_a_reused_channel_binds_one_data_task_that_dispatches_on_its_epoch(): - """Two receives on one channel at one PE share the data task the channel binds. + """Several receives on one channel at one PE share the data task the channel binds. A data task's hardware ID is the color, so binding two of them is not merely wasteful - but rejected by cslc ("task ID '0' bound to more than one task"). The static Batcher - reaches that shape at L=3, where a PE compares against its neighbour in two phases that - the sample gives the same channel. + but rejected by cslc ("task ID '0' bound to more than one task"). Each PE of the chain + receives on the same channel in every one of its phases, so its receives all land in one + dispatcher that has to tell the epochs apart. """ - sample = os.path.join( - os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_1D.sptl') - kernel = parser.parse_file(sample) - kernel = passes.concretize_parameters(kernel, L=3) - kernel = passes.constexpr_propagation(kernel) - - csl_files = lower_spatial_ir_to_csl(kernel) + csl_files = _scalar_exchange_chain() shared = 0 for file in csl_files: @@ -136,7 +184,7 @@ def test_a_reused_channel_binds_one_data_task_that_dispatches_on_its_epoch(): assert f'@block(dtask_{task_index}_id);' in body, ( f'{file.filename}: dtask_{task_index} does not block color {color} when it is done') - assert shared, 'sample no longer reuses a channel for two receives at one PE' + assert shared, 'the chain no longer reuses a channel for several receives at one PE' def test_data_tasks_sharing_a_channel_must_take_turns(): From b0cefbcbc6bd68b0a355da22b98171b13256a7ad Mon Sep 17 00:00:00 2001 From: glukas Date: Tue, 18 Aug 2026 22:33:15 +0200 Subject: [PATCH 28/68] fix data task lowering for wse-3 --- spada/lowering/spatial_ir_to_csl.py | 68 ++++++++++++++++++- tests/spatial_ir/test_shift_bundles.py | 10 ++- .../spatial_ir/test_task_recycling_codegen.py | 40 ++++++----- 3 files changed, 98 insertions(+), 20 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 3fa78d22..d83a952c 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -434,8 +434,9 @@ def generate_rectangle(kernel: spir.Kernel, # Declare each data task ID. Data tasks that share a color are aliases of one hardware ID, and # a state variable selects which of them the shared task runs as. for slot in task_bindings.data_slots: + id_expr = _data_task_id_builtin(rect.metadata, slot, tasks, dsds) for task_index in slot.task_indices: - current_code.write(f'const dtask_{task_index}_id = @get_data_task_id(@get_color({slot.color}));\n') + current_code.write(f'const dtask_{task_index}_id = {id_expr};\n') if slot.recycled: representative = slot.representative_task_index current_code.write(f'var {task_bindings.data_state_var(representative)}: u16 = ' @@ -1439,7 +1440,10 @@ def _write_indented_block(current_code: StringIO, block: str, indent: str) -> No def _data_task_color(rect: PEBlock, task_index: int, task: tdag.CSLTask, color_map: dict[str, int]) -> int: """ - Returns the color a data task listens on, which is also its hardware task ID. + Returns the color a data task listens on. + + On WSE-2 that color is also the hardware task ID. On WSE-3 the ID is the + input queue bound to this color; see ``_data_task_id_builtin``. :param task_index: Only used to name the task in the error message. """ @@ -1455,6 +1459,66 @@ def _data_task_color(rect: PEBlock, task_index: int, task: tdag.CSLTask, color_m raise ValueError(f'Cannot find color for stream "{name_to_csl(sname)}" in data task {task_index}') +def _input_queue_for_data_slot( + rect: PEBlock, + slot: task_recycling.DataTaskSlot, + tasks: list[tdag.CSLTask], + dsds: UniqueDSDDict, +) -> int: + """Return the fabric input queue the receives in ``slot`` share. + + On WSE-3 a data task's hardware ID is that queue, which + ``_declare_queue_initialization`` has already bound to the slot's color. + + :param rect: The PE block being generated. + :param slot: The data-task slot whose color the receives listen on. + :param tasks: All tasks of this PE. + :param dsds: Fabric descriptors of this PE, which record the queue assignment. + :return: The input-queue identifier. + """ + queues: set[int] = set() + for task_index in slot.task_indices: + task = tasks[task_index] + stmt = rect.compute.statements[task.statements[0]] + assert isinstance(stmt, spir.ForeachStatement) + sname = stmt.receive_stream.stream_name + if isinstance(sname, spir.ArraySlice): + sname = sname.array + for _, dsd in dsds.get(sname.as_ir(), []): + if isinstance(dsd, cslstruct.FabricDSD) and dsd.dsd_type == cslstruct.DSDType.fabin: + queues.add(dsd.queue) + if len(queues) != 1: + found = sorted(queues) if queues else 'none' + raise SyntaxError( + f'WSE-3 data task on color {slot.color} needs exactly one input queue, found {found}.\n' + " note: @get_data_task_id takes an input_queue on WSE-3, not a color") + return next(iter(queues)) + + +def _data_task_id_builtin( + rect: PEBlock, + slot: task_recycling.DataTaskSlot, + tasks: list[tdag.CSLTask], + dsds: UniqueDSDDict, +) -> str: + """Return the ``@get_data_task_id(...)`` expression for ``slot``. + + WSE-2 constructs a data-task ID from the color the receive listens on. + WSE-3 constructs it from the input queue already bound to that color; + passing the color is rejected as ``expected 'input_queue' expression, got: 'color'``. + + :param rect: The PE block being generated. + :param slot: The data-task slot, whose color is the receive's fabric color. + :param tasks: All tasks of this PE, indexed as in the slot. + :param dsds: Fabric descriptors of this PE, which record the queue assignment. + :return: A CSL expression of type ``data_task_id``. + """ + if csl.ARCH == 'wse3': + queue = _input_queue_for_data_slot(rect, slot, tasks, dsds) + return f'@get_data_task_id(@get_input_queue({queue}))' + return f'@get_data_task_id(@get_color({slot.color}))' + + def _generate_data_task_slot( rect: PEBlock, slot: task_recycling.DataTaskSlot, diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py index f681e551..4e3e04c9 100644 --- a/tests/spatial_ir/test_shift_bundles.py +++ b/tests/spatial_ir/test_shift_bundles.py @@ -11,7 +11,7 @@ import pytest from spada.lowering.spatial_ir_to_csl import canonicalize_kernel, lower_spatial_ir_to_csl -from spada.syntax.csl import routing as cslrouting +from spada.syntax.csl import constants, routing as cslrouting from spada.syntax.spatial_ir import canonicalization, parser, passes from spada.syntax.spatial_ir.shift_bundles import detect_shift_bundles @@ -330,3 +330,11 @@ def test_a_scalar_receive_lowers_to_a_data_task(): pe0 = next(f.code for f in files if 'code_0_0' in f.filename) assert 'task dtask_' in pe0 assert 'tmp = __x' in pe0 + if constants.ARCH == 'wse3': + assert re.search(r'@get_data_task_id\(@get_input_queue\(\d+\)\)', pe0), pe0 + assert '@get_data_task_id(@get_color(' not in pe0 + queues = re.findall(r'@get_data_task_id\(@get_input_queue\((\d+)\)\)', pe0) + for queue in queues: + assert f'@initialize_queue(@get_input_queue({queue}),' in pe0, pe0 + else: + assert re.search(r'@get_data_task_id\(@get_color\(\d+\)\)', pe0), pe0 diff --git a/tests/spatial_ir/test_task_recycling_codegen.py b/tests/spatial_ir/test_task_recycling_codegen.py index b0f149e4..2546a38b 100644 --- a/tests/spatial_ir/test_task_recycling_codegen.py +++ b/tests/spatial_ir/test_task_recycling_codegen.py @@ -4,7 +4,7 @@ import pytest from spada.lowering.spatial_ir_to_csl import lower_spatial_ir_to_csl -from spada.syntax.csl import task_recycling +from spada.syntax.csl import constants, task_recycling from spada.syntax.csl import tasks as tdag from spada.syntax.spatial_ir import parser, passes @@ -151,38 +151,44 @@ def test_data_tasks_install_the_state_of_a_recycled_successor(): def test_a_reused_channel_binds_one_data_task_that_dispatches_on_its_epoch(): """Several receives on one channel at one PE share the data task the channel binds. - A data task's hardware ID is the color, so binding two of them is not merely wasteful - but rejected by cslc ("task ID '0' bound to more than one task"). Each PE of the chain - receives on the same channel in every one of its phases, so its receives all land in one - dispatcher that has to tell the epochs apart. + A data task's hardware ID is the color on WSE-2 and the input queue on WSE-3, so binding + two of them is not merely wasteful but rejected by cslc ("task ID '0' bound to more than + one task"). Each PE of the chain receives on the same channel in every one of its phases, + so its receives all land in one dispatcher that has to tell the epochs apart. """ csl_files = _scalar_exchange_chain() shared = 0 for file in csl_files: code = file.code - colors: dict[str, list[str]] = {} - for task_index, color in re.findall(r'const dtask_(\d+)_id = @get_data_task_id\(@get_color\((\d+)\)\)', code): - colors.setdefault(color, []).append(task_index) + builtin = r'@get_input_queue' if constants.ARCH == 'wse3' else r'@get_color' + hardware_ids: dict[str, list[str]] = {} + for task_index, hw in re.findall( + rf'const dtask_(\d+)_id = @get_data_task_id\({builtin}\((\d+)\)\)', code): + hardware_ids.setdefault(hw, []).append(task_index) bound = re.findall(r'@bind_data_task\(\w+, dtask_(\d+)_id\);', code) - assert len(bound) == len(colors), ( - f'{file.filename}: binds {len(bound)} data tasks for {len(colors)} colors') + assert len(bound) == len(hardware_ids), ( + f'{file.filename}: binds {len(bound)} data tasks for {len(hardware_ids)} hardware IDs') - for color, task_indices in colors.items(): + dispatchers = re.findall(r'task dtask_color_(\d+)\([^)]*\) void \{(.*?)\n\}', code, re.S) + for hw, task_indices in hardware_ids.items(): if len(task_indices) == 1: continue shared += 1 - dispatcher = re.search(rf'task dtask_color_{color}\([^)]*\) void \{{(.*?)\n\}}', code, re.S) - assert dispatcher, f'{file.filename}: color {color} is reused but has no dispatcher' - body = dispatcher.group(1) - states = re.findall(r'(?:else )?if \(__dtask_color_' + color + r'_state == (\d+)\)', body) + body = next( + (b for _, b in dispatchers if all(f'@block(dtask_{t}_id);' in b for t in task_indices)), + None) + assert body, ( + f'{file.filename}: hardware ID {hw} is reused but has no dispatcher that blocks ' + f'{task_indices}') + states = re.findall(r'(?:else )?if \(__dtask_color_\d+_state == (\d+)\)', body) assert states == [str(state) for state in range(len(task_indices))], ( - f'{file.filename}: color {color} dispatches on {states} for {len(task_indices)} receives') + f'{file.filename}: hardware ID {hw} dispatches on {states} for {len(task_indices)} receives') # A branch that keeps its color live would take the next epoch's wavelets as its own. for task_index in task_indices: assert f'@block(dtask_{task_index}_id);' in body, ( - f'{file.filename}: dtask_{task_index} does not block color {color} when it is done') + f'{file.filename}: dtask_{task_index} does not block when it is done') assert shared, 'the chain no longer reuses a channel for several receives at one PE' From 58d49c1ef52bd7b7e349a98f12bb57a3c24a6301 Mon Sep 17 00:00:00 2001 From: glukas Date: Tue, 18 Aug 2026 23:06:03 +0200 Subject: [PATCH 29/68] Fix WSE-3 task recycling --- spada/lowering/spatial_ir_to_csl.py | 36 ++++++++++++++++--- spada/syntax/csl/constants.py | 21 ++++++++--- tests/spatial_ir/test_task_recycling.py | 18 ++++++++++ .../spatial_ir/test_task_recycling_codegen.py | 25 +++++++++++-- 4 files changed, 90 insertions(+), 10 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index d83a952c..abb9e8da 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -410,7 +410,13 @@ def generate_rectangle(kernel: spir.Kernel, i: _data_task_color(rect.metadata, i, task, color_map) for i, task in enumerate(tasks) if task.task_type == 'data' } - task_bindings = task_recycling.plan_task_bindings(tasks, task_creation_behavior, set(color_map.values()), + # On WSE-2 a data-task ID *is* its color, so a local task must not reuse one. + # On WSE-3 data-task IDs are input queues 0–7; colors and local tasks do not + # share a namespace. memcpy's local tasks are reserved on both generations. + disallowed_task_ids = set(csl.RESERVED_LOCAL_TASK_IDS) + if csl.ARCH != 'wse3': + disallowed_task_ids |= set(color_map.values()) + task_bindings = task_recycling.plan_task_bindings(tasks, task_creation_behavior, disallowed_task_ids, data_task_colors) place_block_bytes = _place_block_storage_bytes(rect.metadata.place) @@ -520,13 +526,13 @@ def generate_rectangle(kernel: spir.Kernel, if tasks[representative].blocked: footer.write(f' @block(dtask_{representative}_id);\n') - max_task_id = max((slot.hardware_task_id for slot in task_bindings.local_slots), default=csl.LOCAL_TASK_IDS[0] - 1) - # Create exit task that unblocks command stream exit_task_sequential = all(typ == tdag.InterTaskEdge.SEQUENCE for t in tasks for n, typ in t.outgoing if n == -1) exit_task_sequential &= not any( t.task_type == 'data' for t in tasks for n, _ in t.outgoing if n == -1) # No data tasks exit_task_blocked = any(n == -1 and typ == tdag.InterTaskEdge.UNBLOCK for t in tasks for n, typ in t.outgoing) + hardware_exit_id = None if exit_task_sequential else _exit_task_hardware_id( + {slot.hardware_task_id for slot in task_bindings.local_slots}, set(color_map.values())) # Bind exit task if not exit_task_sequential: @@ -607,7 +613,7 @@ def generate_rectangle(kernel: spir.Kernel, if not exit_task_sequential: current_code.write(f''' -const exit_task_id = @get_local_task_id({max_task_id + 1}); +const exit_task_id = @get_local_task_id({hardware_exit_id}); task exit_task() void {{ {benchmark_code.kernel_postamble} // On completion, unblock command stream @@ -1438,6 +1444,28 @@ def _write_indented_block(current_code: StringIO, block: str, indent: str) -> No current_code.write(f'{indent}{line}\n') +def _exit_task_hardware_id(used_ids: set[int], color_ids: set[int]) -> int: + """Return a local-task ID for ``exit_task`` that nothing else has bound. + + Walks the activatable range from 8 and skips IDs already taken by program + slots, by memcpy/system reservations, and on WSE-2 by colors (which are also + data-task IDs there). + + :param used_ids: Hardware IDs already assigned to local-task slots. + :param color_ids: Colors allocated to this PE. + :return: A free activatable identifier. + """ + occupied = set(used_ids) | set(csl.RESERVED_LOCAL_TASK_IDS) + if csl.ARCH != 'wse3': + occupied |= set(color_ids) + for tid in range(8, 31): + if tid not in occupied: + return tid + raise SyntaxError( + 'No free local task ID remains for exit_task ' + f'(occupied {sorted(occupied)}).') + + def _data_task_color(rect: PEBlock, task_index: int, task: tdag.CSLTask, color_map: dict[str, int]) -> int: """ Returns the color a data task listens on. diff --git a/spada/syntax/csl/constants.py b/spada/syntax/csl/constants.py index 92566608..6274203e 100644 --- a/spada/syntax/csl/constants.py +++ b/spada/syntax/csl/constants.py @@ -6,12 +6,25 @@ # Cerebras architecture to use. Options: 'wse2', 'wse3' ARCH = os.environ.get('WSE_ARCH', 'wse2') -# From the SDK: IDs 29 and 30 should generally be avoided in programs as they are used for system tasks. -# https://sdk.cerebras.net/csl/language/task-ids?highlight=color#activatable-identifiers -# NOTE: We also avoid task ID 28 as we reserve it for ``exit_task`` +# Activatable local-task IDs: 0–30 on WSE-2, 8–30 on WSE-3. 29 is the teardown +# handler and 30 is the timer; memcpy also binds several of these as local tasks +# (``sys_params.csl``: SYS_EN_MAIN=24, SYS_UBLK_C22=27, SYS_EXIT=28, +# SYS_SEND_CTRL=30). On WSE-3, ``memcpyd2h.csl`` additionally aliases color 21 as +# ``LOCAL_MEMCPYD2H_DATA`` to save an entrypoint, which is the collision cslc +# reports as "task ID '21' bound to more than one task". +# https://sdk.cerebras.ai/csl/language/task-ids +# https://sdk.cerebras.ai/tensor-streaming +_RESERVED_LOCAL_TASK_IDS = { + 'wse2': [24, 27, 28, 29, 30], + 'wse3': [21, 24, 27, 28, 29, 30], +} +RESERVED_LOCAL_TASK_IDS = _RESERVED_LOCAL_TASK_IDS[ARCH] + +# Program-assignable local-task IDs. WSE-3 skips the memcpy holes; ``exit_task`` +# is not in this list and takes the next free ID after the assigned slots. _CSL_LOCAL_TASK_IDS = { 'wse2': list(range(8, 21)), - 'wse3': list(range(8, 28)), + 'wse3': [t for t in range(8, 26) if t not in _RESERVED_LOCAL_TASK_IDS['wse3']], } LOCAL_TASK_IDS = _CSL_LOCAL_TASK_IDS[ARCH] diff --git a/tests/spatial_ir/test_task_recycling.py b/tests/spatial_ir/test_task_recycling.py index 9b026d8b..32c4f0a6 100644 --- a/tests/spatial_ir/test_task_recycling.py +++ b/tests/spatial_ir/test_task_recycling.py @@ -255,3 +255,21 @@ def test_plan_is_deterministic(): plan2 = task_recycling.plan_task_bindings(tasks, tdag.TaskCreationBehavior.STATE_MACHINE_ON_OVERRUN) assert plan1.task_to_local_slot == plan2.task_to_local_slot assert plan1.task_to_local_state == plan2.task_to_local_state + + +def test_local_task_ids_do_not_include_memcpy_reservations(): + """The assignable pool must not contain IDs memcpy already binds.""" + assert set(constants.LOCAL_TASK_IDS).isdisjoint(constants.RESERVED_LOCAL_TASK_IDS) + assert 21 not in constants._CSL_LOCAL_TASK_IDS['wse3'] + assert set(constants._CSL_LOCAL_TASK_IDS['wse3']).isdisjoint(constants._RESERVED_LOCAL_TASK_IDS['wse3']) + assert set(constants._CSL_LOCAL_TASK_IDS['wse2']).isdisjoint(constants._RESERVED_LOCAL_TASK_IDS['wse2']) + + +def test_exit_task_skips_the_first_memcpy_reservation(): + """If every ID below memcpy's first local task is taken, exit_task must hop the hole.""" + first_reserved = min(constants.RESERVED_LOCAL_TASK_IDS) + used = set(range(8, first_reserved)) + exit_id = s2c._exit_task_hardware_id(used, set()) + assert exit_id not in used + assert exit_id not in constants.RESERVED_LOCAL_TASK_IDS + assert exit_id == first_reserved + 1 diff --git a/tests/spatial_ir/test_task_recycling_codegen.py b/tests/spatial_ir/test_task_recycling_codegen.py index 2546a38b..78385bba 100644 --- a/tests/spatial_ir/test_task_recycling_codegen.py +++ b/tests/spatial_ir/test_task_recycling_codegen.py @@ -227,6 +227,27 @@ def test_codegen_avoids_local_task_id_color_overlap(): assert 8 in colors, 'sample should force color 8 to be allocated' assert local_task_ids - assert local_task_ids.isdisjoint(colors), ( - f'local task IDs overlap communication colors: ids={sorted(local_task_ids)}, colors={sorted(colors)}' + assert local_task_ids.isdisjoint(constants.RESERVED_LOCAL_TASK_IDS), ( + f'local task IDs overlap memcpy reservations: ids={sorted(local_task_ids)}' + ) + if constants.ARCH != 'wse3': + # On WSE-2 a data-task ID is its color, so the two sets must be disjoint. + assert local_task_ids.isdisjoint(colors), ( + f'local task IDs overlap communication colors: ids={sorted(local_task_ids)}, ' + f'colors={sorted(colors)}' + ) + + +def test_csl_runtime_task_recycling_sample_avoids_memcpy_local_task_ids(): + """The merge sample's 14 local tasks used to land on memcpy's ID 21 on WSE-3.""" + path = os.path.join(_CSL_RUNTIME_TASK_RECYCLING_SAMPLES, 'task_recycling_merge.sptl') + kernel = parser.parse_file(path) + kernel = passes.constexpr_propagation(kernel) + csl_files = lower_spatial_ir_to_csl( + kernel, task_fusion=False, copy_elision=True, prune_memory=True) + combined = '\n'.join(f.code for f in csl_files) + local_task_ids = {int(v) for v in re.findall(r'@get_local_task_id\((\d+)\)', combined)} + assert local_task_ids + assert local_task_ids.isdisjoint(constants.RESERVED_LOCAL_TASK_IDS), ( + f'generated local task IDs overlap memcpy: {sorted(local_task_ids)}' ) From de357f5f5b2ae6dd3f578eb9ddffd3ac987ebeae Mon Sep 17 00:00:00 2001 From: glukas Date: Tue, 18 Aug 2026 23:28:56 +0200 Subject: [PATCH 30/68] update documentation --- irspec/docs/spatial/routing.md | 188 +----------------- irspec/docs/spatial/routing_wse.md | 115 ++++++++++- .../sort/batcher_oddeven_bundled_1D.sptl | 2 +- spada/syntax/spatial_ir/shift_bundles.py | 2 +- 4 files changed, 117 insertions(+), 190 deletions(-) diff --git a/irspec/docs/spatial/routing.md b/irspec/docs/spatial/routing.md index 734992a8..ee079782 100644 --- a/irspec/docs/spatial/routing.md +++ b/irspec/docs/spatial/routing.md @@ -193,187 +193,7 @@ to receive. especially when considering multiple phases. -## Lowering to Switches - -!!! note "Note: Scope" - Everything in this section describes how the **Cerebras WSE / CSL backend** realizes epochs, not - the semantics of the Spatial IR itself. Colors, routers, switch positions and control messages - are properties of that target; a different backend may implement channel reuse by any means that - preserves the correctness conditions of [Undefined Behavior](#undefined-behavior) above. Where - the two WSE generations differ, the text says so; the compiler selects between them on - `WSE_ARCH`. - -Channels are a scarce resource: each channel that is live at a PE occupies one of the hardware's -routing colors. Epochs are what makes it possible to reuse a channel, and hence a color, for -several streams. - -A stream induces, at each PE of its path, a *route configuration*: the set of directions the PE -receives from and the set of directions it transmits to (where `RAMP` denotes the PE's own compute -element). Consider a fixed channel $C$ and a fixed PE $(i, j)$. Ordering the streams that use $C$ -at $(i, j)$ by their epochs yields a sequence of route configurations -$R_0, R_1, \dotsc, R_{n-1}$, which is realized by the PE's *switch* for the color assigned to $C$: -$R_0$ is the initial configuration and the router *advances* to $R_{k+1}$ at the epoch boundary. - -Whichever side a switch position leaves unspecified keeps the value it currently has, so positions -compose incrementally. - -!!! warning "WSE-2: A Switch Position Carries One Direction" - On WSE-2 a switch position records *either* the input the router receives from or the output it - transmits to, never both — `cslc` rejects a position naming both with *"cannot have both an - input and an output in the same switch position"*. A transition that changes both sides — a PE - that stops receiving on a channel and starts sending on it, as in a systolic chain — therefore - occupies **two** positions, passing through an intermediate configuration that keeps the old - input and takes the new output. The intermediate is a pure relay, occupied only between the two - advances that retire the configuration, and it must keep the old input so that the second - advance still reaches the router. - - WSE-3 accepts both directions in one position, so the same transition costs one position and one - advance there. - -!!! danger "Error: Too Many Route Configurations" - A router holds a bounded number of switch positions per color (four on both WSE-2 and WSE-3). - *If the streams sharing a channel require more positions than that at a single PE, a compile - error is raised.* Assigning a different channel to some of the streams resolves it, at the cost - of an additional color. On WSE-2 a configuration which changes both the input and the output - direction costs two positions, so four positions is fewer than four turnarounds there. - -When the sequence of configurations at a router is periodic — a halo exchange that alternates -between sending and receiving across phases produces $R_0, R_1, R_0, R_1$ — only one period is -stored and the switch wraps around from the last position back to the base one (`ring_mode`). - -Consecutive configurations that are equal do not consume a position and do not require an advance. -This is a common case: two streams declared as `relative_stream(-2, 0)` in successive phases induce -the same configuration at every PE of their paths, so their shared channel needs no switching at -all. - -An advance is driven by the `close` that ends the epoch. The sending PE emits a *switch-advance -control message* on the channel, one per position to be traversed. It follows the stream's path -using the configuration that is being retired, and advances the router of each PE it traverses, -after all data of the epoch. - -!!! warning "WSE: A Control Message's Payload Selects Nothing" - A CSL control wavelet nominally carries up to eight per-router switching commands - (``'s `MAX_CMDS`), which would let one message advance some routers on a path and leave - others alone. **On the WSE hardware the payload does not work that way.** Measured on the - simulator, only command slot 0 is ever executed, and **every** switch-configured router the - wavelet reaches applies it; slots 1–7 had no effect in any topology tested — the sender's own - router, one hop, two hops through a plain relay, and two switch-configured routers in sequence. - A generated kernel built on the opposite assumption, advancing the fourth router of a path with - an `[ADV, NOP, NOP, ADV]` chain, stalled in the fabric. The compiler therefore emits - `encode_single_payload`, which writes slot 0 only. - - The consequence is that a message cannot advance one router while leaving another on the same - path where it is. *If the routers along one path would have to advance by different amounts, a - compile error is raised.* A router that is already on its last position is exempt: outside - `ring_mode` an advance past the last position is a no-op, so a message passing through may - over-advance it harmlessly. (Such a router is not necessarily finished — it may keep relaying the - same configuration for the rest of the kernel, which is exactly what the bundle below relies - on.) - - This is a property of the *payload*, not of switching: a router can still be switched at a time - only it knows, and delivery to a compute element can still be made selective, by the two - mechanisms the next section combines. - -Because the control message travels the path of the retired configuration in order behind the data, -a receiving PE needs to emit nothing to advance its own router: the ordering required by the -[lemma above](#undefined-behavior) is provided by the fabric. A receiver's `close` therefore has no -runtime effect; it exists so that the lifetime of the stream — and hence the number of elements it -carries — is stated by every participant and can be checked. - -## Overlapping Interval Shifts - -The correctness conditions above rule out one shape that occurs constantly: a run of consecutive PEs -all shifting the same distance $d$ along an axis, as in - -``` -dataflow i16 i, i16 j in [0:D + M, 0] { - stream fwd = relative_stream(D, 0) { hops = auto, channel = 0 } -} -``` - -with the PEs in `[0:M)` sending and those in `[D:D+M)` receiving. Source $p$'s word passes through -the routers of sources $p+1, \dotsc, M-1$, so the paths share PEs within one epoch. Written as one -stream per source, that is a channel each, $M$ colors for a shift; the alternative is a chain of -single-hop stores and forwards, which serializes the whole run behind $d$ hops of copying. - -Neither is necessary. The compiler recognizes this pattern — `detect_shift_bundles` — and lowers the -whole run onto **one channel**, giving the routers configurations that are switched only by events -the PE owning them knows locally: - -``` -PE: 0 1 2 3 4 5 (M = 3, D = 3) -role: src0 src1 src2 dst0 dst1 dst2 -sends: 3rd 2nd 1st -- -- -- -routes: R->E R->E R->E W->{R,E} W->{R,E} W->R -pos1: -- W->E W->E -- -- -- -filter: -- -- -- win 2 win 1 win 0 -``` - -**Sources inject, then relay.** A source starts at `rx = RAMP, tx = {EAST}` with `pos1` taking -`rx = WEST`, sends its own words, and its `close` advances its own router into relay mode. The -trigger is local — *"my own send is done"* — which is what makes it expressible at all, given that -the payload of the resulting message [selects nothing](#lowering-to-switches). - -**The order is descending, and enforces itself.** The source nearest the destinations goes first. No -schedule or barrier is needed: a source further away cannot push a word through its neighbour's -router while that neighbour is still injecting from its ramp, so it waits on the link. Backpressure -serializes the run in exactly the order the switches expect. - -**Destinations do not switch; a filter picks their words.** Each destination is statically routed to -`tx = {RAMP, EAST}`, which *duplicates* rather than consumes: every destination's router sees the -entire stream, in one order, and the one the stream reaches last uses `tx = {RAMP}` to take it out of -the network. Which words a destination hands to its compute element is decided by a counter filter -on that color, one linear function of the PE coordinate, so a single `@set_color_config` covers the -whole run. A control message passes such a router without being counted (`count_data = true`) and -without being filtered. - -!!! note "Note: Counter Filter Arithmetic" - As measured on the simulator, a counter filter starts at `init_counter`, increments on every data - wavelet, wraps to zero after `limit1`, and hands a wavelet to the compute element iff the counter - is at most `max_counter`. A window of `words` out of a stream of `length * words` is therefore - `limit1 = length * words - 1`, `max_counter = words - 1`, and an `init_counter` chosen so that - the counter reads zero as the wanted block arrives. - `tests/csl_runtime/test_shift_bundle_filters.sh` is the hand-written layout this was measured - with. - -!!! danger "Error: Too Many Wavelet Filters" - WSE-2 has four filters per PE, of which the `memcpy` module reserves one, so **three** are usable - (`FILTERS_PER_PE`). A PE needs one per color it filters. *If a PE would need more, a compile error - is raised*; the fix is to give some of the streams their own channels, which trades filters for - colors. Three filters therefore means at most three bundled phases per PE, whatever the kernel: - `batcher_oddeven_bundled_1D.sptl` would want ten at $2^4$ PEs and bundles only its three widest - phases, which is where most of the colors are saved anyway. - - Reconfiguring a filter while wavelets are still in flight on its color is a data race, and the - destination that terminates the stream cannot be reconfigured until the stream has drained, - because its router is what removes the wavelets from the network. Filters are consequently set up - once, at layout time, and never reused between phases. - -Bundling applies only when every run it decomposes into has at least two sources and is no longer -than the shift distance, so that no PE is both a source and a destination; a shift of one PE is left -alone, since a chain at distance one is already sequenced by ordinary switch positions. Anything -else falls back to the per-hop lowering, and to the errors above if that conflicts. - -Which shifts are bundled is decided by the channel assignment rather than by an attribute: a bundle -is what several overlapping matchings on *one* channel become, so giving each matching a channel of -its own is how a kernel declines the trade. What it then costs is colors, and those can be won back -by reusing a channel across phases. Two unbundled shifts may share one safely when they agree on -axis and signed distance and their sources agree modulo twice that distance, because a PE's role — -source, relay or destination — is then a function of its position modulo twice the distance alone, -so one static configuration serves every phase in the pool. Sharing on any other basis risks a PE -that sends on the color in one phase and receives on it in another, which needs a two-sided switch -change that a sender cannot drive (see *Lowering to Switches*), and nothing in the compiler -currently rejects it. `batcher_oddeven_bundled_1D.sptl` pools on exactly this rule, and -`batcher_oddeven_1D.sptl` is the same rule written out as arithmetic. - -This arrangement is the one used in Schnyder's *Distributed Sorting on the Cerebras Wafer-Scale -Engine* (fig. 7.6) for the 2D reduce-scatter, and is known to run on WSE-2. - -!!! note "Note: Multiple Rounds on One Color" - Two mechanisms are deliberately left unused, and are what to reach for if the four switch - positions or three filters run out. `SWITCH_RST` restores the initial configuration of every - router a message passes, which retires a whole path with one wavelet. Teardown-based - reconfiguration reprograms the routers between rounds outright, which is the only known way to - put an arbitrary *sequence* of sends and receives on one color: for a long enough sequence there - is a PE for which no fixed cycle of switch positions exists. Until then, splitting the rounds - across channels — as `bitonic_sort_1D.sptl` does — remains the per-kernel fallback. +## Lowering to Cerebras WSE + +How epochs, switch positions, control wavelets, shift bundling, and counter filters are realized on +the Cerebras Wafer-Scale Engine is described in [Routing Semantics on Cerebras WSE](routing_wse.md). diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index 7c6e7fcc..53f479a3 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -76,13 +76,22 @@ after all data of the epoch. others alone. **On the WSE hardware it does not work that way.** Measured on the simulator, only command slot 0 is ever executed, and **every** switch-configured router the wavelet reaches applies it; slots 1–7 had no effect in any topology tested — the sender's own router, one hop, - two hops through a plain relay, and two switch-configured routers in sequence. The compiler - therefore emits `encode_single_payload`, which writes slot 0 only. + two hops through a plain relay, and two switch-configured routers in sequence. A generated kernel + built on the opposite assumption, advancing the fourth router of a path with an `[ADV, NOP, NOP, + ADV]` chain, stalled in the fabric. The compiler therefore emits `encode_single_payload`, which + writes slot 0 only. The consequence is that a message cannot advance one router while leaving another on the same path where it is. *If the routers along one path would have to advance by different amounts, a - compile error is raised.* A router that is already on its last configuration is exempt: it never - routes anything again, so a message passing through may over-advance it harmlessly. + compile error is raised.* A router that is already on its last position is exempt: outside + `ring_mode` an advance past the last position is a no-op, so a message passing through may + over-advance it harmlessly. (Such a router is not necessarily finished — it may keep relaying the + same configuration for the rest of the kernel, which is exactly what the bundle below relies + on.) + + This is a property of the *payload*, not of switching: a router can still be switched at a time + only it knows, and delivery to a compute element can still be made selective, by the two + mechanisms the next section combines. !!! danger "WSE-2: A Two-Advance Turnaround Overshoots the Receiver" A control message stops at the first router whose current output is `RAMP`, and it is routed by @@ -108,3 +117,101 @@ a receiving PE needs to emit nothing to advance its own router: the ordering req [lemma](../routing#undefined-behavior) in the IR semantics is provided by the fabric. A receiver's `close` therefore has no runtime effect; it exists so that the lifetime of the stream — and hence the number of elements it carries — is stated by every participant and can be checked. + +## Overlapping Interval Shifts + +The [correctness conditions](../routing#undefined-behavior) rule out one shape that occurs constantly: +a run of consecutive PEs all shifting the same distance $d$ along an axis, as in + +``` +dataflow i16 i, i16 j in [0:D + M, 0] { + stream fwd = relative_stream(D, 0) { hops = auto, channel = 0 } +} +``` + +with the PEs in `[0:M)` sending and those in `[D:D+M)` receiving. Source $p$'s word passes through +the routers of sources $p+1, \dotsc, M-1$, so the paths share PEs within one epoch. Written as one +stream per source, that is a channel each, $M$ colors for a shift; the alternative is a chain of +single-hop stores and forwards, which serializes the whole run behind $d$ hops of copying. + +Neither is necessary. The compiler recognizes this pattern — `detect_shift_bundles` — and lowers the +whole run onto **one channel**, giving the routers configurations that are switched only by events +the PE owning them knows locally: + +``` +PE: 0 1 2 3 4 5 (M = 3, D = 3) +role: src0 src1 src2 dst0 dst1 dst2 +sends: 3rd 2nd 1st -- -- -- +routes: R->E R->E R->E W->{R,E} W->{R,E} W->R +pos1: -- W->E W->E -- -- -- +filter: -- -- -- win 2 win 1 win 0 +``` + +**Sources inject, then relay.** A source starts at `rx = RAMP, tx = {EAST}` with `pos1` taking +`rx = WEST`, sends its own words, and its `close` advances its own router into relay mode. The +trigger is local — *"my own send is done"* — which is what makes it expressible at all, given that +the payload of the resulting message [selects nothing](#lowering-to-switches). + +**The order is descending, and enforces itself.** The source nearest the destinations goes first. No +schedule or barrier is needed: a source further away cannot push a word through its neighbour's +router while that neighbour is still injecting from its ramp, so it waits on the link. Backpressure +serializes the run in exactly the order the switches expect. + +**Destinations do not switch; a filter picks their words.** Each destination is statically routed to +`tx = {RAMP, EAST}`, which *duplicates* rather than consumes: every destination's router sees the +entire stream, in one order, and the one the stream reaches last uses `tx = {RAMP}` to take it out of +the network. Which words a destination hands to its compute element is decided by a counter filter +on that color, one linear function of the PE coordinate, so a single `@set_color_config` covers the +whole run. A control message passes such a router without being counted (`count_data = true`) and +without being filtered. + +!!! note "Note: Counter Filter Arithmetic" + As measured on the simulator, a counter filter starts at `init_counter`, increments on every data + wavelet, wraps to zero after `limit1`, and hands a wavelet to the compute element iff the counter + is at most `max_counter`. A window of `words` out of a stream of `length * words` is therefore + `limit1 = length * words - 1`, `max_counter = words - 1`, and an `init_counter` chosen so that + the counter reads zero as the wanted block arrives. + `tests/csl_runtime/test_shift_bundle_filters.sh` is the hand-written layout this was measured + with. + +!!! danger "Error: Too Many Wavelet Filters" + WSE-2 has four filters per PE, of which the `memcpy` module reserves one, so **three** are usable + (`FILTERS_PER_PE`). A PE needs one per color it filters. *If a PE would need more, a compile error + is raised*; the fix is to give some of the streams their own channels, which trades filters for + colors. Three filters therefore means at most three bundled phases per PE, whatever the kernel: + `batcher_oddeven_bundled_1D.sptl` would want ten at $2^4$ PEs and bundles only its three widest + phases, which is where most of the colors are saved anyway. + + Reconfiguring a filter while wavelets are still in flight on its color is a data race, and the + destination that terminates the stream cannot be reconfigured until the stream has drained, + because its router is what removes the wavelets from the network. Filters are consequently set up + once, at layout time, and never reused between phases. + +Bundling applies only when every run it decomposes into has at least two sources and is no longer +than the shift distance, so that no PE is both a source and a destination; a shift of one PE is left +alone, since a chain at distance one is already sequenced by ordinary switch positions. Anything +else falls back to the per-hop lowering, and to the errors above if that conflicts. + +Which shifts are bundled is decided by the channel assignment rather than by an attribute: a bundle +is what several overlapping matchings on *one* channel become, so giving each matching a channel of +its own is how a kernel declines the trade. What it then costs is colors, and those can be won back +by reusing a channel across phases. Two unbundled shifts may share one safely when they agree on +axis and signed distance and their sources agree modulo twice that distance, because a PE's role — +source, relay or destination — is then a function of its position modulo twice the distance alone, +so one static configuration serves every phase in the pool. Sharing on any other basis risks a PE +that sends on the color in one phase and receives on it in another, which needs a two-sided switch +change that a sender cannot drive (see [Lowering to Switches](#lowering-to-switches)), and nothing +in the compiler currently rejects it. `batcher_oddeven_bundled_1D.sptl` pools on exactly this rule, +and `batcher_oddeven_1D.sptl` is the same rule written out as arithmetic. + +This arrangement is the one used in Schnyder's *Distributed Sorting on the Cerebras Wafer-Scale +Engine* (fig. 7.6) for the 2D reduce-scatter, and is known to run on WSE-2. + +!!! note "Note: Multiple Rounds on One Color" + Two mechanisms are deliberately left unused, and are what to reach for if the four switch + positions or three filters run out. `SWITCH_RST` restores the initial configuration of every + router a message passes, which retires a whole path with one wavelet. Teardown-based + reconfiguration reprograms the routers between rounds outright, which is the only known way to + put an arbitrary *sequence* of sends and receives on one color: for a long enough sequence there + is a PE for which no fixed cycle of switch positions exists. Until then, splitting the rounds + across channels — as `bitonic_sort_1D.sptl` does — remains the per-kernel fallback. diff --git a/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl b/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl index acf4a6d3..aab1758b 100644 --- a/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl +++ b/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl @@ -38,7 +38,7 @@ * into the next -- an overlapping interval shift, which is what a bundle is. The low PEs of a * block take turns nearest-the-partner-first, handing their routers over to relay mode as they * finish; the high PEs are statically routed and pick their keys out of the stream with a - * counter filter. See irspec/docs/spatial/routing.md. + * counter filter. See irspec/docs/spatial/routing_wse.md. * * Bundling is chosen by the channel assignment, since a bundle is what several overlapping * matchings on one channel become. It costs one filter per participating PE and saves the diff --git a/spada/syntax/spatial_ir/shift_bundles.py b/spada/syntax/spatial_ir/shift_bundles.py index fc6ad9ae..7a660a57 100644 --- a/spada/syntax/spatial_ir/shift_bundles.py +++ b/spada/syntax/spatial_ir/shift_bundles.py @@ -5,7 +5,7 @@ the router of source ``p + 1`` carries source ``p``'s words. One color per direction still suffices, because the routers can be time-multiplexed -- but only if every switch is triggered by something the PE that owns it knows locally, since a control wavelet advances *every* switch-configured router -it reaches (see ``irspec/docs/spatial/routing.md``). +it reaches (see ``irspec/docs/spatial/routing_wse.md``). Send order is what provides that. The sources go nearest-the-destinations first, so a source's router changes from injecting to relaying exactly when that source has finished its own send, which From 6a9584d9c2805706dd33dff4a37e0b073f759e2c Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 00:04:17 +0200 Subject: [PATCH 31/68] Add odd even sort (quadratic work) --- README.md | 3 +- irspec/docs/spatial/routing_wse.md | 12 +- .../sorting/odd_even_sort_1D_looped.sptl | 135 ++++++++++++++++++ .../test_odd_even_sort_1d_looped.sh | 53 +++++++ tests/spatial_ir/test_routing.py | 48 +++++++ 5 files changed, 244 insertions(+), 7 deletions(-) create mode 100644 samples/spatial/sorting/odd_even_sort_1D_looped.sptl create mode 100644 tests/csl_runtime/test_odd_even_sort_1d_looped.sh diff --git a/README.md b/README.md index 78051e35..678292ee 100644 --- a/README.md +++ b/README.md @@ -104,7 +104,8 @@ Sample SPADA programs are in `samples/`: | `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy` | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Sorting networks over `2^L` PEs holding `K` keys each: `batcher_oddeven_1D` (a color per matching) and `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair) | +| `samples/spatial/sort/` | Batcher's odd-even mergesort over `2^L` PEs holding `K` keys each: `batcher_oddeven_1D` (a color per matching) and `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair) | +| `samples/spatial/sorting/` | 1D sorting networks: `bitonic_sort_1D` (channel reuse via router switches, WSE-3) and `odd_even_sort_1D_looped` (N neighbour rounds as a runtime loop, four static channels) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index 53f479a3..1defbca5 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -58,9 +58,9 @@ all. routers never advance. The epoch boundaries around it are then buying only *ordering* — and the fabric already delivers a channel's wavelets in order. Such a sequence of phases can be collapsed into a single epoch with a sequential `for` in the compute blocks, which lowers to a - real loop and so costs code and compile time independent of the number of rounds. Compare - `samples/spatial/sorting/odd_even_sort_1D.sptl` with - `samples/spatial/sorting/odd_even_sort_1D_looped.sptl`. + real loop and so costs code and compile time independent of the number of rounds. + `samples/spatial/sorting/odd_even_sort_1D_looped.sptl` is the example: N odd-even rounds on four + static channels, one CSL loop, no per-round barrier. This does *not* generalize to channels that switch: a router's positions are a static sequence, so the epoch a configuration belongs to has to be visible to the compiler. @@ -108,9 +108,9 @@ after all data of the epoch. eastward channel and every interior PE alternates between sending and receiving on it, so every close is a turnaround and every receiver has a switch-configured neighbour behind it. That kernel deadlocks on WSE-2 and runs on WSE-3, where the turnaround costs one position and one - message and nothing overshoots. `samples/spatial/sorting/odd_even_sort_1D.sptl` avoids it by - giving each round parity its own pair of channels: each PE's role on a channel is then fixed, - no router switches at all, and no close emits a message. + message and nothing overshoots. `samples/spatial/sorting/odd_even_sort_1D_looped.sptl` avoids it + by giving each round parity its own pair of channels: each PE's role on a channel is then + fixed, no router switches at all, and no close emits a message. Because the control message travels the path of the retired configuration in order behind the data, a receiving PE needs to emit nothing to advance its own router: the ordering required by the diff --git a/samples/spatial/sorting/odd_even_sort_1D_looped.sptl b/samples/spatial/sorting/odd_even_sort_1D_looped.sptl new file mode 100644 index 00000000..d30f257e --- /dev/null +++ b/samples/spatial/sorting/odd_even_sort_1D_looped.sptl @@ -0,0 +1,135 @@ +/** + * Odd-even transposition sort over N = 2^L PEs, with the N rounds as a runtime loop. + * + * Each PE holds K keys. Sequence k is element k of every PE; the K sequences are sorted + * independently. After the network, PE i holds the i-th key of each sequence. + * + * Algorithm + * --------- + * Round t compares neighbours (2i, 2i+1) when t is even and (2i+1, 2i+2) when t is odd. + * N rounds suffice. Every comparator is one hop: on a line a longer exchange occupies the + * same links as several one-hop ones and does not cut the round count, which is already + * Theta(N). + * + * One channel per (round parity, direction) -- four in all -- so each PE's role on each + * channel is fixed for the whole run. Sharing a channel across parities would make every + * interior PE alternate send/receive on it; on WSE-2 that turnaround overshoots the + * receiver (see irspec/docs/spatial/routing_wse.md). + * + * Why the rounds are a loop + * ------------------------- + * The same N rounds can be written as a compile-time `for` around two `phase` blocks. + * The compiler unrolls that into N phases, so code and compile time grow with N, and + * each phase boundary is a barrier. Here the rounds are a sequential `for` inside each + * compute block, which lowers to a CSL loop: the body is emitted once, and there is no + * per-round barrier. + * + * A phase boundary is an epoch boundary -- streams close, a channel may change hands, + * routers may advance. This network needs none of that. With roles fixed, no router + * switches and no channel is reassigned. Round t's values arrive before round t+1's + * because a channel is a FIFO. + * + * Role alternation is a per-PE-parity split, which a `compute` subgrid already expresses. + * The two ends sit out every odd round and so run a shorter body: + * + * PE 0 low in even rounds, idle in odd rounds + * PE 2, 4, .. N-2 low in even rounds, high in odd rounds + * PE 1, 3, .. N-3 high in even rounds, low in odd rounds + * PE N-1 high in even rounds, idle in odd rounds + * + * + * Constraints: L >= 1, K >= 1. WSE-2 and WSE-3. + * + **/ +kernel @odd_even_sort_1d_looped(stream[1<[1< east_even = relative_stream(1, 0) { + hops = auto, + channel = 0 + } + stream west_even = relative_stream(-1, 0) { + hops = auto, + channel = 1 + } + stream east_odd = relative_stream(1, 0) { + hops = auto, + channel = 2 + } + stream west_odd = relative_stream(-1, 0) { + hops = auto, + channel = 3 + } + } + + // Load this PE's K keys. + compute i16 i, i16 j in [0:1< val[k] else val[k]) + } + } + await send(val, a_out[i, j]) + } + + // Odd-indexed interior PEs: high in even rounds, low in odd rounds. + compute i16 i, i16 j in [1 : (1< val[k] else val[k]) + } + await send(val, east_odd) + await receive(other, west_odd) + await map i32 k in [0:K] { + val[k] = (other[k] if other[k] < val[k] else val[k]) + } + } + await send(val, a_out[i, j]) + } + + // PE N-1: high partner in every even round, idle in every odd round. + compute i16 i, i16 j in [(1< val[k] else val[k]) + } + } + await send(val, a_out[i, j]) + } +} diff --git a/tests/csl_runtime/test_odd_even_sort_1d_looped.sh b/tests/csl_runtime/test_odd_even_sort_1d_looped.sh new file mode 100644 index 00000000..fb6dbf68 --- /dev/null +++ b/tests/csl_runtime/test_odd_even_sort_1d_looped.sh @@ -0,0 +1,53 @@ +#!/bin/sh +# E2E: odd-even transposition sort on 2^L PEs, N rounds as a runtime loop +# (odd_even_sort_1D_looped.sptl). Each PE holds K keys; sequence k is element k of +# every PE. Reference: OUT_a_out == sort(a_in, axis=0). +# Runs on WSE-2 and WSE-3: four channels, one per (round parity, direction), so no +# router ever switches. L = 1 is two PEs and a single even round; L = 3 is eight PEs +# and exercises every role (ends and both interior parities). + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +. "$SCRIPT_DIR/_lib.sh" + +SAMPLES_DIR="$(cd "$SCRIPT_DIR/../../samples/spatial/sorting" && pwd)" +FOLDER="odd_even_sort_1d_looped_sptl" + +run_sort() { + l=$1 + k=$2 + echo "--- odd_even_sort_1d_looped L=$l K=$k ---" + + sptlc "$SAMPLES_DIR/odd_even_sort_1D_looped.sptl" "$FOLDER" -p L=$l -p K=$k + + python3 - < dict[str, str]: + kernel = parser.parse_file(_ODD_EVEN_LOOPED) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, L=L, K=K)) + return {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_benchmarking=True)} + + +def test_odd_even_sort_looped_uses_four_static_channels(): + """ + One channel per (round parity, direction). Roles never change, so no router switches and the + rounds stay a CSL loop of N/2 iterations rather than N unrolled phases. + """ + files = _lower_odd_even_looped(3) + colors = set(int(c) for c in re.findall(r'@get_color\((\d+)\)', files['layout.csl'])) + assert colors == {0, 1, 2, 3}, sorted(colors) + assert '.switches' not in files['layout.csl'] + + # L = 1 drops the interior rectangles (N = 2 has only the two endpoints). + ends = _lower_odd_even_looped(1, K=1) + assert 'code_0_0.csl' in ends and 'code_1_0.csl' in ends + assert 'code_2_0.csl' not in ends + + interior = files['code_2_0.csl'] + assert 'for (@range(i32, 0, 4, 1))' in interior, interior + # The loop body is emitted once: two even-round transfers and two odd-round transfers, not + # four copies of each for the four even/odd pairs at N = 8. + assert interior.count('fabout_dsd') == 2, interior + assert interior.count('fabin_dsd') == 2, interior + + +def test_odd_even_sort_looped_code_is_independent_of_n(): + """Lowering cost and the interior PE program stay flat as N grows.""" + small = _lower_odd_even_looped(3) + large = _lower_odd_even_looped(6) + assert 'for (@range(i32, 0, 32, 1))' in large['code_2_0.csl'] + # Same four PE roles, so the same number of code files; the loop trip count is the only + # difference that scales with L. + assert len(small) == len(large) + assert abs(len(small['code_2_0.csl']) - len(large['code_2_0.csl'])) < 64 + + if __name__ == '__main__': pytest.main([__file__]) From 6de1d4a2c23ed6acced84ed1f04ec43df998b22d Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 00:07:24 +0200 Subject: [PATCH 32/68] consolidate sort sample folders --- README.md | 3 +-- irspec/docs/spatial/routing_wse.md | 4 ++-- samples/spatial/{sorting => sort}/bitonic_sort_1D.sptl | 0 .../spatial/{sorting => sort}/odd_even_sort_1D_looped.sptl | 0 tests/csl_runtime/test_bitonic_sort_1d.sh | 2 +- tests/csl_runtime/test_odd_even_sort_1d_looped.sh | 2 +- tests/spatial_ir/test_routing.py | 4 ++-- 7 files changed, 7 insertions(+), 8 deletions(-) rename samples/spatial/{sorting => sort}/bitonic_sort_1D.sptl (100%) rename samples/spatial/{sorting => sort}/odd_even_sort_1D_looped.sptl (100%) diff --git a/README.md b/README.md index 678292ee..2600bfbb 100644 --- a/README.md +++ b/README.md @@ -104,8 +104,7 @@ Sample SPADA programs are in `samples/`: | `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy` | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Batcher's odd-even mergesort over `2^L` PEs holding `K` keys each: `batcher_oddeven_1D` (a color per matching) and `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair) | -| `samples/spatial/sorting/` | 1D sorting networks: `bitonic_sort_1D` (channel reuse via router switches, WSE-3) and `odd_even_sort_1D_looped` (N neighbour rounds as a runtime loop, four static channels) | +| `samples/spatial/sort/` | Sorting networks over `2^L` PEs holding `K` keys each: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `bitonic_sort_1D` (channel reuse via router switches, WSE-3), and `odd_even_sort_1D_looped` (N neighbour rounds as a runtime loop, four static channels) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index 1defbca5..a30de64f 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -59,7 +59,7 @@ all. fabric already delivers a channel's wavelets in order. Such a sequence of phases can be collapsed into a single epoch with a sequential `for` in the compute blocks, which lowers to a real loop and so costs code and compile time independent of the number of rounds. - `samples/spatial/sorting/odd_even_sort_1D_looped.sptl` is the example: N odd-even rounds on four + `samples/spatial/sort/odd_even_sort_1D_looped.sptl` is the example: N odd-even rounds on four static channels, one CSL loop, no per-round barrier. This does *not* generalize to channels that switch: a router's positions are a static sequence, @@ -108,7 +108,7 @@ after all data of the epoch. eastward channel and every interior PE alternates between sending and receiving on it, so every close is a turnaround and every receiver has a switch-configured neighbour behind it. That kernel deadlocks on WSE-2 and runs on WSE-3, where the turnaround costs one position and one - message and nothing overshoots. `samples/spatial/sorting/odd_even_sort_1D_looped.sptl` avoids it + message and nothing overshoots. `samples/spatial/sort/odd_even_sort_1D_looped.sptl` avoids it by giving each round parity its own pair of channels: each PE's role on a channel is then fixed, no router switches at all, and no close emits a message. diff --git a/samples/spatial/sorting/bitonic_sort_1D.sptl b/samples/spatial/sort/bitonic_sort_1D.sptl similarity index 100% rename from samples/spatial/sorting/bitonic_sort_1D.sptl rename to samples/spatial/sort/bitonic_sort_1D.sptl diff --git a/samples/spatial/sorting/odd_even_sort_1D_looped.sptl b/samples/spatial/sort/odd_even_sort_1D_looped.sptl similarity index 100% rename from samples/spatial/sorting/odd_even_sort_1D_looped.sptl rename to samples/spatial/sort/odd_even_sort_1D_looped.sptl diff --git a/tests/csl_runtime/test_bitonic_sort_1d.sh b/tests/csl_runtime/test_bitonic_sort_1d.sh index 311b8f43..c9defbb4 100644 --- a/tests/csl_runtime/test_bitonic_sort_1d.sh +++ b/tests/csl_runtime/test_bitonic_sort_1d.sh @@ -29,7 +29,7 @@ L=2 N=4 K=4 FOLDER="bitonic_sort_1d_sptl" -SAMPLES_DIR="$(cd "$SCRIPT_DIR/../../samples/spatial/sorting" && pwd)" +SAMPLES_DIR="$(cd "$SCRIPT_DIR/../../samples/spatial/sort" && pwd)" sptlc "$SAMPLES_DIR/bitonic_sort_1D.sptl" "$FOLDER" -p L=$L -p K=$K diff --git a/tests/csl_runtime/test_odd_even_sort_1d_looped.sh b/tests/csl_runtime/test_odd_even_sort_1d_looped.sh index fb6dbf68..a53f1335 100644 --- a/tests/csl_runtime/test_odd_even_sort_1d_looped.sh +++ b/tests/csl_runtime/test_odd_even_sort_1d_looped.sh @@ -10,7 +10,7 @@ set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" . "$SCRIPT_DIR/_lib.sh" -SAMPLES_DIR="$(cd "$SCRIPT_DIR/../../samples/spatial/sorting" && pwd)" +SAMPLES_DIR="$(cd "$SCRIPT_DIR/../../samples/spatial/sort" && pwd)" FOLDER="odd_even_sort_1d_looped_sptl" run_sort() { diff --git a/tests/spatial_ir/test_routing.py b/tests/spatial_ir/test_routing.py index e07134d1..79dba11c 100644 --- a/tests/spatial_ir/test_routing.py +++ b/tests/spatial_ir/test_routing.py @@ -423,7 +423,7 @@ def test_switch_positions_beyond_capacity_are_rejected(): # bitonic_sort_1D: the heaviest channel reuse in the samples ### -_BITONIC = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sorting', +_BITONIC = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'bitonic_sort_1D.sptl') @@ -474,7 +474,7 @@ def test_bitonic_sort_is_rejected_on_wse2(): ### _ODD_EVEN_LOOPED = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', - 'sorting', 'odd_even_sort_1D_looped.sptl') + 'sort', 'odd_even_sort_1D_looped.sptl') def _lower_odd_even_looped(L: int, K: int = 4) -> dict[str, str]: From 39b4fab2d57f715d6f4fc5915758db6ca591fd70 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 00:18:10 +0200 Subject: [PATCH 33/68] fix wse3 bundle lowering --- spada/lowering/spatial_ir_to_csl.py | 42 ++++++++++++++++++++-- spada/syntax/spatial_ir/stream_lifetime.py | 24 ++++++++++--- tests/spatial_ir/test_shift_bundles.py | 28 +++++++++++++++ tests/spatial_ir/test_stream_lifetime.py | 17 +++++++++ 4 files changed, 105 insertions(+), 6 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index abb9e8da..832d29e0 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -1085,7 +1085,8 @@ def _collect_unique_dsds( # PE -- the hardware rejects "two master input queues for the same color". Queues are therefore # handed out per channel; streams on ``auto`` channels get a color to themselves, so they key on # their own name. Sequential channels may share a queue only when their occupancy spans on this - # PE do not overlap; see ``stream_lifetime.assign_fabric_queues``. + # PE do not overlap; see ``stream_lifetime.assign_fabric_queues``. On WSE-3 a data-task ID is + # that input queue, so inbound colors that bind a data task cannot share one. channel_of_stream = { declaration.stream_name.as_ir(): declaration.stream.routing.resolved_channel for declaration in rect.dataflow.statements @@ -1098,9 +1099,11 @@ def queue_key(stream: spir.Identifier) -> str: input_names = _streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=True) output_names = _streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=False) + exclusive_input_keys = _data_task_queue_keys(tasks, rect.compute, queue_key, input_names) input_queue_of = stream_lifetime.assign_fabric_queues( _queue_spans(rect.compute, input_names, queue_key, inbound=True), csl.INPUT_QUEUE_IDS, - kind='input', architecture=csl.ARCH, location=location) + kind='input', architecture=csl.ARCH, location=location, + exclusive_keys=exclusive_input_keys) output_queue_of = stream_lifetime.assign_fabric_queues( _queue_spans(rect.compute, output_names, queue_key, inbound=False), csl.OUTPUT_QUEUE_IDS, kind='output', architecture=csl.ARCH, location=location) @@ -1466,6 +1469,41 @@ def _exit_task_hardware_id(used_ids: set[int], color_ids: set[int]) -> int: f'(occupied {sorted(occupied)}).') +def _data_task_queue_keys( + tasks: list[tdag.CSLTask], + compute: spir.ComputeBlock, + queue_key, + input_names: set[spir.Identifier], +) -> frozenset[str]: + """Return queue keys whose inbound color will bind a data task. + + Used on WSE-3 so those colors each get their own input queue: the queue is + the data-task hardware ID and the comptime ``@initialize_queue`` bind. WSE-2 + data-task IDs are colors, so occupancy pooling may still share queues there. + + :param tasks: Tasks of this PE, already classified as local or data. + :param compute: The compute block those task statement indices refer to. + :param queue_key: Maps a stream identifier to its grouping key. + :param input_names: Streams that bind a fabric input queue on this PE. + :return: The exclusive keys, or empty when this generation may share queues. + """ + if csl.ARCH != 'wse3': + return frozenset() + keys: set[str] = set() + for task in tasks: + if task.task_type != 'data': + continue + stmt = compute.statements[task.statements[0]] + if not isinstance(stmt, spir.ForeachStatement): + continue + sname = stmt.receive_stream.stream_name + if isinstance(sname, spir.ArraySlice): + sname = sname.array + if sname in input_names: + keys.add(queue_key(sname)) + return frozenset(keys) + + def _data_task_color(rect: PEBlock, task_index: int, task: tdag.CSLTask, color_map: dict[str, int]) -> int: """ Returns the color a data task listens on. diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py index 97ab447e..56cf22ef 100644 --- a/spada/syntax/spatial_ir/stream_lifetime.py +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -741,7 +741,8 @@ def _never_concurrent(first: str, second: str, uses_per_rect: list[dict[spir.Ide def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int], *, kind: str, - architecture: str, location: str) -> dict[str, int]: + architecture: str, location: str, + exclusive_keys: frozenset[str] | None = None) -> dict[str, int]: """ Assigns hardware fabric queues to stream groups from their occupancy spans on one PE. @@ -750,6 +751,10 @@ def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int] arrive in a gap, and remapping the queue onto another color while they sit there is what the hardware rejects. Two groups may share a queue only when those spans do not overlap. + Keys in ``exclusive_keys`` never share a queue, even when their spans are disjoint. That is + required on WSE-3 for inbound colors that bind a data task: the hardware ID *is* the input + queue, and ``@initialize_queue`` is a comptime one-to-one bind. + The spans form an interval graph, so colouring them in start-time order is optimal. :param spans: Mapping of grouping key to an inclusive ``(first_use, last_use)`` statement index @@ -759,6 +764,8 @@ def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int] :param kind: ``'input'`` or ``'output'``, for the diagnostic. :param architecture: The target name, for the diagnostic. :param location: The PE rectangle, for the diagnostic. + :param exclusive_keys: Groups that must each own a queue for the whole PE, typically WSE-3 + data-task colors. :return: Mapping of grouping key to a queue identifier from ``queue_ids``. """ if not spans: @@ -767,13 +774,15 @@ def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int] raise SyntaxError( f'{location} needs {kind} queues, but {architecture} has none that a program may use.') + exclusive_keys = exclusive_keys or frozenset() assigned: dict[str, int] = {} for key in sorted(spans, key=lambda name: (spans[name][0], spans[name][1], name)): start, end = spans[key] used = { assigned[other] for other in assigned - if start <= spans[other][1] and spans[other][0] <= end + if (key in exclusive_keys or other in exclusive_keys + or (start <= spans[other][1] and spans[other][0] <= end)) } for queue in queue_ids: if queue not in used: @@ -782,8 +791,15 @@ def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int] else: overlapping = sorted( other for other, (other_start, other_end) in spans.items() - if other != key and start <= other_end and other_start <= end + if other != key and ( + key in exclusive_keys or other in exclusive_keys + or (start <= other_end and other_start <= end)) ) + extra = '' + if key in exclusive_keys or exclusive_keys.intersection(overlapping): + extra = ( + '\n note: on WSE-3 a data-task ID is its input queue, so two colors that bind ' + 'a data task cannot share one') raise SyntaxError( f'{location} would need {len(used) + 1} concurrent {kind} queues ' f'(live groups {[key] + overlapping}), but a PE can use at most ' @@ -791,7 +807,7 @@ def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int] f' note: a fabric queue is remapped when a new color uses it, and the hardware ' f'rejects that while wavelets remain\n' f' note: a channel keeps one queue for the whole of its lifetime on the PE, ' - f'including gaps between epochs') + f'including gaps between epochs{extra}') return assigned diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py index 4e3e04c9..c50ffd1c 100644 --- a/tests/spatial_ir/test_shift_bundles.py +++ b/tests/spatial_ir/test_shift_bundles.py @@ -338,3 +338,31 @@ def test_a_scalar_receive_lowers_to_a_data_task(): assert f'@initialize_queue(@get_input_queue({queue}),' in pe0, pe0 else: assert re.search(r'@get_data_task_id\(@get_color\(\d+\)\)', pe0), pe0 + + +def test_sequential_data_task_colors_get_distinct_hardware_ids(): + """R=2 binds two data tasks on one PE; they must not share a hardware ID. + + On WSE-3 that ID is the input queue, so occupancy pooling must not remap the + first epoch's queue onto the second color. cslc rejects the shared ID as + "task ID bound to more than one task". + """ + path = os.path.join( + os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'simple', 'exchange_bundle_1D.sptl' + ) + kernel = parser.parse_file(path) + kernel = passes.concretize_parameters(kernel, M=3, D=3, R=2) + kernel = passes.constexpr_propagation(kernel) + files = lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) + pe = next(f.code for f in files if f.filename == 'code_3_0.csl') + ids = re.findall(r'const dtask_\d+_id = (@get_data_task_id\([^;]+);', pe) + assert len(ids) == 2, pe + assert ids[0] != ids[1], pe + if constants.ARCH == 'wse3': + queues = re.findall(r'@get_data_task_id\(@get_input_queue\((\d+)\)\)', pe) + assert len(set(queues)) == 2, pe + inits = re.findall( + r'@initialize_queue\(@get_input_queue\((\d+)\), \.\{ \.color = (\w+) \}\)', pe) + by_queue = {queue: color for queue, color in inits} + for queue in queues: + assert queue in by_queue, pe diff --git a/tests/spatial_ir/test_stream_lifetime.py b/tests/spatial_ir/test_stream_lifetime.py index 6d41639f..d8a82a9a 100644 --- a/tests/spatial_ir/test_stream_lifetime.py +++ b/tests/spatial_ir/test_stream_lifetime.py @@ -496,5 +496,22 @@ def test_three_overlapping_spans_exhaust_two_queues(): kind='input', architecture='wse2', location='PE (0, 0)') +def test_exclusive_keys_do_not_share_a_queue_across_a_gap(): + """WSE-3 data-task colors cannot share a queue even when their spans are disjoint.""" + assigned = stream_lifetime.assign_fabric_queues( + {'channel 0': (0, 2), 'channel 1': (3, 5)}, [2, 3], + kind='input', architecture='wse3', location='PE (0, 0)', + exclusive_keys=frozenset({'channel 0', 'channel 1'})) + assert assigned['channel 0'] != assigned['channel 1'] + + +def test_exclusive_keys_exhaust_queues_when_too_many_data_tasks(): + with pytest.raises(SyntaxError, match='data-task ID'): + stream_lifetime.assign_fabric_queues( + {'channel 0': (0, 1), 'channel 1': (2, 3), 'channel 2': (4, 5)}, [2, 3], + kind='input', architecture='wse3', location='PE (0, 0)', + exclusive_keys=frozenset({'channel 0', 'channel 1', 'channel 2'})) + + if __name__ == '__main__': pytest.main([__file__]) From 79327632365fe5c3f4c11ceea50a6fa71c4a6788 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 09:01:24 +0200 Subject: [PATCH 34/68] Fix fabric shape analysis --- spada/syntax/spatial_ir/analysis.py | 20 +++++++++++++- tests/spatial_ir/test_spatial_ir_analysis.py | 29 ++++++++++++++++++++ 2 files changed, 48 insertions(+), 1 deletion(-) diff --git a/spada/syntax/spatial_ir/analysis.py b/spada/syntax/spatial_ir/analysis.py index b1b093cd..68ae988b 100644 --- a/spada/syntax/spatial_ir/analysis.py +++ b/spada/syntax/spatial_ir/analysis.py @@ -161,12 +161,30 @@ def sends_and_receives(compute: spir.ComputeBlock) -> dict[spir.Identifier, tupl return {k: (k in collector.sends, k in collector.receives) for k in all_identifiers} +def _fabric_shape(shape: list[int]) -> list[int]: + """ + Host tensors and memcpy are always ``(width, height, elem_per_pe)``. + + A 0-D stream occupies one PE; a 1-D array of streams is a row. 2-D shapes + are already a fabric rectangle and are left unchanged. + """ + if len(shape) == 0: + return [1, 1] + if len(shape) == 1: + return [shape[0], 1] + return shape + + def get_kernel_stream_arguments( kernel: spir.Kernel) -> tuple[dict[str, dict[str, list[int] | str]], dict[str, dict[str, list[int] | str]]]: """ Returns two dictionaries: 1. A dictionary mapping input stream names to their data types and shapes. 2. A dictionary mapping output stream names to their data types and shapes. + + Stream argument ``shape`` is the 2D PE rectangle the runtime copies, not the + syntactic rank of the IR type: ``stream`` is ``[1, 1]`` and + ``stream[N]`` is ``[N, 1]``. Compile-time scalars keep ``shape = []``. """ input_streams = {} output_streams = {} @@ -188,7 +206,7 @@ def get_kernel_stream_arguments( arg_as_dict = { "dtype": arg.dtype.element_type.element_type.element_type.element_type.as_ir(), - "shape": shape, + "shape": _fabric_shape(shape), } if isinstance(arg.dtype, spir.StreamType): arg_as_dict["buffer_size"] = arg.dtype.buffer_size.eval() if arg.dtype.buffer_size else None diff --git a/tests/spatial_ir/test_spatial_ir_analysis.py b/tests/spatial_ir/test_spatial_ir_analysis.py index b86fc8e5..eb5cdb1f 100644 --- a/tests/spatial_ir/test_spatial_ir_analysis.py +++ b/tests/spatial_ir/test_spatial_ir_analysis.py @@ -816,6 +816,34 @@ def test_transposed_stream_extents_1D(second_index): assert stream_extents.is_transposed[out_identifier] is False +def test_stream_argument_shapes_are_two_dimensional(): + """ + Metadata shape is the memcpy rectangle ``(w, h)``, plus ``buffer_size`` as the third axis. + 0-D and 1-D stream types are padded; 2-D types and compile-time scalars are not. + """ + collectives = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'collectives') + kernel = parser.parse_file(os.path.join(collectives, 'scalar_reduce_1D.sptl')) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, N=4)) + inputs, outputs = analysis.get_kernel_stream_arguments(kernel) + assert inputs['inp']['shape'] == [4, 1] + assert inputs['inp']['buffer_size'] == 1 + assert outputs['out']['shape'] == [1, 1] + assert outputs['out']['buffer_size'] == 1 + + simple = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'simple') + kernel = parser.parse_file(os.path.join(simple, 'add.sptl')) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, N=8)) + inputs, outputs = analysis.get_kernel_stream_arguments(kernel) + assert inputs['a']['shape'] == [8, 8] + assert outputs['out']['shape'] == [8, 8] + + kernel = parser.parse_file(os.path.join(simple, 'mult_scalar.sptl')) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, N=4)) + inputs, _ = analysis.get_kernel_stream_arguments(kernel) + assert inputs['coeff']['shape'] == [] + assert inputs['a']['shape'] == [4, 4] + + if __name__ == '__main__': test_completion_dag_simple() test_completion_dag_concurrent() @@ -839,3 +867,4 @@ def test_transposed_stream_extents_1D(second_index): test_transposed_stream_extents(True) test_transposed_stream_extents_1D(False) test_transposed_stream_extents_1D(True) + test_stream_argument_shapes_are_two_dimensional() From e530b54f89766fc4aa18b00bee0ffe9c70ba1dea Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 09:02:50 +0200 Subject: [PATCH 35/68] Add WSE-3 simulations to CI --- .github/workflows/python-app.yml | 10 +++++++++- tests/csl_runtime/run_tests.sh | 5 +++++ 2 files changed, 14 insertions(+), 1 deletion(-) diff --git a/.github/workflows/python-app.yml b/.github/workflows/python-app.yml index 368b5f6a..af4fec20 100644 --- a/.github/workflows/python-app.yml +++ b/.github/workflows/python-app.yml @@ -36,7 +36,14 @@ jobs: PYTHONPATH=`pwd` pytest test-csl: + name: test-csl (${{ matrix.wse_arch }}) runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + wse_arch: [wse2, wse3] + env: + WSE_ARCH: ${{ matrix.wse_arch }} steps: - uses: actions/checkout@v4 @@ -89,8 +96,9 @@ jobs: run: | ./cerebras-sdk/cslc -h - - name: Test CSL with simulator + - name: Test CSL with simulator (${{ matrix.wse_arch }}) run: | pip install --no-deps -e . export PATH=$PATH:`pwd`/cerebras-sdk + echo "WSE_ARCH=$WSE_ARCH" ./tests/csl_runtime/run_tests.sh diff --git a/tests/csl_runtime/run_tests.sh b/tests/csl_runtime/run_tests.sh index 98a1948d..20cab76e 100755 --- a/tests/csl_runtime/run_tests.sh +++ b/tests/csl_runtime/run_tests.sh @@ -17,6 +17,7 @@ declare -a FAILED_TESTS echo -e "${BLUE}================================${NC}" echo -e "${BLUE} Running Test Suite${NC}" +echo -e "${BLUE} WSE_ARCH=${WSE_ARCH:-wse2}${NC}" echo -e "${BLUE}================================${NC}" echo "" @@ -36,6 +37,10 @@ NON_TEST_SCRIPTS=("run_tests.sh" "run-in-lima.sh" "sptlc" "_lib.sh") is_non_test() { local name="$1" + # Local debug helpers (zz_*) are not part of the suite. + case "$name" in + zz_*) return 0 ;; + esac for skip in "${NON_TEST_SCRIPTS[@]}"; do [ "$name" = "$skip" ] && return 0 done From 83f92857487d9aaa4af9a021cc0d76eea6e19d52 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 09:28:16 +0200 Subject: [PATCH 36/68] Fixes for WSE3 --- spada/lowering/spatial_ir_to_csl.py | 49 ++++--------------- spada/syntax/spatial_ir/stream_lifetime.py | 5 +- .../csl_runtime/test_shift_bundle_filters.sh | 12 ++++- tests/spatial_ir/test_dsd_ops.py | 30 +++++++++++- tests/spatial_ir/test_shift_bundles.py | 7 ++- 5 files changed, 58 insertions(+), 45 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 832d29e0..3f2bb832 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -1085,8 +1085,9 @@ def _collect_unique_dsds( # PE -- the hardware rejects "two master input queues for the same color". Queues are therefore # handed out per channel; streams on ``auto`` channels get a color to themselves, so they key on # their own name. Sequential channels may share a queue only when their occupancy spans on this - # PE do not overlap; see ``stream_lifetime.assign_fabric_queues``. On WSE-3 a data-task ID is - # that input queue, so inbound colors that bind a data task cannot share one. + # PE do not overlap; see ``stream_lifetime.assign_fabric_queues``. On WSE-3 every inbound color + # keeps its own queue: remapping one that still holds wavelets is a fatal error, and a data-task + # ID is that queue. channel_of_stream = { declaration.stream_name.as_ir(): declaration.stream.routing.resolved_channel for declaration in rect.dataflow.statements @@ -1099,9 +1100,14 @@ def queue_key(stream: spir.Identifier) -> str: input_names = _streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=True) output_names = _streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=False) - exclusive_input_keys = _data_task_queue_keys(tasks, rect.compute, queue_key, input_names) + input_spans = _queue_spans(rect.compute, input_names, queue_key, inbound=True) + # WSE-3 remaps a fabric queue onto the next color at the first transfer that uses it, and + # faults if the queue still holds wavelets. Occupancy in the compute block is not enough + # to prove it is empty, so every inbound color keeps its own queue. That also keeps + # data-task IDs unique, since those IDs *are* the input queues. + exclusive_input_keys = frozenset(input_spans) if csl.ARCH == 'wse3' else frozenset() input_queue_of = stream_lifetime.assign_fabric_queues( - _queue_spans(rect.compute, input_names, queue_key, inbound=True), csl.INPUT_QUEUE_IDS, + input_spans, csl.INPUT_QUEUE_IDS, kind='input', architecture=csl.ARCH, location=location, exclusive_keys=exclusive_input_keys) output_queue_of = stream_lifetime.assign_fabric_queues( @@ -1469,41 +1475,6 @@ def _exit_task_hardware_id(used_ids: set[int], color_ids: set[int]) -> int: f'(occupied {sorted(occupied)}).') -def _data_task_queue_keys( - tasks: list[tdag.CSLTask], - compute: spir.ComputeBlock, - queue_key, - input_names: set[spir.Identifier], -) -> frozenset[str]: - """Return queue keys whose inbound color will bind a data task. - - Used on WSE-3 so those colors each get their own input queue: the queue is - the data-task hardware ID and the comptime ``@initialize_queue`` bind. WSE-2 - data-task IDs are colors, so occupancy pooling may still share queues there. - - :param tasks: Tasks of this PE, already classified as local or data. - :param compute: The compute block those task statement indices refer to. - :param queue_key: Maps a stream identifier to its grouping key. - :param input_names: Streams that bind a fabric input queue on this PE. - :return: The exclusive keys, or empty when this generation may share queues. - """ - if csl.ARCH != 'wse3': - return frozenset() - keys: set[str] = set() - for task in tasks: - if task.task_type != 'data': - continue - stmt = compute.statements[task.statements[0]] - if not isinstance(stmt, spir.ForeachStatement): - continue - sname = stmt.receive_stream.stream_name - if isinstance(sname, spir.ArraySlice): - sname = sname.array - if sname in input_names: - keys.add(queue_key(sname)) - return frozenset(keys) - - def _data_task_color(rect: PEBlock, task_index: int, task: tdag.CSLTask, color_map: dict[str, int]) -> int: """ Returns the color a data task listens on. diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py index 56cf22ef..7f856c66 100644 --- a/spada/syntax/spatial_ir/stream_lifetime.py +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -752,8 +752,9 @@ def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int] hardware rejects. Two groups may share a queue only when those spans do not overlap. Keys in ``exclusive_keys`` never share a queue, even when their spans are disjoint. That is - required on WSE-3 for inbound colors that bind a data task: the hardware ID *is* the input - queue, and ``@initialize_queue`` is a comptime one-to-one bind. + required on WSE-3 for every inbound color: the simulator remaps a queue onto the next color at + the first transfer, and faults if the queue is not empty. It is also required for colors that + bind a data task, because the hardware ID *is* the input queue. The spans form an interval graph, so colouring them in start-time order is optimal. diff --git a/tests/csl_runtime/test_shift_bundle_filters.sh b/tests/csl_runtime/test_shift_bundle_filters.sh index ab3aba92..06422522 100644 --- a/tests/csl_runtime/test_shift_bundle_filters.sh +++ b/tests/csl_runtime/test_shift_bundle_filters.sh @@ -29,11 +29,19 @@ compile_and_run() { k=$2 filter=$3 width=$((2 * m)) + arch="${WSE_ARCH:-wse2}" + if [ "$arch" = "wse3" ]; then + in_queue=2 + init_queues=1 + else + in_queue=0 + init_queues=0 + fi rm -rf "$OUT" - cslc --arch="${WSE_ARCH:-wse2}" "$SRC/layout.csl" -o "$OUT" \ + cslc --arch="$arch" "$SRC/layout.csl" -o "$OUT" \ --fabric-dims=$((7 + width)),3 --fabric-offsets=4,1 --memcpy --channels=1 \ - --params=M:$m,K:$k,FILTER:$filter + --params=M:$m,K:$k,FILTER:$filter,IN_QUEUE:$in_queue,INIT_QUEUES:$init_queues timeout -s 9 240 cs_python "$SRC/run.py" "$OUT" --M "$m" --K "$k" --filter "$filter" } diff --git a/tests/spatial_ir/test_dsd_ops.py b/tests/spatial_ir/test_dsd_ops.py index b680323c..77e9b142 100644 --- a/tests/spatial_ir/test_dsd_ops.py +++ b/tests/spatial_ir/test_dsd_ops.py @@ -4,7 +4,7 @@ from spada.lowering import spatial_ir_to_csl as s2c from spada.syntax.spatial_ir import parser, passes from spada.syntax.spatial_ir.canonicalization import PEBlock -from spada.syntax.csl import dsd_ops +from spada.syntax.csl import constants, dsd_ops def test_dsd_op_detection(): @@ -266,6 +266,34 @@ def test_a_reused_color_keeps_its_input_queue_across_a_gap(): assert queues['fwd__17'] != queues['bwd__30'] +def test_wse3_inbound_colors_do_not_share_an_input_queue(): + """WSE-3 remaps a queue onto the next color and faults if it is not empty. + + Batcher L=2 is the case that hit ``Attempt to remap input queue 2 from C1 to C3`` + when occupancy pooling reused the queue across sequential colors. + """ + if constants.ARCH != 'wse3': + pytest.skip('WSE-2 may remap a drained queue onto the next color') + path = os.path.join( + os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_1D.sptl' + ) + kernel = parser.parse_file(path) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, L=2, K=2)) + files = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_benchmarking=True)} + code = files['code_1_0.csl'] + colors = dict(re.findall(r'const (\w+)_color_in: color = @get_color\((\d+)\);', code)) + queues = dict(re.findall( + r'const (\w+)_in_dsd = @get_dsd\(fabin_dsd, .*?input_queue = @get_input_queue\((\d+)\)', + code)) + by_color: dict[str, set[str]] = {} + for name, color in colors.items(): + if name in queues: + by_color.setdefault(color, set()).add(queues[name]) + assert len(by_color) >= 2, code + used_queues = [next(iter(qs)) for qs in by_color.values()] + assert len(used_queues) == len(set(used_queues)), (by_color, code) + + if __name__ == '__main__': test_dsd_op_detection() test_dsd_op_detection_constant_folding() diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py index c50ffd1c..6a0ee883 100644 --- a/tests/spatial_ir/test_shift_bundles.py +++ b/tests/spatial_ir/test_shift_bundles.py @@ -272,10 +272,15 @@ def test_the_batcher_fits_the_colors_it_has(l: int, colors: int): def test_sixteen_keys_need_three_overlapping_input_queues(): """ At L = 4 a reused inbound color stays live across a gap that already holds two other colors. - WSE-2 has two input queues, so lowering must refuse rather than remap a busy queue. + WSE-2 has two input queues, so occupancy pooling refuses. WSE-3 has six, but remapping a + non-empty queue is illegal, and L = 4 wants seven inbound colors over the kernel. """ from spada.syntax.csl import constants + if constants.ARCH == 'wse3': + with pytest.raises(SyntaxError, match='concurrent input queues'): + _bundled_batcher(4) + return if len(constants.INPUT_QUEUE_IDS) >= 3: pytest.skip(f'{constants.ARCH} has {len(constants.INPUT_QUEUE_IDS)} input queues, enough for L=4') with pytest.raises(SyntaxError, match='concurrent input queues'): From d2f852ed690e78d700103bbdaa02f1cc5e8be50f Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 09:38:39 +0200 Subject: [PATCH 37/68] Improve queue managament for wse-3 --- spada/lowering/spatial_ir_to_csl.py | 28 +++++++++++++++++++--------- 1 file changed, 19 insertions(+), 9 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 3f2bb832..0b7ef25b 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -950,8 +950,15 @@ def _declare_queue_initialization(dsds: UniqueDSDDict, rect: PEBlock, footer: St for statement in rect.metadata.compute.statements: if isinstance(statement, spir.CloseStatement) and statement.switch_advance: - name = cslstmt.name_to_csl(stream_lifetime.underlying_stream(statement.stream_name)) - queue = csl.OUTPUT_QUEUE_IDS[0] + stream = stream_lifetime.underlying_stream(statement.stream_name) + name = cslstmt.name_to_csl(stream) + queue = None + for _, dsd in dsds.get(stream.as_ir(), ()): + if isinstance(dsd, cslstruct.FabricDSD) and dsd.dsd_type == cslstruct.DSDType.fabout: + queue = dsd.queue + break + if queue is None: + queue = csl.OUTPUT_QUEUE_IDS[0] bindings.setdefault(('output_queue', queue), f'@get_color({color_map[name + "_OUT"]})') for (kind, queue), color in sorted(bindings.items()): @@ -1101,18 +1108,21 @@ def queue_key(stream: spir.Identifier) -> str: input_names = _streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=True) output_names = _streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=False) input_spans = _queue_spans(rect.compute, input_names, queue_key, inbound=True) + output_spans = _queue_spans(rect.compute, output_names, queue_key, inbound=False) # WSE-3 remaps a fabric queue onto the next color at the first transfer that uses it, and - # faults if the queue still holds wavelets. Occupancy in the compute block is not enough - # to prove it is empty, so every inbound color keeps its own queue. That also keeps - # data-task IDs unique, since those IDs *are* the input queues. - exclusive_input_keys = frozenset(input_spans) if csl.ARCH == 'wse3' else frozenset() + # faults or stalls if the queue still holds wavelets. Occupancy in the compute block is not + # enough to prove it is empty, so every color keeps its own queue. That also keeps data-task + # IDs unique, since those IDs *are* the input queues. + exclusive = frozenset(input_spans) if csl.ARCH == 'wse3' else frozenset() + exclusive_out = frozenset(output_spans) if csl.ARCH == 'wse3' else frozenset() input_queue_of = stream_lifetime.assign_fabric_queues( input_spans, csl.INPUT_QUEUE_IDS, kind='input', architecture=csl.ARCH, location=location, - exclusive_keys=exclusive_input_keys) + exclusive_keys=exclusive) output_queue_of = stream_lifetime.assign_fabric_queues( - _queue_spans(rect.compute, output_names, queue_key, inbound=False), csl.OUTPUT_QUEUE_IDS, - kind='output', architecture=csl.ARCH, location=location) + output_spans, csl.OUTPUT_QUEUE_IDS, + kind='output', architecture=csl.ARCH, location=location, + exclusive_keys=exclusive_out) def allocate_input_queue(stream: spir.Identifier) -> int: key = queue_key(stream) From 8c641e1b49361f508c5c90815f416ded735b46be Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 10:11:03 +0200 Subject: [PATCH 38/68] Fix handwritten csl tests --- tests/csl_runtime/handwritten/shift_bundle/layout.csl | 10 ++++++++++ .../csl_runtime/handwritten/shift_bundle/receiver.csl | 7 ++++++- tests/csl_runtime/handwritten/shift_bundle/sender.csl | 4 ++++ 3 files changed, 20 insertions(+), 1 deletion(-) diff --git a/tests/csl_runtime/handwritten/shift_bundle/layout.csl b/tests/csl_runtime/handwritten/shift_bundle/layout.csl index 6e027c5d..c94c60e6 100644 --- a/tests/csl_runtime/handwritten/shift_bundle/layout.csl +++ b/tests/csl_runtime/handwritten/shift_bundle/layout.csl @@ -31,6 +31,13 @@ param K: i16; // 1: counter filters; every receiver takes only the K words addressed to it. param FILTER: i16; +// Program input queue for the receivers. WSE-2 may use 0; WSE-3 memcpy already owns 0 and 1, +// so the probe is compiled with IN_QUEUE=2 there. +param IN_QUEUE: i16; + +// 1: emit @initialize_queue (required on WSE-3). 0: omit it (WSE-2). +param INIT_QUEUES: i16; + const memcpy = @import_module("", .{ .width = 2 * M, .height = 1, @@ -57,6 +64,7 @@ layout { .words = K, // The westmost sender is the last to send and has nothing to relay for. .hands_over = x > 0, + .init_queues = INIT_QUEUES, }); } for (@range(i16, 0, M, 1)) |q| { @@ -65,6 +73,8 @@ layout { .stream = STREAM, .words = K, .takes = if (FILTER == 0) STREAM else K, + .in_queue = IN_QUEUE, + .init_queues = INIT_QUEUES, }); } diff --git a/tests/csl_runtime/handwritten/shift_bundle/receiver.csl b/tests/csl_runtime/handwritten/shift_bundle/receiver.csl index 7bb2e07d..caebd77f 100644 --- a/tests/csl_runtime/handwritten/shift_bundle/receiver.csl +++ b/tests/csl_runtime/handwritten/shift_bundle/receiver.csl @@ -8,6 +8,8 @@ param memcpy_params: comptime_struct; param stream: i16; param words: i16; param takes: i16; +param in_queue: i16; +param init_queues: i16; const sys_mod = @import_module("", memcpy_params); @@ -22,7 +24,7 @@ const got_dsd = @get_dsd(mem1d_dsd, .{ .tensor_access = |i|{takes} -> got[i] }); const in_dsd = @get_dsd(fabin_dsd, .{ .extent = takes, .fabric_color = channel, - .input_queue = @get_input_queue(0), + .input_queue = @get_input_queue(in_queue), }); const recv_id = @get_local_task_id(8); @@ -48,5 +50,8 @@ comptime { @export_symbol(__got_ptr, "got"); @bind_local_task(recv_task, recv_id); @bind_local_task(done_task, done_id); + if (init_queues != 0) { + @initialize_queue(@get_input_queue(in_queue), .{ .color = channel }); + } @export_symbol(main, "main"); } diff --git a/tests/csl_runtime/handwritten/shift_bundle/sender.csl b/tests/csl_runtime/handwritten/shift_bundle/sender.csl index 5d8b4e57..8c43c442 100644 --- a/tests/csl_runtime/handwritten/shift_bundle/sender.csl +++ b/tests/csl_runtime/handwritten/shift_bundle/sender.csl @@ -9,6 +9,7 @@ param memcpy_params: comptime_struct; param stream: i16; param words: i16; param hands_over: bool; +param init_queues: i16; const sys_mod = @import_module("", memcpy_params); const ctrl = @import_module(""); @@ -63,5 +64,8 @@ comptime { @export_symbol(__got_ptr, "got"); @bind_local_task(send_task, send_id); @bind_local_task(done_task, done_id); + if (init_queues != 0) { + @initialize_queue(@get_output_queue(2), .{ .color = channel }); + } @export_symbol(main, "main"); } From 8f3528d1d7ddee78710b7c9f1cd35019349b9f85 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 10:32:33 +0200 Subject: [PATCH 39/68] Fix CI for e2e tests --- .github/workflows/python-app.yml | 33 ++++++++++++++++++-------------- 1 file changed, 19 insertions(+), 14 deletions(-) diff --git a/.github/workflows/python-app.yml b/.github/workflows/python-app.yml index af4fec20..6edd04ab 100644 --- a/.github/workflows/python-app.yml +++ b/.github/workflows/python-app.yml @@ -35,15 +35,10 @@ jobs: run: | PYTHONPATH=`pwd` pytest + # WSE-2 and WSE-3 share one runner so Singularity and the SDK are installed + # once. test-csl: - name: test-csl (${{ matrix.wse_arch }}) runs-on: ubuntu-latest - strategy: - fail-fast: false - matrix: - wse_arch: [wse2, wse3] - env: - WSE_ARCH: ${{ matrix.wse_arch }} steps: - uses: actions/checkout@v4 @@ -75,15 +70,15 @@ jobs: sudo apt install ./singularity.deb singularity --version - - name: Cache dependencies - id: cache-deps + - name: Cache Cerebras SDK + id: cache-sdk uses: actions/cache@v4 with: path: cerebras-sdk - key: ${{ runner.os }}-deps + key: ${{ runner.os }}-cerebras-sdk-1.4.0 - name: Install Cerebras SDK v1.4.0 - if: steps.cache-deps.outputs.cache-hit != 'true' + if: steps.cache-sdk.outputs.cache-hit != 'true' run: | mkdir cerebras-sdk cd cerebras-sdk @@ -96,9 +91,19 @@ jobs: run: | ./cerebras-sdk/cslc -h - - name: Test CSL with simulator (${{ matrix.wse_arch }}) + - name: Test CSL with simulator run: | pip install --no-deps -e . export PATH=$PATH:`pwd`/cerebras-sdk - echo "WSE_ARCH=$WSE_ARCH" - ./tests/csl_runtime/run_tests.sh + status=0 + for arch in wse2 wse3; do + echo "::group::WSE_ARCH=$arch" + if WSE_ARCH=$arch ./tests/csl_runtime/run_tests.sh; then + echo "$arch passed" + else + echo "$arch failed" + status=1 + fi + echo "::endgroup::" + done + exit $status From b36abe6737b8d815863503eb1af35a26d6ccc3aa Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 11:46:49 +0200 Subject: [PATCH 40/68] fix tests --- .github/workflows/python-app.yml | 19 +++++++++++++------ tests/csl_runtime/run_tests.sh | 5 +++++ 2 files changed, 18 insertions(+), 6 deletions(-) diff --git a/.github/workflows/python-app.yml b/.github/workflows/python-app.yml index 6edd04ab..8bd16501 100644 --- a/.github/workflows/python-app.yml +++ b/.github/workflows/python-app.yml @@ -39,6 +39,7 @@ jobs: # once. test-csl: runs-on: ubuntu-latest + timeout-minutes: 180 steps: - uses: actions/checkout@v4 @@ -56,18 +57,24 @@ jobs: if [ -f requirements-ci.txt ]; then pip install -r requirements-ci.txt; fi - name: Install CSL dependencies + env: + DEBIAN_FRONTEND: noninteractive run: | + set -euxo pipefail . /etc/os-release echo "Using Ubuntu version $UBUNTU_CODENAME" - # Make dependency setup faster + # Avoid man-db postinst work during package installs. echo 'set man-db/auto-update false' | sudo debconf-communicate >/dev/null - sudo dpkg-reconfigure man-db - sudo apt-get update - sudo apt-get install -y build-essential libssl-dev uuid-dev libgpgme11-dev squashfs-tools - wget -q -O singularity.deb https://github.com/sylabs/singularity/releases/download/v4.2.1/singularity-ce_4.2.1-${UBUNTU_CODENAME}_amd64.deb - sudo apt install ./singularity.deb + sudo apt-get update -y -o Acquire::Retries=3 + sudo apt-get install -y --no-install-recommends \ + build-essential libssl-dev uuid-dev libgpgme11-dev squashfs-tools + + wget -q -O singularity.deb \ + "https://github.com/sylabs/singularity/releases/download/v4.2.1/singularity-ce_4.2.1-${UBUNTU_CODENAME}_amd64.deb" + # Must pass -y: bare "apt install ./singularity.deb" waits for confirmation and hangs CI. + sudo apt-get install -y --no-install-recommends ./singularity.deb singularity --version - name: Cache Cerebras SDK diff --git a/tests/csl_runtime/run_tests.sh b/tests/csl_runtime/run_tests.sh index 20cab76e..20d7b19e 100755 --- a/tests/csl_runtime/run_tests.sh +++ b/tests/csl_runtime/run_tests.sh @@ -1,5 +1,10 @@ #!/bin/bash +# Every test builds in this directory and exchanges data through the fixed names inp.npy and +# OUT_out.npy, removing them once a case is done. Only one run may be active per checkout: two +# concurrent runs overwrite each other's inputs and fail with a shape mismatch. Run architectures +# sequentially, or give each one its own checkout. + # Color codes for output RED='\033[0;31m' GREEN='\033[0;32m' From c269eb0de9f3c8a350c6febbbe122d8b458df5c7 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 12:50:35 +0200 Subject: [PATCH 41/68] Skip batcher L=4 case --- tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh b/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh index 079c6bf0..00cb5025 100755 --- a/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh +++ b/tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh @@ -5,8 +5,10 @@ # three widest phases run on two colors each instead of one per comparator. # L <= 4: three phases is what the filter budget allows, and at L = 5 the phases left unbundled # need more than the 21 colors. On wse2 the ceiling is L = 3: a reused inbound color can stay live -# across a gap that already holds two others, and a PE has only two input queues. wse3 has six, so -# L = 4 still runs there. K does not change either count; it widens each destination's +# across a gap that already holds two others, and a PE has only two input queues. WSE-3 has six +# input queues, but it cannot remap one onto another color while wavelets remain, so a PE needs as +# many queues as inbound colors over the kernel. L = 4 wants seven, which is one more than the +# pool. K does not change either count; it widens each destination's # counter-filter window to K wavelets out of a cycle of M*K, so K > 1 is what exercises that # window on hardware. @@ -55,8 +57,9 @@ run_batcher 1 1 run_batcher 2 2 run_batcher 3 1 run_batcher 3 2 +run_batcher 3 4 if [ "${WSE_ARCH:-wse2}" = "wse3" ]; then - run_batcher 4 2 + echo "Skipping L=4: seven inbound colors, and wse3 cannot remap a non-empty input queue." else echo "Skipping L=4: three inbound channel spans overlap, and wse2 has two input queues." fi From d29176e44b32678042ea9f4a3d58913edc41f707 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 13:36:43 +0200 Subject: [PATCH 42/68] Fixes for WSE3: set .ut_id so microthreads can be shared independently of queues. --- spada/lowering/spatial_ir_to_csl.py | 32 ++++++++++++--- spada/syntax/csl/constants.py | 14 +++++++ spada/syntax/csl/dsd_ops.py | 11 ++++-- spada/syntax/csl/structures.py | 5 +++ spada/syntax/spatial_ir/stream_lifetime.py | 46 ++++++++++++++++++++++ tests/spatial_ir/test_dsd_ops.py | 25 ++++++++++++ tests/spatial_ir/test_stream_lifetime.py | 24 +++++++++++ 7 files changed, 149 insertions(+), 8 deletions(-) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 0b7ef25b..a6150c73 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -1123,6 +1123,18 @@ def queue_key(stream: spir.Identifier) -> str: output_spans, csl.OUTPUT_QUEUE_IDS, kind='output', architecture=csl.ARCH, location=location, exclusive_keys=exclusive_out) + # An asynchronous transfer runs on a microthread, and by default that is the queue ID of the + # operation's highest-priority fabric operand. Since the two directions draw from overlapping + # pools on WSE-3, a receive on input queue N and a send on output queue N would take the same + # microthread and abort with "trying to term ut_instr[N], but it's not ours". Microthreads are + # one resource across both directions, so they are handed out together. + microthread_of = stream_lifetime.assign_microthreads( + {f'in {key}': span for key, span in input_spans.items()} + | {f'out {key}': span for key, span in output_spans.items()}, + csl.MICROTHREAD_IDS, location=location) + + def allocate_microthread(stream: spir.Identifier, inbound: bool) -> int | None: + return microthread_of.get(f'{"in" if inbound else "out"} {queue_key(stream)}') def allocate_input_queue(stream: spir.Identifier) -> int: key = queue_key(stream) @@ -1178,7 +1190,9 @@ def _visit_foreach(stmt: spir.ForeachStatement) -> None: stmt.parameter_range[0].step) extents = (end.eval() - start.eval()) // (step.eval() if step is not None else 1) fabric_color = f'{name_to_csl(stream_name)}_color' - dsd = cslstruct.FabricDSD(cslstruct.DSDType.fabin, fabric_color, extents, allocate_input_queue(stream_name)) + dsd = cslstruct.FabricDSD(cslstruct.DSDType.fabin, fabric_color, extents, + allocate_input_queue(stream_name), + ut=allocate_microthread(stream_name, inbound=True)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) for stmt in rect.compute.statements: @@ -1205,7 +1219,9 @@ def _visit_foreach(stmt: spir.ForeachStatement) -> None: lambda a, b: a * b, [s.eval() if not isinstance(s, int) else s for s in dtypes[stmt.local_array].shape], 1) fabric_color = f'{name_to_csl(stream_name)}_color' - dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, allocate_input_queue(stream_name)) + dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, + allocate_input_queue(stream_name), + ut=allocate_microthread(stream_name, inbound=True)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) elif isinstance(stmt, spir.SendStatement) and stream_name.as_ir() in stream_candidates: dsd_type = cslstruct.DSDType.fabout @@ -1223,7 +1239,9 @@ def _visit_foreach(stmt: spir.ForeachStatement) -> None: lambda a, b: a * b, [s.eval() if not isinstance(s, int) else s for s in dtypes[stmt.local_array].shape], 1) fabric_color = f'{name_to_csl(stream_name)}_color' - dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, allocate_output_queue(stream_name)) + dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, + allocate_output_queue(stream_name), + ut=allocate_microthread(stream_name, inbound=False)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) if isinstance(stmt, spir.SendStatement): @@ -1274,7 +1292,9 @@ def _visit_nested_send(substmt: spir.SendStatement): lambda a, b: a * b, [s.eval() if not isinstance(s, int) else s for s in dtypes[substmt.local_array].shape], 1) fabric_color = f'{name_to_csl(stream_name)}_color' - dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, allocate_output_queue(stream_name)) + dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, + allocate_output_queue(stream_name), + ut=allocate_microthread(stream_name, inbound=False)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) def _visit_nested_receive(substmt: spir.ReceiveStatement): @@ -1300,7 +1320,9 @@ def _visit_nested_receive(substmt: spir.ReceiveStatement): lambda a, b: a * b, [s.eval() if not isinstance(s, int) else s for s in dtypes[local_array].shape], 1) fabric_color = f'{name_to_csl(stream_name)}_color' - dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, allocate_input_queue(stream_name)) + dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, + allocate_input_queue(stream_name), + ut=allocate_microthread(stream_name, inbound=True)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) def _visit_local_array(operand): diff --git a/spada/syntax/csl/constants.py b/spada/syntax/csl/constants.py index 6274203e..ffe37be1 100644 --- a/spada/syntax/csl/constants.py +++ b/spada/syntax/csl/constants.py @@ -63,6 +63,20 @@ } OUTPUT_QUEUE_IDS = _OUTPUT_QUEUE_IDS[ARCH] +# Microthreads that drive in-flight asynchronous DSD operations. Two operations may never run on +# one microthread at the same time. WSE-2 has no say in the matter: the ID is the queue ID of the +# operation's highest-priority fabric operand, which is why its input and output pools above are +# disjoint. WSE-3 keeps that default but lets ``.ut_id`` override it, which it must, since a PE +# there needs an input and an output queue of the same number at once (see +# https://sdk.cerebras.net/csl/language/microthreads_wse3). An empty list means the target cannot +# name microthreads, so the default stands. Queues 0 and 1 belong to memcpy on WSE-3, and so do the +# microthreads it drives them with. +_MICROTHREAD_IDS = { + 'wse2': [], + 'wse3': list(range(2, 8)), +} +MICROTHREAD_IDS = _MICROTHREAD_IDS[ARCH] + _HARDWARE_FABRIC_DIMS = { 'wse2': (757, 996), 'wse3': (762, 1172), diff --git a/spada/syntax/csl/dsd_ops.py b/spada/syntax/csl/dsd_ops.py index f686c50d..3cecf599 100644 --- a/spada/syntax/csl/dsd_ops.py +++ b/spada/syntax/csl/dsd_ops.py @@ -26,9 +26,14 @@ def _append_async_suffix(self, base: str, dsd_objects: list[cslstruct.DataStruct # CSL allows .async only when every operand is a DSD/DSR. A scalar # source (CopyDSDOp.scalar_input) must complete synchronously, then # activate/unblock the next task. - fabric_async = any(isinstance(dsd, cslstruct.FabricDSD) for dsd in dsd_objects) - if fabric_async and not getattr(self, 'scalar_input', False): - return f'{base[:-2]}, .{{ .async = true, .{async_target.inter_task_edge} = {async_target.target_task} }});' + fabric_operands = [dsd for dsd in dsd_objects if isinstance(dsd, cslstruct.FabricDSD)] + if fabric_operands and not getattr(self, 'scalar_input', False): + # ``dsd_objects`` arrives in the order the hardware ranks operands when it picks the + # microthread for the transfer: destination, then the sources left to right. + microthread = fabric_operands[0].ut + ut_id = '' if microthread is None else f' .ut_id = @get_ut_id({microthread}),' + return (f'{base[:-2]}, .{{ .async = true,{ut_id} ' + f'.{async_target.inter_task_edge} = {async_target.target_task} }});') return f'{base}\n@{async_target.inter_task_edge}({async_target.target_task});' def as_csl(self, diff --git a/spada/syntax/csl/structures.py b/spada/syntax/csl/structures.py index ab678ea5..fffc955d 100644 --- a/spada/syntax/csl/structures.py +++ b/spada/syntax/csl/structures.py @@ -61,6 +61,11 @@ class FabricDSD(DataStructureDescriptor): extent: int queue: int control: bool = False + #: Microthread to drive an asynchronous transfer over this descriptor, where the target lets a + #: program name one. ``None`` leaves the hardware default, which is the queue ID. The setting + #: belongs to the operation rather than the descriptor, so ``as_csl`` does not emit it; the + #: operand carries it to whichever operation uses it. + ut: int | None = None def __post_init__(self): assert self.dsd_type in (DSDType.fabin, DSDType.fabout) diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py index 7f856c66..cbb691db 100644 --- a/spada/syntax/spatial_ir/stream_lifetime.py +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -812,6 +812,52 @@ def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int] return assigned +def assign_microthreads(spans: dict[str, tuple[int, int]], microthread_ids: list[int], *, + location: str) -> dict[str, int]: + """ + Assigns microthreads to stream groups from their occupancy spans on one PE. + + A microthread is held only for the lifetime of one asynchronous operation, so unlike a fabric + queue it needs no proof that the hardware has drained: groups whose spans do not overlap take + turns on one microthread. Callers pass the inbound and outbound groups of a PE together, keyed + apart by direction, because a microthread is one resource shared by both directions. + + :param spans: Mapping of grouping key to an inclusive ``(first_use, last_use)`` statement index + pair on this PE, over both directions. + :param microthread_ids: The microthread identifiers a program may name, in the order they should + be handed out. + :param location: The PE rectangle, for the diagnostic. + :return: Mapping of grouping key to a microthread identifier, empty when the target cannot name + microthreads and the hardware default has to stand. + """ + if not spans or not microthread_ids: + return {} + + assigned: dict[str, int] = {} + for key in sorted(spans, key=lambda name: (spans[name][0], spans[name][1], name)): + start, end = spans[key] + used = { + assigned[other] + for other in assigned + if start <= spans[other][1] and spans[other][0] <= end + } + for microthread in microthread_ids: + if microthread not in used: + assigned[key] = microthread + break + else: + overlapping = sorted( + other for other, (other_start, other_end) in spans.items() + if other != key and start <= other_end and other_start <= end) + raise SyntaxError( + f'{location} would need {len(used) + 1} concurrent microthreads ' + f'(live groups {[key] + overlapping}), but a PE may name at most ' + f'{len(microthread_ids)}.\n' + f' note: every asynchronous transfer in flight holds one microthread, counting ' + f'both directions') + return assigned + + ### # Optimization passes ### diff --git a/tests/spatial_ir/test_dsd_ops.py b/tests/spatial_ir/test_dsd_ops.py index 77e9b142..c153bb9a 100644 --- a/tests/spatial_ir/test_dsd_ops.py +++ b/tests/spatial_ir/test_dsd_ops.py @@ -294,6 +294,31 @@ def test_wse3_inbound_colors_do_not_share_an_input_queue(): assert len(used_queues) == len(set(used_queues)), (by_color, code) +def test_wse3_concurrent_transfers_use_distinct_microthreads(): + """Two transfers in flight at once may not share a microthread. + + A laplacian PE receives from one neighbour and forwards to another in the same task. On WSE-3 + the input and output queue pools both start at 2, so leaving the microthread at its default -- + the queue ID of the highest-priority fabric operand -- put both on microthread 2 and aborted the + simulation with ``trying to term ut_instr[2], but it's not ours``. + """ + if constants.ARCH != 'wse3': + pytest.skip('WSE-2 derives the microthread from the queue, and its two pools are disjoint') + path = os.path.join( + os.path.dirname(__file__), '..', '..', 'samples', 'benchmarks', 'laplacian_4_4_4.sptl') + kernel = parser.parse_file(path) + files = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_benchmarking=True)} + code = files['code_2_1.csl'] + + tasks = re.findall(r'task \w+\(\) void \{(.*?)\n\}', code, re.DOTALL) + concurrent = [ + re.findall(r'\.ut_id = @get_ut_id\((\d+)\)', body) for body in tasks + ] + assert any(len(used) > 1 for used in concurrent), code + for used in concurrent: + assert len(used) == len(set(used)), (used, code) + + if __name__ == '__main__': test_dsd_op_detection() test_dsd_op_detection_constant_folding() diff --git a/tests/spatial_ir/test_stream_lifetime.py b/tests/spatial_ir/test_stream_lifetime.py index d8a82a9a..3af2b2f7 100644 --- a/tests/spatial_ir/test_stream_lifetime.py +++ b/tests/spatial_ir/test_stream_lifetime.py @@ -513,5 +513,29 @@ def test_exclusive_keys_exhaust_queues_when_too_many_data_tasks(): exclusive_keys=frozenset({'channel 0', 'channel 1', 'channel 2'})) +def test_microthreads_are_shared_across_directions_when_spans_are_disjoint(): + """A microthread is held only while a transfer is in flight, so turns may be taken.""" + assigned = stream_lifetime.assign_microthreads( + {'in channel 0': (0, 2), 'out channel 1': (3, 5)}, [2, 3], location='PE (0, 0)') + assert assigned == {'in channel 0': 2, 'out channel 1': 2} + + +def test_a_receive_and_a_send_in_flight_together_get_distinct_microthreads(): + assigned = stream_lifetime.assign_microthreads( + {'in channel 0': (0, 4), 'out channel 0': (0, 4)}, [2, 3], location='PE (0, 0)') + assert assigned['in channel 0'] != assigned['out channel 0'] + + +def test_microthreads_run_out_when_too_many_transfers_overlap(): + with pytest.raises(SyntaxError, match='concurrent microthreads'): + stream_lifetime.assign_microthreads( + {'in channel 0': (0, 10), 'out channel 0': (2, 8), 'out channel 1': (4, 6)}, [2, 3], + location='PE (0, 0)') + + +def test_microthreads_are_left_to_the_hardware_when_the_target_cannot_name_them(): + assert stream_lifetime.assign_microthreads({'in channel 0': (0, 2)}, [], location='PE (0, 0)') == {} + + if __name__ == '__main__': pytest.main([__file__]) From 5364a1548170f1b36b0166acca346bdaf1b14021 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 15:13:57 +0200 Subject: [PATCH 43/68] Skip tests marked pending --- tests/csl_runtime/run_tests.sh | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/csl_runtime/run_tests.sh b/tests/csl_runtime/run_tests.sh index 20d7b19e..92870c6f 100755 --- a/tests/csl_runtime/run_tests.sh +++ b/tests/csl_runtime/run_tests.sh @@ -42,9 +42,9 @@ NON_TEST_SCRIPTS=("run_tests.sh" "run-in-lima.sh" "sptlc" "_lib.sh") is_non_test() { local name="$1" - # Local debug helpers (zz_*) are not part of the suite. + # Local debug helpers (zz_*) and not-yet-enabled cases (pending_*) are not part of the suite. case "$name" in - zz_*) return 0 ;; + zz_*|pending_*) return 0 ;; esac for skip in "${NON_TEST_SCRIPTS[@]}"; do [ "$name" = "$skip" ] && return 0 From c89685e8e7fdcdfeddbcc7c2c54d54f1b6a1f7b9 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 17:05:16 +0200 Subject: [PATCH 44/68] Improved WSE-3 sample, improved color allocation, microthread assignment --- README.md | 2 +- irspec/docs/spatial/routing_wse.md | 22 +- .../sort/batcher_oddeven_bundled_1D.sptl | 13 +- .../spatial/sort/batcher_oddeven_wse3_1D.sptl | 261 ++++++++++++++++++ samples/spatial/sort/plot_batcher_routing.py | 122 +++++--- spada/lowering/spatial_ir_to_csl.py | 68 ++++- spada/syntax/spatial_ir/stream_lifetime.py | 38 +-- .../test_batcher_oddeven_wse3_1d.sh | 64 +++++ tests/spatial_ir/test_shift_bundles.py | 66 ++++- tests/spatial_ir/test_stream_lifetime.py | 23 +- 10 files changed, 598 insertions(+), 81 deletions(-) create mode 100644 samples/spatial/sort/batcher_oddeven_wse3_1D.sptl create mode 100755 tests/csl_runtime/test_batcher_oddeven_wse3_1d.sh diff --git a/README.md b/README.md index 2600bfbb..bb12915f 100644 --- a/README.md +++ b/README.md @@ -104,7 +104,7 @@ Sample SPADA programs are in `samples/`: | `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy` | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Sorting networks over `2^L` PEs holding `K` keys each: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `bitonic_sort_1D` (channel reuse via router switches, WSE-3), and `odd_even_sort_1D_looped` (N neighbour rounds as a runtime loop, four static channels) | +| `samples/spatial/sort/` | Sorting networks over `2^L` PEs holding `K` keys each: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), `bitonic_sort_1D` (channel reuse via router switches, WSE-3), and `odd_even_sort_1D_looped` (N neighbour rounds as a runtime loop, four static channels) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index a30de64f..c0d83401 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -198,11 +198,23 @@ its own is how a kernel declines the trade. What it then costs is colors, and th by reusing a channel across phases. Two unbundled shifts may share one safely when they agree on axis and signed distance and their sources agree modulo twice that distance, because a PE's role — source, relay or destination — is then a function of its position modulo twice the distance alone, -so one static configuration serves every phase in the pool. Sharing on any other basis risks a PE -that sends on the color in one phase and receives on it in another, which needs a two-sided switch -change that a sender cannot drive (see [Lowering to Switches](#lowering-to-switches)), and nothing -in the compiler currently rejects it. `batcher_oddeven_bundled_1D.sptl` pools on exactly this rule, -and `batcher_oddeven_1D.sptl` is the same rule written out as arithmetic. +so one static configuration serves every phase in the pool. `batcher_oddeven_bundled_1D.sptl` pools +on exactly this rule, and `batcher_oddeven_1D.sptl` is the same rule written out as arithmetic. + +The agreement on the *sign* of the distance can be dropped without giving that up. Keep the axis, the +magnitude and the source residue modulo twice it, and let the direction of travel vary: the sources +are then the PEs congruent to the residue, the destinations those congruent to residue plus distance, +and the relays the classes strictly between on the one side or the other — three disjoint classes, so +a PE still holds one role on the color for the whole kernel and still never both sends and receives +on it. What varies is the side it faces, which is one switch position either way, since a source only +ever changes where it transmits and a destination only where it receives. This halves the colors a +pooled distance needs, and with them the queues, which is what +`batcher_oddeven_wse3_1D.sptl` is for: on WSE-3 a queue stays bound to its color for the whole +kernel, so what a PE can afford is not how many colors are live at once but how many it ever touches. + +Sharing on a basis looser than either does risk a PE that sends on the color in one phase and receives +on it in another, which needs a two-sided switch change that a sender cannot drive on WSE-2 (see +[Lowering to Switches](#lowering-to-switches)), and nothing in the compiler currently rejects it. This arrangement is the one used in Schnyder's *Distributed Sorting on the Cerebras Wafer-Scale Engine* (fig. 7.6) for the 2D reduce-scatter, and is known to run on WSE-2. diff --git a/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl b/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl index aab1758b..61b2b1c9 100644 --- a/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl +++ b/samples/spatial/sort/batcher_oddeven_bundled_1D.sptl @@ -56,9 +56,10 @@ * The phases left unbundled are routed per matching, as in batcher_oddeven_1D, but their colors * are pooled across phases. Two of them may share when they agree on direction, distance, and * the source residue mod 2d, because a PE's role is then decided by its position mod 2d alone - * and one static configuration serves every phase in the pool. Sharing on any looser rule would - * let a PE send on a color in one phase and receive on it in another, which needs two switch - * advances from a sender that can only emit one. + * and one static configuration serves every phase in the pool -- no router here switches at all, + * which is what makes this the variant to use on wse2. Dropping the agreement on direction pools + * twice as tightly, at the price of switching routers; that is batcher_oddeven_wse3_1D.sptl, which + * is what reaches L = 4 on wse3, where a queue is bound to its color for the whole kernel. * * bundled (4*d >= N) : fwd N + 2*(l*(L+1) + p), bwd fwd + 1 * pooled : fwd 2*((2*d - 2) + c), bwd fwd + 1 @@ -66,8 +67,10 @@ * with c = r for p = 1 and c = d + r for p >= 2; the pooled blocks for successive d are disjoint * because the block for d runs from 2*d-2 to 4*d-3, and unbundled phases have 4*d < N, so every * pooled channel stays below N. Colors: 10 at L = 3 and 18 at L = 4, of the 21 available. L = 5 - * would need 34, which is what caps this kernel. On wse2 an interior PE at L = 4 has three inbound - * colors live at once, and a PE has two input queues, so L <= 3 there; wse3 has six queues. + * would need 34, which is what caps this kernel. Queues cap it earlier than that on either target, + * at L = 3: on wse2 an interior PE at L = 4 has three inbound colors live at once and a PE has two + * input queues, and while wse3 has six, it binds each to its color for the whole kernel, and such a + * PE touches seven. batcher_oddeven_wse3_1D.sptl is the variant that gets L = 4 there. * * Example L=3 (n=8): * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) diff --git a/samples/spatial/sort/batcher_oddeven_wse3_1D.sptl b/samples/spatial/sort/batcher_oddeven_wse3_1D.sptl new file mode 100644 index 00000000..8bbd8b23 --- /dev/null +++ b/samples/spatial/sort/batcher_oddeven_wse3_1D.sptl @@ -0,0 +1,261 @@ +/** + * 1D Batcher odd-even mergesort over N = 2^L PEs, each holding K f32 keys: N*K keys in all. + * Ascending: after the network, PE i holds keys i*K .. i*K + (K-1) of the sorted sequence. + * + * Same network, same compare-split arithmetic and same result as batcher_oddeven_bundled_1D. What + * differs is how the unbundled phases pool their channels, and it exists because that is what stands + * between the bundled variant and L = 4 on WSE-3. + * + * Why the bundled variant stops short on WSE-3 + * ------------------------------------------- + * WSE-3 binds a fabric queue to a color for the whole kernel: a queue may not be remapped onto + * another color, because nothing can prove the wavelets of the old one have drained. A PE therefore + * needs one input queue per color it *ever* receives on and one output queue per color it ever sends + * on, and it has six of each -- five for sending, since memcpy's copy-back of `out` reserves one. + * The bundled variant's busiest PE at L = 4 wants seven of each: + * + * 2 d = 1, one color for the p = 1 phase and one for the p >= 2 phases + * 2 d = 2, likewise + * 3 one per bundled phase: (l=3,p=1), (l=4,p=1), (l=4,p=2) + * + * Bundling cannot take that further. It saves *colors*, and a PE's three filters cap it at three + * phases, but a bundled color cannot be reused by a later phase either -- its sources end in relay + * mode and its destination filters are set at layout time and never reprogrammed -- so each bundled + * phase costs its participants a queue in each direction, exactly as an unbundled one does. + * + * Pooling by origin instead of by direction + * ----------------------------------------- + * What costs two queues per unbundled distance is pooling the forward and backward streams + * separately. Two phases may share a channel there when they agree on direction, distance and source + * residue mod 2d, which keeps every router's configuration fixed for the whole kernel but makes a + * comparator's two messages two colors. + * + * Drop the agreement on direction and keep the rest: a color becomes (distance d, origin residue + * mod 2d), whichever way the message travels. On such a color the three roles fall into disjoint + * residue classes -- sources are the PEs congruent to s, destinations those congruent to s + d, and + * relays the classes strictly between, on one side or the other -- so a PE holds one role on a color + * for the whole kernel and never both sends and receives on it. Only the side it faces varies: + * + * source R->E or R->W two positions, the transmit side alone changes + * destination W->R or E->R two positions, the receive side alone changes + * relay W->E or E->W one position; a relay class is reached from one side only + * + * Two positions of the four a router holds, and never a both-sides change, so this asks for nothing + * WSE-2 lacks either. Every router on a retired path advances once per epoch boundary, and the ones + * with a single position sit on their last, where a control message passing through over-advances + * them harmlessly. Each PE now spends one input and one output queue per unbundled distance instead + * of two, which is what brings L = 4 within budget: + * + * 5 in and 5 out at L = 4, against 6 and 5. + * + * Channels + * -------- + * Bundling still pays where the matchings overlap, and the rule that picks the phases is unchanged: + * 4*d >= N, the three widest, one per filter. The rest pool by origin. + * + * bundled (4*d >= N) : fwd N + 2*(l*(L+1) + p), bwd fwd + 1 + * pooled : fwd (2*d - 2) + c_lo, bwd (2*d - 2) + c_hi + * + * where c_lo is the low partner's residue mod 2d and c_hi the high partner's -- c_lo = r and + * c_hi = r + d for p = 1, the other way round for p >= 2, which is where the two phases at one + * distance meet. The block for d runs from 2*d-2 to 4*d-3 and pooled phases have 4*d < N, so every + * pooled channel stays below N. + * + * Colors: 8 at L = 3 and 12 at L = 4, against 10 and 18 for the bundled variant. All of them switch, + * and WSE-3 implements switches on fifteen of its twenty-one colors, which is the next thing L = 5 + * would run out of -- along with a sixth output queue. + * + * Requires WSE-3 (!) + * ------------------ + * Not for the routing, but for the queues: this kernel is written for a target that reuses none. On + * WSE-2 a queue comes free again once its occupancy span ends, so what binds there is how many colors + * are live at once rather than how many a PE ever touches, and the bundled variant, whose routers + * never switch at all, is the better fit. + * + * Example L=3 (n=8), with the channel each phase's two directions land on: + * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) fwd 0 bwd 1 + * (l=2,p=1) dist 2: (0,2)(1,3)(4,6)(5,7) bundled + * (l=2,p=2) dist 1: (1,2)(5,6) fwd 1 bwd 0 + * (l=3,p=1) dist 4: (0,4)(1,5)(2,6)(3,7) bundled + * (l=3,p=2) dist 2: (2,4)(3,5) bundled + * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) fwd 1 bwd 0 + * + * Constraints: 1 <= L <= 4, K >= 1. WSE-3. + **/ +kernel @batcher_oddeven_wse3_1d( + stream[1<[1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + } + } + + for i16 l in [1:L+1] { + // p = 1: all PEs participate, dist = 1<<(l-1). The low partner of matching r has residue r + // mod 2*dist and the high partner r + dist, which is the channel each of them sends on. + phase { + for i16 r in [0:1<<(l-1)] { + dataflow i16 i, i16 j in [r:1< fwd = relative_stream(1<<(l-1), 0) { + hops = auto, + channel = ((1<= (1< bwd = relative_stream(-(1<<(l-1)), 0) { + hops = auto, + channel = (((1<= (1< fwd = relative_stream(1<<(l-1), 0) { + hops = auto, + channel = ((1<= (1< bwd = relative_stream(-(1<<(l-1)), 0) { + hops = auto, + channel = (((1<= (1<= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + } + + // p = 2 .. l: skip first/last dist PEs of each 2^l box, dist = 1<<(l-p). The low partner's + // residue is dist + r here and the high partner's is r, the reverse of the p = 1 phases, + // which is what lets the two of them share these channels. + for i16 p in [2:l+1] { + phase { + for i16 r in [0:1<<(l-p)] { + for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { + hops = auto, + channel = ((1<= (1< bwd = relative_stream(-(1<<(l-p)), 0) { + hops = auto, + channel = (((1<= (1< fwd = relative_stream(1<<(l-p), 0) { + hops = auto, + channel = ((1<= (1< bwd = relative_stream(-(1<<(l-p)), 0) { + hops = auto, + channel = (((1<= (1<= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + } + } + } + } + + // Write the sorted blocks back to the host. + phase { + compute i16 i, i16 j in [0:1<= n. The rest are routed per matching, on colors - pooled across phases by the source residue mod 2d, so that a PE's - role on a pooled color is the same in every phase that uses it. + pooled across phases by the source residue mod 2d *and* the + direction, so that a PE's role on a pooled color is the same in + every phase that uses it and no router ever switches: + fwd 2*((2*d - 2) + c), bwd fwd + 1, c = r or d + r + wse3 batcher_oddeven_wse3_1D.sptl. Bundles the same three phases, and + pools the rest by the origin's residue alone, dropping the + agreement on direction: + fwd (2*d - 2) + c_lo, bwd (2*d - 2) + c_hi + A comparator's two messages then share a pooled color whenever + two phases at one distance meet on it, which halves how many + colors a PE touches -- what WSE-3 counts, since a queue there + stays bound to its color for the whole kernel. The price is that + pooled routers switch: a source alternates R->E and R->W, a + destination W->R and E->R, both one side at a time. Views ----- @@ -30,8 +42,9 @@ Each comparator is two arrows: down = fwd (east, +d), up = bwd (west, -d), each colored by its channel. table Per-PE @set_color_config, with switch positions resolved the way - WSE-2 stores them (a both-sides change is split through a relay - intermediate). Stacked bands are pos0, pos1, … in that order. + the version's target stores them: on WSE-2 a both-sides change is + split through a relay intermediate, on WSE-3 it is one position. + Stacked bands are pos0, pos1, … in that order. A destination filter is the small ``fN`` in the cell, N being the filter's init_counter; ``--k`` widens the windows the way K keys per PE do, which scales every init_counter by K. @@ -42,6 +55,7 @@ python samples/spatial/sort/plot_batcher_routing.py --n 8 --version bundled --view table python samples/spatial/sort/plot_batcher_routing.py --n 8 --version static bundled --view table python samples/spatial/sort/plot_batcher_routing.py --n 8 --version hybrid --view table --k 4 + python samples/spatial/sort/plot_batcher_routing.py --n 16 --version wse3 --view table """ from __future__ import annotations @@ -59,10 +73,13 @@ from matplotlib.patches import Rectangle -VERSIONS = ("static", "bundled", "hybrid") +VERSIONS = ("static", "bundled", "hybrid", "wse3") MIN_BUNDLE_DISTANCE = 2 MIN_BUNDLE_LENGTH = 2 +# The architecture each version is written for, which decides how a both-sides change is resolved. +TARGET_ARCH = {"static": "WSE-2", "bundled": "WSE-2", "hybrid": "WSE-2", "wse3": "WSE-3"} + TX_ORDER = ("RAMP", "EAST", "WEST") DIR_LETTER = {"RAMP": "R", "EAST": "E", "WEST": "W"} @@ -170,7 +187,7 @@ def is_bundled(phase: Phase, n: int, version: str) -> bool: return False if version == "bundled": return True - return 4 * phase.dist >= n # hybrid: the three widest phases, one per filter + return 4 * phase.dist >= n # hybrid and wse3: the three widest phases, one per filter def assign_channels(phases: list[Phase], version: str, n: int) -> list[Phase]: @@ -182,16 +199,23 @@ def assign_channels(phases: list[Phase], version: str, n: int) -> list[Phase]: for ph in phases: matchings = [] for m in ph.matchings: + low_residue = m.offset if ph.p == 1 else ph.dist + m.offset + high_residue = (low_residue + ph.dist) % (2 * ph.dist) if version == "bundled": - fwd = 2 * ph.index + fwd, bwd = 2 * ph.index, 2 * ph.index + 1 elif is_bundled(ph, n, version): fwd = n + 2 * (ph.l * (log_n + 1) + ph.p) + bwd = fwd + 1 + elif version == "wse3": + # Pooled by the origin's residue mod 2d alone: each direction takes the colour of + # the partner that sends it, so the two phases at one distance meet on both. + fwd = (2 * ph.dist - 2) + low_residue + bwd = (2 * ph.dist - 2) + high_residue else: - # Pooled: the source residue mod 2d decides the colour, so phases that agree on - # it agree on every router configuration and can share. - residue = m.offset if ph.p == 1 else ph.dist + m.offset - fwd = 2 * ((2 * ph.dist - 2) + residue) - matchings.append(replace(m, fwd=fwd, bwd=fwd + 1)) + # Pooled by direction as well, which keeps every router configuration static. + fwd = 2 * ((2 * ph.dist - 2) + low_residue) + bwd = fwd + 1 + matchings.append(replace(m, fwd=fwd, bwd=bwd)) assigned.append(replace(ph, matchings=tuple(matchings))) return _compact_channels(assigned) @@ -211,6 +235,40 @@ def _compact_channels(phases: list[Phase]) -> list[Phase]: for ph in phases] +def _channel_labels(phases: list[Phase], n_channels: int) -> list[str]: + """ + Legend text for each channel: what it carries and which phases put it there. + + Read off the assignment rather than assumed, since which direction a channel carries is what + the versions disagree about -- pooling by direction gives every channel one of them, pooling by + origin gives the shared ones both. + """ + distances: list[set[int]] = [set() for _ in range(n_channels)] + directions: list[set[str]] = [set() for _ in range(n_channels)] + users: list[list[str]] = [[] for _ in range(n_channels)] + for ph in phases: + for m in ph.matchings: + if not m.pairs: + continue + for channel, direction in ((m.fwd, "fwd"), (m.bwd, "bwd")): + distances[channel].add(ph.dist) + directions[channel].add(direction) + if f"l{ph.l}p{ph.p}" not in users[channel]: + users[channel].append(f"l{ph.l}p{ph.p}") + + labels = [] + for ch in range(n_channels): + if not directions[ch]: + labels.append(f" ch {ch} (unused)") + continue + arrow = {("fwd", ): "↓", ("bwd", ): "↑"}.get(tuple(sorted(directions[ch])), "↕") + kind = "+".join(sorted(directions[ch])) + dist = ",".join(str(d) for d in sorted(distances[ch])) + phase_text = " ".join(users[ch]) if len(users[ch]) <= 3 else f"{len(users[ch])} phases" + labels.append(f"{arrow} ch {ch}: d={dist} {kind} ({phase_text})") + return labels + + def channel_count(phases: list[Phase]) -> int: used = [ch for ph in phases for m in ph.matchings for ch in (m.fwd, m.bwd)] return (max(used) + 1) if used else 0 @@ -285,17 +343,10 @@ def _draw_network(ax, phases: list[Phase], n: int, version: str) -> None: ax.plot([x_bwd - cap, x_bwd + cap], [lo, lo], color=bwd_c, lw=1.6, zorder=3) ax.plot([x_bwd - cap, x_bwd + cap], [hi, hi], color=bwd_c, lw=1.6, zorder=3) - handles = [] - for ch in range(n_channels): - arrow = "↓" if ch % 2 == 0 else "↑" - kind = "fwd" if ch % 2 == 0 else "bwd" - if version == "bundled": - label = f"{arrow} ch {ch} (p{ch // 2} {kind})" - else: - label = f"{arrow} ch {ch} ({kind})" - handles.append( - Line2D([0], [0], color=channel_color(ch, n_channels), lw=2.0, label=label) - ) + handles = [ + Line2D([0], [0], color=channel_color(ch, n_channels), lw=2.0, label=label) + for ch, label in enumerate(_channel_labels(phases, n_channels)) + ] ax.legend( handles=handles, title="channel", @@ -329,14 +380,19 @@ def _add(seq: list[Route], config: Route) -> None: seq.append(config) -def _resolve_hardware(configs: list[Route]) -> tuple[Route, ...]: - """WSE-2: a position names one side; a both-sides change is split in two.""" +def _resolve_hardware(configs: list[Route], split_both_sides: bool = True) -> tuple[Route, ...]: + """ + The positions the hardware stores for a sequence of configurations. + + On WSE-2 a position names one side, so a change of both is split through an intermediate that + keeps the old input; WSE-3 takes both in one position and needs no splitting. + """ if not configs: return () positions = [configs[0]] for config in configs[1:]: previous = positions[-1] - if previous.rx != config.rx and previous.tx != config.tx: + if split_both_sides and previous.rx != config.rx and previous.tx != config.tx: positions.append(Route(previous.rx, config.tx)) positions.append(config) return tuple(positions) @@ -439,8 +495,9 @@ def pe_table(phases: list[Phase], n: int, version: str, words: int = 1) -> list[ for lo, hi in m.pairs: _install_ordinary_pair(configs, lo, hi, m.fwd, m.bwd) + split = TARGET_ARCH[version] == "WSE-2" return [ - [Cell(_resolve_hardware(configs[pe][ch]), filters[pe][ch]) for ch in range(n_channels)] + [Cell(_resolve_hardware(configs[pe][ch], split), filters[pe][ch]) for ch in range(n_channels)] for pe in range(n) ] @@ -456,8 +513,8 @@ def _draw_table(ax, phases: list[Phase], n: int, version: str, words: int) -> No ax.set_xlabel("Channel") ax.set_ylabel("PE") ax.set_title( - f"Resolved switch positions ({version}, WSE-2), n={n}, K={words} ({n_channels} colors); " - "stacked bands are pos0, pos1, ...; fN is the filter init_counter" + f"Resolved switch positions ({version}, {TARGET_ARCH[version]}), n={n}, K={words} " + f"({n_channels} colors); stacked bands are pos0, pos1, ...; fN is the filter init_counter" ) ax.set_aspect("equal") @@ -572,9 +629,10 @@ def main() -> None: nargs="+", choices=VERSIONS + ("all",), default=["static"], - help="static (one color per matching), bundled (two per phase) and/or hybrid (the " - "sample: the three widest phases bundled, the rest pooled). 'all' is every one. " - "Repeatable.", + help="static (one color per matching), bundled (two per phase), hybrid (the WSE-2 sample: " + "the three widest phases bundled, the rest pooled by direction and origin) and/or wse3 " + "(the WSE-3 sample: the same bundles, the rest pooled by origin alone, so the pooled " + "routers switch). 'all' is every one. Repeatable.", ) parser.add_argument( "--view", diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index a6150c73..40d45ce9 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -688,22 +688,27 @@ def _collect_colors_globally(kernel: spir.Kernel, rectangles: list[Rectangle[PEB if stream_decl.stream_name in auto_stream_is_read: channel_is_read.add(channel) - # Allocate colors for each channel + # Allocate colors for each channel, the ones whose routers can switch first. WSE-3 implements + # switches on a subset of the colors, and whether a channel needs them is only known once its + # routes are planned, so handing those out first is what lets any channel of a program that stays + # within their number switch. On WSE-2 every color switches and this is the plain order. + allocation_order = csl.SWITCHABLE_COLORS + [color for color in csl.COLORS + if color not in csl.SWITCHABLE_COLORS] color_offset = 0 for channel in range(max_channel + 1): if channel in channel_to_color: continue if channel not in channel_is_read and channel not in channel_is_written: continue # Unused channel - if color_offset >= len(csl.COLORS): + if color_offset >= len(allocation_order): raise SyntaxError( f'Too many communication channels allocated for CSL: channel {channel} cannot be assigned a color') if channel in channel_is_written: - channel_to_color[channel] = csl.COLORS[color_offset] + channel_to_color[channel] = allocation_order[color_offset] color_offset += 1 if channel in channel_is_read: if channel not in channel_to_color: - channel_to_color[channel] = csl.COLORS[color_offset] + channel_to_color[channel] = allocation_order[color_offset] color_offset += 1 return channel_to_color @@ -999,6 +1004,58 @@ def _queue_spans(compute: spir.ComputeBlock, names: set[spir.Identifier], queue_ return spans +def _microthread_intervals(compute: spir.ComputeBlock, input_names: set[spir.Identifier], + output_names: set[spir.Identifier], + queue_key) -> dict[str, list[tuple[int, int]]]: + """ + The intervals over which each stream group holds a microthread on one PE. + + Both directions are numbered in one space, since a microthread is one resource across them. A + transfer that keeps a completion handle is in flight until that handle is awaited, which is where + real concurrency comes from: a receive started before a send is still running while the send is. + A self-awaited transfer is given its own slot and the next one, because the activation that + awaits it also starts what follows. + + :param compute: The compute block being lowered. + :param input_names: Streams that bind an input queue on this PE. + :param output_names: Streams that bind an output queue on this PE. + :param queue_key: Maps a stream identifier to its grouping key (channel, or the name itself). + :return: Mapping of direction-prefixed grouping key to the intervals it is in flight over. + """ + live: dict[str, list[tuple[int, int]]] = {} + pending: dict[str, list[tuple[str, int]]] = {} + index = 0 + + def close(keys: list[tuple[str, int]], end: int) -> None: + for key, start in keys: + live.setdefault(key, []).append((start, end)) + + for statement in compute.statements: + for node in statement.walk(): + if isinstance(node, spir.AwaitAllStatement): + for keys in pending.values(): + close(keys, index) + pending.clear() + continue + if isinstance(node, spir.AwaitCompletionStatement): + close(pending.pop(node.completion_name.as_ir(), []), index) + continue + for inbound, names in ((True, input_names), (False, output_names)): + stream = _fabric_transfer_stream(node, inbound) + if stream is None or stream not in names: + continue + key = f'{"in" if inbound else "out"} {queue_key(stream)}' + completion = getattr(node, 'completion_name', None) + if completion is None: + live.setdefault(key, []).append((index, index + 1)) + else: + pending.setdefault(completion.name.as_ir(), []).append((key, index)) + index += 1 + for keys in pending.values(): + close(keys, index) + return live + + def _fabric_transfer_stream(node: spir.SpatialNode, inbound: bool) -> Optional[spir.Identifier]: """ The stream a node transfers in the requested direction, or ``None``. @@ -1129,8 +1186,7 @@ def queue_key(stream: spir.Identifier) -> str: # microthread and abort with "trying to term ut_instr[N], but it's not ours". Microthreads are # one resource across both directions, so they are handed out together. microthread_of = stream_lifetime.assign_microthreads( - {f'in {key}': span for key, span in input_spans.items()} - | {f'out {key}': span for key, span in output_spans.items()}, + _microthread_intervals(rect.compute, input_names, output_names, queue_key), csl.MICROTHREAD_IDS, location=location) def allocate_microthread(stream: spir.Identifier, inbound: bool) -> int | None: diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py index cbb691db..3b346db9 100644 --- a/spada/syntax/spatial_ir/stream_lifetime.py +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -812,43 +812,43 @@ def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int] return assigned -def assign_microthreads(spans: dict[str, tuple[int, int]], microthread_ids: list[int], *, +def assign_microthreads(live: dict[str, list[tuple[int, int]]], microthread_ids: list[int], *, location: str) -> dict[str, int]: """ - Assigns microthreads to stream groups from their occupancy spans on one PE. + Assigns microthreads to stream groups from the intervals they are in flight over on one PE. A microthread is held only for the lifetime of one asynchronous operation, so unlike a fabric - queue it needs no proof that the hardware has drained: groups whose spans do not overlap take - turns on one microthread. Callers pass the inbound and outbound groups of a PE together, keyed - apart by direction, because a microthread is one resource shared by both directions. - - :param spans: Mapping of grouping key to an inclusive ``(first_use, last_use)`` statement index - pair on this PE, over both directions. + queue it needs no proof that the hardware has drained, and it is not held across the gaps + between a group's transfers: two groups may take turns on one microthread as long as no transfer + of the one is in flight while a transfer of the other is. Callers pass the inbound and outbound + groups of a PE together, keyed apart by direction, because a microthread is one resource shared + by both directions. + + :param live: Mapping of grouping key to the inclusive intervals of statement indices over which + its transfers are in flight, in one index space over both directions. :param microthread_ids: The microthread identifiers a program may name, in the order they should be handed out. :param location: The PE rectangle, for the diagnostic. :return: Mapping of grouping key to a microthread identifier, empty when the target cannot name microthreads and the hardware default has to stand. """ - if not spans or not microthread_ids: + if not live or not microthread_ids: return {} + def concurrent(one: str, other: str) -> bool: + return any(start <= other_end and other_start <= end + for start, end in live[one] + for other_start, other_end in live[other]) + assigned: dict[str, int] = {} - for key in sorted(spans, key=lambda name: (spans[name][0], spans[name][1], name)): - start, end = spans[key] - used = { - assigned[other] - for other in assigned - if start <= spans[other][1] and spans[other][0] <= end - } + for key in sorted(live, key=lambda name: (min(live[name]), name)): + used = {assigned[other] for other in assigned if concurrent(key, other)} for microthread in microthread_ids: if microthread not in used: assigned[key] = microthread break else: - overlapping = sorted( - other for other, (other_start, other_end) in spans.items() - if other != key and start <= other_end and other_start <= end) + overlapping = sorted(other for other in live if other != key and concurrent(key, other)) raise SyntaxError( f'{location} would need {len(used) + 1} concurrent microthreads ' f'(live groups {[key] + overlapping}), but a PE may name at most ' diff --git a/tests/csl_runtime/test_batcher_oddeven_wse3_1d.sh b/tests/csl_runtime/test_batcher_oddeven_wse3_1d.sh new file mode 100755 index 00000000..7695ed14 --- /dev/null +++ b/tests/csl_runtime/test_batcher_oddeven_wse3_1d.sh @@ -0,0 +1,64 @@ +#!/bin/sh +# E2E test: the WSE-3 1D Batcher odd-even mergesort (2^L PEs, a block of K f32 keys per PE). +# Kernel: batcher_oddeven_wse3_1D.sptl params: L, K +# Same result as batcher_oddeven_1D -- OUT_out.reshape(n*k) == sort(inp.reshape(n*k)) -- and the same +# network as the bundled variant. What it adds is L = 4, which the bundled variant cannot reach on +# wse3: there a queue is bound to a color for the whole kernel, so a PE needs one per color it ever +# uses, and the bundled variant wants seven of the six. Pooling the unbundled phases by the origin's +# residue rather than by direction spends one queue per distance in each direction instead of two, +# which brings it to five. The routers switch for it, on twelve colors of the fifteen wse3 can switch. +# This kernel is written for the queue model of wse3 and is only run there; the bundled variant is +# the better fit on wse2, whose routers here would not switch at all. + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +. "$SCRIPT_DIR/_lib.sh" + +SORT_DIR="$(cd "$(dirname "$0")/../../samples/spatial/sort" && pwd)" +FOLDER="batcher_oddeven_wse3_1d_sptl" + +if [ "${WSE_ARCH:-wse2}" != "wse3" ]; then + echo "Skipping: this variant targets wse3; on wse2 use test_batcher_oddeven_bundled_1d.sh." + exit 0 +fi + +run_batcher() { + l=$1 + k=$2 + echo "--- batcher_oddeven_wse3_1d L=$l K=$k ---" + + sptlc "$SORT_DIR/batcher_oddeven_wse3_1D.sptl" "$FOLDER" -p L=$l -p K=$k + + python3 - <= 3: + if len(constants.INPUT_QUEUE_IDS) >= 3 and constants.ARCH != 'wse3': pytest.skip(f'{constants.ARCH} has {len(constants.INPUT_QUEUE_IDS)} input queues, enough for L=4') - with pytest.raises(SyntaxError, match='concurrent input queues'): + with pytest.raises(SyntaxError, match='concurrent (in|out)put queues'): _bundled_batcher(4) +def _wse3_batcher(l: int, k: int = 1): + path = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', + 'batcher_oddeven_wse3_1D.sptl') + kernel = parser.parse_file(path) + kernel = passes.concretize_parameters(kernel, L=l, K=k) + kernel = passes.constexpr_propagation(kernel) + return lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) + + +def _queues_per_pe(files, kind: str) -> int: + return max(len(set(re.findall(rf'@get_{kind}_queue\((\d+)\)', f.code))) + for f in files if f.filename.startswith('code_')) + + +@pytest.mark.parametrize('l, colors', [(3, 8), (4, 12)]) +def test_pooling_by_origin_halves_what_a_pooled_distance_costs(l: int, colors: int): + """ + batcher_oddeven_wse3_1D lets a comparator's two messages share one color instead of taking one + per direction, which is 8 colors at L = 3 and 12 at L = 4 where the bundled variant takes 10 + and 18. The pooled colors now switch, since a PE's role on one is fixed but the side it faces + is not, and a source only changes where it transmits and a destination where it receives. + """ + if l == 4 and len(constants.INPUT_QUEUE_IDS) < 3: + pytest.skip(f'{constants.ARCH} has {len(constants.INPUT_QUEUE_IDS)} input queues') + + layout = next(f.code for f in _wse3_batcher(l) if 'layout' in f.filename) + used, switched = _colors_of(layout), _colors_of(layout, r'\.switches') + assert len(used) == colors + assert used == switched + + # Two positions per router, and never a both-sides change, so nothing here needs WSE-3. + for line in layout.splitlines(): + if '.switches' in line: + assert len(re.findall(r'\.pos\d', line)) <= 2, line + + +def test_sixteen_keys_fit_the_queues_of_a_target_that_reuses_none(): + """ + WSE-3 keeps a queue on its color for the whole kernel, so what a PE can afford is how many + colors it ever touches. One per pooled distance instead of two brings L = 4 inside the six + inbound queues, and inside the five outbound ones left once memcpy has taken its own. + """ + if len(constants.INPUT_QUEUE_IDS) < 3: + # Two queues bind this variant at L = 3 as well, and there the bundled one, whose routers + # never switch, is the better fit anyway. + with pytest.raises(SyntaxError, match='concurrent input queues'): + _wse3_batcher(4) + return + + files = _wse3_batcher(4) + assert _queues_per_pe(files, 'input') <= len(constants.INPUT_QUEUE_IDS) + assert _queues_per_pe(files, 'output') < len(constants.OUTPUT_QUEUE_IDS) + + def test_only_the_widest_batcher_phases_are_bundled(): # Bundling costs one filter at every participating PE and a PE has three, so the sample bundles # every phase that satisfies 4d >= N. At L = 3 that is already three phases (d = 4, 2, 2). diff --git a/tests/spatial_ir/test_stream_lifetime.py b/tests/spatial_ir/test_stream_lifetime.py index 3af2b2f7..f2c989c8 100644 --- a/tests/spatial_ir/test_stream_lifetime.py +++ b/tests/spatial_ir/test_stream_lifetime.py @@ -513,28 +513,39 @@ def test_exclusive_keys_exhaust_queues_when_too_many_data_tasks(): exclusive_keys=frozenset({'channel 0', 'channel 1', 'channel 2'})) -def test_microthreads_are_shared_across_directions_when_spans_are_disjoint(): +def test_microthreads_are_shared_across_directions_when_transfers_do_not_overlap(): """A microthread is held only while a transfer is in flight, so turns may be taken.""" assigned = stream_lifetime.assign_microthreads( - {'in channel 0': (0, 2), 'out channel 1': (3, 5)}, [2, 3], location='PE (0, 0)') + {'in channel 0': [(0, 2)], 'out channel 1': [(3, 5)]}, [2, 3], location='PE (0, 0)') assert assigned == {'in channel 0': 2, 'out channel 1': 2} def test_a_receive_and_a_send_in_flight_together_get_distinct_microthreads(): assigned = stream_lifetime.assign_microthreads( - {'in channel 0': (0, 4), 'out channel 0': (0, 4)}, [2, 3], location='PE (0, 0)') + {'in channel 0': [(0, 4)], 'out channel 0': [(0, 4)]}, [2, 3], location='PE (0, 0)') assert assigned['in channel 0'] != assigned['out channel 0'] +def test_a_group_that_comes_back_after_a_gap_does_not_hold_its_microthread_across_it(): + """ + A channel used again much later is not in flight in between, unlike a fabric queue, which stays + bound to it. The group that runs in the gap may take the same microthread. + """ + assigned = stream_lifetime.assign_microthreads( + {'in channel 0': [(0, 1), (8, 9)], 'out channel 1': [(4, 5)]}, [2, 3], location='PE (0, 0)') + assert assigned == {'in channel 0': 2, 'out channel 1': 2} + + def test_microthreads_run_out_when_too_many_transfers_overlap(): with pytest.raises(SyntaxError, match='concurrent microthreads'): stream_lifetime.assign_microthreads( - {'in channel 0': (0, 10), 'out channel 0': (2, 8), 'out channel 1': (4, 6)}, [2, 3], - location='PE (0, 0)') + {'in channel 0': [(0, 10)], 'out channel 0': [(2, 8)], 'out channel 1': [(4, 6)]}, + [2, 3], location='PE (0, 0)') def test_microthreads_are_left_to_the_hardware_when_the_target_cannot_name_them(): - assert stream_lifetime.assign_microthreads({'in channel 0': (0, 2)}, [], location='PE (0, 0)') == {} + assert stream_lifetime.assign_microthreads({'in channel 0': [(0, 2)]}, [], + location='PE (0, 0)') == {} if __name__ == '__main__': From 6af96ef45cd2964c27649b23f6c865a484a6c8bc Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 17:28:22 +0200 Subject: [PATCH 45/68] Fix odd even sort -- one sequence --- README.md | 2 +- .../spatial/sort/odd_even_sort_1D_looped.sptl | 128 +++++++++++++++--- .../test_odd_even_sort_1d_looped.sh | 25 ++-- 3 files changed, 125 insertions(+), 30 deletions(-) diff --git a/README.md b/README.md index bb12915f..15b37a24 100644 --- a/README.md +++ b/README.md @@ -104,7 +104,7 @@ Sample SPADA programs are in `samples/`: | `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy` | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Sorting networks over `2^L` PEs holding `K` keys each: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), `bitonic_sort_1D` (channel reuse via router switches, WSE-3), and `odd_even_sort_1D_looped` (N neighbour rounds as a runtime loop, four static channels) | +| `samples/spatial/sort/` | Sorting networks over `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), `bitonic_sort_1D` (K independent sequences; channel reuse via router switches, WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | diff --git a/samples/spatial/sort/odd_even_sort_1D_looped.sptl b/samples/spatial/sort/odd_even_sort_1D_looped.sptl index d30f257e..ef198368 100644 --- a/samples/spatial/sort/odd_even_sort_1D_looped.sptl +++ b/samples/spatial/sort/odd_even_sort_1D_looped.sptl @@ -1,11 +1,16 @@ /** * Odd-even transposition sort over N = 2^L PEs, with the N rounds as a runtime loop. * - * Each PE holds K keys. Sequence k is element k of every PE; the K sequences are sorted - * independently. After the network, PE i holds the i-th key of each sequence. + * Each PE holds a block of K f32 keys: N*K keys in all. Ascending: after the network, PE i holds + * keys i*K .. i*K + (K-1) of the sorted sequence. K = 1 is the one-key-per-PE network again. * * Algorithm * --------- + * The load phase sorts this PE's block. Every later comparator is a compare-split of two sorted + * blocks: the partners trade, the lower index keeps the K smallest of the 2K keys and the higher + * index the K largest, and both halves come out ascending. A network that sorts N keys sorts N + * sorted blocks this way, so concatenating the blocks left to right gives the sorted sequence. + * * Round t compares neighbours (2i, 2i+1) when t is even and (2i+1, 2i+2) when t is odd. * N rounds suffice. Every comparator is one hop: on a line a longer exchange occupies the * same links as several one-hop ones and does not cut the round count, which is already @@ -37,6 +42,9 @@ * PE 1, 3, .. N-3 high in even rounds, low in odd rounds * PE N-1 high in even rounds, idle in odd rounds * + * A high PE must send its own block before it merges, which is why it sends after the + * receive but before the walk: the merge overwrites val, and the partner needs the + * pre-merge value. * * Constraints: L >= 1, K >= 1. WSE-2 and WSE-3. * @@ -45,8 +53,15 @@ kernel @odd_even_sort_1d_looped(stream[1<[1<(stream[1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } } // PE 0: low partner in every even round, idle in every odd round. @@ -80,8 +105,19 @@ kernel @odd_even_sort_1d_looped(stream[1<(stream[1< val[k] else val[k]) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] } } await send(val, a_out[i, j]) @@ -109,13 +168,36 @@ kernel @odd_even_sort_1d_looped(stream[1< val[k] else val[k]) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] } await send(val, east_odd) await receive(other, west_odd) - await map i32 k in [0:K] { - val[k] = (other[k] if other[k] < val[k] else val[k]) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] } } await send(val, a_out[i, j]) @@ -126,8 +208,20 @@ kernel @odd_even_sort_1d_looped(stream[1< val[k] else val[k]) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] } } await send(val, a_out[i, j]) diff --git a/tests/csl_runtime/test_odd_even_sort_1d_looped.sh b/tests/csl_runtime/test_odd_even_sort_1d_looped.sh index a53f1335..9136444c 100644 --- a/tests/csl_runtime/test_odd_even_sort_1d_looped.sh +++ b/tests/csl_runtime/test_odd_even_sort_1d_looped.sh @@ -1,10 +1,12 @@ #!/bin/sh # E2E: odd-even transposition sort on 2^L PEs, N rounds as a runtime loop -# (odd_even_sort_1D_looped.sptl). Each PE holds K keys; sequence k is element k of -# every PE. Reference: OUT_a_out == sort(a_in, axis=0). +# (odd_even_sort_1D_looped.sptl). Each PE holds a block of K f32 keys; every comparator is a +# compare-split, so the network sorts all 2^L * K keys and PE i ends up with keys i*K .. i*K + K-1 +# of the sorted sequence. Reference: OUT_a_out.reshape(n*k) == sort(a_in.reshape(n*k)). # Runs on WSE-2 and WSE-3: four channels, one per (round parity, direction), so no # router ever switches. L = 1 is two PEs and a single even round; L = 3 is eight PEs -# and exercises every role (ends and both interior parities). +# and exercises every role (ends and both interior parities). K = 1 is the one-key-per-PE +# network; K = 4 is not a power of two, which is what the K-element merge has to be independent of. set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" @@ -31,17 +33,16 @@ PYEOF python3 - < Date: Wed, 19 Aug 2026 17:28:36 +0200 Subject: [PATCH 46/68] adapt odd even sort test --- tests/csl_runtime/test_odd_even_sort_1d_looped.sh | 0 1 file changed, 0 insertions(+), 0 deletions(-) mode change 100644 => 100755 tests/csl_runtime/test_odd_even_sort_1d_looped.sh diff --git a/tests/csl_runtime/test_odd_even_sort_1d_looped.sh b/tests/csl_runtime/test_odd_even_sort_1d_looped.sh old mode 100644 new mode 100755 From 9fdd803cd386087530655001e5737722caae4665 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 17:51:19 +0200 Subject: [PATCH 47/68] Reject statically overlapping compute and dataflow blocks in one phase at compile time --- spada/syntax/spatial_ir/canonical_subgrids.py | 36 +++++++++- tests/spatial_ir/test_canonical_subgrids.py | 68 +++++++++++++++++++ 2 files changed, 101 insertions(+), 3 deletions(-) create mode 100644 tests/spatial_ir/test_canonical_subgrids.py diff --git a/spada/syntax/spatial_ir/canonical_subgrids.py b/spada/syntax/spatial_ir/canonical_subgrids.py index 9f1027ad..9a3fd9f4 100644 --- a/spada/syntax/spatial_ir/canonical_subgrids.py +++ b/spada/syntax/spatial_ir/canonical_subgrids.py @@ -52,18 +52,48 @@ def visit_PlaceBlock(self, block: spa.PlaceBlock): def visit_DataflowBlock(self, block: spa.DataflowBlock): self.process_block(block) + +def _validate_disjoint_phase_subgrids(subgrids: list[spa.Subgrid]) -> None: + """Reject statically overlapping compute or dataflow blocks in one phase. + + The Spatial IR specification permits at most one compute block per PE in a + phase and requires dataflow subgrids in a phase to be disjoint. Parameter + and metaprogram expressions have already been concretized before this pass, + so these overlaps can be diagnosed exactly. + """ + for index, first in enumerate(subgrids): + first_phase, first_block = first.metadata + if not isinstance(first_block, (ComputeBlock, DataflowBlock)): + continue + + for second in subgrids[index + 1:]: + second_phase, second_block = second.metadata + if first_phase != second_phase or type(first_block) is not type(second_block): + continue + if not first.intersects(second): + continue + + overlap = first.intersection(second) + block_kind = 'compute' if isinstance(first_block, ComputeBlock) else 'dataflow' + raise SyntaxError( + f'Overlapping {block_kind} subgrids in phase {first_phase}: ' + f'PEs x={overlap.x_range}, y={overlap.y_range} belong to multiple ' + f'{block_kind} blocks.' + ) + + def canonicalize_subgrids(kernel: Kernel) -> Kernel: """ This pass ensures that all subgrids either do not intersect or are equal. - Assumes that the subgrids are already correctly defined within each phase. - Specifically, within each phase no two gridpoints may belong to more than one subgrid - of the same block type. + Compute and dataflow overlaps within one phase are rejected before splitting. + Place blocks may overlap because they can declare distinct fields on the same PEs. :param kernel: The kernel to canonicalize. :return: A new kernel with the subgrids canonicalized. """ subgrids = kernel.subgrids() + _validate_disjoint_phase_subgrids(subgrids) # split subgrids so that no two un-equal subgrids overlap print(f"Splitting {len(subgrids)} grids") diff --git a/tests/spatial_ir/test_canonical_subgrids.py b/tests/spatial_ir/test_canonical_subgrids.py new file mode 100644 index 00000000..d679362b --- /dev/null +++ b/tests/spatial_ir/test_canonical_subgrids.py @@ -0,0 +1,68 @@ +import pytest + +from spada.lowering import spatial_ir_to_csl as s2c +from spada.syntax.spatial_ir import parser, passes + + +def _canonicalize(code: str): + kernel = parser.parse_string(code, 'test.sptl') + kernel = passes.constexpr_propagation(kernel) + return s2c.canonicalize_kernel(kernel) + + +def test_overlapping_compute_subgrids_in_one_phase_are_rejected(): + code = """ +kernel @overlapping_compute<>() { + phase { + compute i16 i, i16 j in [0:4, 0] {} + compute i16 i, i16 j in [2:6, 0] {} + } +} +""" + with pytest.raises( + SyntaxError, + match=r'Overlapping compute subgrids in phase 1.*x=\(2, 4, 1\)', + ): + _canonicalize(code) + + +def test_overlapping_dataflow_subgrids_in_one_phase_are_rejected(): + code = """ +kernel @overlapping_dataflow<>() { + phase { + dataflow i16 i, i16 j in [0:4, 0] {} + dataflow i16 i, i16 j in [3:6, 0] {} + } +} +""" + with pytest.raises( + SyntaxError, + match=r'Overlapping dataflow subgrids in phase 1.*x=\(3, 4, 1\)', + ): + _canonicalize(code) + + +def test_disjoint_strided_compute_subgrids_are_accepted(): + code = """ +kernel @disjoint_compute<>() { + phase { + compute i16 i, i16 j in [0:8:2, 0] {} + compute i16 i, i16 j in [1:8:2, 0] {} + } +} +""" + _canonicalize(code) + + +def test_compute_subgrids_may_overlap_across_phases(): + code = """ +kernel @compute_across_phases<>() { + phase { + compute i16 i, i16 j in [0:4, 0] {} + } + phase { + compute i16 i, i16 j in [0:4, 0] {} + } +} +""" + _canonicalize(code) From 1fd930b5703903b8a085a1f3223885666334c3c8 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 17:52:02 +0200 Subject: [PATCH 48/68] Fix odd even sort --- .../spatial/sort/odd_even_sort_1D_looped.sptl | 41 ++++++++++++++++--- tests/spatial_ir/test_dsd_ops.py | 5 +-- 2 files changed, 37 insertions(+), 9 deletions(-) diff --git a/samples/spatial/sort/odd_even_sort_1D_looped.sptl b/samples/spatial/sort/odd_even_sort_1D_looped.sptl index ef198368..9295bac4 100644 --- a/samples/spatial/sort/odd_even_sort_1D_looped.sptl +++ b/samples/spatial/sort/odd_even_sort_1D_looped.sptl @@ -85,8 +85,8 @@ kernel @odd_even_sort_1d_looped(stream[1<(stream[1<(stream[1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } for i32 t in [0 : 1<<(L-1)] { await send(val, east_even) await receive(other, west_even) @@ -165,6 +172,17 @@ kernel @odd_even_sort_1d_looped(stream[1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } for i32 t in [0 : 1<<(L-1)] { await receive(other, east_even) await send(val, west_even) @@ -205,6 +223,17 @@ kernel @odd_even_sort_1d_looped(stream[1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } for i32 t in [0 : 1<<(L-1)] { await receive(other, east_even) await send(val, west_even) diff --git a/tests/spatial_ir/test_dsd_ops.py b/tests/spatial_ir/test_dsd_ops.py index c153bb9a..a580aaa7 100644 --- a/tests/spatial_ir/test_dsd_ops.py +++ b/tests/spatial_ir/test_dsd_ops.py @@ -155,16 +155,15 @@ def test_transfers_inside_a_sequential_for_get_fabric_dsds(): channel = 0 } } - compute i16 i, i16 j in [0:2, 0] { - await receive(val, src[i, j]) - } compute i16 i, i16 j in [0:1, 0] { + await receive(val, src[i, j]) for i32 t in [0:M] { await send(val, east) } await send(val, dst[i, j]) } compute i16 i, i16 j in [1:2, 0] { + await receive(val, src[i, j]) for i32 t in [0:M] { await receive(other, east) } From 1ed4e5dbff8bf82892643ed23da85fc96995ad6e Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 18:35:11 +0200 Subject: [PATCH 49/68] Fix test shift bundles --- tests/spatial_ir/test_shift_bundles.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py index 56c33e19..d5c45134 100644 --- a/tests/spatial_ir/test_shift_bundles.py +++ b/tests/spatial_ir/test_shift_bundles.py @@ -139,9 +139,11 @@ def test_a_shift_of_one_is_not_bundled(): assert _bundles(_SHIFT, M=1, D=1, K=1) == [] -def test_sources_longer_than_the_shift_are_not_bundled(): - # Sources would be destinations of the same bundle, which this arrangement cannot express. - assert _bundles(_SHIFT, M=4, D=2, K=1) == [] +def test_sources_longer_than_the_shift_are_rejected_as_overlapping_compute_subgrids(): + # PEs 2 and 3 would be both sources and destinations in one phase, which violates the + # one-compute-block-per-PE rule before bundle detection is reached. + with pytest.raises(SyntaxError, match='Overlapping compute subgrids'): + _bundles(_SHIFT, M=4, D=2, K=1) def test_sources_inject_then_relay(): From 68784614cce0344942f7200935385f149cccf8e3 Mon Sep 17 00:00:00 2001 From: glukas Date: Wed, 19 Aug 2026 22:53:41 +0200 Subject: [PATCH 50/68] Fix close implementation for wse-2 to fix bundling --- irspec/docs/spatial/routing_wse.md | 21 ++++--- spada/lowering/spatial_ir_to_csl.py | 32 ++++++---- spada/syntax/csl/routing.py | 58 +++++++++++++------ spada/syntax/csl/statements.py | 5 +- spada/syntax/csl/structures.py | 9 ++- spada/syntax/spatial_ir/irnodes.py | 12 +++- .../handwritten/shift_bundle/layout.csl | 9 ++- .../handwritten/shift_bundle/run.py | 24 +++++--- .../handwritten/shift_bundle/sender.csl | 19 +++++- .../csl_runtime/test_shift_bundle_filters.sh | 16 +++-- tests/spatial_ir/test_routing.py | 20 +++---- tests/spatial_ir/test_shift_bundles.py | 14 ++++- 12 files changed, 169 insertions(+), 70 deletions(-) diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index c0d83401..b14a391e 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -65,10 +65,14 @@ all. This does *not* generalize to channels that switch: a router's positions are a static sequence, so the epoch a configuration belongs to has to be visible to the compiler. -An advance is driven by the `close` that ends the epoch. The sending PE emits a *switch-advance -control message* on the channel, one per position to be traversed. It follows the stream's path -using the configuration that is being retired, and advances the router of each PE it traverses, -after all data of the epoch. +An advance is driven by the `close` that ends the epoch. When some *other* router on the path has +to move, the sending PE emits a *switch-advance control message* on the channel, one per position +to be traversed. It follows the stream's path using the configuration that is being retired, and +advances the router of each PE it traverses, after all data of the epoch. When only the sending +PE's own router has to move, and only by one position, the last data wavelet does that itself +(`.advance_switch` on the fabric output DSD). A second operation on the same output queue is what +drops a data wavelet on WSE-2 once a back-pressured send fills it (output queues 2 and 3 hold six +16-bit words; three `f32` values already fill them). !!! warning "WSE: Advances Are Not Selective" A CSL control wavelet nominally carries up to eight per-router switching commands @@ -148,9 +152,12 @@ filter: -- -- -- win 2 win 1 win 0 ``` **Sources inject, then relay.** A source starts at `rx = RAMP, tx = {EAST}` with `pos1` taking -`rx = WEST`, sends its own words, and its `close` advances its own router into relay mode. The -trigger is local — *"my own send is done"* — which is what makes it expressible at all, given that -the payload of the resulting message [selects nothing](#lowering-to-switches). +`rx = WEST`, sends its own words, and the last data wavelet advances its own router into relay +mode (`.advance_switch` on the fabric output DSD). The trigger is local — *"my own send is +done"* — which is what makes it expressible at all, given that the payload of a control message +[selects nothing](#lowering-to-switches). A `SWITCH_ADV` on the same output queue would also +flip that router, but on WSE-2 a back-pressured send of three or more `f32` values fills the +queue and the control wavelet then steals a data word. **The order is descending, and enforces itself.** The source nearest the destinations goes first. No schedule or barrier is needed: a source further away cannot push a word through its neighbour's diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 40d45ce9..f2e35da6 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -151,7 +151,8 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, # Plan the router switch advances, then drop every close no router has to act on cslrouting.plan_switch_advances(rectangles) if close_elision: - stream_lifetime.elide_redundant_closes(rectangles, needs_advance=lambda stmt: bool(stmt.switch_advance)) + stream_lifetime.elide_redundant_closes( + rectangles, needs_advance=lambda stmt: bool(stmt.switch_advance) or stmt.advance_data_switch) for rect in rectangles: # Create a unique CSL code file based on rectangle offset @@ -1206,6 +1207,23 @@ def allocate_output_queue(stream: spir.Identifier) -> int: f'{location}: no output queue was reserved for {key} (stream "{stream.as_ir()}").') return output_queue_of[key] + # A close that only flips this PE's own router does so on the last data wavelet, so the + # outgoing fabric descriptor has to carry ``.advance_switch``. The close itself emits no + # control wavelet and is kept only so this scan can see the flag. + streams_advance_on_send = { + stream_lifetime.underlying_stream(stmt.stream_name).as_ir() + for stmt in rect.compute.statements + if isinstance(stmt, spir.CloseStatement) and stmt.advance_data_switch + } + + def _fabout(stream, extents) -> cslstruct.FabricDSD: + stream = stream_lifetime.underlying_stream(stream) + return cslstruct.FabricDSD( + cslstruct.DSDType.fabout, f'{name_to_csl(stream)}_color', extents, + allocate_output_queue(stream), + ut=allocate_microthread(stream, inbound=False), + advance_switch=stream.as_ir() in streams_advance_on_send) + def _visit_foreach(stmt: spir.ForeachStatement) -> None: """ Registers the fabric input DSD for a ``foreach`` that draws from a stream. @@ -1280,7 +1298,6 @@ def _visit_foreach(stmt: spir.ForeachStatement) -> None: ut=allocate_microthread(stream_name, inbound=True)) dsds[stream_name.as_ir()].append((dsd_name, dsd)) elif isinstance(stmt, spir.SendStatement) and stream_name.as_ir() in stream_candidates: - dsd_type = cslstruct.DSDType.fabout dsd_name = f'{name_to_csl(stream_name)}_out_dsd' extents = stream_candidates[stream_name.as_ir()][1] if extents is not None: # Use buffer size @@ -1294,10 +1311,7 @@ def _visit_foreach(stmt: spir.ForeachStatement) -> None: extents = functools.reduce( lambda a, b: a * b, [s.eval() if not isinstance(s, int) else s for s in dtypes[stmt.local_array].shape], 1) - fabric_color = f'{name_to_csl(stream_name)}_color' - dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, - allocate_output_queue(stream_name), - ut=allocate_microthread(stream_name, inbound=False)) + dsd = _fabout(stream_name, extents) dsds[stream_name.as_ir()].append((dsd_name, dsd)) if isinstance(stmt, spir.SendStatement): @@ -1333,7 +1347,6 @@ def _visit_nested_send(substmt: spir.SendStatement): return # Stream DSD (i.e., await send in a foreach) stream_name = substmt.stream_name - dsd_type = cslstruct.DSDType.fabout dsd_name = f'{name_to_csl(stream_name)}_out_dsd' extents = stream_candidates[stream_name.as_ir()][1] if extents is not None: # Use buffer size @@ -1347,10 +1360,7 @@ def _visit_nested_send(substmt: spir.SendStatement): extents = functools.reduce( lambda a, b: a * b, [s.eval() if not isinstance(s, int) else s for s in dtypes[substmt.local_array].shape], 1) - fabric_color = f'{name_to_csl(stream_name)}_color' - dsd = cslstruct.FabricDSD(dsd_type, fabric_color, extents, - allocate_output_queue(stream_name), - ut=allocate_microthread(stream_name, inbound=False)) + dsd = _fabout(stream_name, extents) dsds[stream_name.as_ir()].append((dsd_name, dsd)) def _visit_nested_receive(substmt: spir.ReceiveStatement): diff --git a/spada/syntax/csl/routing.py b/spada/syntax/csl/routing.py index 2ba5de75..892f5021 100644 --- a/spada/syntax/csl/routing.py +++ b/spada/syntax/csl/routing.py @@ -691,16 +691,20 @@ def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: """ Determines, for every ``close`` statement, how many switch advances it has to emit. - A close only produces code on a PE that *sends* the stream: the control wavelets it emits travel - the path being retired. Every switch-configured router such a wavelet reaches advances -- the - hardware applies the wavelet's single command at each of them rather than indexing a per-router - command array -- so a close cannot move one router while leaving another on its path behind. - All routers on the path that hold switch positions must therefore advance by the same amount, - and that amount is how many wavelets are sent. A close on a receiving PE emits nothing; its - router is advanced by the sending PE's wavelets. + A close only produces a control wavelet on a PE that *sends* the stream, and only when some + *other* router on the path has to move: the wavelet travels the path being retired and every + switch-configured router it reaches advances. A close that only has to flip the sending PE's + own router does that on the last data wavelet (``.advance_switch`` on the fabric output DSD) + instead -- a second operation on the same output queue is what drops a data wavelet on WSE-2 + once a back-pressured send fills it. - The result is recorded on each ``CloseStatement`` as its ``switch_advance`` field; a close that - needs no advance keeps ``None`` there and generates no code. + All routers on the path that hold switch positions and are advanced by a control wavelet must + therefore advance by the same amount, and that amount is how many wavelets are sent. A close on + a receiving PE emits nothing; its router is advanced by the sending PE's wavelets. + + The result is recorded on each ``CloseStatement`` as ``switch_advance`` (control wavelets) or + ``advance_data_switch`` (last-data-wavelet flip). A close that needs neither keeps both unset + and generates no code. :param rectangles: The consolidated PE rectangles of the kernel, annotated in place. :return: The number of closes that retire a route configuration. @@ -738,6 +742,7 @@ def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: if not isinstance(statement, spir.CloseStatement): continue statement.switch_advance = None + statement.advance_data_switch = False name = stream_lifetime.underlying_stream(statement.stream_name) declaration = declarations.get(name) if declaration is None or name not in uses or not uses[name].sent: @@ -749,8 +754,11 @@ def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: group = stream_lifetime.stream_group_key(declaration) # How far each switch-configured router on the path has to move, keyed by the router so - # that a disagreement can name it. + # that a disagreement can name it. The sending PE's own router is tracked separately: + # flipping only that one is done on the last data wavelet, not by a SWITCH_ADV. advances: dict[str, int] = {} + local_advance: int | None = None + remote_advances: dict[str, int] = {} for dx, dy in offsets: site = _find_site(position_of, channel, (rect.x_range[0] + dx, rect.x_range[1] + dx, rect.x_range[2]), @@ -766,13 +774,20 @@ def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: ring, total = wraps[site] position %= len(indices) if position + 1 < len(indices): - advances[site.describe()] = indices[position + 1] - indices[position] + amount = indices[position + 1] - indices[position] elif ring: - advances[site.describe()] = total - indices[position] - # Otherwise this router is on its last position for this color, and outside ring mode - # an advance past it is a no-op, so wavelets passing through over-advance it - # harmlessly. It is not necessarily finished with the color: a shift bundle's source - # keeps relaying its last configuration long after reaching it. + amount = total - indices[position] + else: + # This router is on its last position for this color, and outside ring mode + # an advance past it is a no-op, so wavelets passing through over-advance it + # harmlessly. It is not necessarily finished with the color: a shift bundle's + # source keeps relaying its last configuration long after reaching it. + continue + advances[site.describe()] = amount + if dx == 0 and dy == 0: + local_advance = amount + else: + remote_advances[site.describe()] = amount distinct = set(advances.values()) if not distinct: @@ -788,7 +803,16 @@ def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: ' note: give the streams that disagree separate channels, at the cost of an ' 'additional color') - statement.switch_advance = distinct.pop() + amount = distinct.pop() + # A source that only flips its own router does so on the last data wavelet. Posting a + # SWITCH_ADV into the same output queue afterwards is what drops a data wavelet on + # WSE-2 when a back-pressured send of three or more f32 values fills that queue + # (see tests/csl_runtime/test_shift_bundle_filters.sh). Remote routers still need a + # traveling control wavelet, and a two-position turnaround still needs two of them. + if not remote_advances and local_advance == 1: + statement.advance_data_switch = True + else: + statement.switch_advance = amount planned += 1 return planned diff --git a/spada/syntax/csl/statements.py b/spada/syntax/csl/statements.py index 31761b9a..fb9bbe9a 100644 --- a/spada/syntax/csl/statements.py +++ b/spada/syntax/csl/statements.py @@ -49,8 +49,9 @@ def generate_csl_statement(statement: spir.Statement, # Skip (taken care of when tasks are defined) return "" elif isinstance(statement, spir.CloseStatement): - # Retiring a route configuration means advancing the switches along the stream's path, one - # control wavelet per position, or nothing at all when no router has to move. + # A close that only flips this PE's own router does that on the last data wavelet + # (``.advance_switch`` on the fabric output DSD) and generates no statement here. A close + # that has to move a remote router emits one SWITCH_ADV control wavelet per position. if not statement.switch_advance: return "" stream = statement.stream_name diff --git a/spada/syntax/csl/structures.py b/spada/syntax/csl/structures.py index fffc955d..7a2a97ff 100644 --- a/spada/syntax/csl/structures.py +++ b/spada/syntax/csl/structures.py @@ -61,6 +61,10 @@ class FabricDSD(DataStructureDescriptor): extent: int queue: int control: bool = False + #: When True, the router of this color advances after the last wavelet this descriptor sends. + #: Only meaningful on ``fabout``; it is how a source hands its own router over to relay mode + #: without a second operation on the same output queue. + advance_switch: bool = False #: Microthread to drive an asynchronous transfer over this descriptor, where the target lets a #: program name one. ``None`` leaves the hardware default, which is the queue ID. The setting #: belongs to the operation rather than the descriptor, so ``as_csl`` does not emit it; the @@ -69,14 +73,17 @@ class FabricDSD(DataStructureDescriptor): def __post_init__(self): assert self.dsd_type in (DSDType.fabin, DSDType.fabout) + if self.advance_switch: + assert self.dsd_type == DSDType.fabout def as_csl(self) -> str: direction = "in" if self.dsd_type == DSDType.fabin else "out" queue_type = "input_queue" if self.dsd_type == DSDType.fabin else "output_queue" fabric_color = f' .fabric_color = {self.color}_{direction},' if self.color else '' control = ' .control = true,' if self.control else '' + advance = ' .advance_switch = true,' if self.advance_switch else '' return (f'@get_dsd({self.dsd_type.name}_dsd, .{{ .extent = {self.extent},{fabric_color}' - f'{control} .{queue_type} = @get_{queue_type}({self.queue}) }})') + f'{control}{advance} .{queue_type} = @get_{queue_type}({self.queue}) }})') def __hash__(self): return hash(("FabricDSD", self.as_csl())) diff --git a/spada/syntax/spatial_ir/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index 87943a5c..f83b0152 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -932,9 +932,15 @@ class CloseStatement(Statement): #: How many switch positions the routers along the stream's path move forward when this close #: retires the stream's route configuration. One wavelet is emitted per position, and a #: transition that changes both a router's input and its output direction takes two. Filled in - #: during lowering by ``csl.routing.plan_switch_advances``; ``None`` means no router has to act, - #: in which case the close generates no code. Not part of the surface syntax. + #: during lowering by ``csl.routing.plan_switch_advances``; ``None`` means no control wavelet + #: is sent. Not part of the surface syntax. switch_advance: Optional[int] = None + #: When True, the sending PE flips only its own router, and does so on the last data wavelet + #: (``.advance_switch`` on the fabric output DSD) rather than by a ``SWITCH_ADV`` control + #: wavelet. A control wavelet on the same output queue as the data is what drops a wavelet on + #: WSE-2 once a back-pressured send of three or more f32 values fills the queue. Mutually + #: exclusive with ``switch_advance``. Not part of the surface syntax. + advance_data_switch: bool = False def validate(self) -> None: assert isinstance(self.stream_name, (Identifier, ArraySlice)) @@ -942,6 +948,8 @@ def validate(self) -> None: assert isinstance(self.completion_name, Completion) if self.switch_advance is not None: assert isinstance(self.switch_advance, int) and self.switch_advance > 0 + if self.advance_data_switch: + assert self.switch_advance is None def as_ir(self, indent: int = 0) -> str: indent_str = ' ' * indent diff --git a/tests/csl_runtime/handwritten/shift_bundle/layout.csl b/tests/csl_runtime/handwritten/shift_bundle/layout.csl index c94c60e6..84149769 100644 --- a/tests/csl_runtime/handwritten/shift_bundle/layout.csl +++ b/tests/csl_runtime/handwritten/shift_bundle/layout.csl @@ -6,8 +6,8 @@ // The senders go east to west -- the one closest to the receivers first. That order is what // makes this cheap: a sender injects its own words and only then has to become a relay for // the senders west of it, so the hand-over is triggered by an event it knows locally and can -// signal itself with one SWITCH_ADV. No router has to be switched from a distance, which is -// the thing a control wavelet cannot do selectively. +// signal itself with a switch advance. No router has to be switched from a distance, which +// is the thing a control wavelet cannot do selectively. // // The receivers never switch. Each one transmits to its ramp *and* onward east, so every // receiver's router sees the whole stream and the outermost one terminates it. Which words a @@ -38,6 +38,10 @@ param IN_QUEUE: i16; // 1: emit @initialize_queue (required on WSE-3). 0: omit it (WSE-2). param INIT_QUEUES: i16; +// 1: hand over with ``.advance_switch`` on the data fabout. 0: a SWITCH_ADV control wavelet +// after the data microthread reports completion (see sender.csl). +param ADVANCE_SWITCH: i16; + const memcpy = @import_module("", .{ .width = 2 * M, .height = 1, @@ -65,6 +69,7 @@ layout { // The westmost sender is the last to send and has nothing to relay for. .hands_over = x > 0, .init_queues = INIT_QUEUES, + .advance_switch = ADVANCE_SWITCH, }); } for (@range(i16, 0, M, 1)) |q| { diff --git a/tests/csl_runtime/handwritten/shift_bundle/run.py b/tests/csl_runtime/handwritten/shift_bundle/run.py index a8f99c75..9062ca42 100644 --- a/tests/csl_runtime/handwritten/shift_bundle/run.py +++ b/tests/csl_runtime/handwritten/shift_bundle/run.py @@ -22,6 +22,8 @@ def main() -> int: parser.add_argument('--M', type=int, required=True) parser.add_argument('--K', type=int, default=1) parser.add_argument('--filter', type=int, default=0) + parser.add_argument('--dump-core', action='store_true', + help='write corefile.cs1 before stopping, including on a stall') args = parser.parse_args() m, k = args.M, args.K @@ -37,15 +39,19 @@ def main() -> int: runner.load() runner.run() - runner.memcpy_h2d(val_id, vals.ravel(), 0, 0, width, 1, k, streaming=False, - data_type=crt.MemcpyDataType.MEMCPY_32BIT, - order=crt.MemcpyOrder.ROW_MAJOR, nonblock=False) - runner.launch('main', nonblock=False) - got = np.zeros(width * stream, dtype=np.float32) - runner.memcpy_d2h(got, got_id, 0, 0, width, 1, stream, streaming=False, - data_type=crt.MemcpyDataType.MEMCPY_32BIT, - order=crt.MemcpyOrder.ROW_MAJOR, nonblock=False) - runner.stop() + try: + runner.memcpy_h2d(val_id, vals.ravel(), 0, 0, width, 1, k, streaming=False, + data_type=crt.MemcpyDataType.MEMCPY_32BIT, + order=crt.MemcpyOrder.ROW_MAJOR, nonblock=False) + runner.launch('main', nonblock=False) + got = np.zeros(width * stream, dtype=np.float32) + runner.memcpy_d2h(got, got_id, 0, 0, width, 1, stream, streaming=False, + data_type=crt.MemcpyDataType.MEMCPY_32BIT, + order=crt.MemcpyOrder.ROW_MAJOR, nonblock=False) + finally: + if args.dump_core: + runner.dump_core('corefile.cs1') + runner.stop() got = got.reshape(width, stream) for x in range(width): diff --git a/tests/csl_runtime/handwritten/shift_bundle/sender.csl b/tests/csl_runtime/handwritten/shift_bundle/sender.csl index 8c43c442..feded770 100644 --- a/tests/csl_runtime/handwritten/shift_bundle/sender.csl +++ b/tests/csl_runtime/handwritten/shift_bundle/sender.csl @@ -2,14 +2,25 @@ // // All senders are launched at once and the order sorts itself out in the fabric: a sender // west of us cannot push its words through our router while we still receive from the ramp, -// so it waits on the link until our SWITCH_ADV has moved us to relay mode. That is the whole +// so it waits on the link until our switch has moved us to relay mode. That is the whole // serialisation mechanism -- no barrier, no counting. +// +// Two ways to make that switch, selected by ``advance_switch``: +// +// 0 a SWITCH_ADV control wavelet on the same output queue, after the data microthread +// reports completion. The wavelet has to sit behind the data, which is why it shares +// the queue; WSE-2 output queue 2 holds six 16-bit words, so K f32 values already fill +// it at K = 3, and a sender that cannot drain yet then blocks the compute element on +// the synchronous @mov32. Async completion is not an empty queue (see @queue_flush). +// 1 ``.advance_switch`` on the data fabout itself, which advances this router when the +// last data wavelet is sent. No second operation, no extra wavelet in the stream. param memcpy_params: comptime_struct; param stream: i16; param words: i16; param hands_over: bool; param init_queues: i16; +param advance_switch: i16; const sys_mod = @import_module("", memcpy_params); const ctrl = @import_module(""); @@ -26,6 +37,7 @@ const out_dsd = @get_dsd(fabout_dsd, .{ .extent = words, .fabric_color = channel, .output_queue = @get_output_queue(2), + .advance_switch = if (advance_switch != 0) hands_over else false, }); // The control wavelet has to leave through the same queue as the data, so that it stays // behind our own words and only advances our router once they are out. @@ -44,7 +56,10 @@ task send_task() void { } task done_task() void { - if (hands_over) { + // A separate SWITCH_ADV is only issued when the data DSD does not advance the switch + // itself. Posting it asynchronously is illegal: the source is a scalar, and even as a + // DSD it would share output queue 2 with the data microthread. + if (hands_over and (advance_switch == 0)) { @mov32(switch_dsd, ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0)); } sys_mod.unblock_cmd_stream(); diff --git a/tests/csl_runtime/test_shift_bundle_filters.sh b/tests/csl_runtime/test_shift_bundle_filters.sh index 06422522..d50b3d26 100644 --- a/tests/csl_runtime/test_shift_bundle_filters.sh +++ b/tests/csl_runtime/test_shift_bundle_filters.sh @@ -8,8 +8,10 @@ # Stage 2 turns the counter filters on and checks that each receiver keeps only its own block. # Together they pin down the hardware contract the compiler relies on: # -# * a sender's own SWITCH_ADV advances its own router, and over-advancing a router that is -# already on its last position is harmless, +# * a sender hands its own router over on the last data wavelet (``.advance_switch`` on the +# fabout). A SWITCH_ADV on the same output queue is what drops a wavelet on WSE-2 once a +# back-pressured send of three or more f32 values fills that queue; over-advancing a +# router that is already on its last position is still harmless, # * a receiver transmitting to RAMP and EAST duplicates rather than consumes, # * a filter withholds a wavelet from the compute element without removing it from the # network, and a wavelet nobody keeps is dropped by the terminating router, @@ -41,11 +43,17 @@ compile_and_run() { rm -rf "$OUT" cslc --arch="$arch" "$SRC/layout.csl" -o "$OUT" \ --fabric-dims=$((7 + width)),3 --fabric-offsets=4,1 --memcpy --channels=1 \ - --params=M:$m,K:$k,FILTER:$filter,IN_QUEUE:$in_queue,INIT_QUEUES:$init_queues + --params=M:$m,K:$k,FILTER:$filter,IN_QUEUE:$in_queue,INIT_QUEUES:$init_queues,ADVANCE_SWITCH:1 timeout -s 9 240 cs_python "$SRC/run.py" "$OUT" --M "$m" --K "$k" --filter "$filter" } -for case in "3 1" "4 1" "3 2"; do +# The cases grow in the two directions the contract has to hold in: the number of senders sharing +# the color, and the words each of them ships. K > 1 is what makes the window narrower than the +# cycle, and M*K is the cycle the counter has to wrap at, so "4 4" is the corner where a window of +# four sits inside a cycle of sixteen and the last receiver's counter starts at twelve. That is the +# shape batcher_oddeven_bundled_1D takes at L = 3, K = 4 for its distance-4 phase. Ordered so the +# first failure marks the boundary. +for case in "3 1" "4 1" "3 2" "4 2" "3 4" "4 4"; do set -- $case echo "--- stage 1: arrival order, M=$1 K=$2 (no filters) ---" compile_and_run "$1" "$2" 0 diff --git a/tests/spatial_ir/test_routing.py b/tests/spatial_ir/test_routing.py index 79dba11c..457a8369 100644 --- a/tests/spatial_ir/test_routing.py +++ b/tests/spatial_ir/test_routing.py @@ -216,20 +216,20 @@ def test_two_phase_split_switch_plans(): def test_two_phase_split_emits_two_control_wavelets(): """ Of the six closes in the sample, only the two on PEs that *send* a stream whose path contains a - router that must advance survive elision. + router that must advance survive elision. PE 1 only flips its own router, which the last data + wavelet does; PE 3 has to turn PE 2 around, which is a traveling SWITCH_ADV. """ files = _lower('two_phase_split.sptl', K=32) - emitting = {name: code for name, code in files.items() if 'switch_dsd' in code} - assert sorted(emitting) == ['code_1_0.csl', 'code_3_0.csl'] + + assert '.advance_switch = true' in files['code_1_0.csl'] + assert 'SWITCH_ADV' not in files['code_1_0.csl'] + assert 'switch_dsd' not in files['code_1_0.csl'] payload = 'ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0)' - # PE 1 moves its own router one position; PE 3 retires PE 2's incoming configuration, and PE 2 - # has to turn around, which takes two positions where a switch carries only one direction. - assert emitting['code_1_0.csl'].count(payload) == 1 - assert emitting['code_3_0.csl'].count(payload) == (1 if csl.SWITCH_POSITION_ALLOWS_BOTH else 2) - for code in emitting.values(): - assert 'const ctrl = @import_module("");' in code - assert '.control = true' in code + # PE 2 has to turn around, which takes two positions where a switch carries only one direction. + assert files['code_3_0.csl'].count(payload) == (1 if csl.SWITCH_POSITION_ALLOWS_BOTH else 2) + assert 'const ctrl = @import_module("");' in files['code_3_0.csl'] + assert '.control = true' in files['code_3_0.csl'] def test_two_phase_split_without_switching_falls_back(): diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py index d5c45134..cfeaf869 100644 --- a/tests/spatial_ir/test_shift_bundles.py +++ b/tests/spatial_ir/test_shift_bundles.py @@ -203,12 +203,20 @@ def test_windows_are_as_wide_as_the_stream_bound(): def test_each_source_advances_its_own_switch_once(): codes = _codes(_SHIFT, M=3, D=3, K=1) sending = codes['code_0_0.csl'] - assert sending.count('ctrl.opcode.SWITCH_ADV') == 1 - assert '.control = true' in sending - # The destinations do not switch, so nothing is emitted there. + # A source only has to flip its own router, which the last data wavelet does. A SWITCH_ADV + # on the same output queue is what drops a wavelet on WSE-2 once a back-pressured send fills it. + assert '.advance_switch = true' in sending + assert 'SWITCH_ADV' not in sending + assert '.control = true' not in sending assert 'SWITCH_ADV' not in codes['code_3_0.csl'] +def test_a_wide_payload_still_advances_on_the_last_data_wavelet(): + sending = _codes(_SHIFT, M=3, D=3, K=4)['code_0_0.csl'] + assert '.advance_switch = true' in sending + assert 'SWITCH_ADV' not in sending + + def test_one_color_carries_the_whole_bundle(): layout = _layout(_SHIFT, M=3, D=3, K=1) routes = layout[layout.index('// Routes'):] From 42a4f174c2db8e579ddebcade0bdb15c027a4282 Mon Sep 17 00:00:00 2001 From: glukas Date: Thu, 20 Aug 2026 06:42:50 +0200 Subject: [PATCH 51/68] Fix spmv, gemv --- samples/spatial/blas/gemv.sptl | 9 +++++++-- samples/spatial/blas/gemv_twophase.sptl | 9 +++++++-- samples/spatial/blas/spmv.sptl | 9 +++++++-- 3 files changed, 21 insertions(+), 6 deletions(-) diff --git a/samples/spatial/blas/gemv.sptl b/samples/spatial/blas/gemv.sptl index 396ba772..b109707d 100644 --- a/samples/spatial/blas/gemv.sptl +++ b/samples/spatial/blas/gemv.sptl @@ -42,11 +42,16 @@ kernel @gemv( } // Phase 2: Load x on the j=0 column and y on the i=0 column. + // The three ranges are disjoint: PE(0,0) belongs to both the row and the column. phase { - compute i16 i, i16 j in [0:PX, 0] { + compute i16 i, i16 j in [0, 0] { await receive(x, inp_x[i, j]) + await receive(y_block, inp_y[i, j]) } - compute i16 i, i16 j in [0, 0:PY] { + compute i16 i, i16 j in [1:PX, 0] { + await receive(x, inp_x[i, j]) + } + compute i16 i, i16 j in [0, 1:PY] { await receive(y_block, inp_y[i, j]) } } diff --git a/samples/spatial/blas/gemv_twophase.sptl b/samples/spatial/blas/gemv_twophase.sptl index 10cf36b4..e7b2488d 100644 --- a/samples/spatial/blas/gemv_twophase.sptl +++ b/samples/spatial/blas/gemv_twophase.sptl @@ -47,11 +47,16 @@ kernel @gemv_twophase( } // Phase 2: Load x on the j=0 column and y on the i=0 column. + // The three ranges are disjoint: PE(0,0) belongs to both the row and the column. phase { - compute i16 i, i16 j in [0:G*S, 0] { + compute i16 i, i16 j in [0, 0] { await receive(x, inp_x[i, j]) + await receive(y_block, inp_y[i, j]) } - compute i16 i, i16 j in [0, 0:PY] { + compute i16 i, i16 j in [1:G*S, 0] { + await receive(x, inp_x[i, j]) + } + compute i16 i, i16 j in [0, 1:PY] { await receive(y_block, inp_y[i, j]) } } diff --git a/samples/spatial/blas/spmv.sptl b/samples/spatial/blas/spmv.sptl index db28f519..15144bbb 100644 --- a/samples/spatial/blas/spmv.sptl +++ b/samples/spatial/blas/spmv.sptl @@ -52,11 +52,16 @@ kernel @spmv( } // Phase 2: Load x on the j=0 column and y on the i=0 column. + // The three ranges are disjoint: PE(0,0) belongs to both the row and the column. phase { - compute i16 i, i16 j in [0:PX, 0] { + compute i16 i, i16 j in [0, 0] { await receive(x, inp_x[i, j]) + await receive(y_block, inp_y[i, j]) } - compute i16 i, i16 j in [0, 0:PY] { + compute i16 i, i16 j in [1:PX, 0] { + await receive(x, inp_x[i, j]) + } + compute i16 i, i16 j in [0, 1:PY] { await receive(y_block, inp_y[i, j]) } } From 65de5c3c65240d8e08266fad334c14e2bcb4cfe3 Mon Sep 17 00:00:00 2001 From: glukas Date: Thu, 20 Aug 2026 07:24:31 +0200 Subject: [PATCH 52/68] keep .advance_switch wse-2 only --- irspec/docs/spatial/routing_wse.md | 24 +++++++++++++----------- spada/lowering/spatial_ir_to_csl.py | 4 ++-- spada/syntax/csl/routing.py | 19 +++++++++++-------- spada/syntax/spatial_ir/irnodes.py | 6 +++--- tests/spatial_ir/test_routing.py | 14 +++++++++----- tests/spatial_ir/test_shift_bundles.py | 23 ++++++++++++++++------- 6 files changed, 54 insertions(+), 36 deletions(-) diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index b14a391e..1829f4b4 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -68,11 +68,13 @@ all. An advance is driven by the `close` that ends the epoch. When some *other* router on the path has to move, the sending PE emits a *switch-advance control message* on the channel, one per position to be traversed. It follows the stream's path using the configuration that is being retired, and -advances the router of each PE it traverses, after all data of the epoch. When only the sending -PE's own router has to move, and only by one position, the last data wavelet does that itself -(`.advance_switch` on the fabric output DSD). A second operation on the same output queue is what -drops a data wavelet on WSE-2 once a back-pressured send fills it (output queues 2 and 3 hold six -16-bit words; three `f32` values already fill them). +advances the router of each PE it traverses, after all data of the epoch. On WSE-2, when only the +sending PE's own router has to move, and only by one position, the last data wavelet does that +itself (`.advance_switch` on the fabric output DSD). A second operation on the same output queue +is what drops a data wavelet there once a back-pressured send fills it (output queues 2 and 3 +hold six 16-bit words; three `f32` values already fill them). WSE-3 keeps a traveling `SWITCH_ADV` +for that local flip as well: its output queues hold eight words, and origin-pooled destinations +that also switch are only moved by a control wavelet on the path. !!! warning "WSE: Advances Are Not Selective" A CSL control wavelet nominally carries up to eight per-router switching commands @@ -152,12 +154,12 @@ filter: -- -- -- win 2 win 1 win 0 ``` **Sources inject, then relay.** A source starts at `rx = RAMP, tx = {EAST}` with `pos1` taking -`rx = WEST`, sends its own words, and the last data wavelet advances its own router into relay -mode (`.advance_switch` on the fabric output DSD). The trigger is local — *"my own send is -done"* — which is what makes it expressible at all, given that the payload of a control message -[selects nothing](#lowering-to-switches). A `SWITCH_ADV` on the same output queue would also -flip that router, but on WSE-2 a back-pressured send of three or more `f32` values fills the -queue and the control wavelet then steals a data word. +`rx = WEST`, sends its own words, and then advances its own router into relay mode. The trigger +is local — *"my own send is done"* — which is what makes it expressible at all, given that the +payload of a control message [selects nothing](#lowering-to-switches). On WSE-2 that advance is +`.advance_switch` on the fabric output DSD: a `SWITCH_ADV` on the same output queue would also +flip that router, but a back-pressured send of three or more `f32` values fills the six-word +queue and the control wavelet then steals a data word. On WSE-3 the same close emits `SWITCH_ADV`. **The order is descending, and enforces itself.** The source nearest the destinations goes first. No schedule or barrier is needed: a source further away cannot push a word through its neighbour's diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index f2e35da6..1118d98f 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -1207,8 +1207,8 @@ def allocate_output_queue(stream: spir.Identifier) -> int: f'{location}: no output queue was reserved for {key} (stream "{stream.as_ir()}").') return output_queue_of[key] - # A close that only flips this PE's own router does so on the last data wavelet, so the - # outgoing fabric descriptor has to carry ``.advance_switch``. The close itself emits no + # On WSE-2, a close that only flips this PE's own router does so on the last data wavelet, so + # the outgoing fabric descriptor has to carry ``.advance_switch``. The close itself emits no # control wavelet and is kept only so this scan can see the flag. streams_advance_on_send = { stream_lifetime.underlying_stream(stmt.stream_name).as_ir() diff --git a/spada/syntax/csl/routing.py b/spada/syntax/csl/routing.py index 892f5021..8fcebdda 100644 --- a/spada/syntax/csl/routing.py +++ b/spada/syntax/csl/routing.py @@ -695,8 +695,9 @@ def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: *other* router on the path has to move: the wavelet travels the path being retired and every switch-configured router it reaches advances. A close that only has to flip the sending PE's own router does that on the last data wavelet (``.advance_switch`` on the fabric output DSD) - instead -- a second operation on the same output queue is what drops a data wavelet on WSE-2 - once a back-pressured send fills it. + instead, but only on WSE-2 -- a second operation on the same output queue is what drops a + data wavelet there once a back-pressured send fills it. WSE-3 output queues hold eight words, + so a traveling ``SWITCH_ADV`` is safe, and origin-pooled destinations still need one. All routers on the path that hold switch positions and are advanced by a control wavelet must therefore advance by the same amount, and that amount is how many wavelets are sent. A close on @@ -804,12 +805,14 @@ def plan_switch_advances(rectangles: list[Rectangle[PEBlock]]) -> int: 'additional color') amount = distinct.pop() - # A source that only flips its own router does so on the last data wavelet. Posting a - # SWITCH_ADV into the same output queue afterwards is what drops a data wavelet on - # WSE-2 when a back-pressured send of three or more f32 values fills that queue - # (see tests/csl_runtime/test_shift_bundle_filters.sh). Remote routers still need a - # traveling control wavelet, and a two-position turnaround still needs two of them. - if not remote_advances and local_advance == 1: + # A source that only flips its own router does so on the last data wavelet, but only + # on WSE-2. Posting a SWITCH_ADV into the same output queue afterwards is what drops + # a data wavelet there when a back-pressured send of three or more f32 values fills + # that queue (see tests/csl_runtime/test_shift_bundle_filters.sh). WSE-3 queues hold + # eight words, so SWITCH_ADV is safe; origin-pooled Batcher destinations also switch + # and only a traveling control wavelet moves them. Remote routers still need that + # wavelet, and a two-position turnaround still needs two of them. + if not remote_advances and local_advance == 1 and constants.ARCH == 'wse2': statement.advance_data_switch = True else: statement.switch_advance = amount diff --git a/spada/syntax/spatial_ir/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index f83b0152..9890e3c4 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -937,9 +937,9 @@ class CloseStatement(Statement): switch_advance: Optional[int] = None #: When True, the sending PE flips only its own router, and does so on the last data wavelet #: (``.advance_switch`` on the fabric output DSD) rather than by a ``SWITCH_ADV`` control - #: wavelet. A control wavelet on the same output queue as the data is what drops a wavelet on - #: WSE-2 once a back-pressured send of three or more f32 values fills the queue. Mutually - #: exclusive with ``switch_advance``. Not part of the surface syntax. + #: wavelet. Used only on WSE-2: a control wavelet on the same output queue as the data is + #: what drops a wavelet there once a back-pressured send of three or more f32 values fills + #: the queue. Mutually exclusive with ``switch_advance``. Not part of the surface syntax. advance_data_switch: bool = False def validate(self) -> None: diff --git a/tests/spatial_ir/test_routing.py b/tests/spatial_ir/test_routing.py index 457a8369..66ea9099 100644 --- a/tests/spatial_ir/test_routing.py +++ b/tests/spatial_ir/test_routing.py @@ -216,14 +216,18 @@ def test_two_phase_split_switch_plans(): def test_two_phase_split_emits_two_control_wavelets(): """ Of the six closes in the sample, only the two on PEs that *send* a stream whose path contains a - router that must advance survive elision. PE 1 only flips its own router, which the last data - wavelet does; PE 3 has to turn PE 2 around, which is a traveling SWITCH_ADV. + router that must advance survive elision. PE 1 only flips its own router (WSE-2: last data + wavelet; WSE-3: SWITCH_ADV). PE 3 has to turn PE 2 around, which is a traveling SWITCH_ADV. """ files = _lower('two_phase_split.sptl', K=32) - assert '.advance_switch = true' in files['code_1_0.csl'] - assert 'SWITCH_ADV' not in files['code_1_0.csl'] - assert 'switch_dsd' not in files['code_1_0.csl'] + if csl.ARCH == 'wse2': + assert '.advance_switch = true' in files['code_1_0.csl'] + assert 'SWITCH_ADV' not in files['code_1_0.csl'] + assert 'switch_dsd' not in files['code_1_0.csl'] + else: + assert 'SWITCH_ADV' in files['code_1_0.csl'] + assert '.advance_switch = true' not in files['code_1_0.csl'] payload = 'ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0)' # PE 2 has to turn around, which takes two positions where a switch carries only one direction. diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py index cfeaf869..c3cf8006 100644 --- a/tests/spatial_ir/test_shift_bundles.py +++ b/tests/spatial_ir/test_shift_bundles.py @@ -203,18 +203,27 @@ def test_windows_are_as_wide_as_the_stream_bound(): def test_each_source_advances_its_own_switch_once(): codes = _codes(_SHIFT, M=3, D=3, K=1) sending = codes['code_0_0.csl'] - # A source only has to flip its own router, which the last data wavelet does. A SWITCH_ADV - # on the same output queue is what drops a wavelet on WSE-2 once a back-pressured send fills it. - assert '.advance_switch = true' in sending - assert 'SWITCH_ADV' not in sending - assert '.control = true' not in sending + # A source only has to flip its own router. WSE-2 does that on the last data wavelet: + # a SWITCH_ADV on the same output queue drops a wavelet once a back-pressured send fills it. + # WSE-3 keeps SWITCH_ADV; origin-pooled destinations still need a traveling control wavelet. + if constants.ARCH == 'wse2': + assert '.advance_switch = true' in sending + assert 'SWITCH_ADV' not in sending + assert '.control = true' not in sending + else: + assert 'SWITCH_ADV' in sending + assert '.advance_switch = true' not in sending assert 'SWITCH_ADV' not in codes['code_3_0.csl'] def test_a_wide_payload_still_advances_on_the_last_data_wavelet(): sending = _codes(_SHIFT, M=3, D=3, K=4)['code_0_0.csl'] - assert '.advance_switch = true' in sending - assert 'SWITCH_ADV' not in sending + if constants.ARCH == 'wse2': + assert '.advance_switch = true' in sending + assert 'SWITCH_ADV' not in sending + else: + assert 'SWITCH_ADV' in sending + assert '.advance_switch = true' not in sending def test_one_color_carries_the_whole_bundle(): From dec08a73c80c01574d9c054d673aea3ac4448f51 Mon Sep 17 00:00:00 2001 From: glukas Date: Thu, 20 Aug 2026 07:42:09 +0200 Subject: [PATCH 53/68] deeper testing, remove bitonic sort sample --- samples/spatial/sort/bitonic_sort_1D.sptl | 158 ------------------ tests/csl_runtime/test_batcher_oddeven_1d.sh | 1 + .../test_batcher_oddeven_bundled_1d.sh | 1 + .../test_batcher_oddeven_wse3_1d.sh | 3 + tests/csl_runtime/test_bitonic_sort_1d.sh | 77 --------- .../test_odd_even_sort_1d_looped.sh | 2 + 6 files changed, 7 insertions(+), 235 deletions(-) delete mode 100644 samples/spatial/sort/bitonic_sort_1D.sptl delete mode 100644 tests/csl_runtime/test_bitonic_sort_1d.sh diff --git a/samples/spatial/sort/bitonic_sort_1D.sptl b/samples/spatial/sort/bitonic_sort_1D.sptl deleted file mode 100644 index 08b8e0be..00000000 --- a/samples/spatial/sort/bitonic_sort_1D.sptl +++ /dev/null @@ -1,158 +0,0 @@ -/** - * Batcher's bitonic sorting network over N = 2^L PEs in a row. - * - * Each PE holds K keys, and the network sorts K independent sequences at once: sequence k is made - * of element k of every PE. A compare-exchange is therefore an elementwise min/max over K values, - * which streams as one K-element transfer per epoch. - * - * The point of this sample is *channel economy*. A bitonic network on N keys performs L(L+1)/2 - * compare-exchange steps at distances 1, 2, 4, ..., N/2, and every PE takes part in every step. - * Giving each step its own channel would need one per step; giving each concurrently exchanging - * pair its own channel would need O(N). This kernel instead uses - * - * L = log2(N) channels - * - * -- exactly one per exchange distance -- and reuses each of them across every step, every lane and - * both directions of travel. That reuse is what the router switches pay for. - * - * Why the network needs lanes - * --------------------------- - * At distance J = 2^d every PE with bit d clear is a "low" PE and exchanges with the PE J to its - * east. Those paths overlap: 0 -> 4 and 1 -> 5 both cross PEs 1..4, so they cannot share a channel - * concurrently (see the channel-conflict rule in the routing specification). The step is therefore - * split into J *lanes*: lane c handles the PEs congruent to c modulo 2J, whose paths are exactly - * disjoint. - * - * Why each lane is two phases - * --------------------------- - * A compare-exchange has to move a key in each direction, and both directions use the same channel. - * They cannot be concurrent, so the lane is split into an eastward phase and a westward one; the - * phase boundary closes the stream, which is what frees the channel. Reversing the direction of - * travel is what makes this kernel demanding: every PE on the path -- sender, relay and receiver - * alike -- has to change both its router's input and its output between the two phases. - * - * Structure, generated entirely by compile-time `for` blocks: - * - * for s in [0, L) stage: groups of size 2^(s+1) are sorted - * for e in [0, s] step: exchange distance J = 2^(s-e), halving - * for c in [0, J) lane: two phases, two epochs on channel s-e - * for g in ... the ascending groups, then the descending ones - * - * A group of size 2^(s+1) sorts ascending when bit s+1 of its index is clear and descending - * otherwise, so the ascending groups start at multiples of 2^(s+2) and the descending ones are - * offset by 2^(s+1). In the final stage the descending range is empty, which is what leaves the - * whole array sorted ascending. - * - * One lane (L=3, s=2, e=0, J=4, c=1) on channel 2, as two epochs: - * - * PE: 0 1 2 3 4 5 6 7 - * east *---->-----------------* 1 -> 5, relays 2,3,4 - * west *-----------------<----* 5 -> 1, relays 4,3,2 - * - * (!) Assumes L >= 1. Requires WSE-3: reversing a router between the two phases changes both its - * input and its output direction, and enough of those accumulate on one color to exceed the - * four switch positions a WSE-2 router can hold (where each such reversal costs two). (!) - * - * (!) The router budget caps this at L = 2. At distance 2^d the channel is reused by 2^d lanes in - * two directions each, so an interior router cycles through 2^(d+1) configurations; four - * switch positions run out at d = 2. Sorting more keys needs a channel per direction, which - * doubles the channel count to 2*log2(N) and, since no router then reverses, also runs on - * WSE-2. (!) - * - * Note on syntax: `+`/`-` are right-associative in this grammar, so `s-e+1` would parse as - * `s-(e+1)`. The step expressions below write `(s-e)+1` explicitly. - **/ -kernel @bitonic_sort_1d(stream[1<[1< east = relative_stream(1<<(s-e), 0) { - hops = auto, - channel = s-e - } - } - for i16 g in [0 : 1< west = relative_stream(-(1<<(s-e)), 0) { - hops = auto, - channel = s-e - } - } - for i16 g in [0 : 1< val[k] else val[k]) - } - } - } - for i16 g in [1<<(s+1) : 1< val[k] else val[k]) - } - } - compute i16 i, i16 j in [g + c + (1<<(s-e)) : g + (1<<(s+1)) : 1<<((s-e)+1), 0] { - await send(val, west) - await map i32 k in [0:K] { - val[k] = (other[k] if other[k] < val[k] else val[k]) - } - } - } - } - } - } - } - - // Store the sorted keys. - phase { - compute i16 i, i16 j in [0:1< Date: Thu, 20 Aug 2026 07:47:29 +0200 Subject: [PATCH 54/68] Multi-row sorts --- README.md | 2 +- samples/spatial/sort/batcher_oddeven_1D.sptl | 35 +++++++-------- .../sort/batcher_oddeven_bundled_1D.sptl | 35 +++++++-------- .../spatial/sort/batcher_oddeven_wse3_1D.sptl | 35 +++++++-------- tests/csl_runtime/test_batcher_oddeven_1d.sh | 43 ++++++++++--------- .../test_batcher_oddeven_bundled_1d.sh | 34 ++++++++------- .../test_batcher_oddeven_wse3_1d.sh | 34 ++++++++------- tests/spatial_ir/test_dsd_ops.py | 4 +- tests/spatial_ir/test_shift_bundles.py | 4 +- 9 files changed, 120 insertions(+), 106 deletions(-) diff --git a/README.md b/README.md index 15b37a24..e97567c0 100644 --- a/README.md +++ b/README.md @@ -104,7 +104,7 @@ Sample SPADA programs are in `samples/`: | `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy` | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Sorting networks over `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), `bitonic_sort_1D` (K independent sequences; channel reuse via router switches, WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels) | +| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), `bitonic_sort_1D` (K independent sequences; channel reuse via router switches, WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | diff --git a/samples/spatial/sort/batcher_oddeven_1D.sptl b/samples/spatial/sort/batcher_oddeven_1D.sptl index 560c0e22..023619e4 100644 --- a/samples/spatial/sort/batcher_oddeven_1D.sptl +++ b/samples/spatial/sort/batcher_oddeven_1D.sptl @@ -1,6 +1,7 @@ /** - * 1D Batcher odd-even mergesort over N = 2^L PEs, each holding K f32 keys: N*K keys in all. - * Ascending: after the network, PE i holds keys i*K .. i*K + (K-1) of the sorted sequence. + * 1D Batcher odd-even mergesort over R independent rows of N = 2^L PEs, each holding K f32 keys. + * Ascending: after the network, PE (i, j) holds keys i*K .. i*K + (K-1) of row j. Every hop is + * east-west, so the R rows never mix. * * Merge stages l = 1 .. L, each of width 2^l. Within stage l: * p = 1: every PE in each 2^l box compares at dist = 2^{l-1} @@ -56,13 +57,13 @@ * (l=3,p=2) dist 2: (2,4)(3,5) * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) * - * Constraints: L >= 1, K >= 1 + * Constraints: L >= 1, K >= 1, R >= 1 **/ -kernel @batcher_oddeven_1d( - stream[1<[1<( + stream[1<[1<( // Load this PE's block and sort it, which is the invariant the comparators below rely on. phase { - compute i16 i, i16 j in [0:1<( // Offset r is one disjoint matching; each r has its own fwd/bwd colors. phase { for i16 r in [0:1<<(l-1)] { - dataflow i16 i, i16 j in [r:1< fwd = relative_stream(1<<(l-1), 0) { hops = auto, channel = 2 * (((1<<(l-1)) - 1) + r) @@ -106,7 +107,7 @@ kernel @batcher_oddeven_1d( channel = (2 * (((1<<(l-1)) - 1) + r)) + 1 } } - dataflow i16 i, i16 j in [(r + (1<<(l-1))):1< fwd = relative_stream(1<<(l-1), 0) { hops = auto, channel = 2 * (((1<<(l-1)) - 1) + r) @@ -117,7 +118,7 @@ kernel @batcher_oddeven_1d( } } - compute i16 i, i16 j in [r:1<( val[m] = res[m] } } - compute i16 i, i16 j in [(r + (1<<(l-1))):1<( phase { for i16 r in [0:1<<(l-p)] { for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { hops = auto, channel = (2 * ((1<( channel = ((2 * ((1< fwd = relative_stream(1<<(l-p), 0) { hops = auto, channel = (2 * ((1<( } } - compute i16 i, i16 j in [(b + ((1<<(l-p)) + r)):((b + (1<( val[m] = res[m] } } - compute i16 i, i16 j in [(b + ((2 * (1<<(l-p))) + r)):(b + (1<( // Write the sorted blocks back to the host. phase { - compute i16 i, i16 j in [0:1<= 1 + * Constraints: 1 <= L <= 4 (L <= 3 on wse2), K >= 1, R >= 1 **/ -kernel @batcher_oddeven_1d( - stream[1<[1<( + stream[1<[1<( // Load this PE's block and sort it, which is the invariant the comparators below rely on. phase { - compute i16 i, i16 j in [0:1<( // Offset r is one disjoint matching. The matchings share a channel where the phase is // bundled, and take one each where it is not, which is what turns bundling off. for i16 r in [0:1<<(l-1)] { - dataflow i16 i, i16 j in [r:1< fwd = relative_stream(1<<(l-1), 0) { hops = auto, channel = ((1<= (1<( channel = (((1<= (1< fwd = relative_stream(1<<(l-1), 0) { hops = auto, channel = ((1<= (1<( } } - compute i16 i, i16 j in [r:1<( val[m] = res[m] } } - compute i16 i, i16 j in [(r + (1<<(l-1))):1<( phase { for i16 r in [0:1<<(l-p)] { for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { hops = auto, channel = ((1<= (1<( channel = (((1<= (1< fwd = relative_stream(1<<(l-p), 0) { hops = auto, channel = ((1<= (1<( } } - compute i16 i, i16 j in [(b + ((1<<(l-p)) + r)):((b + (1<( val[m] = res[m] } } - compute i16 i, i16 j in [(b + ((2 * (1<<(l-p))) + r)):(b + (1<( // Write the sorted blocks back to the host. phase { - compute i16 i, i16 j in [0:1<= 1. WSE-3. + * Constraints: 1 <= L <= 4, K >= 1, R >= 1. WSE-3. **/ -kernel @batcher_oddeven_wse3_1d( - stream[1<[1<( + stream[1<[1<( // Load this PE's block and sort it, which is the invariant the comparators below rely on. phase { - compute i16 i, i16 j in [0:1<( // mod 2*dist and the high partner r + dist, which is the channel each of them sends on. phase { for i16 r in [0:1<<(l-1)] { - dataflow i16 i, i16 j in [r:1< fwd = relative_stream(1<<(l-1), 0) { hops = auto, channel = ((1<= (1<( channel = (((1<= (1< fwd = relative_stream(1<<(l-1), 0) { hops = auto, channel = ((1<= (1<( } } - compute i16 i, i16 j in [r:1<( val[m] = res[m] } } - compute i16 i, i16 j in [(r + (1<<(l-1))):1<( phase { for i16 r in [0:1<<(l-p)] { for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { hops = auto, channel = ((1<= (1<( channel = (((1<= (1< fwd = relative_stream(1<<(l-p), 0) { hops = auto, channel = ((1<= (1<( } } - compute i16 i, i16 j in [(b + ((1<<(l-p)) + r)):((b + (1<( val[m] = res[m] } } - compute i16 i, i16 j in [(b + ((2 * (1<<(l-p))) + r)):(b + (1<( // Write the sorted blocks back to the host. phase { - compute i16 i, i16 j in [0:1< Date: Thu, 20 Aug 2026 08:54:07 +0200 Subject: [PATCH 55/68] Fix queue remapping logic --- README.md | 2 +- spada/lowering/spatial_ir_to_csl.py | 56 ++++++++++++++--- tests/spatial_ir/test_routing.py | 96 ++++++++++++++--------------- 3 files changed, 93 insertions(+), 61 deletions(-) diff --git a/README.md b/README.md index e97567c0..6685aeae 100644 --- a/README.md +++ b/README.md @@ -104,7 +104,7 @@ Sample SPADA programs are in `samples/`: | `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy` | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), `bitonic_sort_1D` (K independent sequences; channel reuse via router switches, WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels) | +| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 1118d98f..1d56449e 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -971,14 +971,56 @@ def _declare_queue_initialization(dsds: UniqueDSDDict, rect: PEBlock, footer: St footer.write(f' @initialize_queue(@get_{kind}({queue}), .{{ .color = {color} }});\n') +def _statement_transfer_points(statement: spir.Statement, names: set[spir.Identifier], + inbound: bool) -> list[spir.Identifier]: + """ + Fabric transfers in ``statement``, in source order, with sequential ``for`` bodies repeated. + + Walking a loop body once makes its colours look sequential, so occupancy pooling would give + them one queue. The next iteration of an earlier colour can already occupy the router when a + later colour of the same body remaps that queue -- WSE-2 then aborts with "Attempt to remap + input queue N, from C_i to C_j, but the router is holding wavelets". Appending the body a + second time makes a colour used on both sides of another occupy a span that overlaps it, the + same rule that keeps a reused colour's queue across a gap between unrolled phases. + """ + if isinstance(statement, spir.ForStatement): + body = _transfer_points(statement.body, names, inbound) + return body + body + + nested_skip: set[int] = set() + points: list[spir.Identifier] = [] + for node in statement.walk(): + if id(node) in nested_skip: + continue + if node is not statement and isinstance(node, spir.ForStatement): + for descendant in node.walk(): + nested_skip.add(id(descendant)) + points.extend(_statement_transfer_points(node, names, inbound)) + continue + stream = _fabric_transfer_stream(node, inbound) + if stream is None or stream not in names: + continue + points.append(stream) + return points + + +def _transfer_points(statements: list[spir.Statement], names: set[spir.Identifier], + inbound: bool) -> list[spir.Identifier]: + points: list[spir.Identifier] = [] + for statement in statements: + points.extend(_statement_transfer_points(statement, names, inbound)) + return points + + def _queue_spans(compute: spir.ComputeBlock, names: set[spir.Identifier], queue_key, inbound: bool) -> dict[str, tuple[int, int]]: """ Occupancy of each queue key along the linearized send/receive order of this PE. - Nested transfers in a loop body are ordered by the walk, so sequential halo exchanges in one - ``for`` do not look concurrent. Uses of the same channel still collapse to one span, so a - colour that comes back after a gap keeps its queue for the whole of that span. + Sequential ``for`` bodies are counted twice so a colour that comes back on the next iteration + keeps its queue across the loop-carried gap; see ``_statement_transfer_points``. Uses of the + same channel still collapse to one span, so a colour that comes back after a gap between + unrolled phases keeps its queue for the whole of that span. :param compute: The compute block being lowered. :param names: Streams that actually bind a fabric queue in this direction. @@ -986,13 +1028,7 @@ def _queue_spans(compute: spir.ComputeBlock, names: set[spir.Identifier], queue_ :param inbound: True to walk receives, False to walk sends. :return: Mapping of grouping key to ``(first_use, last_use)`` in linearized order. """ - points: list[spir.Identifier] = [] - for statement in compute.statements: - for node in statement.walk(): - stream = _fabric_transfer_stream(node, inbound) - if stream is None or stream not in names: - continue - points.append(stream) + points = _transfer_points(compute.statements, names, inbound) spans: dict[str, tuple[int, int]] = {} for index, stream in enumerate(points): diff --git a/tests/spatial_ir/test_routing.py b/tests/spatial_ir/test_routing.py index 66ea9099..e946b00c 100644 --- a/tests/spatial_ir/test_routing.py +++ b/tests/spatial_ir/test_routing.py @@ -423,56 +423,6 @@ def test_switch_positions_beyond_capacity_are_rejected(): _lower_string(_stress_kernel(_MAX_STRESS_PHASES + 1), K=4) -### -# bitonic_sort_1D: the heaviest channel reuse in the samples -### - -_BITONIC = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', - 'bitonic_sort_1D.sptl') - - -def _lower_bitonic(L: int, K: int = 4) -> dict[str, str]: - kernel = parser.parse_file(_BITONIC) - kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, L=L, K=K)) - return {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel)} - - -@pytest.mark.skipif(not csl.SWITCH_POSITION_ALLOWS_BOTH, - reason=f'{csl.ARCH} cannot reverse a router within four switch positions') -def test_bitonic_sort_uses_one_channel_per_distance(): - """ - A bitonic network on 2^L keys needs L(L+1)/2 exchange steps but only L channels: one per - exchange distance, reused by every lane, every stage and both directions of travel. - - L is 2 here because that is what the router budget allows: at distance 2^d the channel is - reused by 2^d lanes in two directions each, so an interior router cycles through 2^(d+1) - configurations, and four positions run out at d = 2. - """ - files = _lower_bitonic(2) - colors = set(re.findall(r'@get_color\((\d+)\)', files['layout.csl'])) - assert len(colors) <= 2, sorted(colors) - - layout = files['layout.csl'] - assert '.switches' in layout - # The lane pattern repeats, so the configuration sequence closes into a ring. - assert 'ring_mode' in layout - # Every router stays inside its four positions. - for line in layout.splitlines(): - if '.switches' in line: - assert len(re.findall(r'\.pos\d', line)) < csl.SWITCH_POSITIONS, line - - -@pytest.mark.skipif(csl.SWITCH_POSITION_ALLOWS_BOTH, - reason='this architecture can reverse a router in a single switch position') -def test_bitonic_sort_is_rejected_on_wse2(): - """ - Reversing a router costs two positions where a position carries one direction, and the interior - routers of the network reverse often enough to exhaust them. - """ - with pytest.raises(SyntaxError, match='switch positions'): - _lower_bitonic(2) - - ### # odd_even_sort_1D_looped: N rounds as a runtime loop on four static channels ### @@ -521,5 +471,51 @@ def test_odd_even_sort_looped_code_is_independent_of_n(): assert abs(len(small['code_2_0.csl']) - len(large['code_2_0.csl'])) < 64 +def _fabin_queues(code: str) -> dict[str, str]: + return dict(re.findall( + r'const (\w+)_in_dsd = @get_dsd\(fabin_dsd, .*?input_queue = @get_input_queue\((\d+)\)', + code, flags=re.S)) + + +def test_odd_even_sort_looped_interior_keeps_distinct_input_queues(): + """ + Even-round east and odd-round west are both inbound on an odd interior PE. The west neighbour + can inject the next even-round block on C0 while this PE is already receiving the odd-round + one on C3. Sharing input queue 0 is what the WSE-2 simulator rejects as remapping C0 onto C3 + while the router still holds wavelets (L=2 K=16). + """ + files = _lower_odd_even_looped(2, K=16) + odd_interior = _fabin_queues(files['code_1_0.csl']) + even_interior = _fabin_queues(files['code_2_0.csl']) + assert len(odd_interior) == 2, odd_interior + assert len(set(odd_interior.values())) == 2, odd_interior + assert len(even_interior) == 2, even_interior + assert len(set(even_interior.values())) == 2, even_interior + # Endpoints have one inbound colour and do not need a second queue. + assert len(_fabin_queues(files['code_0_0.csl'])) == 1 + assert len(_fabin_queues(files['code_3_0.csl'])) == 1 + + +def test_unrolled_phases_still_share_a_queue_when_spans_are_disjoint(): + """ + Occupancy pooling across successive phases is still required on WSE-2: a Batcher endpoint + receives on two colours that never overlap, and there is only one input queue to spare. + """ + if len(csl.INPUT_QUEUE_IDS) < 2: + pytest.skip(f'{csl.ARCH} has no pool of input queues to share') + path = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', + 'batcher_oddeven_1D.sptl') + kernel = parser.parse_file(path) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, L=2, K=2, R=1)) + files = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_benchmarking=True)} + queues = _fabin_queues(files['code_0_0.csl']) + assert len(queues) == 2, queues + if csl.ARCH == 'wse3': + assert len(set(queues.values())) == 2, queues + else: + assert len(set(queues.values())) == 1, queues + + + if __name__ == '__main__': pytest.main([__file__]) From 8d86cec56c82ea2ae25644622448a26785130ac4 Mon Sep 17 00:00:00 2001 From: glukas Date: Thu, 20 Aug 2026 08:55:45 +0200 Subject: [PATCH 56/68] larger test for wse3 batcher variant --- tests/csl_runtime/test_batcher_oddeven_wse3_1d.sh | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/csl_runtime/test_batcher_oddeven_wse3_1d.sh b/tests/csl_runtime/test_batcher_oddeven_wse3_1d.sh index 2506124a..1ccd65a7 100755 --- a/tests/csl_runtime/test_batcher_oddeven_wse3_1d.sh +++ b/tests/csl_runtime/test_batcher_oddeven_wse3_1d.sh @@ -68,4 +68,5 @@ run_batcher 2 2 3 run_batcher 4 1 run_batcher 4 2 run_batcher 4 16 +run_batcher 4 32 echo "Skipping L=5: sixteen colors would have to switch, and wse3 switches on fifteen." From 252a37bfb039b00dc5c28099f8057498e9ed7d0d Mon Sep 17 00:00:00 2001 From: glukas Date: Thu, 20 Aug 2026 09:06:51 +0200 Subject: [PATCH 57/68] Add 2D Shearsort sample --- README.md | 2 +- samples/spatial/sort/shearsort_2D_looped.sptl | 1698 +++++++++++++++++ tests/csl_runtime/test_shearsort_2d_looped.sh | 69 + tests/spatial_ir/test_routing.py | 72 + 4 files changed, 1840 insertions(+), 1 deletion(-) create mode 100644 samples/spatial/sort/shearsort_2D_looped.sptl create mode 100755 tests/csl_runtime/test_shearsort_2d_looped.sh diff --git a/README.md b/README.md index 6685aeae..5a649ea1 100644 --- a/README.md +++ b/README.md @@ -104,7 +104,7 @@ Sample SPADA programs are in `samples/`: | `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy` | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels) | +| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels). `shearsort_2D_looped` is the 2D snake-order mesh sort on `N x N` PEs, the same neighbour rounds as a runtime loop on eight static channels (WSE-3) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | diff --git a/samples/spatial/sort/shearsort_2D_looped.sptl b/samples/spatial/sort/shearsort_2D_looped.sptl new file mode 100644 index 00000000..11d89c99 --- /dev/null +++ b/samples/spatial/sort/shearsort_2D_looped.sptl @@ -0,0 +1,1698 @@ +/** + * Shearsort on an N x N mesh, N = 2^L, with the odd-even rounds as runtime loops. + * + * Each PE holds a block of K f32 keys: N*N*K keys in all. The network is snake-order: after + * it, concatenating even rows left to right and odd rows right to left, PE blocks still + * ascending, yields the sorted sequence. K = 1 is the one-key-per-PE network again. + * + * Algorithm + * --------- + * Schnorr-Shamir shearsort: (row-sort, column-sort)^L, then one final row-sort. Ending on a + * column-sort would leave the mesh column-ordered and break the snake. Each 1D sort is the + * neighbour odd-even network of odd_even_sort_1D_looped.sptl, N rounds, as a compare-split of + * two sorted K-blocks. A network that sorts N keys sorts N sorted blocks this way. + * + * Even rows (j even) sort ascending, small keys west. Odd rows sort descending, small keys + * east. Every column sorts with small keys north (j toward 0). Origin is north-west; x grows + * east and y grows south. + * + * Round t of a 1D sort compares neighbours (2i, 2i+1) when t is even and (2i+1, 2i+2) when t + * is odd. N rounds suffice. Every comparator is one hop. + * + * Eight channels + * -------------- + * One channel per (axis, round parity, direction), eight in all, so each PE's role on each + * channel is fixed for the whole run. Four east-west colours would suffice for the rows, but + * the columns need a second quartet: mapping both axes onto the same four would make a PE + * whose row-role and column-role disagree send and then receive on one colour. On WSE-2 that + * turnaround overshoots the receiver (see irspec/docs/spatial/routing_wse.md). + * + * Why the rounds are a loop + * ------------------------- + * The N odd-even rounds of a 1D sort, and the L shearsort iterations, can be written as a + * compile-time `for` around phase blocks. The compiler unrolls that into Theta(L*N) phases, + * so code and compile time grow with N, and each phase boundary is a barrier. Here both are + * sequential `for`s inside each compute block, which lower to CSL loops: each body is emitted + * once, and there is no per-round barrier. + * + * A phase boundary is an epoch boundary -- streams close, a channel may change hands, + * routers may advance. This network needs none of that. With roles fixed, no router + * switches and no channel is reassigned. Round t's values arrive before round t+1's + * because a channel is a FIFO. + * + * Role split + * ---------- + * Communication shape is the product of the 1D odd-even roles on i and on j, which a + * `compute` subgrid already expresses. The four ends of each axis sit out every odd round + * along that axis. Reverse rows only flip keep-smallest against keep-largest; they do not + * change who sends on which channel. + * + * i = 0 low in even row-rounds, idle in odd row-rounds + * i even interior low in even row-rounds, high in odd row-rounds + * i odd interior high in even row-rounds, low in odd row-rounds + * i = N-1 high in even row-rounds, idle in odd row-rounds + * + * and the same four classes on j for the columns, with even j also the ascending rows. + * + * A high PE must send its own block before it merges, which is why it sends after the + * receive but before the walk: the merge overwrites val, and the partner needs the + * pre-merge value. + * + * Queues + * ------ + * A fully interior PE receives on four colours (two row, two column) and sends on the other + * four. WSE-3 binds a queue to its colour for the whole kernel, and has six of each, so + * four inbound and four outbound fit. WSE-2 has two of each and those four inbound colours + * are live in one epoch, so this kernel does not lower there. + * + * Constraints: L >= 1, K >= 1. WSE-3. + * + **/ +kernel @shearsort_2d_looped(stream[1<[1< east_even = relative_stream(1, 0) { + hops = auto, + channel = 0 + } + stream west_even = relative_stream(-1, 0) { + hops = auto, + channel = 1 + } + stream east_odd = relative_stream(1, 0) { + hops = auto, + channel = 2 + } + stream west_odd = relative_stream(-1, 0) { + hops = auto, + channel = 3 + } + stream south_even = relative_stream(0, 1) { + hops = auto, + channel = 4 + } + stream north_even = relative_stream(0, -1) { + hops = auto, + channel = 5 + } + stream south_odd = relative_stream(0, 1) { + hops = auto, + channel = 6 + } + stream north_odd = relative_stream(0, -1) { + hops = auto, + channel = 7 + } + } + + // (i=0, j=0): row low-even ascending, column low-even; idle on both odd rounds. + compute i16 i, i16 j in [0:1, 0:1] { + await receive(val, a_in[i, j]) + for i16 m in [1:K] { + for i16 u in [0:K-1] { + pv = ((m - u) - 1) if ((m - u) - 1) > 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, south_even) + await receive(other, north_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (even-i interior, j=0): row low-even / high-odd ascending; column low-even. + compute i16 i, i16 j in [2 : (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, east_odd) + await send(val, west_odd) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, south_even) + await receive(other, north_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, east_odd) + await send(val, west_odd) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (odd-i interior, j=0): row high-even / low-odd ascending; column low-even. + compute i16 i, i16 j in [1 : (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, east_odd) + await receive(other, west_odd) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, south_even) + await receive(other, north_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, east_odd) + await receive(other, west_odd) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (i=N-1, j=0): row high-even ascending; column low-even. + compute i16 i, i16 j in [(1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, south_even) + await receive(other, north_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (i=0, even-j interior): row low-even ascending; column low-even / high-odd. + compute i16 i, i16 j in [0:1, 2 : (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, south_even) + await receive(other, north_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, south_odd) + await send(val, north_odd) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (even-i interior, even-j interior): row low-even / high-odd ascending; + // column low-even / high-odd. Fully interior, four inbound colours. + compute i16 i, i16 j in [2 : (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, east_odd) + await send(val, west_odd) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, south_even) + await receive(other, north_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, south_odd) + await send(val, north_odd) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, east_odd) + await send(val, west_odd) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (odd-i interior, even-j interior): row high-even / low-odd ascending; + // column low-even / high-odd. + compute i16 i, i16 j in [1 : (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, east_odd) + await receive(other, west_odd) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, south_even) + await receive(other, north_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, south_odd) + await send(val, north_odd) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, east_odd) + await receive(other, west_odd) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (i=N-1, even-j interior): row high-even ascending; column low-even / high-odd. + compute i16 i, i16 j in [(1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, south_even) + await receive(other, north_even) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, south_odd) + await send(val, north_odd) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (i=0, odd-j interior): row low-even descending; column high-even / low-odd. + compute i16 i, i16 j in [0:1, 1 : (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, south_even) + await send(val, north_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, south_odd) + await receive(other, north_odd) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (even-i interior, odd-j interior): row low-even / high-odd descending; + // column high-even / low-odd. + compute i16 i, i16 j in [2 : (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, east_odd) + await send(val, west_odd) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, south_even) + await send(val, north_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, south_odd) + await receive(other, north_odd) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, east_odd) + await send(val, west_odd) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (odd-i interior, odd-j interior): row high-even / low-odd descending; + // column high-even / low-odd. + compute i16 i, i16 j in [1 : (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, east_odd) + await receive(other, west_odd) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, south_even) + await send(val, north_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, south_odd) + await receive(other, north_odd) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, east_odd) + await receive(other, west_odd) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (i=N-1, odd-j interior): row high-even descending; column high-even / low-odd. + compute i16 i, i16 j in [(1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, south_even) + await send(val, north_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, south_odd) + await receive(other, north_odd) + // Keep the K smallest: merge upwards from the two block starts. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (i=0, j=N-1): row low-even descending; column high-even. + compute i16 i, i16 j in [0:1, (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, south_even) + await send(val, north_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (even-i interior, j=N-1): row low-even / high-odd descending; column high-even. + compute i16 i, i16 j in [2 : (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, east_odd) + await send(val, west_odd) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, south_even) + await send(val, north_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await send(val, east_even) + await receive(other, west_even) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await receive(other, east_odd) + await send(val, west_odd) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (odd-i interior, j=N-1): row high-even / low-odd descending; column high-even. + compute i16 i, i16 j in [1 : (1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, east_odd) + await receive(other, west_odd) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, south_even) + await send(val, north_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + await send(val, east_odd) + await receive(other, west_odd) + // Keep the K largest: reverse row, west partner holds the large keys. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } + + // (i=N-1, j=N-1): row high-even descending; column high-even. + compute i16 i, i16 j in [(1< 0 else 0 + pt = pv + 1 + x = val[pv] + y = val[pt] + val[pv] = x if x < y else y + val[pt] = y if x < y else x + } + } + for i32 s in [0:L] { + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, south_even) + await send(val, north_even) + // Keep the K largest: merge downwards from the two block ends. + pv = K - 1 + pt = K - 1 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x >= y else 0 + pr = (K - 1) - m + res[pr] = x if take == 1 else y + pv = pv - take + pt = pt - (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + } + for i32 t in [0 : 1<<(L-1)] { + await receive(other, east_even) + await send(val, west_even) + // Keep the K smallest: reverse row, east partner holds the small keys. + pv = 0 + pt = 0 + for i16 m in [0:K] { + x = val[pv] + y = other[pt] + take = 1 if x <= y else 0 + res[m] = x if take == 1 else y + pv = pv + take + pt = pt + (1 - take) + } + await map i16 m in [0:K] { + val[m] = res[m] + } + } + await send(val, a_out[i, j]) + } +} diff --git a/tests/csl_runtime/test_shearsort_2d_looped.sh b/tests/csl_runtime/test_shearsort_2d_looped.sh new file mode 100755 index 00000000..6f6a90bc --- /dev/null +++ b/tests/csl_runtime/test_shearsort_2d_looped.sh @@ -0,0 +1,69 @@ +#!/bin/sh +# E2E: shearsort on an N x N mesh, N = 2^L, N neighbour odd-even rounds as a runtime loop +# (shearsort_2D_looped.sptl). Each PE holds a block of K f32 keys; every comparator is a +# compare-split, so the network sorts all N*N*K keys into snake order: even rows left to +# right, odd rows right to left, each block still ascending. +# Reference: flatten(OUT_a_out in snake order) == sort(a_in.reshape(n*n*k)). +# WSE-3 only: a fully interior PE receives on four colours in one epoch, and WSE-3 binds a +# queue to its colour for the whole kernel (six of each). WSE-2 has two input queues. + +set -e +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +. "$SCRIPT_DIR/_lib.sh" + +if [ "${WSE_ARCH:-wse2}" != "wse3" ]; then + echo "Skipping shearsort_2d_looped: four inbound colours live in one epoch, and wse2 has two input queues." + exit 0 +fi + +SAMPLES_DIR="$(cd "$SCRIPT_DIR/../../samples/spatial/sort" && pwd)" +FOLDER="shearsort_2d_looped_sptl" + +run_sort() { + l=$1 + k=$2 + echo "--- shearsort_2d_looped L=$l K=$k ---" + + sptlc "$SAMPLES_DIR/shearsort_2D_looped.sptl" "$FOLDER" -p L=$l -p K=$k + + python3 - < dict[str, str]: + kernel = parser.parse_file(_SHEARSORT_LOOPED) + kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, L=L, K=K)) + return {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_benchmarking=True)} + + +def _require_wse3_shearsort(): + if csl.ARCH != 'wse3': + pytest.skip('shearsort_2D_looped needs four inbound queues; WSE-2 has two') + + +def test_shearsort_looped_uses_eight_static_channels(): + """ + One channel per (axis, round parity, direction). Roles never change, so no router switches + and both the L shearsort iterations and the N odd-even rounds stay CSL loops. + """ + _require_wse3_shearsort() + files = _lower_shearsort_looped(2) + colors = set(int(c) for c in re.findall(r'@get_color\((\d+)\)', files['layout.csl'])) + assert colors == {0, 1, 2, 3, 4, 5, 6, 7}, sorted(colors) + assert '.switches' not in files['layout.csl'] + + # L = 1 drops the interior rectangles (N = 2 has only the four corners). + ends = _lower_shearsort_looped(1, K=1) + assert 'code_0_0.csl' in ends and 'code_1_1.csl' in ends + assert 'code_2_2.csl' not in ends + + interior = files['code_2_2.csl'] + assert 'for (@range(i32, 0, 2, 1))' in interior, interior + # Four outbound colours (even-row east, odd-row west, even-column south, odd-column north) + # and four inbound, each emitted once rather than unrolled over L or N. + assert interior.count('fabout_dsd') == 4, interior + assert interior.count('fabin_dsd') == 4, interior + + +def test_shearsort_looped_code_is_independent_of_n(): + """Lowering cost and the interior PE program stay flat as N grows.""" + _require_wse3_shearsort() + small = _lower_shearsort_looped(2) + large = _lower_shearsort_looped(3) + assert 'for (@range(i32, 0, 3, 1))' in large['code_2_2.csl'] + assert 'for (@range(i32, 0, 4, 1))' in large['code_2_2.csl'] + # Same sixteen PE roles, so the same number of code files; the loop trip counts are the + # only difference that scales with L. + assert len(small) == len(large) + assert abs(len(small['code_2_2.csl']) - len(large['code_2_2.csl'])) < 64 + + +def test_shearsort_looped_interior_keeps_distinct_input_queues(): + """ + A fully interior PE receives on two row colours and two column colours. WSE-3 binds each + inbound colour to its own queue for the whole kernel, so those four must be distinct. + """ + _require_wse3_shearsort() + files = _lower_shearsort_looped(2, K=1) + even_even = _fabin_queues(files['code_2_2.csl']) + odd_odd = _fabin_queues(files['code_1_1.csl']) + assert len(even_even) == 4, even_even + assert len(set(even_even.values())) == 4, even_even + assert len(odd_odd) == 4, odd_odd + assert len(set(odd_odd.values())) == 4, odd_odd + # The north-west corner only receives even-round west and even-round north. + assert len(_fabin_queues(files['code_0_0.csl'])) == 2 + if __name__ == '__main__': pytest.main([__file__]) From d03a906992dfa3b76f8c527175f359b8aefcea84 Mon Sep 17 00:00:00 2001 From: glukas Date: Thu, 20 Aug 2026 11:01:45 +0200 Subject: [PATCH 58/68] Tune shearsort test size --- tests/csl_runtime/test_shearsort_2d_looped.sh | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/tests/csl_runtime/test_shearsort_2d_looped.sh b/tests/csl_runtime/test_shearsort_2d_looped.sh index 6f6a90bc..a9069b6a 100755 --- a/tests/csl_runtime/test_shearsort_2d_looped.sh +++ b/tests/csl_runtime/test_shearsort_2d_looped.sh @@ -59,11 +59,9 @@ PYEOF } run_sort 1 1 -run_sort 1 4 +run_sort 1 8 run_sort 2 1 -run_sort 2 8 -run_sort 2 16 +run_sort 2 4 run_sort 3 1 run_sort 3 8 -run_sort 3 16 run_sort 4 2 \ No newline at end of file From 2449cfc78733b13fde60d080a4cf0f4a05899a9d Mon Sep 17 00:00:00 2001 From: glukas Date: Thu, 20 Aug 2026 11:13:08 +0200 Subject: [PATCH 59/68] Document WSE-3, shearsort, and looped sort channel splits. --- README.md | 8 +++++--- irspec/docs/spatial/routing_wse.md | 13 +++++++------ 2 files changed, 12 insertions(+), 9 deletions(-) diff --git a/README.md b/README.md index 5a649ea1..76f44814 100644 --- a/README.md +++ b/README.md @@ -101,7 +101,7 @@ Sample SPADA programs are in `samples/`: | `samples/stencils.py` | GT4Py stencil definitions (Laplacian, vertical advection, UVBKE, …) | | `samples/advanced_stencils.py` | GT4Py definitions for horizontal diffusion kernels | | `samples/benchmarks/` | Pre-compiled `.spst`/`.sptl` pairs for five kernels at five domain sizes | -| `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy` | +| `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy`, plus the one-color `shift_bundle_1D` and `exchange_bundle_1D` microkernels | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | | `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels). `shearsort_2D_looped` is the 2D snake-order mesh sort on `N x N` PEs, the same neighbour rounds as a runtime loop on eight static channels (WSE-3) | @@ -112,7 +112,7 @@ Sample SPADA programs are in `samples/`: ## SDK Version and WSE compatibility -The code has been tested for CSL SDK 1.4 and WSE-2. +The code has been tested for CSL SDK 1.4 on WSE-2 and WSE-3. The compiler and the CSL runtime tests read `WSE_ARCH` (`wse2`, the default, or `wse3`). CI runs the simulator suite for both. A kernel that a generation cannot express skips in that architecture's run: `shearsort_2D_looped` needs four inbound queues in one epoch, which WSE-2 does not have, and `batcher_oddeven_wse3_1D` is the origin-pooled Batcher that reaches `L = 4` only on WSE-3. ## Testing @@ -155,6 +155,7 @@ make -C tests/csl_runtime check-sdk ```bash make -C tests/csl_runtime test +make -C tests/csl_runtime test WSE_ARCH=wse3 ``` **Run a single test:** @@ -185,13 +186,14 @@ tests/csl_runtime/run-in-lima.sh --sdk-url ``` This creates the Lima VM on first use (~5–10 min), downloads and extracts the SDK to `tests/csl_runtime/cerebras-sdk/`, installs Python dependencies inside the VM, and runs the full test suite. -If the SDK tarball is already downloaded or extracted, use `--sdk /path/to/cs_sdk` instead of `--sdk-url`. +If the SDK tarball is already downloaded or extracted, use `--sdk /path/to/cs_sdk` instead of `--sdk-url`. Pass `--arch wse3` to compile and simulate for WSE-3. Other modes: ```bash # Run a single test tests/csl_runtime/run-in-lima.sh --sdk --test test_add.sh +tests/csl_runtime/run-in-lima.sh --sdk --arch wse3 --test test_shearsort_2d_looped.sh # Verify the SDK toolchain only tests/csl_runtime/run-in-lima.sh --sdk --check diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index 1829f4b4..dac7a136 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -60,7 +60,10 @@ all. collapsed into a single epoch with a sequential `for` in the compute blocks, which lowers to a real loop and so costs code and compile time independent of the number of rounds. `samples/spatial/sort/odd_even_sort_1D_looped.sptl` is the example: N odd-even rounds on four - static channels, one CSL loop, no per-round barrier. + static channels, one CSL loop, no per-round barrier. The 2D analogue is + `samples/spatial/sort/shearsort_2D_looped.sptl`: eight static neighbour channels, nested + loops, no switches. A fully interior PE there receives on four colours in one epoch, which + fits WSE-3's six exclusive queues and not WSE-2's two. This does *not* generalize to channels that switch: a router's positions are a static sequence, so the epoch a configuration belongs to has to be visible to the compiler. @@ -216,10 +219,7 @@ are then the PEs congruent to the residue, the destinations those congruent to r and the relays the classes strictly between on the one side or the other — three disjoint classes, so a PE still holds one role on the color for the whole kernel and still never both sends and receives on it. What varies is the side it faces, which is one switch position either way, since a source only -ever changes where it transmits and a destination only where it receives. This halves the colors a -pooled distance needs, and with them the queues, which is what -`batcher_oddeven_wse3_1D.sptl` is for: on WSE-3 a queue stays bound to its color for the whole -kernel, so what a PE can afford is not how many colors are live at once but how many it ever touches. +ever changes where it transmits and a destination only where it receives. Sharing on a basis looser than either does risk a PE that sends on the color in one phase and receives on it in another, which needs a two-sided switch change that a sender cannot drive on WSE-2 (see @@ -235,4 +235,5 @@ Engine* (fig. 7.6) for the 2D reduce-scatter, and is known to run on WSE-2. reconfiguration reprograms the routers between rounds outright, which is the only known way to put an arbitrary *sequence* of sends and receives on one color: for a long enough sequence there is a PE for which no fixed cycle of switch positions exists. Until then, splitting the rounds - across channels — as `bitonic_sort_1D.sptl` does — remains the per-kernel fallback. + across channels — as `odd_even_sort_1D_looped.sptl` and `shearsort_2D_looped.sptl` do — + remains the per-kernel fallback. From 7949f88069f9aaa979f688094febba9ebb3c3f4d Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Tue, 25 Aug 2026 12:11:38 -0700 Subject: [PATCH 60/68] Update python-app.yml --- .github/workflows/python-app.yml | 16 ++++------------ 1 file changed, 4 insertions(+), 12 deletions(-) diff --git a/.github/workflows/python-app.yml b/.github/workflows/python-app.yml index 8bd16501..acecfceb 100644 --- a/.github/workflows/python-app.yml +++ b/.github/workflows/python-app.yml @@ -40,6 +40,9 @@ jobs: test-csl: runs-on: ubuntu-latest timeout-minutes: 180 + strategy: + matrix: + wse-arch: ["wse2", "wse3"] steps: - uses: actions/checkout@v4 @@ -102,15 +105,4 @@ jobs: run: | pip install --no-deps -e . export PATH=$PATH:`pwd`/cerebras-sdk - status=0 - for arch in wse2 wse3; do - echo "::group::WSE_ARCH=$arch" - if WSE_ARCH=$arch ./tests/csl_runtime/run_tests.sh; then - echo "$arch passed" - else - echo "$arch failed" - status=1 - fi - echo "::endgroup::" - done - exit $status + WSE_ARCH=${{ matrix.wse-arch }} ./tests/csl_runtime/run_tests.sh From 45b32149a44f59912c9324a6a5846b45a0c90dec Mon Sep 17 00:00:00 2001 From: Lux Gianinazzi Date: Tue, 25 Aug 2026 22:16:29 +0200 Subject: [PATCH 61/68] Update routing_wse.md --- irspec/docs/spatial/routing_wse.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index dac7a136..bf3f0421 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -225,8 +225,8 @@ Sharing on a basis looser than either does risk a PE that sends on the color in on it in another, which needs a two-sided switch change that a sender cannot drive on WSE-2 (see [Lowering to Switches](#lowering-to-switches)), and nothing in the compiler currently rejects it. -This arrangement is the one used in Schnyder's *Distributed Sorting on the Cerebras Wafer-Scale -Engine* (fig. 7.6) for the 2D reduce-scatter, and is known to run on WSE-2. +This arrangement is the one used in Luis Schnyder's Bachelor Thesis *Distributed Sorting on the Cerebras Wafer-Scale +Engine* for the 2D reduce-scatter. !!! note "Note: Multiple Rounds on One Color" Two mechanisms are deliberately left unused, and are what to reach for if the four switch From 3ccd974bf2764fc4b546d6262adf4d4daf5b1c0c Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 5 Oct 2026 19:21:34 +0200 Subject: [PATCH 62/68] Remove unecessary tests --- README.md | 4 +- irspec/docs/spatial/routing_wse.md | 9 +- .../spatial/collectives/scalar_reduce_1D.sptl | 2 +- samples/spatial/sort/batcher_oddeven_1D.sptl | 236 ---------------- .../sort/batcher_oddeven_bundled_1D.sptl | 261 ------------------ spada/syntax/csl/routing.py | 4 +- spada/syntax/spatial_ir/shift_bundles.py | 3 +- .../handwritten/shift_bundle/layout.csl | 137 --------- .../handwritten/shift_bundle/receiver.csl | 57 ---- .../handwritten/shift_bundle/run.py | 82 ------ .../handwritten/shift_bundle/sender.csl | 86 ------ tests/csl_runtime/pending_scalar_reduce_1d.sh | 52 ---- .../csl_runtime/samples}/shift_bundle_1D.sptl | 0 tests/csl_runtime/test_batcher_oddeven_1d.sh | 58 ---- .../test_batcher_oddeven_bundled_1d.sh | 70 ----- .../test_batcher_oddeven_wse3_1d.sh | 13 +- tests/csl_runtime/test_shift_bundle_1d.sh | 2 +- .../csl_runtime/test_shift_bundle_filters.sh | 65 ----- tests/spatial_ir/test_dsd_ops.py | 53 ---- tests/spatial_ir/test_routing.py | 20 -- tests/spatial_ir/test_shift_bundles.py | 68 ----- 21 files changed, 12 insertions(+), 1270 deletions(-) delete mode 100644 samples/spatial/sort/batcher_oddeven_1D.sptl delete mode 100644 samples/spatial/sort/batcher_oddeven_bundled_1D.sptl delete mode 100644 tests/csl_runtime/handwritten/shift_bundle/layout.csl delete mode 100644 tests/csl_runtime/handwritten/shift_bundle/receiver.csl delete mode 100644 tests/csl_runtime/handwritten/shift_bundle/run.py delete mode 100644 tests/csl_runtime/handwritten/shift_bundle/sender.csl delete mode 100755 tests/csl_runtime/pending_scalar_reduce_1d.sh rename {samples/spatial/simple => tests/csl_runtime/samples}/shift_bundle_1D.sptl (100%) delete mode 100644 tests/csl_runtime/test_batcher_oddeven_1d.sh delete mode 100755 tests/csl_runtime/test_batcher_oddeven_bundled_1d.sh delete mode 100644 tests/csl_runtime/test_shift_bundle_filters.sh diff --git a/README.md b/README.md index 76f44814..ef29388b 100644 --- a/README.md +++ b/README.md @@ -101,10 +101,10 @@ Sample SPADA programs are in `samples/`: | `samples/stencils.py` | GT4Py stencil definitions (Laplacian, vertical advection, UVBKE, …) | | `samples/advanced_stencils.py` | GT4Py definitions for horizontal diffusion kernels | | `samples/benchmarks/` | Pre-compiled `.spst`/`.sptl` pairs for five kernels at five domain sizes | -| `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy`, plus the one-color `shift_bundle_1D` and `exchange_bundle_1D` microkernels | +| `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy`, plus the one-color `exchange_bundle_1D` microkernel | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_1D` (a color per matching), `batcher_oddeven_bundled_1D` (the widest phases bundled onto one color pair), `batcher_oddeven_wse3_1D` (the same, pooling the rest by origin rather than by direction, which reaches `L = 4` on WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels). `shearsort_2D_looped` is the 2D snake-order mesh sort on `N x N` PEs, the same neighbour rounds as a runtime loop on eight static channels (WSE-3) | +| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_wse3_1D` (the widest phases bundled onto one color, the rest pooled by origin, which reaches `L = 4` on WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels). `shearsort_2D_looped` is the 2D snake-order mesh sort on `N x N` PEs, the same neighbour rounds as a runtime loop on eight static channels (WSE-3) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index bf3f0421..a4a85cee 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -183,16 +183,14 @@ without being filtered. is at most `max_counter`. A window of `words` out of a stream of `length * words` is therefore `limit1 = length * words - 1`, `max_counter = words - 1`, and an `init_counter` chosen so that the counter reads zero as the wanted block arrives. - `tests/csl_runtime/test_shift_bundle_filters.sh` is the hand-written layout this was measured - with. !!! danger "Error: Too Many Wavelet Filters" WSE-2 has four filters per PE, of which the `memcpy` module reserves one, so **three** are usable (`FILTERS_PER_PE`). A PE needs one per color it filters. *If a PE would need more, a compile error is raised*; the fix is to give some of the streams their own channels, which trades filters for colors. Three filters therefore means at most three bundled phases per PE, whatever the kernel: - `batcher_oddeven_bundled_1D.sptl` would want ten at $2^4$ PEs and bundles only its three widest - phases, which is where most of the colors are saved anyway. + a Batcher sort would want ten at $2^4$ PEs, so `batcher_oddeven_wse3_1D.sptl` bundles only its + widest phases, which is where most of the colors are saved anyway. Reconfiguring a filter while wavelets are still in flight on its color is a data race, and the destination that terminates the stream cannot be reconfigured until the stream has drained, @@ -210,8 +208,7 @@ its own is how a kernel declines the trade. What it then costs is colors, and th by reusing a channel across phases. Two unbundled shifts may share one safely when they agree on axis and signed distance and their sources agree modulo twice that distance, because a PE's role — source, relay or destination — is then a function of its position modulo twice the distance alone, -so one static configuration serves every phase in the pool. `batcher_oddeven_bundled_1D.sptl` pools -on exactly this rule, and `batcher_oddeven_1D.sptl` is the same rule written out as arithmetic. +so one static configuration serves every phase in the pool. The agreement on the *sign* of the distance can be dropped without giving that up. Keep the axis, the magnitude and the source residue modulo twice it, and let the direction of travel vary: the sources diff --git a/samples/spatial/collectives/scalar_reduce_1D.sptl b/samples/spatial/collectives/scalar_reduce_1D.sptl index 1f51ab3f..bf592456 100644 --- a/samples/spatial/collectives/scalar_reduce_1D.sptl +++ b/samples/spatial/collectives/scalar_reduce_1D.sptl @@ -3,7 +3,7 @@ * N is the number of PEs in the first row. * Root is 0,0 * Receiving and sending share channel 0, so every middle PE's router switches between the - * two configurations. See tests/csl_runtime/pending_scalar_reduce_1d.sh + * two configurations. * * WARNING: this sample does not compile yet, for a reason unrelated to routing. `await * receive(rcv_val, westwards)` targets a scalar, and scalars in place blocks get no DSD, so diff --git a/samples/spatial/sort/batcher_oddeven_1D.sptl b/samples/spatial/sort/batcher_oddeven_1D.sptl deleted file mode 100644 index 023619e4..00000000 --- a/samples/spatial/sort/batcher_oddeven_1D.sptl +++ /dev/null @@ -1,236 +0,0 @@ -/** - * 1D Batcher odd-even mergesort over R independent rows of N = 2^L PEs, each holding K f32 keys. - * Ascending: after the network, PE (i, j) holds keys i*K .. i*K + (K-1) of row j. Every hop is - * east-west, so the R rows never mix. - * - * Merge stages l = 1 .. L, each of width 2^l. Within stage l: - * p = 1: every PE in each 2^l box compares at dist = 2^{l-1} - * p = 2 .. l: skip box endpoints, compare at dist = 2^{l-p} - * - * Each drawn comparator is two messages (low -> high, then high -> low), a block of K keys each. - * Lower index keeps the small keys; higher index keeps the large ones. - * - * A block per PE - * -------------- - * Every PE keeps its block sorted. The load phase establishes that and every comparator preserves - * it, because a comparator becomes a *compare-split*, which is the standard block form of one: the - * partners trade blocks, the lower index keeps the K smallest of the 2K keys and the higher index - * the K largest, and both halves come out ascending since each is a merge of two ascending runs. - * A network that sorts N keys sorts N sorted blocks this way, so concatenating the blocks left to - * right gives the sorted sequence. K = 1 is the one-key-per-PE network again, and K need not be a - * power of two. - * - * Two ascending runs are merged by a K-step walk with one index into each. No bounds guard is - * needed: the walk takes K steps and each step advances exactly one index, so the low PE reads - * val[pv] and tmp[pt] with pv + pt = m <= K-1, and the high PE, walking down from the two block - * ends, keeps (K-1-pv) + (K-1-pt) = m <= K-1. The walk cannot run in place -- it writes the m-th - * result while still reading val at an index it has already passed -- so it writes res and copies - * back. - * - * The load phase sorts the block with an insertion network. Its inner trip count is K-1 rather - * than the m of a plain insertion sort, because a compute-level loop bound has to be a - * compile-time expression; the surplus trips clamp onto the (0, 1) pair, and a compare-exchange of - * an ordered pair is a no-op. - * - * Static channel assignment (no router reconfiguration): - * At distance d, the d interleaved matchings (offset r = 0 .. d-1) overlap - * on the 1D mesh, so each (d, r) pair gets its own colors. A channel may only - * be reused by a later stage if every PE keeps the same role on it: a PE that - * sends on a channel and later receives on it would have to swap both the - * input and the output of its router, and on WSE-2 that costs two switch - * advances, of which a sender can only ever emit one (the first advance takes - * RAMP off its input, so the second never leaves the PE). - * - * The p = 1 stages each get their own block, since a PE's role there depends - * on the stage. The p >= 2 stages of any distance d agree on roles -- sender - * iff index = d + r (mod 2d) -- and so share one pair of channels: - * p = 1 : fwd 2*((2^(l-1) - 1) + r), bwd fwd + 1 - * p >= 2 : fwd 2*(N - 1) + 2*((d - 1) + r), bwd fwd + 1 - * Total colors = 3*N - 4, whatever K is: a wider block is more wavelets on a - * channel, not more channels. Intended for small N (e.g. L <= 3). - * - * Example L=3 (n=8): - * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) - * (l=2,p=1) dist 2: (0,2)(1,3)(4,6)(5,7) - * (l=2,p=2) dist 1: (1,2)(5,6) - * (l=3,p=1) dist 4: (0,4)(1,5)(2,6)(3,7) - * (l=3,p=2) dist 2: (2,4)(3,5) - * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) - * - * Constraints: L >= 1, K >= 1, R >= 1 - **/ -kernel @batcher_oddeven_1d( - stream[1<[1< 0 else 0 - pt = pv + 1 - x = val[pv] - y = val[pt] - val[pv] = x if x < y else y - val[pt] = y if x < y else x - } - } - } - } - - for i16 l in [1:L+1] { - // p = 1: all PEs participate, dist = 1<<(l-1). - // Offset r is one disjoint matching; each r has its own fwd/bwd colors. - phase { - for i16 r in [0:1<<(l-1)] { - dataflow i16 i, i16 j in [r:1< fwd = relative_stream(1<<(l-1), 0) { - hops = auto, - channel = 2 * (((1<<(l-1)) - 1) + r) - } - stream bwd = relative_stream(-(1<<(l-1)), 0) { - hops = auto, - channel = (2 * (((1<<(l-1)) - 1) + r)) + 1 - } - } - dataflow i16 i, i16 j in [(r + (1<<(l-1))):1< fwd = relative_stream(1<<(l-1), 0) { - hops = auto, - channel = 2 * (((1<<(l-1)) - 1) + r) - } - stream bwd = relative_stream(-(1<<(l-1)), 0) { - hops = auto, - channel = (2 * (((1<<(l-1)) - 1) + r)) + 1 - } - } - - compute i16 i, i16 j in [r:1<= y else 0 - pr = (K - 1) - m - res[pr] = x if take == 1 else y - pv = pv - take - pt = pt - (1 - take) - } - await map i16 m in [0:K] { - val[m] = res[m] - } - } - } - } - - // p = 2 .. l: skip first/last dist PEs of each 2^l box, dist = 1<<(l-p). - for i16 p in [2:l+1] { - phase { - for i16 r in [0:1<<(l-p)] { - for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { - hops = auto, - channel = (2 * ((1< bwd = relative_stream(-(1<<(l-p)), 0) { - hops = auto, - channel = ((2 * ((1< fwd = relative_stream(1<<(l-p), 0) { - hops = auto, - channel = (2 * ((1< bwd = relative_stream(-(1<<(l-p)), 0) { - hops = auto, - channel = ((2 * ((1<= y else 0 - pr = (K - 1) - m - res[pr] = x if take == 1 else y - pv = pv - take - pt = pt - (1 - take) - } - await map i16 m in [0:K] { - val[m] = res[m] - } - } - } - } - } - } - } - - // Write the sorted blocks back to the host. - phase { - compute i16 i, i16 j in [0:1< high, then high -> low), a block of K keys each. - * Lower index keeps the small keys; higher index keeps the large ones. - * - * A block per PE - * -------------- - * Every PE keeps its block sorted. The load phase establishes that and every comparator preserves - * it, because a comparator becomes a *compare-split*, which is the standard block form of one: the - * partners trade blocks, the lower index keeps the K smallest of the 2K keys and the higher index - * the K largest, and both halves come out ascending since each is a merge of two ascending runs. - * A network that sorts N keys sorts N sorted blocks this way, so concatenating the blocks left to - * right gives the sorted sequence. K = 1 is the one-key-per-PE network again, and K need not be a - * power of two. - * - * Two ascending runs are merged by a K-step walk with one index into each. No bounds guard is - * needed: the walk takes K steps and each step advances exactly one index, so the low PE reads - * val[pv] and tmp[pt] with pv + pt = m <= K-1, and the high PE, walking down from the two block - * ends, keeps (K-1-pv) + (K-1-pt) = m <= K-1. The walk cannot run in place -- it writes the m-th - * result while still reading val at an index it has already passed -- so it writes res and copies - * back. - * - * The load phase sorts the block with an insertion network. Its inner trip count is K-1 rather - * than the m of a plain insertion sort, because a compute-level loop bound has to be a - * compile-time expression; the surplus trips clamp onto the (0, 1) pair, and a compare-exchange of - * an ordered pair is a no-op. - * - * Bundling - * -------- - * This is the bundled variant of batcher_oddeven_1D. The comparators of a phase at distance d - * partition the line into alternating blocks of d PEs, and each block ships its keys d steps - * into the next -- an overlapping interval shift, which is what a bundle is. The low PEs of a - * block take turns nearest-the-partner-first, handing their routers over to relay mode as they - * finish; the high PEs are statically routed and pick their keys out of the stream with a - * counter filter. See irspec/docs/spatial/routing_wse.md. - * - * Bundling is chosen by the channel assignment, since a bundle is what several overlapping - * matchings on one channel become. It costs one filter per participating PE and saves the - * colors the matchings would otherwise need, and a PE has only three filters, so only the three - * widest phases are bundled -- the rule 4*d >= N picks exactly those for any L >= 3, and they - * are where the saving is largest. A bundled color cannot be reused by a later phase: its - * sources advance to pos1 and nothing resets them. - * - * K only widens the wavelet windows, never the color or filter count: a bundle of M sources - * carries M*K wavelets per epoch instead of M, and each destination keeps the K of them that its - * counter filter windows out (limit1 = M*K - 1, max_counter = K - 1). Streams are bounded here, - * unlike in batcher_oddeven_1D, because a bundle's source hands its router on when its stream - * closes, and a bound is what closes it as soon as the K-th wavelet is out. - * - * The phases left unbundled are routed per matching, as in batcher_oddeven_1D, but their colors - * are pooled across phases. Two of them may share when they agree on direction, distance, and - * the source residue mod 2d, because a PE's role is then decided by its position mod 2d alone - * and one static configuration serves every phase in the pool -- no router here switches at all, - * which is what makes this the variant to use on wse2. Dropping the agreement on direction pools - * twice as tightly, at the price of switching routers; that is batcher_oddeven_wse3_1D.sptl, which - * is what reaches L = 4 on wse3, where a queue is bound to its color for the whole kernel. - * - * bundled (4*d >= N) : fwd N + 2*(l*(L+1) + p), bwd fwd + 1 - * pooled : fwd 2*((2*d - 2) + c), bwd fwd + 1 - * - * with c = r for p = 1 and c = d + r for p >= 2; the pooled blocks for successive d are disjoint - * because the block for d runs from 2*d-2 to 4*d-3, and unbundled phases have 4*d < N, so every - * pooled channel stays below N. Colors: 10 at L = 3 and 18 at L = 4, of the 21 available. L = 5 - * would need 34, which is what caps this kernel. Queues cap it earlier than that on either target, - * at L = 3: on wse2 an interior PE at L = 4 has three inbound colors live at once and a PE has two - * input queues, and while wse3 has six, it binds each to its color for the whole kernel, and such a - * PE touches seven. batcher_oddeven_wse3_1D.sptl is the variant that gets L = 4 there. - * - * Example L=3 (n=8): - * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) - * (l=2,p=1) dist 2: (0,2)(1,3)(4,6)(5,7) - * (l=2,p=2) dist 1: (1,2)(5,6) - * (l=3,p=1) dist 4: (0,4)(1,5)(2,6)(3,7) - * (l=3,p=2) dist 2: (2,4)(3,5) - * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) - * - * Constraints: 1 <= L <= 4 (L <= 3 on wse2), K >= 1, R >= 1 - **/ -kernel @batcher_oddeven_1d( - stream[1<[1< 0 else 0 - pt = pv + 1 - x = val[pv] - y = val[pt] - val[pv] = x if x < y else y - val[pt] = y if x < y else x - } - } - } - } - - for i16 l in [1:L+1] { - // p = 1: all PEs participate, dist = 1<<(l-1). - phase { - // Offset r is one disjoint matching. The matchings share a channel where the phase is - // bundled, and take one each where it is not, which is what turns bundling off. - for i16 r in [0:1<<(l-1)] { - dataflow i16 i, i16 j in [r:1< fwd = relative_stream(1<<(l-1), 0) { - hops = auto, - channel = ((1<= (1< bwd = relative_stream(-(1<<(l-1)), 0) { - hops = auto, - channel = (((1<= (1< fwd = relative_stream(1<<(l-1), 0) { - hops = auto, - channel = ((1<= (1< bwd = relative_stream(-(1<<(l-1)), 0) { - hops = auto, - channel = (((1<= (1<= y else 0 - pr = (K - 1) - m - res[pr] = x if take == 1 else y - pv = pv - take - pt = pt - (1 - take) - } - await map i16 m in [0:K] { - val[m] = res[m] - } - } - } - } - - // p = 2 .. l: skip first/last dist PEs of each 2^l box, dist = 1<<(l-p). - for i16 p in [2:l+1] { - phase { - for i16 r in [0:1<<(l-p)] { - for i16 b in [0:1< fwd = relative_stream(1<<(l-p), 0) { - hops = auto, - channel = ((1<= (1< bwd = relative_stream(-(1<<(l-p)), 0) { - hops = auto, - channel = (((1<= (1< fwd = relative_stream(1<<(l-p), 0) { - hops = auto, - channel = ((1<= (1< bwd = relative_stream(-(1<<(l-p)), 0) { - hops = auto, - channel = (((1<= (1<= y else 0 - pr = (K - 1) - m - res[pr] = x if take == 1 else y - pv = pv - take - pt = pt - (1 - take) - } - await map i16 m in [0:K] { - val[m] = res[m] - } - } - } - } - } - } - } - - // Write the sorted blocks back to the host. - phase { - compute i16 i, i16 j in [0:1< int: # A source that only flips its own router does so on the last data wavelet, but only # on WSE-2. Posting a SWITCH_ADV into the same output queue afterwards is what drops # a data wavelet there when a back-pressured send of three or more f32 values fills - # that queue (see tests/csl_runtime/test_shift_bundle_filters.sh). WSE-3 queues hold + # that queue. WSE-3 queues hold # eight words, so SWITCH_ADV is safe; origin-pooled Batcher destinations also switch # and only a traveling control wavelet moves them. Remote routers still need that # wavelet, and a two-position turnaround still needs two of them. diff --git a/spada/syntax/spatial_ir/shift_bundles.py b/spada/syntax/spatial_ir/shift_bundles.py index 7a660a57..1cc248ab 100644 --- a/spada/syntax/spatial_ir/shift_bundles.py +++ b/spada/syntax/spatial_ir/shift_bundles.py @@ -18,8 +18,7 @@ its ramp alone and thereby takes the stream out of the network. This is the arrangement Schnyder's 2D reduce-scatter uses ("Distributed Sorting on the Cerebras -Wafer-Scale Engine", fig. 7.6), and ``tests/csl_runtime/test_shift_bundle_filters.sh`` is a -hand-written version of it that pins down the hardware behaviour relied on here. +Wafer-Scale Engine", fig. 7.6). """ from __future__ import annotations diff --git a/tests/csl_runtime/handwritten/shift_bundle/layout.csl b/tests/csl_runtime/handwritten/shift_bundle/layout.csl deleted file mode 100644 index 84149769..00000000 --- a/tests/csl_runtime/handwritten/shift_bundle/layout.csl +++ /dev/null @@ -1,137 +0,0 @@ -// Counted switching by send order, with counter filters at the receivers. -// -// M senders at x in [0, M) each ship K f32 words M steps east, to the receiver at x + M. A -// single switchable color carries all of it, even though the paths overlap pairwise. -// -// The senders go east to west -- the one closest to the receivers first. That order is what -// makes this cheap: a sender injects its own words and only then has to become a relay for -// the senders west of it, so the hand-over is triggered by an event it knows locally and can -// signal itself with a switch advance. No router has to be switched from a distance, which -// is the thing a control wavelet cannot do selectively. -// -// The receivers never switch. Each one transmits to its ramp *and* onward east, so every -// receiver's router sees the whole stream and the outermost one terminates it. Which words a -// receiver keeps is decided by a counter filter, when FILTER != 0. -// -// Counter filter semantics, as measured on the simulator (the manual's "(exclusive)" aside -// notwithstanding, and matching its "reject all wavelets whose active counter is greater -// than max_counter"): -// -// * the counter starts at init_counter and increments on every data wavelet, -// * it wraps to zero after limit1, so it cycles through limit1 + 1 values, -// * a wavelet reaches the compute element iff counter <= max_counter. -// -// So a window of K words out of a stream of M*K is limit1 = M*K - 1, max_counter = K - 1, -// and init_counter placing the wanted word at counter zero. - -param M: i16; -param K: i16; - -// 0: no filter; every receiver takes all M*K words. Establishes the arrival order. -// 1: counter filters; every receiver takes only the K words addressed to it. -param FILTER: i16; - -// Program input queue for the receivers. WSE-2 may use 0; WSE-3 memcpy already owns 0 and 1, -// so the probe is compiled with IN_QUEUE=2 there. -param IN_QUEUE: i16; - -// 1: emit @initialize_queue (required on WSE-3). 0: omit it (WSE-2). -param INIT_QUEUES: i16; - -// 1: hand over with ``.advance_switch`` on the data fabout. 0: a SWITCH_ADV control wavelet -// after the data microthread reports completion (see sender.csl). -param ADVANCE_SWITCH: i16; - -const memcpy = @import_module("", .{ - .width = 2 * M, - .height = 1, -}); - -const CHANNEL: i16 = 1; - -const STREAM: i16 = M * K; - -// Sends run east to west, so the words from sender q arrive as block M - 1 - q of the -// stream. Receiver M + q wants that block, so its counter must read zero when the block -// starts: shift the start of the cycle back by as many words as precede the block. -fn window_start(q: i16) i16 { - return @as(i16, ((q + 1) * K) % STREAM); -} - -layout { - @set_rectangle(2 * M, 1); - - for (@range(i16, 0, M, 1)) |x| { - @set_tile_code(x, 0, "sender.csl", .{ - .memcpy_params = memcpy.get_params(x), - .stream = STREAM, - .words = K, - // The westmost sender is the last to send and has nothing to relay for. - .hands_over = x > 0, - .init_queues = INIT_QUEUES, - .advance_switch = ADVANCE_SWITCH, - }); - } - for (@range(i16, 0, M, 1)) |q| { - @set_tile_code(M + q, 0, "receiver.csl", .{ - .memcpy_params = memcpy.get_params(M + q), - .stream = STREAM, - .words = K, - .takes = if (FILTER == 0) STREAM else K, - .in_queue = IN_QUEUE, - .init_queues = INIT_QUEUES, - }); - } - - // Senders: inject from the ramp, then relay from the west. - @set_color_config(0, 0, @get_color(CHANNEL), .{ .routes = .{ .rx = .{RAMP}, .tx = .{EAST} } }); - for (@range(i16, 1, M, 1)) |x| { - @set_color_config(x, 0, @get_color(CHANNEL), .{ - .routes = .{ .rx = .{RAMP}, .tx = .{EAST} }, - .switches = .{ .pos1 = .{ .rx = WEST } }, - }); - } - - // Receivers: statically routed, duplicating to the ramp and onward east. The outermost - // one terminates the stream instead of running off the east edge, which also makes its - // router the thing that removes the words nobody kept. - for (@range(i16, 0, M, 1)) |q| { - if (FILTER == 0) { - if (q == M - 1) { - @set_color_config(M + q, 0, @get_color(CHANNEL), - .{ .routes = .{ .rx = .{WEST}, .tx = .{RAMP} } }); - } else { - @set_color_config(M + q, 0, @get_color(CHANNEL), - .{ .routes = .{ .rx = .{WEST}, .tx = .{RAMP, EAST} } }); - } - } else { - if (q == M - 1) { - @set_color_config(M + q, 0, @get_color(CHANNEL), .{ - .routes = .{ .rx = .{WEST}, .tx = .{RAMP} }, - .filter = .{ - .kind = .{ .counter = true }, - .count_data = true, - .init_counter = window_start(q), - .limit1 = STREAM - 1, - .max_counter = K - 1, - }, - }); - } else { - @set_color_config(M + q, 0, @get_color(CHANNEL), .{ - .routes = .{ .rx = .{WEST}, .tx = .{RAMP, EAST} }, - .filter = .{ - .kind = .{ .counter = true }, - .count_data = true, - .init_counter = window_start(q), - .limit1 = STREAM - 1, - .max_counter = K - 1, - }, - }); - } - } - } - - @export_name("val", [*]f32, true); - @export_name("got", [*]f32, true); - @export_name("main", fn()void); -} diff --git a/tests/csl_runtime/handwritten/shift_bundle/receiver.csl b/tests/csl_runtime/handwritten/shift_bundle/receiver.csl deleted file mode 100644 index caebd77f..00000000 --- a/tests/csl_runtime/handwritten/shift_bundle/receiver.csl +++ /dev/null @@ -1,57 +0,0 @@ -// One receiver of the bundle. Its router never switches; a counter filter in the layout -// decides which of the words passing through are handed to this compute element. -// -// ``takes`` is how many that is: the whole stream when no filter is configured (used to -// establish the arrival order), otherwise the ``words`` addressed to this receiver. - -param memcpy_params: comptime_struct; -param stream: i16; -param words: i16; -param takes: i16; -param in_queue: i16; -param init_queues: i16; - -const sys_mod = @import_module("", memcpy_params); - -const channel: color = @get_color(1); - -var val: [words]f32; -var got: [stream]f32; -var __val_ptr: [*]f32 = &val; -var __got_ptr: [*]f32 = &got; - -const got_dsd = @get_dsd(mem1d_dsd, .{ .tensor_access = |i|{takes} -> got[i] }); -const in_dsd = @get_dsd(fabin_dsd, .{ - .extent = takes, - .fabric_color = channel, - .input_queue = @get_input_queue(in_queue), -}); - -const recv_id = @get_local_task_id(8); -const done_id = @get_local_task_id(9); - -task recv_task() void { - @fmovs(got_dsd, in_dsd, .{ .async = true, .activate = done_id }); -} - -task done_task() void { - sys_mod.unblock_cmd_stream(); -} - -fn main() void { - for (@range(i16, 0, stream, 1)) |i| { - got[i] = -1.0; - } - @activate(recv_id); -} - -comptime { - @export_symbol(__val_ptr, "val"); - @export_symbol(__got_ptr, "got"); - @bind_local_task(recv_task, recv_id); - @bind_local_task(done_task, done_id); - if (init_queues != 0) { - @initialize_queue(@get_input_queue(in_queue), .{ .color = channel }); - } - @export_symbol(main, "main"); -} diff --git a/tests/csl_runtime/handwritten/shift_bundle/run.py b/tests/csl_runtime/handwritten/shift_bundle/run.py deleted file mode 100644 index 9062ca42..00000000 --- a/tests/csl_runtime/handwritten/shift_bundle/run.py +++ /dev/null @@ -1,82 +0,0 @@ -#!/usr/bin/env cs_python -""" -Host side of the counted-switch prototype. - -Sender x holds K words 10*(x+1) + w and ships them M steps east. With --filter 0 every -receiver takes the whole stream, which reports the order the words arrive in; with ---filter 1 the counter filters are active and every receiver should end up with exactly -the block addressed to it. - -Exits non-zero on a mismatch, after printing the full picture either way. -""" -import argparse -import sys - -import numpy as np -from cerebras.sdk.runtime import sdkruntimepybind as crt - - -def main() -> int: - parser = argparse.ArgumentParser() - parser.add_argument('outdir') - parser.add_argument('--M', type=int, required=True) - parser.add_argument('--K', type=int, default=1) - parser.add_argument('--filter', type=int, default=0) - parser.add_argument('--dump-core', action='store_true', - help='write corefile.cs1 before stopping, including on a stall') - args = parser.parse_args() - - m, k = args.M, args.K - width = 2 * m - stream = m * k - - runner = crt.SdkRuntime(args.outdir, suppress_simfab_trace=True) - val_id = runner.get_id('val') - got_id = runner.get_id('got') - - vals = np.array([[10 * (x + 1) + w for w in range(k)] for x in range(width)], - dtype=np.float32) - - runner.load() - runner.run() - try: - runner.memcpy_h2d(val_id, vals.ravel(), 0, 0, width, 1, k, streaming=False, - data_type=crt.MemcpyDataType.MEMCPY_32BIT, - order=crt.MemcpyOrder.ROW_MAJOR, nonblock=False) - runner.launch('main', nonblock=False) - got = np.zeros(width * stream, dtype=np.float32) - runner.memcpy_d2h(got, got_id, 0, 0, width, 1, stream, streaming=False, - data_type=crt.MemcpyDataType.MEMCPY_32BIT, - order=crt.MemcpyOrder.ROW_MAJOR, nonblock=False) - finally: - if args.dump_core: - runner.dump_core('corefile.cs1') - runner.stop() - - got = got.reshape(width, stream) - for x in range(width): - role = 'sender ' if x < m else 'receiver' - print(f' PE {x} ({role}) val={vals[x]} got={got[x]}') - - if args.filter == 0: - # Sends run east to west, so sender m-1's block arrives first. - expected = np.concatenate([vals[x] for x in range(m - 1, -1, -1)]) - bad = [x for x in range(m, width) if not np.allclose(got[x], expected, atol=1e-6)] - if bad: - print(f'FAILED: receivers {bad} did not see {expected}') - return 1 - print(f'Passed M={m} K={k}: every receiver saw the whole stream as {expected}.') - return 0 - - # Receiver m + q is the destination of sender q. - bad = [q for q in range(m) if not np.allclose(got[m + q][:k], vals[q], atol=1e-6)] - if bad: - for q in bad: - print(f'FAILED: receiver {m + q} wanted {vals[q]}, kept {got[m + q][:k]}') - return 1 - print(f'Passed M={m} K={k}: every receiver kept exactly the block addressed to it.') - return 0 - - -if __name__ == '__main__': - sys.exit(main()) diff --git a/tests/csl_runtime/handwritten/shift_bundle/sender.csl b/tests/csl_runtime/handwritten/shift_bundle/sender.csl deleted file mode 100644 index feded770..00000000 --- a/tests/csl_runtime/handwritten/shift_bundle/sender.csl +++ /dev/null @@ -1,86 +0,0 @@ -// One sender of the bundle: inject this PE's words, then hand the router over to relay mode. -// -// All senders are launched at once and the order sorts itself out in the fabric: a sender -// west of us cannot push its words through our router while we still receive from the ramp, -// so it waits on the link until our switch has moved us to relay mode. That is the whole -// serialisation mechanism -- no barrier, no counting. -// -// Two ways to make that switch, selected by ``advance_switch``: -// -// 0 a SWITCH_ADV control wavelet on the same output queue, after the data microthread -// reports completion. The wavelet has to sit behind the data, which is why it shares -// the queue; WSE-2 output queue 2 holds six 16-bit words, so K f32 values already fill -// it at K = 3, and a sender that cannot drain yet then blocks the compute element on -// the synchronous @mov32. Async completion is not an empty queue (see @queue_flush). -// 1 ``.advance_switch`` on the data fabout itself, which advances this router when the -// last data wavelet is sent. No second operation, no extra wavelet in the stream. - -param memcpy_params: comptime_struct; -param stream: i16; -param words: i16; -param hands_over: bool; -param init_queues: i16; -param advance_switch: i16; - -const sys_mod = @import_module("", memcpy_params); -const ctrl = @import_module(""); - -const channel: color = @get_color(1); - -var val: [words]f32; -var got: [stream]f32; -var __val_ptr: [*]f32 = &val; -var __got_ptr: [*]f32 = &got; - -const val_dsd = @get_dsd(mem1d_dsd, .{ .tensor_access = |i|{words} -> val[i] }); -const out_dsd = @get_dsd(fabout_dsd, .{ - .extent = words, - .fabric_color = channel, - .output_queue = @get_output_queue(2), - .advance_switch = if (advance_switch != 0) hands_over else false, -}); -// The control wavelet has to leave through the same queue as the data, so that it stays -// behind our own words and only advances our router once they are out. -const switch_dsd = @get_dsd(fabout_dsd, .{ - .extent = 1, - .fabric_color = channel, - .control = true, - .output_queue = @get_output_queue(2), -}); - -const send_id = @get_local_task_id(8); -const done_id = @get_local_task_id(9); - -task send_task() void { - @fmovs(out_dsd, val_dsd, .{ .async = true, .activate = done_id }); -} - -task done_task() void { - // A separate SWITCH_ADV is only issued when the data DSD does not advance the switch - // itself. Posting it asynchronously is illegal: the source is a scalar, and even as a - // DSD it would share output queue 2 with the data microthread. - if (hands_over and (advance_switch == 0)) { - @mov32(switch_dsd, ctrl.encode_single_payload(ctrl.opcode.SWITCH_ADV, true, {}, 0)); - } - sys_mod.unblock_cmd_stream(); -} - -fn main() void { - // Senders keep nothing, so the host sees the sentinel and cannot mistake a stray - // delivery here for a real one. - for (@range(i16, 0, stream, 1)) |i| { - got[i] = -1.0; - } - @activate(send_id); -} - -comptime { - @export_symbol(__val_ptr, "val"); - @export_symbol(__got_ptr, "got"); - @bind_local_task(send_task, send_id); - @bind_local_task(done_task, done_id); - if (init_queues != 0) { - @initialize_queue(@get_output_queue(2), .{ .color = channel }); - } - @export_symbol(main, "main"); -} diff --git a/tests/csl_runtime/pending_scalar_reduce_1d.sh b/tests/csl_runtime/pending_scalar_reduce_1d.sh deleted file mode 100755 index 54b8e4f7..00000000 --- a/tests/csl_runtime/pending_scalar_reduce_1d.sh +++ /dev/null @@ -1,52 +0,0 @@ -#!/bin/sh -# NOT RUN YET (named "pending_" so run_tests.sh does not collect it). -# -# The routing this exercises lowers correctly -- every middle PE turns its router around on one -# channel, which now compiles and runs. What blocks the sample is unrelated: `await -# receive(rcv_val, westwards)` targets a scalar, and scalars in place blocks get no DSD, so -# emit_copy falls through to a plain assignment and emits `rcv_val = westwards;`, which cslc -# rejects as an undeclared identifier. Rename back to test_*.sh once scalar receives lower. -# E2E test: 1D scalar chain reduction over a single channel (scalar_reduce_1D.sptl). -# -# Every PE but the last receives a partial sum from the east and sends the accumulated value west, -# both on channel 0. Receiving and sending need incompatible router configurations, so each middle -# PE holds two switch positions and is advanced onto the second by the control wavelet that retires -# its upstream neighbour's stream. This is the sample that could not be lowered before router -# switching existed. -# -# Reference: OUT_out == sum over all input rows. - -set -e -SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" -. "$SCRIPT_DIR/_lib.sh" - -N=4 -FOLDER="scalar_reduce_1d_sptl" - -sptlc "$COLLECTIVES_DIR/scalar_reduce_1D.sptl" "$FOLDER" -p N=$N - -grep -q '\.switches' "$FOLDER/layout.csl" || { - echo "Test failed: no router switch configuration was generated." - exit 1 -} - -python3 - < 1 is what exercises that -# window on hardware. - -set -e -SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" -. "$SCRIPT_DIR/_lib.sh" - -SORT_DIR="$(cd "$(dirname "$0")/../../samples/spatial/sort" && pwd)" -FOLDER="batcher_oddeven_bundled_1d_sptl" - -run_batcher() { - l=$1 - k=$2 - r=${3:-1} - echo "--- batcher_oddeven_bundled_1d L=$l K=$k R=$r ---" - - sptlc "$SORT_DIR/batcher_oddeven_bundled_1D.sptl" "$FOLDER" -p L=$l -p K=$k -p R=$r - - python3 - < 1 is what makes the window narrower than the -# cycle, and M*K is the cycle the counter has to wrap at, so "4 4" is the corner where a window of -# four sits inside a cycle of sixteen and the last receiver's counter starts at twelve. That is the -# shape batcher_oddeven_bundled_1D takes at L = 3, K = 4 for its distance-4 phase. Ordered so the -# first failure marks the boundary. -for case in "3 1" "4 1" "3 2" "4 2" "3 4" "4 4"; do - set -- $case - echo "--- stage 1: arrival order, M=$1 K=$2 (no filters) ---" - compile_and_run "$1" "$2" 0 - echo "--- stage 2: filtered delivery, M=$1 K=$2 ---" - compile_and_run "$1" "$2" 1 -done - -rm -rf "$OUT" -echo "Passed: counted switch by send order with filtered delivery." diff --git a/tests/spatial_ir/test_dsd_ops.py b/tests/spatial_ir/test_dsd_ops.py index 12528799..06226e24 100644 --- a/tests/spatial_ir/test_dsd_ops.py +++ b/tests/spatial_ir/test_dsd_ops.py @@ -240,59 +240,6 @@ def test_an_array_of_one_element_still_gets_a_dsd(k: int): assert destination.endswith('_dsd') and source.endswith('_dsd'), code -def test_a_reused_color_keeps_its_input_queue_across_a_gap(): - """ - On the interior Batcher PE, one inbound color is used on both sides of another. Sharing the - queue across that gap is what the simulator rejects: remapping it onto the middle color while - wavelets of the outer color remain. The outer color must keep the queue for its whole span. - """ - path = os.path.join( - os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_1D.sptl' - ) - kernel = parser.parse_file(path) - kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, L=3, K=1, R=1)) - files = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_benchmarking=True)} - code = files['code_2_0.csl'] - - colors = dict(re.findall(r'const (\w+)_color_in: color = @get_color\((\d+)\);', code)) - queues = dict(re.findall( - r'const (\w+)_in_dsd = @get_dsd\(fabin_dsd, .*?input_queue = @get_input_queue\((\d+)\)', - code)) - # fwd__17 and fwd__37 share a color and straddle bwd__30. - assert colors['fwd__17'] == colors['fwd__37'] - assert colors['fwd__17'] != colors['bwd__30'] - assert queues['fwd__17'] == queues['fwd__37'] - assert queues['fwd__17'] != queues['bwd__30'] - - -def test_wse3_inbound_colors_do_not_share_an_input_queue(): - """WSE-3 remaps a queue onto the next color and faults if it is not empty. - - Batcher L=2 is the case that hit ``Attempt to remap input queue 2 from C1 to C3`` - when occupancy pooling reused the queue across sequential colors. - """ - if constants.ARCH != 'wse3': - pytest.skip('WSE-2 may remap a drained queue onto the next color') - path = os.path.join( - os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_1D.sptl' - ) - kernel = parser.parse_file(path) - kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, L=2, K=2, R=1)) - files = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_benchmarking=True)} - code = files['code_1_0.csl'] - colors = dict(re.findall(r'const (\w+)_color_in: color = @get_color\((\d+)\);', code)) - queues = dict(re.findall( - r'const (\w+)_in_dsd = @get_dsd\(fabin_dsd, .*?input_queue = @get_input_queue\((\d+)\)', - code)) - by_color: dict[str, set[str]] = {} - for name, color in colors.items(): - if name in queues: - by_color.setdefault(color, set()).add(queues[name]) - assert len(by_color) >= 2, code - used_queues = [next(iter(qs)) for qs in by_color.values()] - assert len(used_queues) == len(set(used_queues)), (by_color, code) - - def test_wse3_concurrent_transfers_use_distinct_microthreads(): """Two transfers in flight at once may not share a microthread. diff --git a/tests/spatial_ir/test_routing.py b/tests/spatial_ir/test_routing.py index bdbd288a..d8485389 100644 --- a/tests/spatial_ir/test_routing.py +++ b/tests/spatial_ir/test_routing.py @@ -496,26 +496,6 @@ def test_odd_even_sort_looped_interior_keeps_distinct_input_queues(): assert len(_fabin_queues(files['code_3_0.csl'])) == 1 -def test_unrolled_phases_still_share_a_queue_when_spans_are_disjoint(): - """ - Occupancy pooling across successive phases is still required on WSE-2: a Batcher endpoint - receives on two colours that never overlap, and there is only one input queue to spare. - """ - if len(csl.INPUT_QUEUE_IDS) < 2: - pytest.skip(f'{csl.ARCH} has no pool of input queues to share') - path = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', - 'batcher_oddeven_1D.sptl') - kernel = parser.parse_file(path) - kernel = passes.constexpr_propagation(passes.concretize_parameters(kernel, L=2, K=2, R=1)) - files = {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_benchmarking=True)} - queues = _fabin_queues(files['code_0_0.csl']) - assert len(queues) == 2, queues - if csl.ARCH == 'wse3': - assert len(set(queues.values())) == 2, queues - else: - assert len(set(queues.values())) == 1, queues - - ### # shearsort_2D_looped: (RC)^L R neighbour rounds as runtime loops on eight static channels ### diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py index 40f37eb4..c24e0059 100644 --- a/tests/spatial_ir/test_shift_bundles.py +++ b/tests/spatial_ir/test_shift_bundles.py @@ -1,9 +1,5 @@ """ Tests for bundling overlapping 1D interval shifts onto one color. - -The mechanism these check the lowering against is measured in -``tests/csl_runtime/test_shift_bundle_filters.sh``, which runs a hand-written version of the same -layout on the simulator. """ import os import re @@ -263,47 +259,11 @@ def entry(): cslrouting._check_filter_budget(entries) -def _bundled_batcher(l: int, k: int = 1): - path = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', - 'batcher_oddeven_bundled_1D.sptl') - kernel = parser.parse_file(path) - kernel = passes.concretize_parameters(kernel, L=l, K=k, R=1) - kernel = passes.constexpr_propagation(kernel) - return lower_spatial_ir_to_csl(kernel, disable_benchmarking=True) - - def _colors_of(layout: str, pattern: str = '') -> set[int]: routes = layout[layout.index('// Routes'):] return {int(color) for color in re.findall(r'@get_color\((\d+)\)[^;]*' + pattern, routes)} -@pytest.mark.parametrize('l, colors', [(2, 6), (3, 10)]) -def test_the_batcher_fits_the_colors_it_has(l: int, colors: int): - # A bundled phase puts all of its matchings on one color pair; the phases left unbundled take a - # pair per matching, but share those across phases wherever their sources agree mod 2d. - from spada.syntax.csl import constants - - used = _colors_of(next(f.code for f in _bundled_batcher(l) if 'layout' in f.filename)) - assert len(used) == colors - assert len(used) <= len(constants.COLORS) - - -def test_sixteen_keys_need_three_overlapping_input_queues(): - """ - At L = 4 a reused inbound color stays live across a gap that already holds two other colors. - WSE-2 has two input queues, so occupancy pooling refuses. WSE-3 has six, but remapping a - non-empty queue is illegal, and L = 4 wants seven colors in each direction over the kernel -- - seven outbound once memcpy's copy-back of `out` is counted, which is what it reports first. - batcher_oddeven_wse3_1D is the variant that fits there. - """ - from spada.syntax.csl import constants - - if len(constants.INPUT_QUEUE_IDS) >= 3 and constants.ARCH != 'wse3': - pytest.skip(f'{constants.ARCH} has {len(constants.INPUT_QUEUE_IDS)} input queues, enough for L=4') - with pytest.raises(SyntaxError, match='concurrent (in|out)put queues'): - _bundled_batcher(4) - - def _wse3_batcher(l: int, k: int = 1): path = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sort', 'batcher_oddeven_wse3_1D.sptl') @@ -358,34 +318,6 @@ def test_sixteen_keys_fit_the_queues_of_a_target_that_reuses_none(): assert _queues_per_pe(files, 'output') < len(constants.OUTPUT_QUEUE_IDS) -def test_only_the_widest_batcher_phases_are_bundled(): - # Bundling costs one filter at every participating PE and a PE has three, so the sample bundles - # every phase that satisfies 4d >= N. At L = 3 that is already three phases (d = 4, 2, 2). - from spada.syntax.csl import constants - - layout = next(f.code for f in _bundled_batcher(3) if 'layout' in f.filename) - used, filtered, switched = _colors_of(layout), _colors_of(layout, r'\.filter'), _colors_of(layout, r'\.switches') - - assert len(filtered) == 2 * constants.FILTERS_PER_PE # three phases, two directions each - assert filtered == switched # a bundled color is one whose sources hand over to relay mode - assert not (used - filtered) & switched # the pooled ones hold a single static configuration - - -def test_a_wider_block_costs_wavelets_not_colors(): - # The Batcher trades K keys per comparator instead of one. A bundle of M sources then carries - # M*K wavelets per epoch, of which each destination keeps the K its filter windows out -- the - # colors and the filters stay as they are, only the counters grow. - narrow = next(f.code for f in _bundled_batcher(3, 1) if 'layout' in f.filename) - wide = next(f.code for f in _bundled_batcher(3, 4) if 'layout' in f.filename) - - assert _colors_of(wide) == _colors_of(narrow) - assert _colors_of(wide, r'\.filter') == _colors_of(narrow, r'\.filter') - # At L=3 the bundled phases have two and four sources, so cycles of 8 and 16 words. - assert '.limit1 = 7, .max_counter = 3' in wide - assert '.limit1 = 15, .max_counter = 3' in wide - assert '.init_counter = (pe_x - 1) * 4' in wide - - def test_a_scalar_receive_lowers_to_a_data_task(): # A bundle destination that keeps a single key holds it in a scalar, which arrives as the # argument of a data task rather than as a move out of the fabric. From e1f191caedfb80ed1cf7e3990a4ca5a11d717d80 Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 5 Oct 2026 19:27:26 +0200 Subject: [PATCH 63/68] Additional cleanup --- README.md | 4 +- irspec/docs/spatial/routing_wse.md | 4 +- samples/spatial/sort/plot_batcher_routing.py | 661 ------------------ ...rsort_2D_looped.sptl => shearsort_2D.sptl} | 0 tests/csl_runtime/run_tests.sh | 4 - tests/csl_runtime/test_shearsort_2d_looped.sh | 4 +- tests/spatial_ir/test_routing.py | 2 +- 7 files changed, 7 insertions(+), 672 deletions(-) delete mode 100644 samples/spatial/sort/plot_batcher_routing.py rename samples/spatial/sort/{shearsort_2D_looped.sptl => shearsort_2D.sptl} (100%) diff --git a/README.md b/README.md index ce6b528c..25dcfc37 100644 --- a/README.md +++ b/README.md @@ -109,7 +109,7 @@ Sample SpaDA programs are in `samples/`: | `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy`, plus the one-color `exchange_bundle_1D` microkernel | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_wse3_1D` (the widest phases bundled onto one color, the rest pooled by origin, which reaches `L = 4` on WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels). `shearsort_2D_looped` is the 2D snake-order mesh sort on `N x N` PEs, the same neighbour rounds as a runtime loop on eight static channels (WSE-3) | +| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_wse3_1D` (the widest phases bundled onto one color, the rest pooled by origin, which reaches `L = 4` on WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels). `shearsort_2D` is the 2D snake-order mesh sort on `N x N` PEs, the same neighbour rounds as a runtime loop on eight static channels (WSE-3) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | @@ -117,7 +117,7 @@ Sample SpaDA programs are in `samples/`: ## SDK Version and WSE compatibility -The code has been tested for CSL SDK 1.4 on WSE-2 and WSE-3. The compiler and the CSL runtime tests read `WSE_ARCH` (`wse2`, the default, or `wse3`). CI runs the simulator suite for both. A kernel that a generation cannot express skips in that architecture's run: `shearsort_2D_looped` needs four inbound queues in one epoch, which WSE-2 does not have, and `batcher_oddeven_wse3_1D` is the origin-pooled Batcher that reaches `L = 4` only on WSE-3. +The code has been tested for CSL SDK 1.4 on WSE-2 and WSE-3. The compiler and the CSL runtime tests read `WSE_ARCH` (`wse2`, the default, or `wse3`). CI runs the simulator suite for both. A kernel that a generation cannot express skips in that architecture's run: `shearsort_2D` needs four inbound queues in one epoch, which WSE-2 does not have, and `batcher_oddeven_wse3_1D` is the origin-pooled Batcher that reaches `L = 4` only on WSE-3. ## Testing diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index a4a85cee..e94a4994 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -61,7 +61,7 @@ all. real loop and so costs code and compile time independent of the number of rounds. `samples/spatial/sort/odd_even_sort_1D_looped.sptl` is the example: N odd-even rounds on four static channels, one CSL loop, no per-round barrier. The 2D analogue is - `samples/spatial/sort/shearsort_2D_looped.sptl`: eight static neighbour channels, nested + `samples/spatial/sort/shearsort_2D.sptl`: eight static neighbour channels, nested loops, no switches. A fully interior PE there receives on four colours in one epoch, which fits WSE-3's six exclusive queues and not WSE-2's two. @@ -232,5 +232,5 @@ Engine* for the 2D reduce-scatter. reconfiguration reprograms the routers between rounds outright, which is the only known way to put an arbitrary *sequence* of sends and receives on one color: for a long enough sequence there is a PE for which no fixed cycle of switch positions exists. Until then, splitting the rounds - across channels — as `odd_even_sort_1D_looped.sptl` and `shearsort_2D_looped.sptl` do — + across channels — as `odd_even_sort_1D_looped.sptl` and `shearsort_2D.sptl` do — remains the per-kernel fallback. diff --git a/samples/spatial/sort/plot_batcher_routing.py b/samples/spatial/sort/plot_batcher_routing.py deleted file mode 100644 index 210512fa..00000000 --- a/samples/spatial/sort/plot_batcher_routing.py +++ /dev/null @@ -1,661 +0,0 @@ -#!/usr/bin/env python3 -""" -Visualize the channel assignment of the 1D Batcher samples. - -Versions --------- - static batcher_oddeven_1D.sptl. No router reconfiguration at all: a - channel is only reused where every PE keeps its role on it, so each - cell of the table holds a single switch position. The p = 1 stages - get their own block, since a PE's role there depends on the stage; - the p >= 2 stages of one distance d agree on roles and share: - p = 1 : fwd 2*((d - 1) + r), bwd fwd + 1 - p >= 2 : fwd 2*(n - 1) + 2*((d - 1) + r), bwd fwd + 1 - bundled Each phase uses two colors, one per direction. Phases with d >= 2 - whose comparators form a run of at least two sources are a shift - bundle: sources inject then relay (pos0 / pos1), destinations stay - put and pick their word with a counter filter. Colors are not - reused across phases. This is what the sample did before the filter - budget capped it at L = 3; kept here for comparison. - hybrid batcher_oddeven_bundled_1D.sptl as it stands. Only the three widest - phases bundle, which is exactly the three filters a PE has: the - rule is 4*d >= n. The rest are routed per matching, on colors - pooled across phases by the source residue mod 2d *and* the - direction, so that a PE's role on a pooled color is the same in - every phase that uses it and no router ever switches: - fwd 2*((2*d - 2) + c), bwd fwd + 1, c = r or d + r - wse3 batcher_oddeven_wse3_1D.sptl. Bundles the same three phases, and - pools the rest by the origin's residue alone, dropping the - agreement on direction: - fwd (2*d - 2) + c_lo, bwd (2*d - 2) + c_hi - A comparator's two messages then share a pooled color whenever - two phases at one distance meet on it, which halves how many - colors a PE touches -- what WSE-3 counts, since a queue there - stays bound to its color for the whole kernel. The price is that - pooled routers switch: a source alternates R->E and R->W, a - destination W->R and E->R, both one side at a time. - -Views ------ - network Knuth-style sorting network. Concurrent matchings (same phase) - occupy adjacent sub-columns so overlapping spans stay visible. - Each comparator is two arrows: down = fwd (east, +d), up = bwd - (west, -d), each colored by its channel. - table Per-PE @set_color_config, with switch positions resolved the way - the version's target stores them: on WSE-2 a both-sides change is - split through a relay intermediate, on WSE-3 it is one position. - Stacked bands are pos0, pos1, … in that order. - A destination filter is the small ``fN`` in the cell, N being - the filter's init_counter; ``--k`` widens the windows the way K - keys per PE do, which scales every init_counter by K. - -Examples --------- - python samples/spatial/sort/plot_batcher_routing.py --n 8 - python samples/spatial/sort/plot_batcher_routing.py --n 8 --version bundled --view table - python samples/spatial/sort/plot_batcher_routing.py --n 8 --version static bundled --view table - python samples/spatial/sort/plot_batcher_routing.py --n 8 --version hybrid --view table --k 4 - python samples/spatial/sort/plot_batcher_routing.py --n 16 --version wse3 --view table -""" - -from __future__ import annotations - -import argparse -import math -from dataclasses import dataclass, replace - -import matplotlib - -matplotlib.rcParams["pdf.fonttype"] = 42 -matplotlib.rcParams["ps.fonttype"] = 42 -import matplotlib.pyplot as plt -from matplotlib.lines import Line2D -from matplotlib.patches import Rectangle - - -VERSIONS = ("static", "bundled", "hybrid", "wse3") -MIN_BUNDLE_DISTANCE = 2 -MIN_BUNDLE_LENGTH = 2 - -# The architecture each version is written for, which decides how a both-sides change is resolved. -TARGET_ARCH = {"static": "WSE-2", "bundled": "WSE-2", "hybrid": "WSE-2", "wse3": "WSE-3"} - -TX_ORDER = ("RAMP", "EAST", "WEST") -DIR_LETTER = {"RAMP": "R", "EAST": "E", "WEST": "W"} - - -@dataclass(frozen=True) -class Matching: - l: int - p: int - dist: int - offset: int - pairs: tuple[tuple[int, int], ...] - fwd: int - bwd: int - - -@dataclass(frozen=True) -class Phase: - index: int - l: int - p: int - dist: int - matchings: tuple[Matching, ...] - - -@dataclass(frozen=True) -class Route: - """One router configuration: receive from ``rx``, transmit to ``tx``.""" - - rx: str - tx: tuple[str, ...] - - def label(self) -> str: - tx = "".join(DIR_LETTER[d] for d in self.tx) - return f"{DIR_LETTER[self.rx]}→{tx}" - - -@dataclass(frozen=True) -class Cell: - """Resolved hardware switch positions of one (PE, channel), plus its filter.""" - - positions: tuple[Route, ...] - filter_init: int | None = None - - -def fwd_channel_static(dist: int, offset: int, p: int, n: int) -> int: - """Eastbound channel of matching ``(dist, offset)`` of a stage with the given ``p``. - - The p = 1 stages live in their own block of ``2*(n - 1)`` channels because a - PE's role on such a channel depends on the stage; the p >= 2 stages of one - distance agree on roles and so share a channel above that block. - """ - block = 0 if p == 1 else 2 * (n - 1) - return block + 2 * ((dist - 1) + offset) - - -def bwd_channel_static(dist: int, offset: int, p: int, n: int) -> int: - return fwd_channel_static(dist, offset, p, n) + 1 - - -def batcher_phases(n: int) -> list[Phase]: - if n < 2 or n & (n - 1): - raise ValueError(f"n must be a power of two >= 2, got {n}") - log_n = int(math.log2(n)) - phases: list[Phase] = [] - index = 0 - for l in range(1, log_n + 1): - dist = 1 << (l - 1) - matchings = [] - for r in range(dist): - pairs = tuple((i, i + dist) for i in range(r, n, 1 << l)) - matchings.append( - Matching( - l, 1, dist, r, pairs, - fwd_channel_static(dist, r, 1, n), - bwd_channel_static(dist, r, 1, n), - ) - ) - phases.append(Phase(index, l, 1, dist, tuple(matchings))) - index += 1 - for p in range(2, l + 1): - dist = 1 << (l - p) - matchings = [] - for r in range(dist): - pairs = [] - for b in range(0, n, 1 << l): - start = b + dist + r - stop = b + (1 << l) - dist + r - for lo in range(start, stop, 2 * dist): - pairs.append((lo, lo + dist)) - matchings.append( - Matching( - l, p, dist, r, tuple(pairs), - fwd_channel_static(dist, r, p, n), - bwd_channel_static(dist, r, p, n), - ) - ) - phases.append(Phase(index, l, p, dist, tuple(matchings))) - index += 1 - return phases - - -def is_bundled(phase: Phase, n: int, version: str) -> bool: - """Whether ``version`` serializes a phase's matchings onto one colour pair.""" - if version == "static" or phase.dist < MIN_BUNDLE_DISTANCE: - return False - if version == "bundled": - return True - return 4 * phase.dist >= n # hybrid and wse3: the three widest phases, one per filter - - -def assign_channels(phases: list[Phase], version: str, n: int) -> list[Phase]: - """Rewrite matching channels to match the sample of ``version``.""" - if version == "static": - return phases - log_n = int(math.log2(n)) - assigned = [] - for ph in phases: - matchings = [] - for m in ph.matchings: - low_residue = m.offset if ph.p == 1 else ph.dist + m.offset - high_residue = (low_residue + ph.dist) % (2 * ph.dist) - if version == "bundled": - fwd, bwd = 2 * ph.index, 2 * ph.index + 1 - elif is_bundled(ph, n, version): - fwd = n + 2 * (ph.l * (log_n + 1) + ph.p) - bwd = fwd + 1 - elif version == "wse3": - # Pooled by the origin's residue mod 2d alone: each direction takes the colour of - # the partner that sends it, so the two phases at one distance meet on both. - fwd = (2 * ph.dist - 2) + low_residue - bwd = (2 * ph.dist - 2) + high_residue - else: - # Pooled by direction as well, which keeps every router configuration static. - fwd = 2 * ((2 * ph.dist - 2) + low_residue) - bwd = fwd + 1 - matchings.append(replace(m, fwd=fwd, bwd=bwd)) - assigned.append(replace(ph, matchings=tuple(matchings))) - return _compact_channels(assigned) - - -def _compact_channels(phases: list[Phase]) -> list[Phase]: - """ - Renumbers the channels densely, as the compiler's colour allocation does. - - A hand-written channel formula generally leaves gaps, and only the channels a kernel actually - uses are given a colour, in ascending order. Renumbering here is what makes the drawn channel - axis the colour axis of the emitted layout. - """ - used = sorted({ch for ph in phases for m in ph.matchings for ch in (m.fwd, m.bwd)}) - color_of = {channel: index for index, channel in enumerate(used)} - return [replace(ph, matchings=tuple(replace(m, fwd=color_of[m.fwd], bwd=color_of[m.bwd]) - for m in ph.matchings)) - for ph in phases] - - -def _channel_labels(phases: list[Phase], n_channels: int) -> list[str]: - """ - Legend text for each channel: what it carries and which phases put it there. - - Read off the assignment rather than assumed, since which direction a channel carries is what - the versions disagree about -- pooling by direction gives every channel one of them, pooling by - origin gives the shared ones both. - """ - distances: list[set[int]] = [set() for _ in range(n_channels)] - directions: list[set[str]] = [set() for _ in range(n_channels)] - users: list[list[str]] = [[] for _ in range(n_channels)] - for ph in phases: - for m in ph.matchings: - if not m.pairs: - continue - for channel, direction in ((m.fwd, "fwd"), (m.bwd, "bwd")): - distances[channel].add(ph.dist) - directions[channel].add(direction) - if f"l{ph.l}p{ph.p}" not in users[channel]: - users[channel].append(f"l{ph.l}p{ph.p}") - - labels = [] - for ch in range(n_channels): - if not directions[ch]: - labels.append(f" ch {ch} (unused)") - continue - arrow = {("fwd", ): "↓", ("bwd", ): "↑"}.get(tuple(sorted(directions[ch])), "↕") - kind = "+".join(sorted(directions[ch])) - dist = ",".join(str(d) for d in sorted(distances[ch])) - phase_text = " ".join(users[ch]) if len(users[ch]) <= 3 else f"{len(users[ch])} phases" - labels.append(f"{arrow} ch {ch}: d={dist} {kind} ({phase_text})") - return labels - - -def channel_count(phases: list[Phase]) -> int: - used = [ch for ph in phases for m in ph.matchings for ch in (m.fwd, m.bwd)] - return (max(used) + 1) if used else 0 - - -def channel_color(channel: int, n_channels: int, cmap_name: str = "tab20"): - cmap = plt.get_cmap(cmap_name) - if n_channels <= 20: - return cmap(channel % 20) - return plt.get_cmap("gist_ncar")(channel / max(n_channels - 1, 1)) - - -def _dir_arrow(ax, x: float, y_from: float, y_to: float, color) -> None: - ax.annotate( - "", - xy=(x, y_to), - xytext=(x, y_from), - arrowprops=dict(arrowstyle="-|>", color=color, lw=1.6, mutation_scale=9, shrinkA=1.5, shrinkB=1.5), - zorder=2, - ) - - -def _draw_network(ax, phases: list[Phase], n: int, version: str) -> None: - """Draw concurrent matchings as sub-columns; fwd and bwd as offset arrows.""" - n_channels = channel_count(phases) - slot = 0.38 - gap = 0.55 - dx = 0.06 - origins = [] - x = 0.0 - for ph in phases: - origins.append(x) - x += max(len(ph.matchings), 1) * slot + gap - x_end = x - gap - - ax.set_xlim(-0.7, x_end + 0.4) - ax.set_ylim(n - 0.5, -0.5) - ax.set_yticks(range(n)) - ax.set_yticklabels([str(i) for i in range(n)]) - ax.set_ylabel("PE index") - ax.set_xlabel("Phase (left arrow ↓ fwd, right arrow ↑ bwd)") - ax.set_title(f"Batcher network ({version}), n={n}: downward = fwd (east), upward = bwd (west)") - - tick_pos = [] - tick_lab = [] - for ph, x0 in zip(phases, origins): - n_m = max(len(ph.matchings), 1) - width = n_m * slot - ax.axvspan(x0 - 0.10, x0 + width - slot + 0.10, color="0.93", zorder=0) - tick_pos.append(x0 + (n_m - 1) * slot / 2) - tick_lab.append(f"l={ph.l}\np={ph.p}\nd={ph.dist}") - - ax.set_xticks(tick_pos) - ax.set_xticklabels(tick_lab, fontsize=8) - - for pe in range(n): - ax.plot([-0.5, x_end + 0.2], [pe, pe], color="0.78", lw=0.8, zorder=1) - - cap = 0.05 - for ph, x0 in zip(phases, origins): - for k, m in enumerate(ph.matchings): - x = x0 + k * slot - fwd_c = channel_color(m.fwd, n_channels) - bwd_c = channel_color(m.bwd, n_channels) - for lo, hi in m.pairs: - x_fwd = x - dx - x_bwd = x + dx - _dir_arrow(ax, x_fwd, lo, hi, fwd_c) - _dir_arrow(ax, x_bwd, hi, lo, bwd_c) - ax.plot([x_fwd - cap, x_fwd + cap], [lo, lo], color=fwd_c, lw=1.6, zorder=3) - ax.plot([x_fwd - cap, x_fwd + cap], [hi, hi], color=fwd_c, lw=1.6, zorder=3) - ax.plot([x_bwd - cap, x_bwd + cap], [lo, lo], color=bwd_c, lw=1.6, zorder=3) - ax.plot([x_bwd - cap, x_bwd + cap], [hi, hi], color=bwd_c, lw=1.6, zorder=3) - - handles = [ - Line2D([0], [0], color=channel_color(ch, n_channels), lw=2.0, label=label) - for ch, label in enumerate(_channel_labels(phases, n_channels)) - ] - ax.legend( - handles=handles, - title="channel", - loc="upper left", - bbox_to_anchor=(1.02, 1), - fontsize=7, - ncol=1 if n_channels <= 16 else 2, - ) - - -ROUTE_COLOR = { - "R→E": "#1f77b4", - "W→E": "#ff7f0e", - "W→R": "#2ca02c", - "W→RE": "#9467bd", - "R→W": "#5fa8d3", - "E→W": "#f4a261", - "E→R": "#6dce6d", - "E→RW": "#c77dff", -} -ROUTE_ORDER = ("R→E", "W→E", "W→RE", "W→R", "R→W", "E→W", "E→RW", "E→R") - - -def _tx(*dirs: str) -> tuple[str, ...]: - wanted = set(dirs) - return tuple(d for d in TX_ORDER if d in wanted) - - -def _add(seq: list[Route], config: Route) -> None: - if not seq or seq[-1] != config: - seq.append(config) - - -def _resolve_hardware(configs: list[Route], split_both_sides: bool = True) -> tuple[Route, ...]: - """ - The positions the hardware stores for a sequence of configurations. - - On WSE-2 a position names one side, so a change of both is split through an intermediate that - keeps the old input; WSE-3 takes both in one position and needs no splitting. - """ - if not configs: - return () - positions = [configs[0]] - for config in configs[1:]: - previous = positions[-1] - if split_both_sides and previous.rx != config.rx and previous.tx != config.tx: - positions.append(Route(previous.rx, config.tx)) - positions.append(config) - return tuple(positions) - - -def _shift_runs(pairs: tuple[tuple[int, int], ...], dist: int) -> list[tuple[int, int]]: - """Group ``(lo, lo+dist)`` pairs into maximal consecutive source runs ``(start, length)``.""" - sources = sorted(lo for lo, hi in pairs if hi - lo == dist) - runs: list[tuple[int, int]] = [] - i = 0 - while i < len(sources): - start = sources[i] - length = 1 - while i + length < len(sources) and sources[i + length] == start + length: - length += 1 - runs.append((start, length)) - i += length - return runs - - -def _install_hop(configs: list[list[list[Route]]], pe: int, ch: int, route: Route) -> None: - _add(configs[pe][ch], route) - - -def _install_ordinary_pair(configs: list[list[list[Route]]], lo: int, hi: int, fwd: int, bwd: int) -> None: - _install_hop(configs, lo, fwd, Route("RAMP", _tx("EAST"))) - _install_hop(configs, hi, fwd, Route("WEST", _tx("RAMP"))) - for mid in range(lo + 1, hi): - _install_hop(configs, mid, fwd, Route("WEST", _tx("EAST"))) - _install_hop(configs, hi, bwd, Route("RAMP", _tx("WEST"))) - _install_hop(configs, lo, bwd, Route("EAST", _tx("RAMP"))) - for mid in range(lo + 1, hi): - _install_hop(configs, mid, bwd, Route("EAST", _tx("WEST"))) - - -def _window_init(pe: int, first: int, step: int, words: int) -> int: - """``init_counter`` of a destination filter keeping ``words`` of them, as ``_window_start`` emits it.""" - offset = 1 - first if step > 0 else first + 1 - return (pe + offset if step > 0 else offset - pe) * words - - -def _install_bundle( - configs: list[list[list[Route]]], - filters: list[list[int | None]], - start: int, - length: int, - dist: int, - fwd: int, - bwd: int, - words: int, -) -> None: - """Eastbound inject-then-relay plus westbound mirror, with destination filters.""" - src_lo, src_hi = start, start + length - dst_lo, dst_hi = start + dist, start + dist + length - - for pe in range(src_lo, src_hi): - _install_hop(configs, pe, fwd, Route("RAMP", _tx("EAST"))) - _install_hop(configs, pe, fwd, Route("WEST", _tx("EAST"))) - for pe in range(src_hi, dst_lo): - _install_hop(configs, pe, fwd, Route("WEST", _tx("EAST"))) - # Stream travels east: last destination terminates, the others copy-and-forward. - for pe in range(dst_lo, dst_hi - 1): - _install_hop(configs, pe, fwd, Route("WEST", _tx("RAMP", "EAST"))) - filters[pe][fwd] = _window_init(pe, dst_lo, 1, words) - _install_hop(configs, dst_hi - 1, fwd, Route("WEST", _tx("RAMP"))) - filters[dst_hi - 1][fwd] = 0 - - for pe in range(dst_lo, dst_hi): - _install_hop(configs, pe, bwd, Route("RAMP", _tx("WEST"))) - _install_hop(configs, pe, bwd, Route("EAST", _tx("WEST"))) - for pe in range(src_hi, dst_lo): - _install_hop(configs, pe, bwd, Route("EAST", _tx("WEST"))) - # Stream travels west: lowest destination terminates. - for pe in range(src_lo + 1, src_hi): - _install_hop(configs, pe, bwd, Route("EAST", _tx("RAMP", "WEST"))) - filters[pe][bwd] = _window_init(pe, src_hi - 1, -1, words) - _install_hop(configs, src_lo, bwd, Route("EAST", _tx("RAMP"))) - filters[src_lo][bwd] = 0 - - -def pe_table(phases: list[Phase], n: int, version: str, words: int = 1) -> list[list[Cell]]: - """Build the resolved per-PE switch table of ``version``, for ``words`` keys per PE.""" - n_channels = channel_count(phases) - configs: list[list[list[Route]]] = [[[] for _ in range(n_channels)] for _ in range(n)] - filters: list[list[int | None]] = [[None] * n_channels for _ in range(n)] - - for ph in phases: - pairs = tuple(pair for m in ph.matchings for pair in m.pairs) - if not pairs: - continue - if is_bundled(ph, n, version): - fwd, bwd = ph.matchings[0].fwd, ph.matchings[0].bwd - for start, length in _shift_runs(pairs, ph.dist): - if length >= MIN_BUNDLE_LENGTH: - _install_bundle(configs, filters, start, length, ph.dist, fwd, bwd, words) - else: - _install_ordinary_pair(configs, start, start + ph.dist, fwd, bwd) - else: - for m in ph.matchings: - for lo, hi in m.pairs: - _install_ordinary_pair(configs, lo, hi, m.fwd, m.bwd) - - split = TARGET_ARCH[version] == "WSE-2" - return [ - [Cell(_resolve_hardware(configs[pe][ch], split), filters[pe][ch]) for ch in range(n_channels)] - for pe in range(n) - ] - - -def _draw_table(ax, phases: list[Phase], n: int, version: str, words: int) -> None: - """Resolved per-PE switch positions; destination filters as a small ``fN``.""" - table = pe_table(phases, n, version, words) - n_channels = channel_count(phases) - ax.set_xlim(-0.5, n_channels - 0.5) - ax.set_ylim(n - 0.5, -0.5) - ax.set_xticks(range(n_channels)) - ax.set_yticks(range(n)) - ax.set_xlabel("Channel") - ax.set_ylabel("PE") - ax.set_title( - f"Resolved switch positions ({version}, {TARGET_ARCH[version]}), n={n}, K={words} " - f"({n_channels} colors); stacked bands are pos0, pos1, ...; fN is the filter init_counter" - ) - ax.set_aspect("equal") - - for pe in range(n): - for ch in range(n_channels): - cell = table[pe][ch] - if not cell.positions: - ax.add_patch( - Rectangle( - (ch - 0.45, pe - 0.45), - 0.9, - 0.9, - facecolor="#f4f4f4", - edgecolor="0.85", - lw=0.4, - ) - ) - continue - n_pos = len(cell.positions) - # Leave a thin strip at the bottom when a filter is present. - usable = 0.78 if cell.filter_init is not None else 0.9 - band = usable / n_pos - fontsize = 6 if n_pos == 1 else 5 - for i, route in enumerate(cell.positions): - y0 = pe - 0.45 + i * band - label = route.label() - ax.add_patch( - Rectangle( - (ch - 0.45, y0), - 0.9, - band, - facecolor=ROUTE_COLOR.get(label, "#bbbbbb"), - edgecolor="0.85", - lw=0.4, - ) - ) - text = f"{i}:{label}" if n_pos > 1 else label - ax.text(ch, y0 + band / 2, text, ha="center", va="center", fontsize=fontsize, color="0.1") - if cell.filter_init is not None: - ax.text( - ch, - pe + 0.38, - f"f{cell.filter_init}", - ha="center", - va="center", - fontsize=5, - color="0.15", - fontweight="bold", - ) - - handles = [ - Line2D([0], [0], marker="s", color="w", markerfacecolor=ROUTE_COLOR[p], markersize=10, label=p) - for p in ROUTE_ORDER - if any(route.label() == p for row in table for cell in row for route in cell.positions) - ] - if any(cell.filter_init is not None for row in table for cell in row): - handles.append( - Line2D([0], [0], marker="$f$", color="0.15", markerfacecolor="w", markersize=10, label="fN filter init") - ) - ax.legend(handles=handles, title="rx→tx", loc="upper left", bbox_to_anchor=(1.02, 1), fontsize=8) - - -def plot(n: int, view: str, version: str, words: int, outfile: str | None, show: bool) -> None: - phases = assign_channels(batcher_phases(n), version, n) - if view == "network": - n_slots = sum(max(len(ph.matchings), 1) for ph in phases) - fig, ax = plt.subplots(figsize=(max(8, 0.7 * n_slots + 0.8 * len(phases)), max(4, 0.45 * n))) - _draw_network(ax, phases, n, version) - elif view == "table": - n_channels = channel_count(phases) - fig, ax = plt.subplots(figsize=(max(8, 0.45 * n_channels), max(4, 0.5 * n))) - _draw_table(ax, phases, n, version, words) - else: - raise ValueError(f"unknown view {view}") - - fig.tight_layout() - if outfile: - fig.savefig(outfile, bbox_inches="tight") - if outfile.endswith(".pdf"): - fig.savefig(outfile[:-4] + ".png", dpi=160, bbox_inches="tight") - print(f"wrote {outfile}") - if show: - plt.show() - plt.close(fig) - - -def _expand_versions(requested: list[str]) -> list[str]: - if "all" in requested: - return list(VERSIONS) - # Preserve order, drop duplicates. - seen: set[str] = set() - versions = [] - for version in requested: - if version not in seen: - seen.add(version) - versions.append(version) - return versions - - -def main() -> None: - parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) - parser.add_argument("--n", type=int, default=8, help="number of PEs, power of two (default 8)") - parser.add_argument( - "--k", - type=int, - default=1, - help="keys per PE (default 1). A comparator trades K of them, so a bundle's cycle is " - "K times longer and every destination filter starts K times further back.", - ) - parser.add_argument( - "--version", - nargs="+", - choices=VERSIONS + ("all",), - default=["static"], - help="static (one color per matching), bundled (two per phase), hybrid (the WSE-2 sample: " - "the three widest phases bundled, the rest pooled by direction and origin) and/or wse3 " - "(the WSE-3 sample: the same bundles, the rest pooled by origin alone, so the pooled " - "routers switch). 'all' is every one. Repeatable.", - ) - parser.add_argument( - "--view", - choices=("network", "table"), - default="network", - help="network: sorting-network comparators; table: resolved switch positions and filters", - ) - parser.add_argument("--out", default=None, help="output PDF path (version is inserted before the extension if several)") - parser.add_argument("--show", action="store_true", help="open an interactive window") - args = parser.parse_args() - versions = _expand_versions(args.version) - for version in versions: - outfile = args.out - if outfile is None and not args.show: - suffix = "" if args.k == 1 else f"_k{args.k}" - outfile = f"samples/spatial/sort/batcher_routing_{version}_n{args.n}{suffix}_{args.view}.pdf" - elif outfile is not None and len(versions) > 1: - if outfile.endswith(".pdf"): - outfile = f"{outfile[:-4]}_{version}.pdf" - else: - outfile = f"{outfile}_{version}" - plot(args.n, args.view, version, args.k, outfile, args.show) - - -if __name__ == "__main__": - main() diff --git a/samples/spatial/sort/shearsort_2D_looped.sptl b/samples/spatial/sort/shearsort_2D.sptl similarity index 100% rename from samples/spatial/sort/shearsort_2D_looped.sptl rename to samples/spatial/sort/shearsort_2D.sptl diff --git a/tests/csl_runtime/run_tests.sh b/tests/csl_runtime/run_tests.sh index 92870c6f..2e7df16a 100755 --- a/tests/csl_runtime/run_tests.sh +++ b/tests/csl_runtime/run_tests.sh @@ -42,10 +42,6 @@ NON_TEST_SCRIPTS=("run_tests.sh" "run-in-lima.sh" "sptlc" "_lib.sh") is_non_test() { local name="$1" - # Local debug helpers (zz_*) and not-yet-enabled cases (pending_*) are not part of the suite. - case "$name" in - zz_*|pending_*) return 0 ;; - esac for skip in "${NON_TEST_SCRIPTS[@]}"; do [ "$name" = "$skip" ] && return 0 done diff --git a/tests/csl_runtime/test_shearsort_2d_looped.sh b/tests/csl_runtime/test_shearsort_2d_looped.sh index a9069b6a..4290222d 100755 --- a/tests/csl_runtime/test_shearsort_2d_looped.sh +++ b/tests/csl_runtime/test_shearsort_2d_looped.sh @@ -1,6 +1,6 @@ #!/bin/sh # E2E: shearsort on an N x N mesh, N = 2^L, N neighbour odd-even rounds as a runtime loop -# (shearsort_2D_looped.sptl). Each PE holds a block of K f32 keys; every comparator is a +# (shearsort_2D.sptl). Each PE holds a block of K f32 keys; every comparator is a # compare-split, so the network sorts all N*N*K keys into snake order: even rows left to # right, odd rows right to left, each block still ascending. # Reference: flatten(OUT_a_out in snake order) == sort(a_in.reshape(n*n*k)). @@ -24,7 +24,7 @@ run_sort() { k=$2 echo "--- shearsort_2d_looped L=$l K=$k ---" - sptlc "$SAMPLES_DIR/shearsort_2D_looped.sptl" "$FOLDER" -p L=$l -p K=$k + sptlc "$SAMPLES_DIR/shearsort_2D.sptl" "$FOLDER" -p L=$l -p K=$k python3 - < dict[str, str]: From 156dd5d0035e8e84b2b660a784f11aaffdf7387b Mon Sep 17 00:00:00 2001 From: Lux Gianinazzi Date: Mon, 5 Oct 2026 19:31:49 +0200 Subject: [PATCH 64/68] Update spada/syntax/csl/constants.py Co-authored-by: Tal Ben-Nun --- spada/syntax/csl/constants.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/spada/syntax/csl/constants.py b/spada/syntax/csl/constants.py index ffe37be1..0457c481 100644 --- a/spada/syntax/csl/constants.py +++ b/spada/syntax/csl/constants.py @@ -24,10 +24,10 @@ # is not in this list and takes the next free ID after the assigned slots. _CSL_LOCAL_TASK_IDS = { 'wse2': list(range(8, 21)), - 'wse3': [t for t in range(8, 26) if t not in _RESERVED_LOCAL_TASK_IDS['wse3']], + 'wse3': list(range(8, 26)), } -LOCAL_TASK_IDS = _CSL_LOCAL_TASK_IDS[ARCH] +LOCAL_TASK_IDS = [t for t in _CSL_LOCAL_TASK_IDS[ARCH] if t not in RESERVED_LOCAL_TASK_IDS] _CSL_CONTROL_TASK_IDS = { 'wse2': list(range(0, 64)), From 2662e37ee406a3b17614421b5da0c7e99e2f88f0 Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 5 Oct 2026 19:44:57 +0200 Subject: [PATCH 65/68] Run pytest for both architectures and fix WSE-3 data tasks. A data task reading a host stream now reserves an input queue on WSE-3, and tests that cannot recycle or fit in six queues skip there. --- .github/workflows/python-app.yml | 5 ++++ spada/lowering/spatial_ir_to_csl.py | 21 +++++++++++---- .../test_lowering_spatial_ir_to_csl.py | 22 ++++++++++++++- tests/spatial_ir/test_task_recycling.py | 5 ++-- .../spatial_ir/test_task_recycling_codegen.py | 27 +++++++++++++++++-- 5 files changed, 70 insertions(+), 10 deletions(-) diff --git a/.github/workflows/python-app.yml b/.github/workflows/python-app.yml index 998fb98a..5e07b831 100644 --- a/.github/workflows/python-app.yml +++ b/.github/workflows/python-app.yml @@ -15,6 +15,9 @@ permissions: jobs: test-pytest: runs-on: ubuntu-latest + strategy: + matrix: + wse-arch: ["wse2", "wse3"] steps: - uses: actions/checkout@v4 @@ -32,6 +35,8 @@ jobs: pip install -e ".[dev]" - name: Test with pytest + env: + WSE_ARCH: ${{ matrix.wse-arch }} run: | pytest diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index fb612364..332453cf 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -1077,6 +1077,8 @@ def close(keys: list[tuple[str, int]], end: int) -> None: if isinstance(node, spir.AwaitCompletionStatement): close(pending.pop(node.completion_name.as_ir(), []), index) continue + if isinstance(node, spir.ForeachStatement) and not node.parameter_range: + continue # A data task runs on no microthread. for inbound, names in ((True, input_names), (False, output_names)): stream = _fabric_transfer_stream(node, inbound) if stream is None or stream not in names: @@ -1096,12 +1098,15 @@ def close(keys: list[tuple[str, int]], end: int) -> None: def _fabric_transfer_stream(node: spir.SpatialNode, inbound: bool) -> Optional[spir.Identifier]: """ The stream a node transfers in the requested direction, or ``None``. + + A data-task receive (``foreach`` with no range) counts only on WSE-3, where its task ID is an + input queue and so needs one reserved. """ if inbound: if isinstance(node, spir.ReceiveStatement): return stream_lifetime.underlying_stream(node.stream_name) - if (isinstance(node, spir.ForeachStatement) and node.parameter_range - and node.receive_stream is not None): + if (isinstance(node, spir.ForeachStatement) and node.receive_stream is not None + and (node.parameter_range or csl.ARCH == 'wse3')): return stream_lifetime.underlying_stream(node.receive_stream.stream_name) return None if isinstance(node, spir.SendStatement): @@ -1114,8 +1119,8 @@ def _streams_with_fabric_dsds(compute: spir.ComputeBlock, memcpy_mode: bool, """ Streams that lower to a fabric DSD in one direction, so they need a hardware queue. - A data-task receive (``foreach`` with no range) binds the color itself and does not take a - queue. Memcpy arguments are already in local memory, so they do not either. + A data-task receive (``foreach`` with no range) binds the color itself and takes a queue only on + WSE-3. Memcpy arguments are already in local memory, so they do not take one either. :param compute: The compute block being lowered. :param memcpy_mode: Whether memcpy mode is used. @@ -1278,7 +1283,13 @@ def _visit_foreach(stmt: spir.ForeachStatement) -> None: raise SyntaxError(f'Foreach generator "{stream_name.as_ir()}" without a defined ' f'range must only be used with a kernel argument or extern_stream.' f'\n In line {stmt.lineinfo}') - # A data task will be created instead (handled in _generate_data_task) + # A data task will be created instead (handled in _generate_data_task_slot). On WSE-3 + # its ID is an input queue, which only a fabric descriptor records. + if (csl.ARCH == 'wse3' and not (memcpy_mode and stream_name in stream_args) + and stream_name.as_ir() in stream_candidates): + dsd = cslstruct.FabricDSD(cslstruct.DSDType.fabin, f'{name_to_csl(stream_name)}_color', 1, + allocate_input_queue(stream_name)) + dsds[stream_name.as_ir()].append((f'{name_to_csl(stream_name)}_in_dsd', dsd)) return if stream_name.as_ir() not in stream_candidates: return diff --git a/tests/spatial_ir/test_lowering_spatial_ir_to_csl.py b/tests/spatial_ir/test_lowering_spatial_ir_to_csl.py index 75a769a8..5e23d99a 100644 --- a/tests/spatial_ir/test_lowering_spatial_ir_to_csl.py +++ b/tests/spatial_ir/test_lowering_spatial_ir_to_csl.py @@ -1,9 +1,29 @@ import os from spada.lowering.spatial_ir_to_csl import lower_spatial_ir_to_csl +from spada.syntax.csl import constants from spada.syntax.spatial_ir import parser, passes import pytest +def _lower_or_skip_queue_limit(kernel, **kwargs): + """ + Lowers ``kernel``, skipping the test where a PE needs more input queues than WSE-3 has. + + WSE-3 keeps one input queue per inbound color for the whole kernel, so a PE that receives on + more than six channels cannot be lowered there. + + :param kernel: A concretized kernel. + :param kwargs: Passed on to ``lower_spatial_ir_to_csl``. + :return: The generated CSL files. + """ + try: + return lower_spatial_ir_to_csl(kernel, **kwargs) + except SyntaxError as error: + if constants.ARCH == 'wse3' and 'concurrent input queues' in str(error): + pytest.skip(f'needs more input queues than WSE-3 has: {str(error).splitlines()[0]}') + raise + + def test_non_concrete_program(): file = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'simple', 'add.sptl') kernel = parser.parse_file(file) @@ -67,7 +87,7 @@ def test_tree_reduce_1d_compiles_512_pes(): # K must be >= 2: K=1 breaks foreach/receive lowering (empty DSD slot for __x). kernel = passes.concretize_parameters(kernel, L=9, K=2) kernel = passes.constexpr_propagation(kernel) - csl_files = lower_spatial_ir_to_csl(kernel, copy_elision=True, prune_memory=True) + csl_files = _lower_or_skip_queue_limit(kernel, copy_elision=True, prune_memory=True) assert csl_files, 'expected at least one generated CSL file' assert all(f.code.strip() for f in csl_files), 'expected non-empty CSL bodies' diff --git a/tests/spatial_ir/test_task_recycling.py b/tests/spatial_ir/test_task_recycling.py index 32c4f0a6..a898da58 100644 --- a/tests/spatial_ir/test_task_recycling.py +++ b/tests/spatial_ir/test_task_recycling.py @@ -72,8 +72,9 @@ def test_task_recycling_all_tasks_assigned(): def test_task_recycling_plan_reuses_local_slots(): tasks = _create_unfused_tasks() local_task_count = sum(1 for task in tasks if task.task_type == 'local') - - assert local_task_count > len(constants.LOCAL_TASK_IDS) + if local_task_count <= len(constants.LOCAL_TASK_IDS): + pytest.skip(f'{local_task_count} local tasks fit the {len(constants.LOCAL_TASK_IDS)} IDs of ' + f'{constants.ARCH}, so nothing is recycled') plan = task_recycling.plan_task_bindings(tasks, tdag.TaskCreationBehavior.STATE_MACHINE_ON_OVERRUN) diff --git a/tests/spatial_ir/test_task_recycling_codegen.py b/tests/spatial_ir/test_task_recycling_codegen.py index 78385bba..763a1648 100644 --- a/tests/spatial_ir/test_task_recycling_codegen.py +++ b/tests/spatial_ir/test_task_recycling_codegen.py @@ -66,6 +66,25 @@ _CHAIN_PHASES = 12 +def _lower_or_skip_queue_limit(kernel, **kwargs): + """ + Lowers ``kernel``, skipping the test where a PE needs more input queues than WSE-3 has. + + WSE-3 keeps one input queue per inbound color for the whole kernel, so a PE that receives on + more than six channels cannot be lowered there. + + :param kernel: A concretized kernel. + :param kwargs: Passed on to ``lower_spatial_ir_to_csl``. + :return: The generated CSL files. + """ + try: + return lower_spatial_ir_to_csl(kernel, **kwargs) + except SyntaxError as error: + if constants.ARCH == 'wse3' and 'concurrent input queues' in str(error): + pytest.skip(f'needs more input queues than WSE-3 has: {str(error).splitlines()[0]}') + raise + + def _scalar_exchange_chain(phases: int = _CHAIN_PHASES): kernel = parser.parse_string(_SCALAR_EXCHANGE_CHAIN) kernel = passes.concretize_parameters(kernel, R=phases) @@ -80,7 +99,7 @@ def test_task_recycling_codegen_uses_else_if_dispatch_for_recycled_slots(): kernel = passes.concretize_parameters(kernel, LX=8, LY=8, K=16) kernel = passes.constexpr_propagation(kernel) - csl_files = lower_spatial_ir_to_csl(kernel, task_fusion=False) + csl_files = _lower_or_skip_queue_limit(kernel, task_fusion=False) code = next(file.code for file in csl_files if file.filename == 'code_0_0.csl') task_id_occurrences: dict[str, int] = {} @@ -115,6 +134,10 @@ def test_csl_runtime_task_recycling_sample_lowers(filename: str): assert csl_files, 'expected at least one generated CSL file' combined = '\n'.join(f.code for f in csl_files) assert combined.strip(), 'expected non-empty CSL' + local_tasks = max(len(re.findall(r'const task_\d+_id = ', f.code)) for f in csl_files) + if local_tasks <= len(constants.LOCAL_TASK_IDS): + pytest.skip(f'{local_tasks} local tasks fit the {len(constants.LOCAL_TASK_IDS)} IDs of ' + f'{constants.ARCH}, so nothing is recycled') assert '__task_slot_' in combined, 'expected task-ID recycling in generated CSL' @@ -218,7 +241,7 @@ def test_codegen_avoids_local_task_id_color_overlap(): kernel = parser.parse_file(path) kernel = passes.constexpr_propagation(kernel) - csl_files = lower_spatial_ir_to_csl( + csl_files = _lower_or_skip_queue_limit( kernel, task_fusion=False, copy_elision=True, prune_memory=True) combined = '\n'.join(f.code for f in csl_files) From 0d58bfb7ff12596e171fff7eca683a69f84b1ff4 Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 5 Oct 2026 19:56:17 +0200 Subject: [PATCH 66/68] Refactor spatial_ir_to_csl --- spada/lowering/fabric_occupancy.py | 215 +++++++++++++++ spada/lowering/spatial_ir_to_csl.py | 339 +----------------------- spada/lowering/wse3.py | 192 ++++++++++++++ spada/syntax/csl/routing.py | 2 +- tests/spatial_ir/test_task_recycling.py | 16 +- 5 files changed, 430 insertions(+), 334 deletions(-) create mode 100644 spada/lowering/fabric_occupancy.py create mode 100644 spada/lowering/wse3.py diff --git a/spada/lowering/fabric_occupancy.py b/spada/lowering/fabric_occupancy.py new file mode 100644 index 00000000..90408378 --- /dev/null +++ b/spada/lowering/fabric_occupancy.py @@ -0,0 +1,215 @@ +""" +Occupancy of fabric transfers on one processing element. + +``_collect_unique_dsds`` uses these spans to decide which transfers may share a fabric queue or a +microthread. A sequential ``for`` body is counted twice, so a color that returns on the next +iteration overlaps whatever sits between its two uses and keeps its own queue. +""" + +from typing import Optional + +from spada.syntax.csl import constants as csl +from spada.syntax.spatial_ir import irnodes as spir +from spada.syntax.spatial_ir import stream_lifetime + + +def statement_transfer_points( + statement: spir.Statement, names: set[spir.Identifier], inbound: bool +) -> list[spir.Identifier]: + """ + Fabric transfers in ``statement``, in source order, with sequential ``for`` bodies repeated. + + Walking a loop body once makes its colours look sequential, so occupancy pooling would give + them one queue. The next iteration of an earlier colour can already occupy the router when a + later colour of the same body remaps that queue -- WSE-2 then aborts with "Attempt to remap + input queue N, from C_i to C_j, but the router is holding wavelets". Appending the body a + second time makes a colour used on both sides of another occupy a span that overlaps it, the + same rule that keeps a reused colour's queue across a gap between unrolled phases. + + :param statement: The statement to walk. + :param names: Streams that bind a fabric queue in this direction. + :param inbound: True to collect receives, False to collect sends. + :return: The transferred streams, with each ``for`` body listed twice. + """ + if isinstance(statement, spir.ForStatement): + body = transfer_points(statement.body, names, inbound) + return body + body + + nested_skip: set[int] = set() + points: list[spir.Identifier] = [] + for node in statement.walk(): + if id(node) in nested_skip: + continue + if node is not statement and isinstance(node, spir.ForStatement): + for descendant in node.walk(): + nested_skip.add(id(descendant)) + points.extend(statement_transfer_points(node, names, inbound)) + continue + stream = fabric_transfer_stream(node, inbound) + if stream is None or stream not in names: + continue + points.append(stream) + return points + + +def transfer_points( + statements: list[spir.Statement], names: set[spir.Identifier], inbound: bool +) -> list[spir.Identifier]: + """ + Fabric transfers across ``statements``, in source order. + + :param statements: The statements to walk, in execution order. + :param names: Streams that bind a fabric queue in this direction. + :param inbound: True to collect receives, False to collect sends. + :return: The transferred streams. + """ + points: list[spir.Identifier] = [] + for statement in statements: + points.extend(statement_transfer_points(statement, names, inbound)) + return points + + +def queue_spans( + compute: spir.ComputeBlock, names: set[spir.Identifier], queue_key, inbound: bool +) -> dict[str, tuple[int, int]]: + """ + Occupancy of each queue key along the linearized send/receive order of this PE. + + Sequential ``for`` bodies are counted twice so a colour that comes back on the next iteration + keeps its queue across the loop-carried gap; see ``statement_transfer_points``. Uses of the + same channel still collapse to one span, so a colour that comes back after a gap between + unrolled phases keeps its queue for the whole of that span. + + :param compute: The compute block being lowered. + :param names: Streams that actually bind a fabric queue in this direction. + :param queue_key: Maps a stream identifier to its grouping key (channel, or the name itself). + :param inbound: True to walk receives, False to walk sends. + :return: Mapping of grouping key to ``(first_use, last_use)`` in linearized order. + """ + points = transfer_points(compute.statements, names, inbound) + + spans: dict[str, tuple[int, int]] = {} + for index, stream in enumerate(points): + key = queue_key(stream) + if key in spans: + start, _ = spans[key] + spans[key] = (start, index) + else: + spans[key] = (index, index) + return spans + + +def microthread_intervals( + compute: spir.ComputeBlock, + input_names: set[spir.Identifier], + output_names: set[spir.Identifier], + queue_key, +) -> dict[str, list[tuple[int, int]]]: + """ + The intervals over which each stream group holds a microthread on one PE. + + Both directions are numbered in one space, since a microthread is one resource across them. A + transfer that keeps a completion handle is in flight until that handle is awaited, which is where + real concurrency comes from: a receive started before a send is still running while the send is. + A self-awaited transfer is given its own slot and the next one, because the activation that + awaits it also starts what follows. + + :param compute: The compute block being lowered. + :param input_names: Streams that bind an input queue on this PE. + :param output_names: Streams that bind an output queue on this PE. + :param queue_key: Maps a stream identifier to its grouping key (channel, or the name itself). + :return: Mapping of direction-prefixed grouping key to the intervals it is in flight over. + """ + live: dict[str, list[tuple[int, int]]] = {} + pending: dict[str, list[tuple[str, int]]] = {} + index = 0 + + def close(keys: list[tuple[str, int]], end: int) -> None: + for key, start in keys: + live.setdefault(key, []).append((start, end)) + + for statement in compute.statements: + for node in statement.walk(): + if isinstance(node, spir.AwaitAllStatement): + for keys in pending.values(): + close(keys, index) + pending.clear() + continue + if isinstance(node, spir.AwaitCompletionStatement): + close(pending.pop(node.completion_name.as_ir(), []), index) + continue + if isinstance(node, spir.ForeachStatement) and not node.parameter_range: + continue # A data task runs on no microthread. + for inbound, names in ((True, input_names), (False, output_names)): + stream = fabric_transfer_stream(node, inbound) + if stream is None or stream not in names: + continue + key = f"{'in' if inbound else 'out'} {queue_key(stream)}" + completion = getattr(node, "completion_name", None) + if completion is None: + live.setdefault(key, []).append((index, index + 1)) + else: + pending.setdefault(completion.name.as_ir(), []).append((key, index)) + index += 1 + for keys in pending.values(): + close(keys, index) + return live + + +def fabric_transfer_stream( + node: spir.SpatialNode, inbound: bool +) -> Optional[spir.Identifier]: + """ + The stream a node transfers in the requested direction, or ``None``. + + A data-task receive (``foreach`` with no range) counts only on WSE-3, where its task ID is an + input queue and so needs one reserved. + + :param node: A node of the compute block. + :param inbound: True for a receive, False for a send. + :return: The underlying stream, or ``None`` when the node is not such a transfer. + """ + if inbound: + if isinstance(node, spir.ReceiveStatement): + return stream_lifetime.underlying_stream(node.stream_name) + if ( + isinstance(node, spir.ForeachStatement) + and node.receive_stream is not None + and (node.parameter_range or csl.ARCH == "wse3") + ): + return stream_lifetime.underlying_stream(node.receive_stream.stream_name) + return None + if isinstance(node, spir.SendStatement): + return stream_lifetime.underlying_stream(node.stream_name) + return None + + +def streams_with_fabric_dsds( + compute: spir.ComputeBlock, + memcpy_mode: bool, + stream_args: set[spir.Identifier], + inbound: bool, +) -> set[spir.Identifier]: + """ + Streams that lower to a fabric DSD in one direction, so they need a hardware queue. + + A data-task receive (``foreach`` with no range) binds the color itself and takes a queue only on + WSE-3. Memcpy arguments are already in local memory, so they do not take one either. + + :param compute: The compute block being lowered. + :param memcpy_mode: Whether memcpy mode is used. + :param stream_args: Kernel-argument streams, which memcpy has already copied. + :param inbound: True for receives, False for sends. + :return: The stream identifiers that need a queue in that direction. + """ + result: set[spir.Identifier] = set() + argument_names = {name.as_ir() for name in stream_args} + for statement in compute.statements: + for node in statement.walk(): + name = fabric_transfer_stream(node, inbound) + if name is None: + continue + if memcpy_mode and name.as_ir() in argument_names: + continue + result.add(name) + return result diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 332453cf..17b2f7c7 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -7,7 +7,7 @@ import functools from io import StringIO import textwrap -from typing import Optional +from spada.lowering import fabric_occupancy, wse3 from spada.syntax.common.types import BIT_WIDTH from spada.syntax.spatial_ir import irnodes as spir, canonicalization, analysis, passes from spada.syntax.spatial_ir import copy_elimination @@ -394,7 +394,7 @@ def generate_rectangle(kernel: spir.Kernel, raise cslrouting.declare_switch_advances(rect, header, color_map, dsds) - _declare_queue_initialization(dsds, rect, footer, color_map) + wse3.declare_queue_initialization(dsds, rect, footer, color_map) # Fuse tasks as much as possible to reduce number of resources if task_fusion: @@ -408,7 +408,7 @@ def generate_rectangle(kernel: spir.Kernel, print(f'P{rect.x_range[0]},{rect.y_range[0]}: Reduced from {len_for_reporting} to {len(tasks)} tasks.') data_task_colors = { - i: _data_task_color(rect.metadata, i, task, color_map) + i: wse3.data_task_color(rect.metadata, i, task, color_map) for i, task in enumerate(tasks) if task.task_type == 'data' } # On WSE-2 a data-task ID *is* its color, so a local task must not reuse one. @@ -441,7 +441,7 @@ def generate_rectangle(kernel: spir.Kernel, # Declare each data task ID. Data tasks that share a color are aliases of one hardware ID, and # a state variable selects which of them the shared task runs as. for slot in task_bindings.data_slots: - id_expr = _data_task_id_builtin(rect.metadata, slot, tasks, dsds) + id_expr = wse3.data_task_id_builtin(rect.metadata, slot, tasks, dsds) for task_index in slot.task_indices: current_code.write(f'const dtask_{task_index}_id = {id_expr};\n') if slot.recycled: @@ -532,7 +532,7 @@ def generate_rectangle(kernel: spir.Kernel, exit_task_sequential &= not any( t.task_type == 'data' for t in tasks for n, _ in t.outgoing if n == -1) # No data tasks exit_task_blocked = any(n == -1 and typ == tdag.InterTaskEdge.UNBLOCK for t in tasks for n, typ in t.outgoing) - hardware_exit_id = None if exit_task_sequential else _exit_task_hardware_id( + hardware_exit_id = None if exit_task_sequential else wse3.exit_task_hardware_id( {slot.hardware_task_id for slot in task_bindings.local_slots}, set(color_map.values())) # Bind exit task @@ -925,222 +925,6 @@ def _dsd_from_stream(stream_candidates: dict[str, tuple[spir.StreamDeclaration | return cslstruct.MemoryDSD(dsd_type, name, extents, idxvars, indices) -def _declare_queue_initialization(dsds: UniqueDSDDict, rect: PEBlock, footer: StringIO, - color_map: dict[str, int]) -> None: - """ - Binds every fabric queue this PE uses to its color, which WSE-3 requires. - - On WSE-2 a fabric queue picks its color up from the descriptor that uses it. WSE-3 does not: - a queue must be tied to a color with ``@initialize_queue`` before any transfer over it will - proceed, and a program that omits it simply hangs. Queues are handed out per channel (see - ``_collect_unique_dsds``), so each one is named by exactly one color here. - - :param dsds: The descriptors collected for this rectangle. - :param rect: The PE block being generated, used for the switch-advance descriptors. - :param footer: The ``comptime`` block to write the bindings into. - :param color_map: Stream name to color number, for the switch-advance descriptors. - """ - if not csl.ARCH == 'wse3': - return - - # (queue kind, queue id) -> color expression. Both the data descriptors and the control - # descriptors that carry switch advances need their queue bound. - bindings: dict[tuple[str, int], str] = {} - for entries in dsds.values(): - for _, dsd in entries: - if not isinstance(dsd, cslstruct.FabricDSD) or not dsd.color: - continue - direction = 'in' if dsd.dsd_type == cslstruct.DSDType.fabin else 'out' - kind = 'input_queue' if dsd.dsd_type == cslstruct.DSDType.fabin else 'output_queue' - bindings.setdefault((kind, dsd.queue), f'{dsd.color}_{direction}') - - for statement in rect.metadata.compute.statements: - if isinstance(statement, spir.CloseStatement) and statement.switch_advance: - stream = stream_lifetime.underlying_stream(statement.stream_name) - name = cslstmt.name_to_csl(stream) - queue = None - for _, dsd in dsds.get(stream.as_ir(), ()): - if isinstance(dsd, cslstruct.FabricDSD) and dsd.dsd_type == cslstruct.DSDType.fabout: - queue = dsd.queue - break - if queue is None: - queue = csl.OUTPUT_QUEUE_IDS[0] - bindings.setdefault(('output_queue', queue), f'@get_color({color_map[name + "_OUT"]})') - - for (kind, queue), color in sorted(bindings.items()): - footer.write(f' @initialize_queue(@get_{kind}({queue}), .{{ .color = {color} }});\n') - - -def _statement_transfer_points(statement: spir.Statement, names: set[spir.Identifier], - inbound: bool) -> list[spir.Identifier]: - """ - Fabric transfers in ``statement``, in source order, with sequential ``for`` bodies repeated. - - Walking a loop body once makes its colours look sequential, so occupancy pooling would give - them one queue. The next iteration of an earlier colour can already occupy the router when a - later colour of the same body remaps that queue -- WSE-2 then aborts with "Attempt to remap - input queue N, from C_i to C_j, but the router is holding wavelets". Appending the body a - second time makes a colour used on both sides of another occupy a span that overlaps it, the - same rule that keeps a reused colour's queue across a gap between unrolled phases. - """ - if isinstance(statement, spir.ForStatement): - body = _transfer_points(statement.body, names, inbound) - return body + body - - nested_skip: set[int] = set() - points: list[spir.Identifier] = [] - for node in statement.walk(): - if id(node) in nested_skip: - continue - if node is not statement and isinstance(node, spir.ForStatement): - for descendant in node.walk(): - nested_skip.add(id(descendant)) - points.extend(_statement_transfer_points(node, names, inbound)) - continue - stream = _fabric_transfer_stream(node, inbound) - if stream is None or stream not in names: - continue - points.append(stream) - return points - - -def _transfer_points(statements: list[spir.Statement], names: set[spir.Identifier], - inbound: bool) -> list[spir.Identifier]: - points: list[spir.Identifier] = [] - for statement in statements: - points.extend(_statement_transfer_points(statement, names, inbound)) - return points - - -def _queue_spans(compute: spir.ComputeBlock, names: set[spir.Identifier], queue_key, - inbound: bool) -> dict[str, tuple[int, int]]: - """ - Occupancy of each queue key along the linearized send/receive order of this PE. - - Sequential ``for`` bodies are counted twice so a colour that comes back on the next iteration - keeps its queue across the loop-carried gap; see ``_statement_transfer_points``. Uses of the - same channel still collapse to one span, so a colour that comes back after a gap between - unrolled phases keeps its queue for the whole of that span. - - :param compute: The compute block being lowered. - :param names: Streams that actually bind a fabric queue in this direction. - :param queue_key: Maps a stream identifier to its grouping key (channel, or the name itself). - :param inbound: True to walk receives, False to walk sends. - :return: Mapping of grouping key to ``(first_use, last_use)`` in linearized order. - """ - points = _transfer_points(compute.statements, names, inbound) - - spans: dict[str, tuple[int, int]] = {} - for index, stream in enumerate(points): - key = queue_key(stream) - if key in spans: - start, _ = spans[key] - spans[key] = (start, index) - else: - spans[key] = (index, index) - return spans - - -def _microthread_intervals(compute: spir.ComputeBlock, input_names: set[spir.Identifier], - output_names: set[spir.Identifier], - queue_key) -> dict[str, list[tuple[int, int]]]: - """ - The intervals over which each stream group holds a microthread on one PE. - - Both directions are numbered in one space, since a microthread is one resource across them. A - transfer that keeps a completion handle is in flight until that handle is awaited, which is where - real concurrency comes from: a receive started before a send is still running while the send is. - A self-awaited transfer is given its own slot and the next one, because the activation that - awaits it also starts what follows. - - :param compute: The compute block being lowered. - :param input_names: Streams that bind an input queue on this PE. - :param output_names: Streams that bind an output queue on this PE. - :param queue_key: Maps a stream identifier to its grouping key (channel, or the name itself). - :return: Mapping of direction-prefixed grouping key to the intervals it is in flight over. - """ - live: dict[str, list[tuple[int, int]]] = {} - pending: dict[str, list[tuple[str, int]]] = {} - index = 0 - - def close(keys: list[tuple[str, int]], end: int) -> None: - for key, start in keys: - live.setdefault(key, []).append((start, end)) - - for statement in compute.statements: - for node in statement.walk(): - if isinstance(node, spir.AwaitAllStatement): - for keys in pending.values(): - close(keys, index) - pending.clear() - continue - if isinstance(node, spir.AwaitCompletionStatement): - close(pending.pop(node.completion_name.as_ir(), []), index) - continue - if isinstance(node, spir.ForeachStatement) and not node.parameter_range: - continue # A data task runs on no microthread. - for inbound, names in ((True, input_names), (False, output_names)): - stream = _fabric_transfer_stream(node, inbound) - if stream is None or stream not in names: - continue - key = f'{"in" if inbound else "out"} {queue_key(stream)}' - completion = getattr(node, 'completion_name', None) - if completion is None: - live.setdefault(key, []).append((index, index + 1)) - else: - pending.setdefault(completion.name.as_ir(), []).append((key, index)) - index += 1 - for keys in pending.values(): - close(keys, index) - return live - - -def _fabric_transfer_stream(node: spir.SpatialNode, inbound: bool) -> Optional[spir.Identifier]: - """ - The stream a node transfers in the requested direction, or ``None``. - - A data-task receive (``foreach`` with no range) counts only on WSE-3, where its task ID is an - input queue and so needs one reserved. - """ - if inbound: - if isinstance(node, spir.ReceiveStatement): - return stream_lifetime.underlying_stream(node.stream_name) - if (isinstance(node, spir.ForeachStatement) and node.receive_stream is not None - and (node.parameter_range or csl.ARCH == 'wse3')): - return stream_lifetime.underlying_stream(node.receive_stream.stream_name) - return None - if isinstance(node, spir.SendStatement): - return stream_lifetime.underlying_stream(node.stream_name) - return None - - -def _streams_with_fabric_dsds(compute: spir.ComputeBlock, memcpy_mode: bool, - stream_args: set[spir.Identifier], inbound: bool) -> set[spir.Identifier]: - """ - Streams that lower to a fabric DSD in one direction, so they need a hardware queue. - - A data-task receive (``foreach`` with no range) binds the color itself and takes a queue only on - WSE-3. Memcpy arguments are already in local memory, so they do not take one either. - - :param compute: The compute block being lowered. - :param memcpy_mode: Whether memcpy mode is used. - :param stream_args: Kernel-argument streams, which memcpy has already copied. - :param inbound: True for receives, False for sends. - :return: The stream identifiers that need a queue in that direction. - """ - result: set[spir.Identifier] = set() - argument_names = {name.as_ir() for name in stream_args} - for statement in compute.statements: - for node in statement.walk(): - name = _fabric_transfer_stream(node, inbound) - if name is None: - continue - if memcpy_mode and name.as_ir() in argument_names: - continue - result.add(name) - return result - - def _collect_unique_dsds( tasks: list[tdag.CSLTask], rect: PEBlock, @@ -1204,10 +988,10 @@ def queue_key(stream: spir.Identifier) -> str: channel = channel_of_stream.get(stream.as_ir(), 'auto') return stream.as_ir() if channel == 'auto' else f'channel {channel}' - input_names = _streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=True) - output_names = _streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=False) - input_spans = _queue_spans(rect.compute, input_names, queue_key, inbound=True) - output_spans = _queue_spans(rect.compute, output_names, queue_key, inbound=False) + input_names = fabric_occupancy.streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=True) + output_names = fabric_occupancy.streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=False) + input_spans = fabric_occupancy.queue_spans(rect.compute, input_names, queue_key, inbound=True) + output_spans = fabric_occupancy.queue_spans(rect.compute, output_names, queue_key, inbound=False) # WSE-3 remaps a fabric queue onto the next color at the first transfer that uses it, and # faults or stalls if the queue still holds wavelets. Occupancy in the compute block is not # enough to prove it is empty, so every color keeps its own queue. That also keeps data-task @@ -1228,7 +1012,7 @@ def queue_key(stream: spir.Identifier) -> str: # microthread and abort with "trying to term ut_instr[N], but it's not ours". Microthreads are # one resource across both directions, so they are handed out together. microthread_of = stream_lifetime.assign_microthreads( - _microthread_intervals(rect.compute, input_names, output_names, queue_key), + fabric_occupancy.microthread_intervals(rect.compute, input_names, output_names, queue_key), csl.MICROTHREAD_IDS, location=location) def allocate_microthread(stream: spir.Identifier, inbound: bool) -> int | None: @@ -1598,109 +1382,6 @@ def _write_indented_block(current_code: StringIO, block: str, indent: str) -> No current_code.write(f'{indent}{line}\n') -def _exit_task_hardware_id(used_ids: set[int], color_ids: set[int]) -> int: - """Return a local-task ID for ``exit_task`` that nothing else has bound. - - Walks the activatable range from 8 and skips IDs already taken by program - slots, by memcpy/system reservations, and on WSE-2 by colors (which are also - data-task IDs there). - - :param used_ids: Hardware IDs already assigned to local-task slots. - :param color_ids: Colors allocated to this PE. - :return: A free activatable identifier. - """ - occupied = set(used_ids) | set(csl.RESERVED_LOCAL_TASK_IDS) - if csl.ARCH != 'wse3': - occupied |= set(color_ids) - for tid in range(8, 31): - if tid not in occupied: - return tid - raise SyntaxError( - 'No free local task ID remains for exit_task ' - f'(occupied {sorted(occupied)}).') - - -def _data_task_color(rect: PEBlock, task_index: int, task: tdag.CSLTask, color_map: dict[str, int]) -> int: - """ - Returns the color a data task listens on. - - On WSE-2 that color is also the hardware task ID. On WSE-3 the ID is the - input queue bound to this color; see ``_data_task_id_builtin``. - - :param task_index: Only used to name the task in the error message. - """ - stmt = rect.compute.statements[task.statements[0]] - assert isinstance(stmt, spir.ForeachStatement) - sname = stmt.receive_stream.stream_name - if isinstance(sname, spir.ArraySlice): - sname = sname.array - if name_to_csl(sname) + '_H2D' in color_map: - return color_map[name_to_csl(sname) + '_H2D'] - if name_to_csl(sname) + '_IN' in color_map: - return color_map[name_to_csl(sname) + '_IN'] - raise ValueError(f'Cannot find color for stream "{name_to_csl(sname)}" in data task {task_index}') - - -def _input_queue_for_data_slot( - rect: PEBlock, - slot: task_recycling.DataTaskSlot, - tasks: list[tdag.CSLTask], - dsds: UniqueDSDDict, -) -> int: - """Return the fabric input queue the receives in ``slot`` share. - - On WSE-3 a data task's hardware ID is that queue, which - ``_declare_queue_initialization`` has already bound to the slot's color. - - :param rect: The PE block being generated. - :param slot: The data-task slot whose color the receives listen on. - :param tasks: All tasks of this PE. - :param dsds: Fabric descriptors of this PE, which record the queue assignment. - :return: The input-queue identifier. - """ - queues: set[int] = set() - for task_index in slot.task_indices: - task = tasks[task_index] - stmt = rect.compute.statements[task.statements[0]] - assert isinstance(stmt, spir.ForeachStatement) - sname = stmt.receive_stream.stream_name - if isinstance(sname, spir.ArraySlice): - sname = sname.array - for _, dsd in dsds.get(sname.as_ir(), []): - if isinstance(dsd, cslstruct.FabricDSD) and dsd.dsd_type == cslstruct.DSDType.fabin: - queues.add(dsd.queue) - if len(queues) != 1: - found = sorted(queues) if queues else 'none' - raise SyntaxError( - f'WSE-3 data task on color {slot.color} needs exactly one input queue, found {found}.\n' - " note: @get_data_task_id takes an input_queue on WSE-3, not a color") - return next(iter(queues)) - - -def _data_task_id_builtin( - rect: PEBlock, - slot: task_recycling.DataTaskSlot, - tasks: list[tdag.CSLTask], - dsds: UniqueDSDDict, -) -> str: - """Return the ``@get_data_task_id(...)`` expression for ``slot``. - - WSE-2 constructs a data-task ID from the color the receive listens on. - WSE-3 constructs it from the input queue already bound to that color; - passing the color is rejected as ``expected 'input_queue' expression, got: 'color'``. - - :param rect: The PE block being generated. - :param slot: The data-task slot, whose color is the receive's fabric color. - :param tasks: All tasks of this PE, indexed as in the slot. - :param dsds: Fabric descriptors of this PE, which record the queue assignment. - :return: A CSL expression of type ``data_task_id``. - """ - if csl.ARCH == 'wse3': - queue = _input_queue_for_data_slot(rect, slot, tasks, dsds) - return f'@get_data_task_id(@get_input_queue({queue}))' - return f'@get_data_task_id(@get_color({slot.color}))' - - def _generate_data_task_slot( rect: PEBlock, slot: task_recycling.DataTaskSlot, diff --git a/spada/lowering/wse3.py b/spada/lowering/wse3.py new file mode 100644 index 00000000..a3ee9ba5 --- /dev/null +++ b/spada/lowering/wse3.py @@ -0,0 +1,192 @@ +""" +CSL details that differ on WSE-3. + +WSE-3 binds a fabric queue to one color for the whole kernel, and a data task's hardware ID is +that input queue rather than the color. WSE-2 takes the color from the descriptor and uses it as +the data-task ID, so the same helpers emit both forms. +""" + +from io import StringIO + +from spada.syntax.csl import constants as csl +from spada.syntax.csl import structures as cslstruct +from spada.syntax.csl import task_recycling, tasks as tdag +from spada.syntax.csl.statements import name_to_csl +from spada.syntax.spatial_ir import irnodes as spir +from spada.syntax.spatial_ir import stream_lifetime +from spada.syntax.spatial_ir.canonicalization import PEBlock + +UniqueDSDDict = cslstruct.UniqueDSDDict + + +def declare_queue_initialization( + dsds: UniqueDSDDict, rect: PEBlock, footer: StringIO, color_map: dict[str, int] +) -> None: + """ + Binds every fabric queue this PE uses to its color, which WSE-3 requires. + + On WSE-2 a fabric queue picks its color up from the descriptor that uses it. WSE-3 does not: + a queue must be tied to a color with ``@initialize_queue`` before any transfer over it will + proceed, and a program that omits it simply hangs. Queues are handed out per channel (see + ``_collect_unique_dsds``), so each one is named by exactly one color here. + + :param dsds: The descriptors collected for this rectangle. + :param rect: The PE block being generated, used for the switch-advance descriptors. + :param footer: The ``comptime`` block to write the bindings into. + :param color_map: Stream name to color number, for the switch-advance descriptors. + """ + if not csl.ARCH == "wse3": + return + + # (queue kind, queue id) -> color expression. Both the data descriptors and the control + # descriptors that carry switch advances need their queue bound. + bindings: dict[tuple[str, int], str] = {} + for entries in dsds.values(): + for _, dsd in entries: + if not isinstance(dsd, cslstruct.FabricDSD) or not dsd.color: + continue + direction = "in" if dsd.dsd_type == cslstruct.DSDType.fabin else "out" + kind = ( + "input_queue" + if dsd.dsd_type == cslstruct.DSDType.fabin + else "output_queue" + ) + bindings.setdefault((kind, dsd.queue), f"{dsd.color}_{direction}") + + for statement in rect.metadata.compute.statements: + if isinstance(statement, spir.CloseStatement) and statement.switch_advance: + stream = stream_lifetime.underlying_stream(statement.stream_name) + name = name_to_csl(stream) + queue = None + for _, dsd in dsds.get(stream.as_ir(), ()): + if ( + isinstance(dsd, cslstruct.FabricDSD) + and dsd.dsd_type == cslstruct.DSDType.fabout + ): + queue = dsd.queue + break + if queue is None: + queue = csl.OUTPUT_QUEUE_IDS[0] + bindings.setdefault( + ("output_queue", queue), f"@get_color({color_map[name + '_OUT']})" + ) + + for (kind, queue), color in sorted(bindings.items()): + footer.write( + f" @initialize_queue(@get_{kind}({queue}), .{{ .color = {color} }});\n" + ) + + +def exit_task_hardware_id(used_ids: set[int], color_ids: set[int]) -> int: + """Return a local-task ID for ``exit_task`` that nothing else has bound. + + Walks the activatable range from 8 and skips IDs already taken by program + slots, by memcpy/system reservations, and on WSE-2 by colors (which are also + data-task IDs there). + + :param used_ids: Hardware IDs already assigned to local-task slots. + :param color_ids: Colors allocated to this PE. + :return: A free activatable identifier. + """ + occupied = set(used_ids) | set(csl.RESERVED_LOCAL_TASK_IDS) + if csl.ARCH != "wse3": + occupied |= set(color_ids) + for tid in range(8, 31): + if tid not in occupied: + return tid + raise SyntaxError( + f"No free local task ID remains for exit_task (occupied {sorted(occupied)})." + ) + + +def data_task_color( + rect: PEBlock, task_index: int, task: tdag.CSLTask, color_map: dict[str, int] +) -> int: + """ + Returns the color a data task listens on. + + On WSE-2 that color is also the hardware task ID. On WSE-3 the ID is the + input queue bound to this color; see ``data_task_id_builtin``. + + :param rect: The PE block the task belongs to. + :param task_index: Only used to name the task in the error message. + :param task: The data task, whose first statement is the receiving ``foreach``. + :param color_map: Stream name to color number. + :return: The color number. + """ + stmt = rect.compute.statements[task.statements[0]] + assert isinstance(stmt, spir.ForeachStatement) + sname = stmt.receive_stream.stream_name + if isinstance(sname, spir.ArraySlice): + sname = sname.array + if name_to_csl(sname) + "_H2D" in color_map: + return color_map[name_to_csl(sname) + "_H2D"] + if name_to_csl(sname) + "_IN" in color_map: + return color_map[name_to_csl(sname) + "_IN"] + raise ValueError( + f'Cannot find color for stream "{name_to_csl(sname)}" in data task {task_index}' + ) + + +def input_queue_for_data_slot( + rect: PEBlock, + slot: task_recycling.DataTaskSlot, + tasks: list[tdag.CSLTask], + dsds: UniqueDSDDict, +) -> int: + """Return the fabric input queue the receives in ``slot`` share. + + On WSE-3 a data task's hardware ID is that queue, which + ``declare_queue_initialization`` has already bound to the slot's color. + + :param rect: The PE block being generated. + :param slot: The data-task slot whose color the receives listen on. + :param tasks: All tasks of this PE. + :param dsds: Fabric descriptors of this PE, which record the queue assignment. + :return: The input-queue identifier. + """ + queues: set[int] = set() + for task_index in slot.task_indices: + task = tasks[task_index] + stmt = rect.compute.statements[task.statements[0]] + assert isinstance(stmt, spir.ForeachStatement) + sname = stmt.receive_stream.stream_name + if isinstance(sname, spir.ArraySlice): + sname = sname.array + for _, dsd in dsds.get(sname.as_ir(), []): + if ( + isinstance(dsd, cslstruct.FabricDSD) + and dsd.dsd_type == cslstruct.DSDType.fabin + ): + queues.add(dsd.queue) + if len(queues) != 1: + found = sorted(queues) if queues else "none" + raise SyntaxError( + f"WSE-3 data task on color {slot.color} needs exactly one input queue, found {found}.\n" + " note: @get_data_task_id takes an input_queue on WSE-3, not a color" + ) + return next(iter(queues)) + + +def data_task_id_builtin( + rect: PEBlock, + slot: task_recycling.DataTaskSlot, + tasks: list[tdag.CSLTask], + dsds: UniqueDSDDict, +) -> str: + """Return the ``@get_data_task_id(...)`` expression for ``slot``. + + WSE-2 constructs a data-task ID from the color the receive listens on. + WSE-3 constructs it from the input queue already bound to that color; + passing the color is rejected as ``expected 'input_queue' expression, got: 'color'``. + + :param rect: The PE block being generated. + :param slot: The data-task slot, whose color is the receive's fabric color. + :param tasks: All tasks of this PE, indexed as in the slot. + :param dsds: Fabric descriptors of this PE, which record the queue assignment. + :return: A CSL expression of type ``data_task_id``. + """ + if csl.ARCH == "wse3": + queue = input_queue_for_data_slot(rect, slot, tasks, dsds) + return f"@get_data_task_id(@get_input_queue({queue}))" + return f"@get_data_task_id(@get_color({slot.color}))" diff --git a/spada/syntax/csl/routing.py b/spada/syntax/csl/routing.py index 5efd89b4..9537c7e1 100644 --- a/spada/syntax/csl/routing.py +++ b/spada/syntax/csl/routing.py @@ -293,7 +293,7 @@ def declare_switch_advances(rect: Rectangle[PEBlock], header: StringIO, color_ma ``switch_advance`` field; this only has to provide the descriptor it is sent through. The descriptor reuses the *same* output queue as the stream's data, which is mandatory rather - than tidy: a queue is bound to one color on WSE-3 (see ``_declare_queue_initialization``), and + than tidy: a queue is bound to one color on WSE-3 (see :func:`spada.lowering.wse3.declare_queue_initialization`), and queues are handed out per channel, so sending a control wavelet for one color through the queue that belongs to another silently fails to advance anything. A close only ever runs on a PE that sends the stream, so the outgoing descriptor always exists. diff --git a/tests/spatial_ir/test_task_recycling.py b/tests/spatial_ir/test_task_recycling.py index a898da58..ef362e75 100644 --- a/tests/spatial_ir/test_task_recycling.py +++ b/tests/spatial_ir/test_task_recycling.py @@ -2,6 +2,7 @@ import pytest from spada.lowering import spatial_ir_to_csl as s2c +from spada.lowering import wse3 from spada.syntax.csl import constants, task_recycling, tasks as tdag from spada.syntax.spatial_ir import analysis, parser, passes from spada.syntax.spatial_ir.canonicalization import PEBlock @@ -259,10 +260,17 @@ def test_plan_is_deterministic(): def test_local_task_ids_do_not_include_memcpy_reservations(): - """The assignable pool must not contain IDs memcpy already binds.""" + """The assignable pool must not contain IDs memcpy already binds. + + The WSE-3 range includes task 21, which memcpy binds; ``LOCAL_TASK_IDS`` drops the reserved IDs. + """ assert set(constants.LOCAL_TASK_IDS).isdisjoint(constants.RESERVED_LOCAL_TASK_IDS) - assert 21 not in constants._CSL_LOCAL_TASK_IDS['wse3'] - assert set(constants._CSL_LOCAL_TASK_IDS['wse3']).isdisjoint(constants._RESERVED_LOCAL_TASK_IDS['wse3']) + wse3_assignable = [ + task_id for task_id in constants._CSL_LOCAL_TASK_IDS['wse3'] + if task_id not in constants._RESERVED_LOCAL_TASK_IDS['wse3'] + ] + assert 21 in constants._CSL_LOCAL_TASK_IDS['wse3'] + assert 21 not in wse3_assignable assert set(constants._CSL_LOCAL_TASK_IDS['wse2']).isdisjoint(constants._RESERVED_LOCAL_TASK_IDS['wse2']) @@ -270,7 +278,7 @@ def test_exit_task_skips_the_first_memcpy_reservation(): """If every ID below memcpy's first local task is taken, exit_task must hop the hole.""" first_reserved = min(constants.RESERVED_LOCAL_TASK_IDS) used = set(range(8, first_reserved)) - exit_id = s2c._exit_task_hardware_id(used, set()) + exit_id = wse3.exit_task_hardware_id(used, set()) assert exit_id not in used assert exit_id not in constants.RESERVED_LOCAL_TASK_IDS assert exit_id == first_reserved + 1 From 9765b85be434010645a4fbaf408f7646820a3157 Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 5 Oct 2026 20:10:02 +0200 Subject: [PATCH 67/68] Improve documentation and comment style --- README.md | 6 +- irspec/docs/spatial/routing_wse.md | 346 +++++++----------- .../spatial/simple/exchange_bundle_1D.sptl | 18 +- .../spatial/sort/batcher_oddeven_wse3_1D.sptl | 117 +++--- .../spatial/sort/odd_even_sort_1D_looped.sptl | 51 +-- samples/spatial/sort/shearsort_2D.sptl | 83 +---- spada/lowering/fabric_occupancy.py | 38 +- spada/lowering/spatial_ir_to_csl.py | 27 +- spada/lowering/wse3.py | 5 +- spada/runtime/runtime.py | 41 +-- spada/syntax/csl/constants.py | 51 +-- spada/syntax/csl/routing.py | 73 ++-- spada/syntax/csl/task_recycling.py | 30 +- spada/syntax/spatial_ir/canonicalization.py | 3 +- spada/syntax/spatial_ir/shift_bundles.py | 52 +-- spada/syntax/spatial_ir/stream_lifetime.py | 64 ++-- tests/csl_runtime/run_tests.sh | 6 +- .../samples/data_task_two_epochs.sptl | 12 +- .../csl_runtime/samples/shift_bundle_1D.sptl | 15 +- .../csl_runtime/test_data_task_two_epochs.sh | 5 +- tests/csl_runtime/test_exchange_bundle_1d.sh | 6 +- .../test_odd_even_sort_1d_looped.sh | 9 +- tests/csl_runtime/test_shearsort_2d_looped.sh | 10 +- tests/csl_runtime/test_shift_bundle_1d.sh | 5 +- tests/csl_runtime/test_spmv.sh | 9 +- tests/spatial_ir/test_dsd_ops.py | 2 +- tests/spatial_ir/test_routing.py | 12 +- 27 files changed, 408 insertions(+), 688 deletions(-) diff --git a/README.md b/README.md index 25dcfc37..55d1ec78 100644 --- a/README.md +++ b/README.md @@ -106,10 +106,10 @@ Sample SpaDA programs are in `samples/`: | `samples/stencils.py` | GT4Py stencil definitions (Laplacian, vertical advection, UVBKE, …) | | `samples/advanced_stencils.py` | GT4Py definitions for horizontal diffusion kernels | | `samples/benchmarks/` | Pre-compiled `.spst`/`.sptl` pairs for five kernels at five domain sizes | -| `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy`, plus the one-color `exchange_bundle_1D` microkernel | +| `samples/spatial/simple/` | Basic single-PE and streaming operations: `add`, `copy`, `forward_sum`, `backward_sum`, `mult_scalar`, `streaming_copy`, and `exchange_bundle_1D` | | `samples/spatial/blas/` | Linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase`, `spmv` | | `samples/spatial/collectives/` | Reductions (`scalar`, `chain`, `tree`, `twophase` in 1D/2D) and broadcasts (`broadcast_1D`, `broadcast_2D`, and multicast variants) | -| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each, one contiguous sequence of `N*K` keys per row: `batcher_oddeven_wse3_1D` (the widest phases bundled onto one color, the rest pooled by origin, which reaches `L = 4` on WSE-3), and `odd_even_sort_1D_looped` (N neighbour compare-splits as a runtime loop, four static channels). `shearsort_2D` is the 2D snake-order mesh sort on `N x N` PEs, the same neighbour rounds as a runtime loop on eight static channels (WSE-3) | +| `samples/spatial/sort/` | Sorting networks over `R` independent rows of `2^L` PEs holding `K` keys each: `batcher_oddeven_wse3_1D` (Batcher odd-even mergesort), `odd_even_sort_1D_looped` (1D looped odd-even transposition sort), and `shearsort_2D` (2D snake-order shearsort on WSE-3) | | `samples/spatial/stencils/` | Stencil examples: `laplacian` (high-level) and `laplacian_routed` (explicit routing) | | `samples/spst/` | Stencil IR examples | @@ -117,7 +117,7 @@ Sample SpaDA programs are in `samples/`: ## SDK Version and WSE compatibility -The code has been tested for CSL SDK 1.4 on WSE-2 and WSE-3. The compiler and the CSL runtime tests read `WSE_ARCH` (`wse2`, the default, or `wse3`). CI runs the simulator suite for both. A kernel that a generation cannot express skips in that architecture's run: `shearsort_2D` needs four inbound queues in one epoch, which WSE-2 does not have, and `batcher_oddeven_wse3_1D` is the origin-pooled Batcher that reaches `L = 4` only on WSE-3. +The code has been tested with CSL SDK 1.4 on WSE-2 and WSE-3. The target architecture is selected via `WSE_ARCH` (`wse2`, the default, or `wse3`). CI runs simulator tests for both architectures. Kernels targeting features unique to WSE-3 (such as `shearsort_2D`, which requires four concurrent inbound queues per PE, and `batcher_oddeven_wse3_1D`) automatically skip during WSE-2 test runs. ## Testing diff --git a/irspec/docs/spatial/routing_wse.md b/irspec/docs/spatial/routing_wse.md index e94a4994..31450d32 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -2,135 +2,96 @@ This page describes how the **Cerebras WSE / CSL backend** realizes routing concepts, such as epochs, on the target hardware architecture. The SpaDA IR abstracts away details such as colors, routers, -switch positions, and control wavelets, and the code generation process lowers potentially differently -to specific WSE architecture (selected via the `WSE_ARCH` environment variable). Where -the two WSE generations differ, the text says so; the compiler selects between them on -`WSE_ARCH`. The code generator still preserves the correctness conditions of -[Undefined Behavior](../routing#undefined-behavior). - -## Lowering to Switches - -Channels are a scarce resource: each channel that is live at a PE occupies one of the hardware's -routing colors. Epochs are what makes it possible to reuse a channel, and hence a color, for -several streams. - -A stream induces, at each PE of its path, a *route configuration*: the set of directions the PE -receives from and the set of directions it transmits to (where `RAMP` denotes the PE's own compute -element). Consider a fixed channel $C$ and a fixed PE $(i, j)$. Ordering the streams that use $C$ -at $(i, j)$ by their epochs yields a sequence of route configurations -$R_0, R_1, \dotsc, R_{n-1}$, which is realized by the PE's *switch* for the color assigned to $C$: -$R_0$ is the initial configuration and the router *advances* to $R_{k+1}$ at the epoch boundary. - -Whichever side a switch position leaves unspecified keeps the value it currently has, so positions -compose incrementally. - -!!! warning "WSE-2: A Switch Position Carries One Direction" - On WSE-2 a switch position records *either* the input the router receives from or the output it - transmits to, never both — `cslc` rejects a position naming both with *"cannot have both an - input and an output in the same switch position"*. A transition that changes both sides — a PE - that stops receiving on a channel and starts sending on it, as in a systolic chain — therefore - occupies **two** positions, passing through an intermediate configuration that keeps the old - input and takes the new output. The intermediate is a pure relay, occupied only between the two - advances that retire the configuration, and it must keep the old input so that the second - advance still reaches the router. - - WSE-3 accepts both directions in one position, so the same transition costs one position and one - advance there. - -!!! danger "Error: Too Many Route Configurations" - A router holds a bounded number of switch positions per color (four on both WSE-2 and WSE-3). - *If the streams sharing a channel require more positions than that at a single PE, a compile - error is raised.* Assigning a different channel to some of the streams resolves it, at the cost - of an additional color. On WSE-2 a configuration which changes both the input and the output - direction costs two positions, so four positions is fewer than four turnarounds there. - -When the sequence of configurations at a router is periodic — a halo exchange that alternates -between sending and receiving across phases produces $R_0, R_1, R_0, R_1$ — only one period is -stored and the switch wraps around from the last position back to the base one (`ring_mode`). - -Consecutive configurations that are equal do not consume a position and do not require an advance. -This is a common case: two streams declared as `relative_stream(-2, 0)` in successive phases induce -the same configuration at every PE of their paths, so their shared channel needs no switching at -all. - -!!! tip "When the Configuration Never Changes, the Phases May Not Be Needed" - A channel whose route configuration is the same in every epoch is never reassigned, and its - routers never advance. The epoch boundaries around it are then buying only *ordering* — and the - fabric already delivers a channel's wavelets in order. Such a sequence of phases can be - collapsed into a single epoch with a sequential `for` in the compute blocks, which lowers to a - real loop and so costs code and compile time independent of the number of rounds. - `samples/spatial/sort/odd_even_sort_1D_looped.sptl` is the example: N odd-even rounds on four - static channels, one CSL loop, no per-round barrier. The 2D analogue is - `samples/spatial/sort/shearsort_2D.sptl`: eight static neighbour channels, nested - loops, no switches. A fully interior PE there receives on four colours in one epoch, which - fits WSE-3's six exclusive queues and not WSE-2's two. - - This does *not* generalize to channels that switch: a router's positions are a static sequence, - so the epoch a configuration belongs to has to be visible to the compiler. - -An advance is driven by the `close` that ends the epoch. When some *other* router on the path has -to move, the sending PE emits a *switch-advance control message* on the channel, one per position -to be traversed. It follows the stream's path using the configuration that is being retired, and -advances the router of each PE it traverses, after all data of the epoch. On WSE-2, when only the -sending PE's own router has to move, and only by one position, the last data wavelet does that -itself (`.advance_switch` on the fabric output DSD). A second operation on the same output queue -is what drops a data wavelet there once a back-pressured send fills it (output queues 2 and 3 -hold six 16-bit words; three `f32` values already fill them). WSE-3 keeps a traveling `SWITCH_ADV` -for that local flip as well: its output queues hold eight words, and origin-pooled destinations -that also switch are only moved by a control wavelet on the path. - -!!! warning "WSE: Advances Are Not Selective" - A CSL control wavelet nominally carries up to eight per-router switching commands - (``'s `MAX_CMDS`), which would let one message advance some routers on a path and leave - others alone. **On the WSE hardware it does not work that way.** Measured on the simulator, only - command slot 0 is ever executed, and **every** switch-configured router the wavelet reaches - applies it; slots 1–7 had no effect in any topology tested — the sender's own router, one hop, - two hops through a plain relay, and two switch-configured routers in sequence. A generated kernel - built on the opposite assumption, advancing the fourth router of a path with an `[ADV, NOP, NOP, - ADV]` chain, stalled in the fabric. The compiler therefore emits `encode_single_payload`, which - writes slot 0 only. - - The consequence is that a message cannot advance one router while leaving another on the same - path where it is. *If the routers along one path would have to advance by different amounts, a - compile error is raised.* A router that is already on its last position is exempt: outside - `ring_mode` an advance past the last position is a no-op, so a message passing through may - over-advance it harmlessly. (Such a router is not necessarily finished — it may keep relaying the - same configuration for the rest of the kernel, which is exactly what the bundle below relies - on.) - - This is a property of the *payload*, not of switching: a router can still be switched at a time - only it knows, and delivery to a compute element can still be made selective, by the two - mechanisms the next section combines. - -!!! danger "WSE-2: A Two-Advance Turnaround Overshoots the Receiver" - A control message stops at the first router whose current output is `RAMP`, and it is routed by - the position that router holds *when the message arrives*. The two messages of a WSE-2 - turnaround therefore do not travel the same distance: the first stops at the receiver and moves - it onto the intermediate position, whose output is a real direction rather than `RAMP`, so the - second is *forwarded past the receiver* to the next PE on the line. - - That is harmless when the next PE is not switch-configured on the same color. When it is, the - stray message reaches a router that no epoch of its own is retiring. **The compiler does not - detect this**, and a kernel that trips it hangs rather than failing to compile. - - An odd-even transposition sort makes the hazard concrete: put both round parities on one - eastward channel and every interior PE alternates between sending and receiving on it, so every - close is a turnaround and every receiver has a switch-configured neighbour behind it. That - kernel deadlocks on WSE-2 and runs on WSE-3, where the turnaround costs one position and one - message and nothing overshoots. `samples/spatial/sort/odd_even_sort_1D_looped.sptl` avoids it - by giving each round parity its own pair of channels: each PE's role on a channel is then - fixed, no router switches at all, and no close emits a message. - -Because the control message travels the path of the retired configuration in order behind the data, -a receiving PE needs to emit nothing to advance its own router: the ordering required by the -[lemma](../routing#undefined-behavior) in the IR semantics is provided by the fabric. A receiver's `close` -therefore has no runtime effect; it exists so that the lifetime of the stream — and hence the number -of elements it carries — is stated by every participant and can be checked. +switch positions, and control wavelets, and the code generation process lowers to the specific +WSE architecture selected via the `WSE_ARCH` environment variable (`wse2` or `wse3`). The code +generator preserves the correctness conditions defined in +[Undefined Behavior](../routing#undefined-behavior). + +## Lowering to Switches + +Fabric channels correspond to hardware routing colors. Epochs enable the sequential reuse of channels +and colors across multiple communication phases. + +A stream induces a *route configuration* at each PE along its path, defined by an input port set and an +output port set (where `RAMP` denotes the local compute element). For a fixed channel $C$ and PE +$(i, j)$, ordering the active streams across successive epochs yields a sequence of configurations +$R_0, R_1, \dotsc, R_{n-1}$. This sequence maps directly to the router's hardware switch positions +for the assigned color: $R_0$ is the base configuration, and the router advances to $R_{k+1}$ at +epoch boundaries. + +Unspecified directions in a switch position retain their previous configuration. + +!!! warning "WSE-2: Single-Direction Switch Positions" + On WSE-2, each switch position may configure either an input direction or an output direction, + not both simultaneously. A transition changing both directions (such as a PE transitioning from + receiver to sender) requires two switch positions and an intermediate relay configuration. + + WSE-3 supports configuring both input and output directions within a single switch position, + requiring only one position and one advance for bidirectional transitions. + +!!! danger "Error: Switch Position Capacity Exceeded" + Routers provide up to four switch positions per color on both WSE-2 and WSE-3. If the stream + sequence on a channel requires more than four positions at any PE, compilation fails. Such + cases must be resolved by assigning additional channels to partition the traffic. + +When the sequence of configurations at a router is periodic, only one period is stored and the router +is configured with `ring_mode = true` to wrap back to the initial position. + +Consecutive configurations that are identical do not consume switch positions and require no advance. + +!!! tip "Phase Elimination for Static Configurations" + When a channel's route configuration remains constant across all epochs, its routers never + advance. In such cases, epoch boundaries enforce only relative ordering. A sequence of + identical-configuration phases can be collapsed into a single epoch containing sequential loops + (`for`) within compute blocks, lowering directly to CSL loops and eliminating phase synchronization + barriers. Examples include: + - `samples/spatial/sort/odd_even_sort_1D_looped.sptl`: $N$ odd-even rounds over four static + channels within a single CSL loop. + - `samples/spatial/sort/shearsort_2D.sptl`: 2D mesh sort using eight static channels across + nested loops without dynamic switches. Interior PEs receive on four colors within a single + epoch, requiring WSE-3 (six input queues). + + This optimization applies only to non-switching channels, as dynamic switch sequences require + compile-time epoch boundaries. + +Router advances are triggered by stream `close` operations at epoch boundaries: +- **Remote advance**: When downstream routers along the path must advance, the sender transmits a + switch-advance control wavelet along the channel using the outgoing configuration. The control + wavelet advances each traversed router after all data wavelets have cleared. +- **Local advance (WSE-2)**: When only the sender's own router must advance by a single position, + the final data wavelet triggers the transition via `.advance_switch` on the fabric output DSD, + avoiding control wavelet transmission and preventing output queue contention. +- **Local advance (WSE-3)**: WSE-3 emits a `SWITCH_ADV` control wavelet, as queue depth and + queue-to-color binding semantics accommodate explicit control wavelets. + +!!! warning "Uniform Switch Advances Along a Path" + CSL control wavelets contain command slots for router reconfiguration (`'s MAX_CMDS`). + In hardware execution and simulation, only command slot 0 is executed by traversed routers. + Every switch-configured router reached by the control wavelet applies the command in slot 0. + + Consequently, a single control wavelet cannot selectively advance a subset of routers along a + path while leaving others unchanged. All advancing routers along an active path must advance by + an identical number of positions; otherwise, a compile error is raised. Routers that have reached + their final switch position are exempt: outside `ring_mode`, advances past the final position are + no-ops. + +!!! danger "WSE-2: Turnaround Control Wavelet Propagation" + A switch-advance control wavelet terminates at the first router whose active routing configuration + targets `RAMP`. In WSE-2 bidirectional turnarounds (where a receiver transitions to become a + sender), the two required switch positions can cause the second control message to route past + the intermediate receiver to downstream PEs. + + To avoid this issue on WSE-2, communication topologies where PEs alternate sending and receiving + roles on a shared color should be partitioned into separate directional channels (e.g., + `samples/spatial/sort/odd_even_sort_1D_looped.sptl`). + +Because switch-advance control wavelets traverse the path behind the payload data, receiving PEs do +not need to emit messages to advance their routers. A receiver's `close` statement has no runtime +overhead; it serves to validate stream lifetimes and element transfer counts statically. ## Overlapping Interval Shifts -The [correctness conditions](../routing#undefined-behavior) rule out one shape that occurs constantly: -a run of consecutive PEs all shifting the same distance $d$ along an axis, as in +When multiple consecutive PEs shift data by uniform distance $d$ along an axis: ``` dataflow i16 i, i16 j in [0:D + M, 0] { @@ -138,14 +99,12 @@ dataflow i16 i, i16 j in [0:D + M, 0] { } ``` -with the PEs in `[0:M)` sending and those in `[D:D+M)` receiving. Source $p$'s word passes through -the routers of sources $p+1, \dotsc, M-1$, so the paths share PEs within one epoch. Written as one -stream per source, that is a channel each, $M$ colors for a shift; the alternative is a chain of -single-hop stores and forwards, which serializes the whole run behind $d$ hops of copying. +with sources in $[0:M)$ and destinations in $[D:D+M)$, communication paths overlap along intermediate +routers. Without optimization, this pattern requires either $M$ distinct channels or a serialized +store-and-forward chain. -Neither is necessary. The compiler recognizes this pattern — `detect_shift_bundles` — and lowers the -whole run onto **one channel**, giving the routers configurations that are switched only by events -the PE owning them knows locally: +The compiler detects this pattern (`detect_shift_bundles`) and lowers the entire shift onto **a single +channel** using coordinated router switching and destination filtering: ``` PE: 0 1 2 3 4 5 (M = 3, D = 3) @@ -156,81 +115,46 @@ pos1: -- W->E W->E -- -- -- filter: -- -- -- win 2 win 1 win 0 ``` -**Sources inject, then relay.** A source starts at `rx = RAMP, tx = {EAST}` with `pos1` taking -`rx = WEST`, sends its own words, and then advances its own router into relay mode. The trigger -is local — *"my own send is done"* — which is what makes it expressible at all, given that the -payload of a control message [selects nothing](#lowering-to-switches). On WSE-2 that advance is -`.advance_switch` on the fabric output DSD: a `SWITCH_ADV` on the same output queue would also -flip that router, but a back-pressured send of three or more `f32` values fills the six-word -queue and the control wavelet then steals a data word. On WSE-3 the same close emits `SWITCH_ADV`. - -**The order is descending, and enforces itself.** The source nearest the destinations goes first. No -schedule or barrier is needed: a source further away cannot push a word through its neighbour's -router while that neighbour is still injecting from its ramp, so it waits on the link. Backpressure -serializes the run in exactly the order the switches expect. - -**Destinations do not switch; a filter picks their words.** Each destination is statically routed to -`tx = {RAMP, EAST}`, which *duplicates* rather than consumes: every destination's router sees the -entire stream, in one order, and the one the stream reaches last uses `tx = {RAMP}` to take it out of -the network. Which words a destination hands to its compute element is decided by a counter filter -on that color, one linear function of the PE coordinate, so a single `@set_color_config` covers the -whole run. A control message passes such a router without being counted (`count_data = true`) and -without being filtered. - -!!! note "Note: Counter Filter Arithmetic" - As measured on the simulator, a counter filter starts at `init_counter`, increments on every data - wavelet, wraps to zero after `limit1`, and hands a wavelet to the compute element iff the counter - is at most `max_counter`. A window of `words` out of a stream of `length * words` is therefore - `limit1 = length * words - 1`, `max_counter = words - 1`, and an `init_counter` chosen so that - the counter reads zero as the wanted block arrives. - -!!! danger "Error: Too Many Wavelet Filters" - WSE-2 has four filters per PE, of which the `memcpy` module reserves one, so **three** are usable - (`FILTERS_PER_PE`). A PE needs one per color it filters. *If a PE would need more, a compile error - is raised*; the fix is to give some of the streams their own channels, which trades filters for - colors. Three filters therefore means at most three bundled phases per PE, whatever the kernel: - a Batcher sort would want ten at $2^4$ PEs, so `batcher_oddeven_wse3_1D.sptl` bundles only its - widest phases, which is where most of the colors are saved anyway. - - Reconfiguring a filter while wavelets are still in flight on its color is a data race, and the - destination that terminates the stream cannot be reconfigured until the stream has drained, - because its router is what removes the wavelets from the network. Filters are consequently set up - once, at layout time, and never reused between phases. - -Bundling applies only when every run it decomposes into has at least two sources and is no longer -than the shift distance, so that no PE is both a source and a destination; a shift of one PE is left -alone, since a chain at distance one is already sequenced by ordinary switch positions. Anything -else falls back to the per-hop lowering, and to the errors above if that conflicts. - -Which shifts are bundled is decided by the channel assignment rather than by an attribute: a bundle -is what several overlapping matchings on *one* channel become, so giving each matching a channel of -its own is how a kernel declines the trade. What it then costs is colors, and those can be won back -by reusing a channel across phases. Two unbundled shifts may share one safely when they agree on -axis and signed distance and their sources agree modulo twice that distance, because a PE's role — -source, relay or destination — is then a function of its position modulo twice the distance alone, -so one static configuration serves every phase in the pool. - -The agreement on the *sign* of the distance can be dropped without giving that up. Keep the axis, the -magnitude and the source residue modulo twice it, and let the direction of travel vary: the sources -are then the PEs congruent to the residue, the destinations those congruent to residue plus distance, -and the relays the classes strictly between on the one side or the other — three disjoint classes, so -a PE still holds one role on the color for the whole kernel and still never both sends and receives -on it. What varies is the side it faces, which is one switch position either way, since a source only -ever changes where it transmits and a destination only where it receives. - -Sharing on a basis looser than either does risk a PE that sends on the color in one phase and receives -on it in another, which needs a two-sided switch change that a sender cannot drive on WSE-2 (see -[Lowering to Switches](#lowering-to-switches)), and nothing in the compiler currently rejects it. - -This arrangement is the one used in Luis Schnyder's Bachelor Thesis *Distributed Sorting on the Cerebras Wafer-Scale -Engine* for the 2D reduce-scatter. - -!!! note "Note: Multiple Rounds on One Color" - Two mechanisms are deliberately left unused, and are what to reach for if the four switch - positions or three filters run out. `SWITCH_RST` restores the initial configuration of every - router a message passes, which retires a whole path with one wavelet. Teardown-based - reconfiguration reprograms the routers between rounds outright, which is the only known way to - put an arbitrary *sequence* of sends and receives on one color: for a long enough sequence there - is a PE for which no fixed cycle of switch positions exists. Until then, splitting the rounds - across channels — as `odd_even_sort_1D_looped.sptl` and `shearsort_2D.sptl` do — - remains the per-kernel fallback. +### Mechanism + +1. **Descending Transmission Order**: Sources transmit in descending order of distance to the + destination range (nearest source first). Hardware link backpressure automatically serializes + transfers without software coordination. +2. **Local Source Advance**: Each source initializes with transmission from the local ramp + (`rx = RAMP, tx = {EAST}`). Upon completing its own transmission, the source locally advances + its router to relay mode (`rx = WEST, tx = {EAST}`). On WSE-2, this advance is executed via + `.advance_switch` on the final data DSD; on WSE-3, via `SWITCH_ADV`. +3. **Static Destination Filtering**: Destination routers do not switch during the epoch. Intermediate + destinations duplicate traffic to both the local ramp and downstream neighbors (`tx = {RAMP, EAST}`), + while the terminal destination consumes the stream (`tx = {RAMP}`). Each destination isolates its + designated slice of data using a hardware counter filter. + +!!! note "Counter Filter Configuration" + A hardware counter filter initializes at `init_counter`, increments on every counted data wavelet, + and resets to zero after reaching `limit1`. A wavelet is delivered to the local compute element + if and only if `counter <= max_counter`. For a window of $W$ words within a total stream of + $L \times W$ words: + $$\text{limit1} = L \cdot W - 1, \quad \text{max\_counter} = W - 1$$ + `init_counter` is configured as an affine function of PE coordinates such that the counter + reaches zero at the arrival of the target block. + +!!! danger "Hardware Filter Capacity" + WSE-2 and WSE-3 provide four hardware filters per PE, of which one is reserved by the `memcpy` + runtime module, leaving **three** available for application kernels (`FILTERS_PER_PE`). Because + hardware filters cannot be safely reconfigured while traffic is active, filters are configured + once at layout time and cannot be reused across phases. Kernels requiring more than three filtered + phases must allocate separate channels. + +### Bundling Preconditions and Channel Reuse + +Shift bundling applies when every decomposed contiguous segment contains at least two sources and +does not exceed the shift distance ($2 \le M \le D$). Shifts of distance 1 are mapped directly to +standard switch chains. + +Unbundled shifts may safely share a channel across phases when they share axis, distance $d$, and +source coordinates modulo $2d$. Because source, destination, and relay roles form disjoint residue +classes modulo $2d$, router configurations remain invariant across all pooled phases. + +Channel sharing can be extended across opposite directions along the same axis by pooling by distance +and residue modulo $2d$. Each PE maintains a single role on the color across all phases, requiring at +most two switch positions (one for transmitting direction, one for receiving direction). diff --git a/samples/spatial/simple/exchange_bundle_1D.sptl b/samples/spatial/simple/exchange_bundle_1D.sptl index 41abc8d4..09672549 100644 --- a/samples/spatial/simple/exchange_bundle_1D.sptl +++ b/samples/spatial/simple/exchange_bundle_1D.sptl @@ -1,17 +1,15 @@ /** - * Two overlapping 1D interval shifts, in opposite directions, repeated R times. + * Repeated pairwise exchange between two blocks of PEs. * - * D + M PEs on a line. The PEs in [0:M) and those in [D:D+M) swap values pairwise: PE i - * trades with PE i + D. Each direction is a shift bundle of its own on a color of its own, - * so one exchange costs two colors however many pairs there are. M <= D, so the two halves - * are disjoint. + * Given D + M PEs on a line, PEs [0:M) swap values with PEs [D:D+M) (PE i swaps + * with PE i + D). The exchange repeats over R phases. * - * This is the skeleton of a sorting network's phase (see the bundled Batcher), reduced to - * one operation: every PE is a source on one color and a filtered destination on the other. - * R repeats it in R phases, none of which shares a color with another, so every PE ends up - * with R wavelet filters and R is what the per-PE filter budget limits. + * Each phase uses two channels (one east, one west) bundled with counter filters + * at the receivers. Because hardware filters cannot be reconfigured during a run, + * each repeat consumes one filter per PE, limiting R to the per-PE filter limit + * (R <= 3 on WSE-2 and WSE-3). * - * Constraints: 2 <= M <= D, R >= 1. + * Constraints: 2 <= M <= D, 1 <= R <= 3. **/ kernel @exchange_bundle_1d( stream[D + M, 1] readonly inp, diff --git a/samples/spatial/sort/batcher_oddeven_wse3_1D.sptl b/samples/spatial/sort/batcher_oddeven_wse3_1D.sptl index 46eb4ce0..f14fcdbc 100644 --- a/samples/spatial/sort/batcher_oddeven_wse3_1D.sptl +++ b/samples/spatial/sort/batcher_oddeven_wse3_1D.sptl @@ -1,87 +1,56 @@ /** - * 1D Batcher odd-even mergesort over R independent rows of N = 2^L PEs, each holding K f32 keys. - * Ascending: after the network, PE (i, j) holds keys i*K .. i*K + (K-1) of row j. Every hop is - * east-west, so the R rows never mix. - * - * Same network, same compare-split arithmetic and same result as batcher_oddeven_bundled_1D. What - * differs is how the unbundled phases pool their channels, and it exists because that is what stands - * between the bundled variant and L = 4 on WSE-3. - * - * Why the bundled variant stops short on WSE-3 - * ------------------------------------------- - * WSE-3 binds a fabric queue to a color for the whole kernel: a queue may not be remapped onto - * another color, because nothing can prove the wavelets of the old one have drained. A PE therefore - * needs one input queue per color it *ever* receives on and one output queue per color it ever sends - * on, and it has six of each -- five for sending, since memcpy's copy-back of `out` reserves one. - * The bundled variant's busiest PE at L = 4 wants seven of each: - * - * 2 d = 1, one color for the p = 1 phase and one for the p >= 2 phases - * 2 d = 2, likewise - * 3 one per bundled phase: (l=3,p=1), (l=4,p=1), (l=4,p=2) - * - * Bundling cannot take that further. It saves *colors*, and a PE's three filters cap it at three - * phases, but a bundled color cannot be reused by a later phase either -- its sources end in relay - * mode and its destination filters are set at layout time and never reprogrammed -- so each bundled - * phase costs its participants a queue in each direction, exactly as an unbundled one does. - * - * Pooling by origin instead of by direction - * ----------------------------------------- - * What costs two queues per unbundled distance is pooling the forward and backward streams - * separately. Two phases may share a channel there when they agree on direction, distance and source - * residue mod 2d, which keeps every router's configuration fixed for the whole kernel but makes a - * comparator's two messages two colors. - * - * Drop the agreement on direction and keep the rest: a color becomes (distance d, origin residue - * mod 2d), whichever way the message travels. On such a color the three roles fall into disjoint - * residue classes -- sources are the PEs congruent to s, destinations those congruent to s + d, and - * relays the classes strictly between, on one side or the other -- so a PE holds one role on a color - * for the whole kernel and never both sends and receives on it. Only the side it faces varies: - * - * source R->E or R->W two positions, the transmit side alone changes - * destination W->R or E->R two positions, the receive side alone changes - * relay W->E or E->W one position; a relay class is reached from one side only + * Batcher odd-even mergesort on R independent rows of N = 2^L PEs. + * + * Each PE holds a block of K f32 keys. After the kernel, every row is sorted in ascending order: + * PE (i, j) holds keys i*K .. i*K + (K-1) of row j. All communication is along x, so rows do not + * interact. + * + * Algorithm + * --------- + * The first phase loads each PE's block and sorts it locally. The network then runs merge levels + * l = 1 .. L, each consisting of sub-phases p = 1 .. l that compare PEs at distance d = 2^(l-p). + * In sub-phase p = 1, PE i is paired with PE i + d within each box of 2^l PEs. In sub-phases + * p >= 2, the first and last d PEs of each box do not participate. Every comparator is a + * compare-split: the partners exchange their blocks, the lower PE keeps the K smallest and the + * higher PE the K largest of the 2K keys. + * + * Channel assignment + * ------------------ + * Long-distance phases (4*d >= N) have overlapping paths and are bundled onto one channel per + * direction using hardware counter filters (see irspec/docs/spatial/routing_wse.md). For L <= 4, + * there are at most three such phases, matching the three available hardware filters per PE. * - * Two positions of the four a router holds, and never a both-sides change, so this asks for nothing - * WSE-2 lacks either. Every router on a retired path advances once per epoch boundary, and the ones - * with a single position sit on their last, where a control message passing through over-advances - * them harmlessly. Each PE now spends one input and one output queue per unbundled distance instead - * of two, which is what brings L = 4 within budget: + * bundled (4*d >= N): fwd = N + 2*(l*(L+1) + p), bwd = fwd + 1 * - * 5 in and 5 out at L = 4, against 6 and 5. + * The remaining short-distance phases pool channels by distance d and the residue (coordinate mod 2d), + * regardless of send direction: * - * Channels - * -------- - * Bundling still pays where the matchings overlap, and the rule that picks the phases is unchanged: - * 4*d >= N, the three widest, one per filter. The rest pool by origin. + * pooled (4*d < N): fwd = 2*d - 2 + c_lo, bwd = 2*d - 2 + c_hi * - * bundled (4*d >= N) : fwd N + 2*(l*(L+1) + p), bwd fwd + 1 - * pooled : fwd (2*d - 2) + c_lo, bwd (2*d - 2) + c_hi + * where c_lo and c_hi are the residues modulo 2d of the lower and the higher partner. For p = 1, + * c_lo = r and c_hi = r + d; for p >= 2 the two are swapped, so all phases at one distance use the + * same 2d channels, numbered 2d - 2 .. 4d - 3. On such a channel, sources, destinations and relays + * form disjoint residue classes. Each PE therefore keeps one role on a channel for the whole kernel, + * and its router alternates between at most two switch positions: * - * where c_lo is the low partner's residue mod 2d and c_hi the high partner's -- c_lo = r and - * c_hi = r + d for p = 1, the other way round for p >= 2, which is where the two phases at one - * distance meet. The block for d runs from 2*d-2 to 4*d-3 and pooled phases have 4*d < N, so every - * pooled channel stays below N. + * source R->E or R->W + * destination W->R or E->R + * relay W->E or E->W * - * Colors: 8 at L = 3 and 12 at L = 4, against 10 and 18 for the bundled variant. All of them switch, - * and WSE-3 implements switches on fifteen of its twenty-one colors, which is the next thing L = 5 - * would run out of -- along with a sixth output queue. + * Channels per phase for L = 3: * - * Requires WSE-3 (!) - * ------------------ - * Not for the routing, but for the queues: this kernel is written for a target that reuses none. On - * WSE-2 a queue comes free again once its occupancy span ends, so what binds there is how many colors - * are live at once rather than how many a PE ever touches, and the bundled variant, whose routers - * never switch at all, is the better fit. + * (l=1,p=1) d=1: (0,1)(2,3)(4,5)(6,7) fwd 0 bwd 1 + * (l=2,p=1) d=2: (0,2)(1,3)(4,6)(5,7) bundled + * (l=2,p=2) d=1: (1,2)(5,6) fwd 1 bwd 0 + * (l=3,p=1) d=4: (0,4)(1,5)(2,6)(3,7) bundled + * (l=3,p=2) d=2: (2,4)(3,5) bundled + * (l=3,p=3) d=1: (1,2)(3,4)(5,6) fwd 1 bwd 0 * - * Example L=3 (n=8), with the channel each phase's two directions land on: - * (l=1,p=1) dist 1: (0,1)(2,3)(4,5)(6,7) fwd 0 bwd 1 - * (l=2,p=1) dist 2: (0,2)(1,3)(4,6)(5,7) bundled - * (l=2,p=2) dist 1: (1,2)(5,6) fwd 1 bwd 0 - * (l=3,p=1) dist 4: (0,4)(1,5)(2,6)(3,7) bundled - * (l=3,p=2) dist 2: (2,4)(3,5) bundled - * (l=3,p=3) dist 1: (1,2)(3,4)(5,6) fwd 1 bwd 0 + * The kernel uses 8 colors for L = 3 and 12 colors for L = 4. * - * Constraints: 1 <= L <= 4, K >= 1, R >= 1. WSE-3. + * Constraints: 1 <= L <= 4, K >= 1, R >= 1. Requires WSE-3, which binds each fabric queue to one + * color for the whole kernel; the pooling above keeps the busiest PE at five input and five output + * queues for L = 4. **/ kernel @batcher_oddeven_wse3_1d( stream[1<= 1, K >= 1. WSE-2 and WSE-3. - * **/ kernel @odd_even_sort_1d_looped(stream[1<[1<= 1, K >= 1. WSE-3. + * Routing and channels + * -------------------- + * The kernel uses 8 static channels, one per (axis, round parity, direction): + * - Channels 0-3: row exchanges (even/odd round, east/west). + * - Channels 4-7: column exchanges (even/odd round, south/north). + * Because roles on each channel remain fixed throughout the entire execution, no routers switch + * and no channel remapping is needed. The rounds are executed as sequential loops within compute + * blocks rather than separate compiler phases, avoiding barrier overhead. * + * Constraints: L >= 1, K >= 1. Requires WSE-3, as interior PEs receive on 4 concurrent inbound + * channels in a single epoch. **/ kernel @shearsort_2d_looped(stream[1<[1<(stream[1< list[spir.Identifier]: - """ - Fabric transfers in ``statement``, in source order, with sequential ``for`` bodies repeated. + """Collect fabric transfers in ``statement`` in source order, duplicating sequential ``for`` bodies. - Walking a loop body once makes its colours look sequential, so occupancy pooling would give - them one queue. The next iteration of an earlier colour can already occupy the router when a - later colour of the same body remaps that queue -- WSE-2 then aborts with "Attempt to remap - input queue N, from C_i to C_j, but the router is holding wavelets". Appending the body a - second time makes a colour used on both sides of another occupy a span that overlaps it, the - same rule that keeps a reused colour's queue across a gap between unrolled phases. + Repeating loop bodies captures loop-carried reuse in occupancy intervals, ensuring that + a color used across iterations cannot be prematurely remapped to another queue while + in-flight wavelets remain. :param statement: The statement to walk. :param names: Streams that bind a fabric queue in this direction. :param inbound: True to collect receives, False to collect sends. - :return: The transferred streams, with each ``for`` body listed twice. + :return: Transferred stream identifiers with each ``for`` body duplicated. """ if isinstance(statement, spir.ForStatement): body = transfer_points(statement.body, names, inbound) @@ -72,19 +66,13 @@ def transfer_points( def queue_spans( compute: spir.ComputeBlock, names: set[spir.Identifier], queue_key, inbound: bool ) -> dict[str, tuple[int, int]]: - """ - Occupancy of each queue key along the linearized send/receive order of this PE. - - Sequential ``for`` bodies are counted twice so a colour that comes back on the next iteration - keeps its queue across the loop-carried gap; see ``statement_transfer_points``. Uses of the - same channel still collapse to one span, so a colour that comes back after a gap between - unrolled phases keeps its queue for the whole of that span. + """Compute the occupancy span (first_use, last_use) of each queue key along the transfer order. :param compute: The compute block being lowered. - :param names: Streams that actually bind a fabric queue in this direction. - :param queue_key: Maps a stream identifier to its grouping key (channel, or the name itself). - :param inbound: True to walk receives, False to walk sends. - :return: Mapping of grouping key to ``(first_use, last_use)`` in linearized order. + :param names: Streams that bind a fabric queue in this direction. + :param queue_key: Function mapping a stream identifier to its grouping key. + :param inbound: True to inspect receives, False for sends. + :return: Map from grouping key to inclusive ``(first_use, last_use)`` indices. """ points = transfer_points(compute.statements, names, inbound) diff --git a/spada/lowering/spatial_ir_to_csl.py b/spada/lowering/spatial_ir_to_csl.py index 17b2f7c7..6629082e 100644 --- a/spada/lowering/spatial_ir_to_csl.py +++ b/spada/lowering/spatial_ir_to_csl.py @@ -970,14 +970,12 @@ def _collect_unique_dsds( array_candidates[place_statement.field_name.as_ir()] = (place_statement, place_statement.dtype.shape) - # Find used DSDs in compute block - # Streams that share a channel share a color, and a color binds to exactly one fabric queue per - # PE -- the hardware rejects "two master input queues for the same color". Queues are therefore - # handed out per channel; streams on ``auto`` channels get a color to themselves, so they key on - # their own name. Sequential channels may share a queue only when their occupancy spans on this - # PE do not overlap; see ``stream_lifetime.assign_fabric_queues``. On WSE-3 every inbound color - # keeps its own queue: remapping one that still holds wavelets is a fatal error, and a data-task - # ID is that queue. + # Find used DSDs in compute block. + # Streams sharing a channel share a hardware color. A color binds to at most one fabric queue + # per PE. Queues are assigned per channel; streams with auto channels receive unique colors. + # Channels may share queues across disjoint occupancy intervals (see assign_fabric_queues). + # On WSE-3, inbound colors receive dedicated queues because data task IDs correspond directly + # to fabric input queue IDs. channel_of_stream = { declaration.stream_name.as_ir(): declaration.stream.routing.resolved_channel for declaration in rect.dataflow.statements @@ -992,10 +990,7 @@ def queue_key(stream: spir.Identifier) -> str: output_names = fabric_occupancy.streams_with_fabric_dsds(rect.compute, memcpy_mode, stream_args, inbound=False) input_spans = fabric_occupancy.queue_spans(rect.compute, input_names, queue_key, inbound=True) output_spans = fabric_occupancy.queue_spans(rect.compute, output_names, queue_key, inbound=False) - # WSE-3 remaps a fabric queue onto the next color at the first transfer that uses it, and - # faults or stalls if the queue still holds wavelets. Occupancy in the compute block is not - # enough to prove it is empty, so every color keeps its own queue. That also keeps data-task - # IDs unique, since those IDs *are* the input queues. + # On WSE-3, inbound and outbound colors are given dedicated queues to avoid runtime remapping. exclusive = frozenset(input_spans) if csl.ARCH == 'wse3' else frozenset() exclusive_out = frozenset(output_spans) if csl.ARCH == 'wse3' else frozenset() input_queue_of = stream_lifetime.assign_fabric_queues( @@ -1006,11 +1001,9 @@ def queue_key(stream: spir.Identifier) -> str: output_spans, csl.OUTPUT_QUEUE_IDS, kind='output', architecture=csl.ARCH, location=location, exclusive_keys=exclusive_out) - # An asynchronous transfer runs on a microthread, and by default that is the queue ID of the - # operation's highest-priority fabric operand. Since the two directions draw from overlapping - # pools on WSE-3, a receive on input queue N and a send on output queue N would take the same - # microthread and abort with "trying to term ut_instr[N], but it's not ours". Microthreads are - # one resource across both directions, so they are handed out together. + # By default, an asynchronous DSD operation uses the queue ID of its highest-priority fabric + # operand as its microthread ID. On WSE-3, input and output queue IDs overlap, so microthreads + # are assigned explicitly to prevent collisions between concurrent sends and receives. microthread_of = stream_lifetime.assign_microthreads( fabric_occupancy.microthread_intervals(rect.compute, input_names, output_names, queue_key), csl.MICROTHREAD_IDS, location=location) diff --git a/spada/lowering/wse3.py b/spada/lowering/wse3.py index a3ee9ba5..a0ecc1d4 100644 --- a/spada/lowering/wse3.py +++ b/spada/lowering/wse3.py @@ -176,9 +176,8 @@ def data_task_id_builtin( ) -> str: """Return the ``@get_data_task_id(...)`` expression for ``slot``. - WSE-2 constructs a data-task ID from the color the receive listens on. - WSE-3 constructs it from the input queue already bound to that color; - passing the color is rejected as ``expected 'input_queue' expression, got: 'color'``. + On WSE-2, data task IDs are constructed from the color. + On WSE-3, data task IDs are constructed from the bound input queue. :param rect: The PE block being generated. :param slot: The data-task slot, whose color is the receive's fabric color. diff --git a/spada/runtime/runtime.py b/spada/runtime/runtime.py index ce9615c4..25d0b1fb 100644 --- a/spada/runtime/runtime.py +++ b/spada/runtime/runtime.py @@ -95,11 +95,10 @@ def from_json(cls, json_data: Union[str, Dict[str, Any]]) -> "ProgramMetadata": def memcpy_data_type(dtype: np.dtype) -> "crt.MemcpyDataType": - """ - Pick the transfer width for a kernel argument of ``dtype``. + """Return the transfer width for a kernel argument of ``dtype``. - :param dtype: The dtype the kernel declared for the argument. - :return: The ``MemcpyDataType`` to pass alongside the buffer. + :param dtype: The declared argument dtype. + :return: The corresponding ``MemcpyDataType``. """ if dtype.itemsize == 4: return crt.MemcpyDataType.MEMCPY_32BIT @@ -110,28 +109,24 @@ def memcpy_data_type(dtype: np.dtype) -> "crt.MemcpyDataType": def memcpy_word_dtype(dtype: np.dtype) -> np.dtype: - """ - Give the host-buffer dtype for a kernel argument of ``dtype``, one element per 32-bit word. + """Return the host buffer dtype for a kernel argument of ``dtype``, aligned to 32-bit words. - ``memcpy_h2d`` and ``memcpy_d2h`` reject a buffer whose elements are not 32 bits ("Internal - data type of any memcpy_d2h() or memcpy_h2d() operation should be 32 bit") even when the - device-side array is 16-bit: ``MEMCPY_16BIT`` means only the low half of each word travels. + Cerebras SDK memcpy operations require 32-bit word alignment on the host even for 16-bit + transfers (MEMCPY_16BIT transfers the lower 16 bits of each 32-bit word). - :param dtype: The dtype the kernel declared for the argument. - :return: ``dtype`` itself when it is already 32 bits wide, else a 32-bit word dtype. + :param dtype: The declared argument dtype. + :return: The original dtype if 32 bits wide, otherwise uint32. """ return dtype if dtype.itemsize == 4 else np.dtype(np.uint32) def as_memcpy_words(data: np.ndarray) -> np.ndarray: - """ - Widen a 16-bit array into the 32-bit words ``memcpy_h2d`` expects. + """Widen a 16-bit array into 32-bit words expected by host memcpy. - The widening is bit-for-bit rather than by value, so that a negative ``i16`` and an ``f16`` - both arrive on the device unchanged. + 16-bit values are bitcast without sign extension to preserve exact binary representations. - :param data: A contiguous array in the dtype the kernel declared. - :return: ``data`` itself when it is already 32 bits wide, else a widened copy. + :param data: Contiguous input array. + :return: The array itself if 32 bits wide, otherwise a widened copy. """ if data.dtype.itemsize == 4: return data @@ -139,12 +134,11 @@ def as_memcpy_words(data: np.ndarray) -> np.ndarray: def from_memcpy_words(words: np.ndarray, dtype: np.dtype) -> np.ndarray: - """ - Undo :func:`as_memcpy_words` for data copied back from the device. + """Narrow 32-bit memcpy words back to the target output dtype. - :param words: The buffer ``memcpy_d2h`` filled. - :param dtype: The dtype the kernel declared for the output. - :return: ``words`` reinterpreted in ``dtype``, keeping the shape. + :param words: Buffer populated by ``memcpy_d2h``. + :param dtype: Target output dtype. + :return: Array reinterpreted in ``dtype``. """ if dtype.itemsize == 4: return words @@ -345,8 +339,7 @@ def __init__( cmaddr = cm_addr or os.environ.get("CM_ADDR", None) self.simulator = cmaddr is None print("SIMULATOR?", self.simulator) - # Fabric traces are large, so they are only written when asked for: they are what tells a - # stalled or faulting simulator run apart, per tile and per color. + # Enable fabric tracing when SPADA_SIMFAB_TRACE is set (useful for debugging simulator stalls). trace = os.environ.get("SPADA_SIMFAB_TRACE") is not None self.runtime = crt.SdkRuntime(str(self.out_folder), suppress_simfab_trace=not trace, cmaddr=cmaddr) diff --git a/spada/syntax/csl/constants.py b/spada/syntax/csl/constants.py index 0457c481..2163e68c 100644 --- a/spada/syntax/csl/constants.py +++ b/spada/syntax/csl/constants.py @@ -48,29 +48,25 @@ } MEMCPY_COLORS = _MEMCPY_COLORS[ARCH] -# See https://sdk.cerebras.net/csl/language/dsds#fabric-queues +# Fabric queue IDs available for application communication. +# See https://sdk.cerebras.ai/csl/language/dsds#fabric-queues _INPUT_QUEUE_IDS = { - 'wse2': list(range(0, 2)), # Ignoring 2-7 as they are smaller in capacity - # On WSE-3 a data task's ID *is* its input queue, and memcpy takes 0 and 1 for its own; binding - # either of them with ``@initialize_queue`` is rejected as "already been set". - 'wse3': list(range(2, 8)), + 'wse2': list(range(0, 2)), # Queues 0-1 provide full capacity on WSE-2 + 'wse3': list(range(2, 8)), # Queues 2-7; queues 0-1 are reserved by memcpy } INPUT_QUEUE_IDS = _INPUT_QUEUE_IDS[ARCH] _OUTPUT_QUEUE_IDS = { - 'wse2': list(range(2, 4)), # Ignoring 0-1,4-5 as they are smaller in capacity - 'wse3': list(range(2, 8)), # All queues are equivalent, but memcpy reserves 0 and 1 + 'wse2': list(range(2, 4)), # Queues 2-3 provide full capacity on WSE-2 + 'wse3': list(range(2, 8)), # Queues 2-7; queues 0-1 are reserved by memcpy } OUTPUT_QUEUE_IDS = _OUTPUT_QUEUE_IDS[ARCH] -# Microthreads that drive in-flight asynchronous DSD operations. Two operations may never run on -# one microthread at the same time. WSE-2 has no say in the matter: the ID is the queue ID of the -# operation's highest-priority fabric operand, which is why its input and output pools above are -# disjoint. WSE-3 keeps that default but lets ``.ut_id`` override it, which it must, since a PE -# there needs an input and an output queue of the same number at once (see -# https://sdk.cerebras.net/csl/language/microthreads_wse3). An empty list means the target cannot -# name microthreads, so the default stands. Queues 0 and 1 belong to memcpy on WSE-3, and so do the -# microthreads it drives them with. +# Hardware microthread IDs for asynchronous DSD operations. +# On WSE-2, the microthread ID is implicitly tied to the queue ID of the highest-priority +# fabric operand. On WSE-3, microthreads 2-7 can be explicitly assigned via the `.ut_id` +# DSD field (queues 0-1 and their corresponding microthreads are reserved by memcpy). +# See https://sdk.cerebras.ai/csl/language/microthreads_wse3 _MICROTHREAD_IDS = { 'wse2': [], 'wse3': list(range(2, 8)), @@ -88,16 +84,10 @@ # See https://sdk.cerebras.ai/csl/language/builtins#switching-configuration-semantics SWITCH_POSITIONS = 4 -# Number of switching command slots a control wavelet carries (````'s MAX_CMDS). -# -# NOTE: only slot 0 is ever executed. Measured on the simulator, every switch-configured router a -# wavelet reaches applies the command in slot 0; slots 1-7 had no effect in any topology tested -# (the sender's own router, one hop, two hops through a plain relay, and two switch-configured -# routers in sequence). A wavelet therefore cannot advance one router while skipping another on its -# path, which is why ``routing.plan_switch_advances`` requires the routers along a path to agree. -# ````'s ``encode_payload`` also loops over all eight slots regardless of the array length -# it is given, so it must be passed exactly eight; ``encode_single_payload`` writes slot 0 only and -# is what the compiler emits. +# Number of switching command slots in a CSL control wavelet ('s MAX_CMDS). +# Hardware routers execute command slot 0 across all traversed switch-configured routers; +# remaining slots are ignored. Control messages therefore advance all routers along their +# path uniformly, using encode_single_payload. MAX_CONTROL_COMMANDS = 8 # Colors whose routers support switches. WSE-3 only implements switches on a subset of colors. @@ -113,10 +103,9 @@ _SWITCH_POSITION_ALLOWS_BOTH = {'wse2': False, 'wse3': True} SWITCH_POSITION_ALLOWS_BOTH = _SWITCH_POSITION_ALLOWS_BOTH[ARCH] -# Router filters usable per PE, over all colors. The hardware has four, but the memcpy module -# reserves one, so a program may configure three (Schnyder, "Distributed Sorting on the Cerebras -# Wafer-Scale Engine", ch. 7). -# -# NOTE: a filter must not be reconfigured while wavelets it counts are still in flight. Filters are -# therefore set once in the layout and never rewritten between phases. +# Maximum hardware counter filters configurable per PE across all colors. +# While the hardware provides four filters per PE, the memcpy runtime module reserves one, +# leaving three available for application routing. +# Hardware filters cannot be safely reconfigured while traffic is active, so they are +# initialized at layout time. FILTERS_PER_PE = 3 diff --git a/spada/syntax/csl/routing.py b/spada/syntax/csl/routing.py index 9537c7e1..94311776 100644 --- a/spada/syntax/csl/routing.py +++ b/spada/syntax/csl/routing.py @@ -67,23 +67,18 @@ def changes_both_sides(self, previous: 'RouteConfig') -> bool: @dataclass(frozen=True) class FilterConfig: - """ - A counter filter: which of the wavelets passing a router are handed to its compute element. - - The router keeps a counter per filter. It starts at ``init_counter``, advances on every wavelet - the filter counts, and wraps to zero after ``limit1``, so it cycles through ``limit1 + 1`` - values. A wavelet is delivered iff the counter is at most ``max_counter``, and *withheld* - otherwise -- withheld is not the same as consumed: the wavelet carries on along the router's - ``tx`` directions, so PEs further along still see it. Only a router that transmits to the ramp - alone drops what it withholds, which is what takes a wavelet out of the network. - (Measured on the simulator. The manual describes - ``max_counter`` as exclusive, but a wavelet arriving at ``counter == max_counter`` is delivered.) + """Hardware counter filter configuration for a router color. - The fields are expression strings rather than integers because a filter's window generally - depends on where the PE sits: a shift bundle's receivers share one ``@set_color_config`` whose - ``init_counter`` is a function of ``pe_x``. + A counter filter monitors incoming data wavelets on a color. The counter starts at + ``init_counter``, increments on each counted wavelet, and wraps to zero after reaching + ``limit1``. A wavelet is delivered to the local compute element (RAMP) when + ``counter <= max_counter``; otherwise, it is forwarded along the router's transmit + directions without delivery. - A PE can hold only ``constants.FILTERS_PER_PE`` of these across all of its colors. + Attributes: + init_counter: Expression string for the initial counter value. + limit1: Expression string for the counter wrap limit. + max_counter: Expression string for the maximum inclusive delivery threshold. """ init_counter: str limit1: str @@ -129,7 +124,7 @@ def expand_positions(configs: list[RouteConfig], ring: bool = False) -> tuple[li ``constants.SWITCH_POSITION_ALLOWS_BOTH``) keep the transition as a single position. The intermediate keeps the *old* input direction, so a switch-advance wavelet arriving from the - same neighbour as before is still accepted once the router has taken the intermediate position; + same neighbor as before is still accepted once the router has taken the intermediate position; the second wavelet would never reach the router otherwise. :param configs: The logical configurations, in the order the router takes them. @@ -397,16 +392,18 @@ def _bundle_ports(bundle: shift_bundles.ShiftBundle) -> tuple[str, str]: def _window_start(variable: str, first: int, step: int, words: int) -> str: - """ - Returns the counter value a destination's filter starts at, as an expression in the loop variable. + """Compute the initial counter value for a destination filter as an affine expression. - Every destination sees the whole stream, in one order, so which words a destination keeps is - decided by where it sits: the ``p``-th destination along the direction of travel keeps the block - the ``p``-th-from-last source sent, which begins ``(p + 1) * words`` short of the end of the - cycle. Starting the counter there brings it to zero just as that block arrives. + Each destination observes the aggregated stream. The p-th destination along the travel + direction receives the block from the p-th from last source, starting (p + 1) * words + before the end of the count cycle. Initializing the counter at this offset aligns the zero + point with the arrival of the destination's block. - :param first: The coordinate of the destination the stream reaches first, where ``p`` is zero. - :param step: ``+1`` or ``-1``, the direction the coordinate grows in as ``p`` grows. + :param variable: Coordinate variable ('pe_x' or 'pe_y'). + :param first: Coordinate of the first destination reached by the stream. + :param step: Direction of coordinate progression (+1 or -1). + :param words: Word count per source transfer. + :return: Expression string for init_counter. """ offset = 1 - first if step > 0 else first + 1 if step > 0: @@ -420,21 +417,19 @@ def _window_start(variable: str, first: int, step: int, words: int) -> str: def _bundle_route_entries(bundle: shift_bundles.ShiftBundle, color: int, order: tuple[int, int, int], rect_index: int, stream_name: spir.Identifier) -> list[tuple['_RouteSite', '_RouteEntry']]: - """ - Returns the route configurations of one shift bundle, replacing the per-hop ones. - - Four kinds of router take part, and none of them is a compute rectangle shifted as a whole -- the - relays in between are as many as the shift distance minus the run length, which is neither half's - width -- so every site here is a standalone one. - - The sources all get the same pair of configurations, injecting and then relaying, including the - one furthest from the destinations which has nothing to relay for. Giving it a switch position it - never uses costs nothing and keeps one ``@set_color_config`` for the whole run; its own - switch-advance wavelet moves it into a configuration that never carries anything, and the - wavelets of the sources behind it pass routers already sitting on their last position, which is a - no-op. - - :param order: The switch-position order key of the send that this bundle carries. + """Generate route configurations for a shift bundle. + + Sources configure an initial injection route followed by a relay switch position. + Intermediate PEs configure plain relay routes. Destinations configure static duplicate + routes (RAMP and forward) with associated counter filter configurations, while the + final destination configures a terminal RAMP route. + + :param bundle: The shift bundle descriptor. + :param color: The hardware color assigned to this bundle. + :param order: Total ordering key for switch position resolution. + :param rect_index: Source PE rectangle index. + :param stream_name: Spatial IR stream identifier. + :return: List of (site, entry) pairs for layout emission. """ incoming, outgoing = _bundle_ports(bundle) variable = 'pe_x' if bundle.axis == 'x' else 'pe_y' diff --git a/spada/syntax/csl/task_recycling.py b/spada/syntax/csl/task_recycling.py index e57742c6..2b278edf 100644 --- a/spada/syntax/csl/task_recycling.py +++ b/spada/syntax/csl/task_recycling.py @@ -363,7 +363,7 @@ def plan_task_bindings( When recycling is needed, all tasks are colored together using load-balanced greedy coloring in degeneracy order, distributing tasks - evenly across hardware slots to minimise dispatcher state machine size. + evenly across hardware slots to minimize dispatcher state machine size. :param data_task_colors: The color each data task listens on, keyed by task index. Data tasks are grouped by it unconditionally; @@ -381,21 +381,19 @@ def plan_data_task_slots( tasks: list[tdag.CSLTask], data_task_colors: dict[int, int], ) -> tuple[tuple[DataTaskSlot, ...], dict[int, int], dict[int, int]]: - """Group the data tasks by the color they listen on. - - Sharing a color is sound only if the receives take it in turns, which is the - same criterion local slots use: every trigger source of the later task must - be reachable from the earlier one. That much orders the *installation* of the - later branch, but not the arrival of its wavelets, which the fabric may - deliver while the earlier branch is still installed. Codegen closes that gap - by having each branch of a recycled slot ``@block`` its own color once its - last wavelet has arrived, so wavelets of the next epoch wait in the queue - until their branch is installed and unblocked. - - :param data_task_colors: The color each data task listens on, keyed by task index. - :return: ``(slots, task_to_slot, task_to_state)``, the last two mapping a task - index to its slot number and to its state within that slot. - :raises SyntaxError: If two data tasks share a color without being ordered. + """Group logical data tasks by the hardware color they receive on. + + Multiple data tasks may share a hardware color slot sequentially if every + trigger source of a subsequent task is reachable from the preceding task in + the task dependency graph. To prevent incoming wavelets of a subsequent epoch + from being processed prematurely, each branch in a shared slot blocks its + associated color upon receiving its final expected wavelet. The color is + subsequently unblocked when the corresponding logical task is activated. + + :param tasks: List of CSL tasks for a PE. + :param data_task_colors: Map from data task index to hardware color. + :return: A tuple of (slots, task_to_slot, task_to_state). + :raises SyntaxError: If unordered data tasks share a color. """ by_color: dict[int, list[int]] = {} for task_index, task in enumerate(tasks): diff --git a/spada/syntax/spatial_ir/canonicalization.py b/spada/syntax/spatial_ir/canonicalization.py index a89b3443..48e51531 100644 --- a/spada/syntax/spatial_ir/canonicalization.py +++ b/spada/syntax/spatial_ir/canonicalization.py @@ -421,12 +421,11 @@ def inline_phases(kernel: spir.Kernel) -> spir.Kernel: raise TypeError(f'Unexpected block type "{type(block).__name__}" in kernel. Was ``canonicalize_phases`` ' 'called?') - new_kernel = spir.Kernel( + return spir.Kernel( name=kernel.name, parameters=copy.deepcopy(kernel.parameters), arguments=copy.deepcopy(kernel.arguments), body=list(rect_place.values()) + list(rect_dataflow.values()) + list(rect_compute.values())) - return new_kernel @dataclass diff --git a/spada/syntax/spatial_ir/shift_bundles.py b/spada/syntax/spatial_ir/shift_bundles.py index 1cc248ab..d492b1d3 100644 --- a/spada/syntax/spatial_ir/shift_bundles.py +++ b/spada/syntax/spatial_ir/shift_bundles.py @@ -1,24 +1,17 @@ -""" -Overlapping 1D interval shifts on one color. - -A run of consecutive PEs each shifting the same distance ``d`` along an axis has overlapping paths: -the router of source ``p + 1`` carries source ``p``'s words. One color per direction still suffices, -because the routers can be time-multiplexed -- but only if every switch is triggered by something -the PE that owns it knows locally, since a control wavelet advances *every* switch-configured router -it reaches (see ``irspec/docs/spatial/routing_wse.md``). - -Send order is what provides that. The sources go nearest-the-destinations first, so a source's -router changes from injecting to relaying exactly when that source has finished its own send, which -it can signal itself with one switch-advance wavelet. Nothing needs to be told from a distance, and -the order enforces itself: a source further from the destinations cannot push a word through its -neighbour's router while that neighbour is still injecting, so it waits on the link. - -The destinations do not switch at all. Each transmits to its ramp *and* onward, so all of them see -the whole stream and a counter filter decides which words each one keeps; the last one transmits to -its ramp alone and thereby takes the stream out of the network. - -This is the arrangement Schnyder's 2D reduce-scatter uses ("Distributed Sorting on the Cerebras -Wafer-Scale Engine", fig. 7.6). +"""1D interval shift bundling for multiplexing overlapping communication paths onto one color. + +When consecutive PEs execute a uniform relative shift along an axis (e.g., each PE +in [0:M) sending to PE i + d), their transmission paths overlap across intermediate routers. +This pattern can be multiplexed onto a single fabric channel using hardware switch advances +and destination counter filters: + +1. Sources transmit in descending order of distance to destinations (nearest destination first). + Link-level backpressure naturally serializes transfers without software coordination. +2. After transmitting its elements, each source router locally advances from injection mode + to relay mode. +3. Destination routers statically forward wavelets to both the local ramp and downstream neighbors + (or to the ramp only for the final destination). Hardware counter filters at each destination + select the designated slice of data. """ from __future__ import annotations @@ -151,19 +144,14 @@ def _consecutive_runs(values: set[int]) -> list[tuple[int, int]]: def detect_shift_bundles(rectangles: list[Rectangle[PEBlock]]) -> list[ShiftBundle]: - """ - Finds the interval shifts in a kernel whose paths overlap, and which therefore need bundling. - - Sources are collected per channel and shift, across rectangles: one logical shift is often - declared by several compute blocks -- a sorting network's matchings at successive offsets, for - instance -- and only their union shows which PEs form a consecutive run. + """Detect interval shifts with overlapping router paths suitable for bundling. - A shift is bundled only if *every* run it decomposes into can be: at least two sources, and no - longer than the shift distance, so that sources and destinations stay disjoint. A shift that - fails this is left to the ordinary per-hop lowering, which reports the conflict if there is one. + Collects sources across PE blocks. A shift qualifies for bundling if all + decomposed contiguous segments contain at least two sources and do not exceed + the shift distance, ensuring source and destination intervals remain disjoint. - :param rectangles: The consolidated PE rectangles of the kernel, with channels already resolved. - :return: The bundles, in a deterministic order. + :param rectangles: Consolidated PE blocks of the kernel with resolved channels. + :return: A list of detected ShiftBundle descriptors. """ # (channel, axis, signed distance, cross-axis range) -> (source coordinates, words, group) groups: dict[tuple[int, str, int, tuple[int, int, int]], tuple[set[int], int, str]] = {} diff --git a/spada/syntax/spatial_ir/stream_lifetime.py b/spada/syntax/spatial_ir/stream_lifetime.py index 3b346db9..ec39a4a5 100644 --- a/spada/syntax/spatial_ir/stream_lifetime.py +++ b/spada/syntax/spatial_ir/stream_lifetime.py @@ -743,31 +743,23 @@ def _never_concurrent(first: str, second: str, uses_per_rect: list[dict[spir.Ide def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int], *, kind: str, architecture: str, location: str, exclusive_keys: frozenset[str] | None = None) -> dict[str, int]: - """ - Assigns hardware fabric queues to stream groups from their occupancy spans on one PE. - - A group is one channel (or one ``auto`` stream). It occupies a single interval from its first - use on this PE to its last, including gaps between epochs: wavelets of that colour can still - arrive in a gap, and remapping the queue onto another color while they sit there is what the - hardware rejects. Two groups may share a queue only when those spans do not overlap. + """Assign fabric queues to stream groups based on occupancy spans. - Keys in ``exclusive_keys`` never share a queue, even when their spans are disjoint. That is - required on WSE-3 for every inbound color: the simulator remaps a queue onto the next color at - the first transfer, and faults if the queue is not empty. It is also required for colors that - bind a data task, because the hardware ID *is* the input queue. + Each group spans from its first use on the PE to its last. Because incoming wavelets + may arrive during gaps between epochs, queues cannot be safely remapped during a gap; + two groups may share a queue only if their active intervals are disjoint. - The spans form an interval graph, so colouring them in start-time order is optimal. + Groups in ``exclusive_keys`` are assigned dedicated queues that are never shared + across the entire kernel. This is required on WSE-3 for inbound colors and for + colors binding data tasks (where the task ID is the input queue ID). - :param spans: Mapping of grouping key to an inclusive ``(first_use, last_use)`` statement index - pair on this PE. - :param queue_ids: The hardware queue identifiers this direction may use, in the order they - should be handed out. - :param kind: ``'input'`` or ``'output'``, for the diagnostic. - :param architecture: The target name, for the diagnostic. - :param location: The PE rectangle, for the diagnostic. - :param exclusive_keys: Groups that must each own a queue for the whole PE, typically WSE-3 - data-task colors. - :return: Mapping of grouping key to a queue identifier from ``queue_ids``. + :param spans: Map from stream group key to inclusive (first_use, last_use) statement indices. + :param queue_ids: Available hardware queue IDs. + :param kind: 'input' or 'output', used in diagnostics. + :param architecture: Target architecture name, used in diagnostics. + :param location: PE coordinate description, used in diagnostics. + :param exclusive_keys: Stream groups requiring dedicated, unshared queues. + :return: Map from stream group key to assigned queue ID. """ if not spans: return {} @@ -814,23 +806,17 @@ def assign_fabric_queues(spans: dict[str, tuple[int, int]], queue_ids: list[int] def assign_microthreads(live: dict[str, list[tuple[int, int]]], microthread_ids: list[int], *, location: str) -> dict[str, int]: - """ - Assigns microthreads to stream groups from the intervals they are in flight over on one PE. - - A microthread is held only for the lifetime of one asynchronous operation, so unlike a fabric - queue it needs no proof that the hardware has drained, and it is not held across the gaps - between a group's transfers: two groups may take turns on one microthread as long as no transfer - of the one is in flight while a transfer of the other is. Callers pass the inbound and outbound - groups of a PE together, keyed apart by direction, because a microthread is one resource shared - by both directions. - - :param live: Mapping of grouping key to the inclusive intervals of statement indices over which - its transfers are in flight, in one index space over both directions. - :param microthread_ids: The microthread identifiers a program may name, in the order they should - be handed out. - :param location: The PE rectangle, for the diagnostic. - :return: Mapping of grouping key to a microthread identifier, empty when the target cannot name - microthreads and the hardware default has to stand. + """Assign microthreads to asynchronous transfer intervals on a PE. + + Microthreads are occupied only while an asynchronous DSD operation is in flight. + Two transfers may share a microthread as long as their active intervals do not overlap. + Inbound and outbound groups are considered jointly because microthreads are a shared + resource across both transfer directions. + + :param live: Map from stream group key to list of inclusive (start, end) statement intervals. + :param microthread_ids: Available microthread IDs. + :param location: PE coordinate description, used in diagnostics. + :return: Map from stream group key to assigned microthread ID. """ if not live or not microthread_ids: return {} diff --git a/tests/csl_runtime/run_tests.sh b/tests/csl_runtime/run_tests.sh index 2e7df16a..88fc0d95 100755 --- a/tests/csl_runtime/run_tests.sh +++ b/tests/csl_runtime/run_tests.sh @@ -1,9 +1,7 @@ #!/bin/bash -# Every test builds in this directory and exchanges data through the fixed names inp.npy and -# OUT_out.npy, removing them once a case is done. Only one run may be active per checkout: two -# concurrent runs overwrite each other's inputs and fail with a shape mismatch. Run architectures -# sequentially, or give each one its own checkout. +# Tests share fixed temporary files (inp.npy, OUT_out.npy) within this directory. +# Execute architecture test suites sequentially to avoid file collisions. # Color codes for output RED='\033[0;31m' diff --git a/tests/csl_runtime/samples/data_task_two_epochs.sptl b/tests/csl_runtime/samples/data_task_two_epochs.sptl index c923d289..6ee917ed 100644 --- a/tests/csl_runtime/samples/data_task_two_epochs.sptl +++ b/tests/csl_runtime/samples/data_task_two_epochs.sptl @@ -1,12 +1,12 @@ /** - * One channel, two epochs, and nothing else. + * Reusing a channel and its data task across two phases. * - * PE0 sends a word to PE1 in one phase and another word in the next, both on channel 0. - * The channel is a reusable resource, so this is legal -- but PE1 receives twice on it, and - * the data task a channel binds is the color itself. The two receives therefore have to - * share one hardware task and take turns, which is what this sample exercises. + * PE 0 sends to PE 1 in phase 1, and again in phase 2, both on channel 0. + * PE 1 receives in both phases on channel 0. In CSL, each data task binds to + * a hardware color/queue, so the two receives share a single data task slot + * driven by an alternating state variable. * - * Constraints: R >= 1 (repeats the pair of epochs R times). + * Constraints: R >= 1 (repeats the two-phase sequence R times). **/ kernel @data_task_two_epochs( stream[2, 1] readonly inp, diff --git a/tests/csl_runtime/samples/shift_bundle_1D.sptl b/tests/csl_runtime/samples/shift_bundle_1D.sptl index 1db7dee7..1d900b6e 100644 --- a/tests/csl_runtime/samples/shift_bundle_1D.sptl +++ b/tests/csl_runtime/samples/shift_bundle_1D.sptl @@ -1,15 +1,10 @@ /** - * An overlapping 1D interval shift, and nothing else. + * Shift values from PEs [0:M) to PEs [D:D+M) on a single channel. * - * D + M PEs on a line. Sources [0:M) each send one f32 to the PE D steps east, which - * overwrites its own value with it. Sources keep theirs. M <= D, so no PE is both a - * source and a destination. - * - * The paths overlap: source i's word passes through the routers of sources i+1 .. M-1. - * A color holds one (rx, tx) pair at a time, so the sources take turns -- nearest the - * destinations first, each handing its router over to relay mode once its own word is - * out. The destinations are statically routed and pick their word out of the stream - * with a counter filter. All of it on one color. + * Sources [0:M) each send one value to destination PEs [D:D+M) at offset D. + * The paths overlap along intermediate routers. To share a single channel: + * - Sources transmit in reverse order (nearest destination first), then switch to relay. + * - Destinations stay in relay mode and use hardware counter filters to pick their value. * * Constraints: 2 <= M <= D. **/ diff --git a/tests/csl_runtime/test_data_task_two_epochs.sh b/tests/csl_runtime/test_data_task_two_epochs.sh index 49f4f33c..2dab573d 100755 --- a/tests/csl_runtime/test_data_task_two_epochs.sh +++ b/tests/csl_runtime/test_data_task_two_epochs.sh @@ -1,8 +1,5 @@ #!/bin/sh -# E2E: one channel carrying two epochs between the same pair of PEs, and nothing else. -# Kernel: samples/data_task_two_epochs.sptl params: R (repeats of the pair of epochs). -# PE1 receives twice on channel 0, so both receives share the data task the channel binds. -# After the run both PEs hold PE0's two keys. +# E2E test: samples/data_task_two_epochs.sptl (params: R) set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" diff --git a/tests/csl_runtime/test_exchange_bundle_1d.sh b/tests/csl_runtime/test_exchange_bundle_1d.sh index 32f2dd9a..1c5b139e 100755 --- a/tests/csl_runtime/test_exchange_bundle_1d.sh +++ b/tests/csl_runtime/test_exchange_bundle_1d.sh @@ -1,9 +1,5 @@ #!/bin/sh -# E2E: two opposite shift bundles per phase, one color each, repeated R times. -# Kernel: exchange_bundle_1D.sptl params: M (pairs), D (distance), R (repeats), M <= D. -# The PEs in [0:M) and [D:D+M) swap pairwise once per repeat, so after an odd R -# OUT_out[i] == inp[i + D] and OUT_out[i + D] == inp[i], and after an even R nothing moved. -# R also sets how many wavelet filters each PE needs, which is what caps it at three. +# E2E test: exchange_bundle_1D.sptl (params: M, D, R) set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" diff --git a/tests/csl_runtime/test_odd_even_sort_1d_looped.sh b/tests/csl_runtime/test_odd_even_sort_1d_looped.sh index 427b3509..200c22a5 100755 --- a/tests/csl_runtime/test_odd_even_sort_1d_looped.sh +++ b/tests/csl_runtime/test_odd_even_sort_1d_looped.sh @@ -1,12 +1,5 @@ #!/bin/sh -# E2E: odd-even transposition sort on 2^L PEs, N rounds as a runtime loop -# (odd_even_sort_1D_looped.sptl). Each PE holds a block of K f32 keys; every comparator is a -# compare-split, so the network sorts all 2^L * K keys and PE i ends up with keys i*K .. i*K + K-1 -# of the sorted sequence. Reference: OUT_a_out.reshape(n*k) == sort(a_in.reshape(n*k)). -# Runs on WSE-2 and WSE-3: four channels, one per (round parity, direction), so no -# router ever switches. L = 1 is two PEs and a single even round; L = 3 is eight PEs -# and exercises every role (ends and both interior parities). K = 1 is the one-key-per-PE -# network; K = 4 is not a power of two, which is what the K-element merge has to be independent of. +# E2E test: odd_even_sort_1D_looped.sptl (params: L, K) set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" diff --git a/tests/csl_runtime/test_shearsort_2d_looped.sh b/tests/csl_runtime/test_shearsort_2d_looped.sh index 4290222d..d5863ac0 100755 --- a/tests/csl_runtime/test_shearsort_2d_looped.sh +++ b/tests/csl_runtime/test_shearsort_2d_looped.sh @@ -1,18 +1,12 @@ #!/bin/sh -# E2E: shearsort on an N x N mesh, N = 2^L, N neighbour odd-even rounds as a runtime loop -# (shearsort_2D.sptl). Each PE holds a block of K f32 keys; every comparator is a -# compare-split, so the network sorts all N*N*K keys into snake order: even rows left to -# right, odd rows right to left, each block still ascending. -# Reference: flatten(OUT_a_out in snake order) == sort(a_in.reshape(n*n*k)). -# WSE-3 only: a fully interior PE receives on four colours in one epoch, and WSE-3 binds a -# queue to its colour for the whole kernel (six of each). WSE-2 has two input queues. +# E2E test: shearsort_2D.sptl on WSE-3 (params: L, K) set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" . "$SCRIPT_DIR/_lib.sh" if [ "${WSE_ARCH:-wse2}" != "wse3" ]; then - echo "Skipping shearsort_2d_looped: four inbound colours live in one epoch, and wse2 has two input queues." + echo "Skipping shearsort_2d_looped: four inbound colors live in one epoch, and wse2 has two input queues." exit 0 fi diff --git a/tests/csl_runtime/test_shift_bundle_1d.sh b/tests/csl_runtime/test_shift_bundle_1d.sh index 8d2555cc..0f1f2bfb 100755 --- a/tests/csl_runtime/test_shift_bundle_1d.sh +++ b/tests/csl_runtime/test_shift_bundle_1d.sh @@ -1,8 +1,5 @@ #!/bin/sh -# E2E: an overlapping eastbound interval shift on one color, and nothing else. -# Kernel: shift_bundle_1D.sptl params: M (sources), D (shift distance), M <= D. -# After the shift, OUT_out[D:D+M] == inp[0:M], and every other PE keeps its own value. -# D > M is the case with pure relays between the two halves. +# E2E test: shift_bundle_1D.sptl (params: M, D) set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" diff --git a/tests/csl_runtime/test_spmv.sh b/tests/csl_runtime/test_spmv.sh index e8a907c7..fd514871 100644 --- a/tests/csl_runtime/test_spmv.sh +++ b/tests/csl_runtime/test_spmv.sh @@ -1,12 +1,5 @@ #!/bin/sh -# E2E test: distributed sparse GEMV y = alpha * A * x + beta * y. -# Grid PX × PY; each PE holds a padded COO block of A with bound NZ. -# Phase 1: load COO A; Phase 2: load x (j=0) and y (i=0); Phase 3: broadcast x in Y; -# Phase 4: local COO SpMV; Phase 5: pipelined chain reduce z in X, -# root applies alpha*z + beta*y and outputs result. -# Reference: OUT_out.npy[0, j, :] == (alpha * A_full @ x_flat + beta * y_flat)[j*K:(j+1)*K] -# Tested with (PX, PY) ∈ {(2,2), (2,3), (3,2), (3,4)}, K=2, NZ=K*K, -# and a sparser case (2,2) with NZ=2 < K*K and explicit zero padding. +# E2E test: spmv.sptl (params: PX, PY, K, NZ) set -e SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" diff --git a/tests/spatial_ir/test_dsd_ops.py b/tests/spatial_ir/test_dsd_ops.py index 06226e24..a5ef8eb4 100644 --- a/tests/spatial_ir/test_dsd_ops.py +++ b/tests/spatial_ir/test_dsd_ops.py @@ -243,7 +243,7 @@ def test_an_array_of_one_element_still_gets_a_dsd(k: int): def test_wse3_concurrent_transfers_use_distinct_microthreads(): """Two transfers in flight at once may not share a microthread. - A laplacian PE receives from one neighbour and forwards to another in the same task. On WSE-3 + A laplacian PE receives from one neighbor and forwards to another in the same task. On WSE-3 the input and output queue pools both start at 2, so leaving the microthread at its default -- the queue ID of the highest-priority fabric operand -- put both on microthread 2 and aborted the simulation with ``trying to term ut_instr[2], but it's not ours``. diff --git a/tests/spatial_ir/test_routing.py b/tests/spatial_ir/test_routing.py index 0fd6119b..0f761f06 100644 --- a/tests/spatial_ir/test_routing.py +++ b/tests/spatial_ir/test_routing.py @@ -479,7 +479,7 @@ def _fabin_queues(code: str) -> dict[str, str]: def test_odd_even_sort_looped_interior_keeps_distinct_input_queues(): """ - Even-round east and odd-round west are both inbound on an odd interior PE. The west neighbour + Even-round east and odd-round west are both inbound on an odd interior PE. The west neighbor can inject the next even-round block on C0 while this PE is already receiving the odd-round one on C3. Sharing input queue 0 is what the WSE-2 simulator rejects as remapping C0 onto C3 while the router still holds wavelets (L=2 K=16). @@ -491,13 +491,13 @@ def test_odd_even_sort_looped_interior_keeps_distinct_input_queues(): assert len(set(odd_interior.values())) == 2, odd_interior assert len(even_interior) == 2, even_interior assert len(set(even_interior.values())) == 2, even_interior - # Endpoints have one inbound colour and do not need a second queue. + # Endpoints have one inbound color and do not need a second queue. assert len(_fabin_queues(files['code_0_0.csl'])) == 1 assert len(_fabin_queues(files['code_3_0.csl'])) == 1 ### -# shearsort_2D_looped: (RC)^L R neighbour rounds as runtime loops on eight static channels +# shearsort_2D_looped: (RC)^L R neighbor rounds as runtime loops on eight static channels ### _SHEARSORT_LOOPED = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', @@ -533,7 +533,7 @@ def test_shearsort_looped_uses_eight_static_channels(): interior = files['code_2_2.csl'] assert 'for (@range(i32, 0, 2, 1))' in interior, interior - # Four outbound colours (even-row east, odd-row west, even-column south, odd-column north) + # Four outbound colors (even-row east, odd-row west, even-column south, odd-column north) # and four inbound, each emitted once rather than unrolled over L or N. assert interior.count('fabout_dsd') == 4, interior assert interior.count('fabin_dsd') == 4, interior @@ -554,8 +554,8 @@ def test_shearsort_looped_code_is_independent_of_n(): def test_shearsort_looped_interior_keeps_distinct_input_queues(): """ - A fully interior PE receives on two row colours and two column colours. WSE-3 binds each - inbound colour to its own queue for the whole kernel, so those four must be distinct. + A fully interior PE receives on two row colors and two column colors. WSE-3 binds each + inbound color to its own queue for the whole kernel, so those four must be distinct. """ _require_wse3_shearsort() files = _lower_shearsort_looped(2, K=1) From 94204a219b7c50206b32ff00e04df101603a96ca Mon Sep 17 00:00:00 2001 From: glukas Date: Mon, 5 Oct 2026 20:16:27 +0200 Subject: [PATCH 68/68] Add comment on how to get thesis copy --- spada/syntax/spatial_ir/shift_bundles.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/spada/syntax/spatial_ir/shift_bundles.py b/spada/syntax/spatial_ir/shift_bundles.py index d492b1d3..9b1b81a3 100644 --- a/spada/syntax/spatial_ir/shift_bundles.py +++ b/spada/syntax/spatial_ir/shift_bundles.py @@ -12,6 +12,9 @@ 3. Destination routers statically forward wavelets to both the local ramp and downstream neighbors (or to the ramp only for the final destination). Hardware counter filters at each destination select the designated slice of data. + +This is the arrangement in Louis Schnyders Bachelor thesis, "Distributed Sorting on the Cerebras Wafer-Scale Engine", +fig. 7.6. Unpublished; reach out to the author or L. Gianinazzi for a private copy. """ from __future__ import annotations