diff --git a/.github/workflows/python-app.yml b/.github/workflows/python-app.yml index da167b43..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,11 +35,19 @@ jobs: pip install -e ".[dev]" - name: Test with pytest + env: + WSE_ARCH: ${{ matrix.wse-arch }} run: | pytest + # WSE-2 and WSE-3 share one runner so Singularity and the SDK are installed + # once. test-csl: runs-on: ubuntu-latest + timeout-minutes: 180 + strategy: + matrix: + wse-arch: ["wse2", "wse3"] steps: - uses: actions/checkout@v4 @@ -54,29 +65,35 @@ jobs: pip install -e ".[dev]" - 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 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 @@ -92,7 +109,7 @@ jobs: - name: Test CSL with simulator run: | export PATH=$PATH:`pwd`/cerebras-sdk - ./tests/csl_runtime/run_tests.sh + WSE_ARCH=${{ matrix.wse-arch }} ./tests/csl_runtime/run_tests.sh build-package: runs-on: ubuntu-latest diff --git a/README.md b/README.md index eb106154..55d1ec78 100644 --- a/README.md +++ b/README.md @@ -106,9 +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` | -| `samples/spatial/blas/` | Dense linear algebra: `axpy`, `matvec`, `gemv`, `gemv_twophase` | +| `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: `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 | @@ -116,7 +117,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 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 @@ -159,6 +160,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:** @@ -189,13 +191,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.md b/irspec/docs/spatial/routing.md index 21c9b485..ee079782 100644 --- a/irspec/docs/spatial/routing.md +++ b/irspec/docs/spatial/routing.md @@ -192,3 +192,8 @@ to receive. If multicasting is used, the correctness conditions must be adapted accordingly, especially when considering multiple phases. + +## 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..31450d32 100644 --- a/irspec/docs/spatial/routing_wse.md +++ b/irspec/docs/spatial/routing_wse.md @@ -2,109 +2,159 @@ 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. Compare - `samples/spatial/sorting/odd_even_sort_1D.sptl` with - `samples/spatial/sorting/odd_even_sort_1D_looped.sptl`. - - 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. - -!!! 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. - -!!! 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/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. - -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 + +When multiple consecutive PEs shift data by uniform distance $d$ along an axis: + +``` +dataflow i16 i, i16 j in [0:D + M, 0] { + stream fwd = relative_stream(D, 0) { hops = auto, channel = 0 } +} +``` + +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. + +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) +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 +``` + +### 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/irspec/docs/spatial/spatial.md b/irspec/docs/spatial/spatial.md index d62ab384..79749a8f 100644 --- a/irspec/docs/spatial/spatial.md +++ b/irspec/docs/spatial/spatial.md @@ -499,6 +499,9 @@ stream stream_name = relative_stream(dx, dy) { 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. +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`. @@ -857,8 +860,9 @@ an explicit `close` on a bounded stream is redundant but legal. An unbounded str by an explicit `close`. !!! note "Note: Verification of Bounds" - When the compiler can infer stream bounds statically, it may generate a compiler error. - Otherwise, no diagnostic is emitted. + 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* 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 new file mode 100644 index 00000000..15144bbb --- /dev/null +++ b/samples/spatial/blas/spmv.sptl @@ -0,0 +1,149 @@ +/** + * 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 + 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 + } + + // 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. + // The three ranges are disjoint: PE(0,0) belongs to both the row and the column. + phase { + 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 [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]) + } + } + + // 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] { + await map i16 k in [0:K] { + z[k] = 0.0 + } + for i16 p in [0:NZ] { + tmp = A_row[p] + tmp2 = A_col[p] + z[tmp] = z[tmp] + A_val[p] * x[tmp2] + } + } + } + + // 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 + } + await map 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/collectives/scalar_reduce_1D.sptl b/samples/spatial/collectives/scalar_reduce_1D.sptl index 0ea362da..bf592456 100644 --- a/samples/spatial/collectives/scalar_reduce_1D.sptl +++ b/samples/spatial/collectives/scalar_reduce_1D.sptl @@ -3,10 +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 - * - * This sample does not compile for the Cerebras WSE because of the reuse of - * the same channel for both sending and receiving in the middle PEs. + * 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 + * `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/samples/spatial/simple/exchange_bundle_1D.sptl b/samples/spatial/simple/exchange_bundle_1D.sptl new file mode 100644 index 00000000..09672549 --- /dev/null +++ b/samples/spatial/simple/exchange_bundle_1D.sptl @@ -0,0 +1,59 @@ +/** + * Repeated pairwise exchange between two blocks of PEs. + * + * 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. + * + * 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, 1 <= R <= 3. + **/ +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/sort/batcher_oddeven_wse3_1D.sptl b/samples/spatial/sort/batcher_oddeven_wse3_1D.sptl new file mode 100644 index 00000000..f14fcdbc --- /dev/null +++ b/samples/spatial/sort/batcher_oddeven_wse3_1D.sptl @@ -0,0 +1,231 @@ +/** + * 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. + * + * bundled (4*d >= N): fwd = N + 2*(l*(L+1) + p), bwd = fwd + 1 + * + * The remaining short-distance phases pool channels by distance d and the residue (coordinate mod 2d), + * regardless of send direction: + * + * pooled (4*d < N): 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: + * + * source R->E or R->W + * destination W->R or E->R + * relay W->E or E->W + * + * Channels per phase for L = 3: + * + * (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 + * + * The kernel uses 8 colors for L = 3 and 12 colors for L = 4. + * + * 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< 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<= 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 + } + } + + // PE 0: low partner in every even round, idle in every odd round. + compute i16 i, i16 j in [0:1, 0] { + await receive(val, a_in[i, j]) + for i16 m in [1:K] { + for i16 s in [0:K-1] { + pv = ((m - s) - 1) if ((m - s) - 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) + // 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-indexed interior PEs: low in even rounds, high in odd rounds. + 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 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-indexed interior PEs: high in even rounds, low in odd rounds. + 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 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]) + } + + // PE N-1: high partner in every even round, idle in every odd round. + 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 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]) + } +} diff --git a/samples/spatial/sort/shearsort_2D.sptl b/samples/spatial/sort/shearsort_2D.sptl new file mode 100644 index 00000000..06cfdcfa --- /dev/null +++ b/samples/spatial/sort/shearsort_2D.sptl @@ -0,0 +1,1655 @@ +/** + * 2D Shearsort on an N x N mesh (N = 2^L), sorting N*N*K keys into snake order. + * + * Each PE holds a block of K f32 keys. After sorting, the mesh is ordered in snake order: + * even rows are sorted ascending left-to-right, odd rows descending right-to-left, and PE blocks + * are internally sorted ascending. + * + * Algorithm + * --------- + * The algorithm executes L iterations of (row sort, column sort) followed by a final row sort: + * - Even rows sort ascending (west to east). + * - Odd rows sort descending (east to west). + * - Columns sort ascending (north to south). + * Each 1D sort executes N rounds of neighbor compare-split operations on K-element blocks. + * + * 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< 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 colors. + 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/samples/spatial/sorting/bitonic_sort_1D.sptl b/samples/spatial/sorting/bitonic_sort_1D.sptl deleted file mode 100644 index 08b8e0be..00000000 --- a/samples/spatial/sorting/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< list[spir.Identifier]: + """Collect fabric transfers in ``statement`` in source order, duplicating sequential ``for`` bodies. + + 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: Transferred stream identifiers with each ``for`` body duplicated. + """ + 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]]: + """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 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) + + 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 4c15a246..6629082e 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 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 @@ -150,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 @@ -176,7 +178,8 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, rect_size = x1 - x0, y1 - y0 # Collect unique routes for all rectangles - routes_per_rectangle = cslrouting.collect_routes(rectangles, color_maps, disable_switching) + routes_per_rectangle, standalone_routes = cslrouting.collect_routes(rectangles, color_maps, + disable_switching, rect_offset) if use_memcpy_mode: layout_code.write(f''' @@ -272,6 +275,10 @@ def lower_spatial_ir_to_csl(kernel: spir.Kernel, }} }}\n''') + # Sites that are no rectangle shifted as a whole bring their own loop. + for loop in standalone_routes: + layout_code.write('\n' + loop) + for rinst in routing_instructions: layout_code.write(rinst + '\n') @@ -378,14 +385,16 @@ 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()}\".") 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: @@ -398,14 +407,26 @@ 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: 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. + # 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) 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. @@ -417,24 +438,16 @@ 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: + 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: + 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: @@ -505,22 +518,22 @@ 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) - footer.write(f' @bind_data_task(dtask_{i}, dtask_{i}_id);\n') - if task.blocked: - footer.write(f' @block(dtask_{i}_id);\n') - - max_task_id = max((slot.hardware_task_id for slot in task_bindings.local_slots), default=csl.LOCAL_TASK_IDS[0] - 1) + # 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') # 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 wse3.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: @@ -553,6 +566,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": @@ -563,7 +584,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') @@ -590,7 +614,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 @@ -642,34 +666,50 @@ def _collect_colors_globally(kernel: spir.Kernel, rectangles: list[Rectangle[PEB max_channel = max(channel_is_read.union(channel_is_written), default=-1) - # Assign all "auto" channels + # Assign all "auto" channels, one per stream of a phase rather than one per declaration. A stream + # declared for a grid that consolidation splits shows up once per rectangle, and ``inline_phases`` + # gives each copy a name of its own, so the copies are recognised by what they route rather than + # by their name: same phase, same offsets, same hops. Sharing them is what puts the matchings of + # a sorting network's phase on two colors instead of two per matching. + 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: - channel_is_written.add(max_channel + 1) - if stream_decl.stream_name in auto_stream_is_read: - channel_is_read.add(max_channel + 1) + if stream_decl.stream.routing.resolved_channel != "auto": + continue + key = (stream_decl.phase, stream_lifetime.stream_group_key(stream_decl)) + channel = auto_channels.get(key) + if channel is None: max_channel += 1 - - # Allocate colors for each channel + channel = max_channel + auto_channels[key] = channel + stream_decl.stream.routing.channel = channel + if stream_decl.stream_name in auto_stream_is_written: + channel_is_written.add(channel) + if stream_decl.stream_name in auto_stream_is_read: + channel_is_read.add(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 @@ -885,45 +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: - 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, @@ -931,6 +932,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. @@ -958,27 +960,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. + # 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 @@ -989,25 +986,62 @@ 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 = 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) + # 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( + input_spans, csl.INPUT_QUEUE_IDS, + kind='input', architecture=csl.ARCH, location=location, + exclusive_keys=exclusive) + output_queue_of = stream_lifetime.assign_fabric_queues( + output_spans, csl.OUTPUT_QUEUE_IDS, + kind='output', architecture=csl.ARCH, location=location, + exclusive_keys=exclusive_out) + # 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) + + 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: - 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] + # 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() + 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. @@ -1026,7 +1060,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 @@ -1048,7 +1088,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: @@ -1075,10 +1117,11 @@ 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 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 @@ -1092,8 +1135,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)) + dsd = _fabout(stream_name, extents) dsds[stream_name.as_ir()].append((dsd_name, dsd)) if isinstance(stmt, spir.SendStatement): @@ -1129,7 +1171,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 @@ -1143,8 +1184,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)) + dsd = _fabout(stream_name, extents) dsds[stream_name.as_ir()].append((dsd_name, dsd)) def _visit_nested_receive(substmt: spir.ReceiveStatement): @@ -1170,7 +1210,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): @@ -1333,85 +1375,151 @@ def _write_indented_block(current_code: StringIO, block: str, indent: str) -> No current_code.write(f'{indent}{line}\n') -def _generate_data_task( +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 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 "" - next_task_code = f'@{itedge_code}({prefix}task_{next_task}_id);' + 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);') - 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, @@ -1470,23 +1578,31 @@ 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() - # DSD operation or async call - if any(dsdop in line for line in lines for dsdop in dsd_ops.DSD_ASSIGNMENT_MAPPING): - # Asynchronous DSD op. DSD line already contains activation or unblocking + # A DSD operation given an async target performs the transition itself, either as the + # operation's ``.activate``/``.unblock`` field or as a call right after it. Naming the + # target is what distinguishes those from a statement that merely happens to use a DSD + # instruction -- a ``close``, whose control wavelets go out through a plain ``@mov32``. + 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') @@ -1506,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/lowering/wse3.py b/spada/lowering/wse3.py new file mode 100644 index 00000000..a0ecc1d4 --- /dev/null +++ b/spada/lowering/wse3.py @@ -0,0 +1,191 @@ +""" +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``. + + 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. + :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/runtime/runtime.py b/spada/runtime/runtime.py index 9ef3adc6..25d0b1fb 100644 --- a/spada/runtime/runtime.py +++ b/spada/runtime/runtime.py @@ -94,6 +94,57 @@ def from_json(cls, json_data: Union[str, Dict[str, Any]]) -> "ProgramMetadata": ######################################################## +def memcpy_data_type(dtype: np.dtype) -> "crt.MemcpyDataType": + """Return the transfer width for a kernel argument of ``dtype``. + + :param dtype: The declared argument dtype. + :return: The corresponding ``MemcpyDataType``. + """ + 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: + """Return the host buffer dtype for a kernel argument of ``dtype``, aligned to 32-bit words. + + 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 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 32-bit words expected by host memcpy. + + 16-bit values are bitcast without sign extension to preserve exact binary representations. + + :param data: Contiguous input array. + :return: The array itself if 32 bits wide, otherwise 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: + """Narrow 32-bit memcpy words back to the target output dtype. + + :param words: Buffer populated by ``memcpy_d2h``. + :param dtype: Target output dtype. + :return: Array reinterpreted in ``dtype``. + """ + 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 +168,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 +196,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 +208,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]: @@ -288,7 +339,10 @@ def __init__( cmaddr = cm_addr or os.environ.get("CM_ADDR", None) self.simulator = cmaddr is None print("SIMULATOR?", self.simulator) - self.runtime = crt.SdkRuntime(str(self.out_folder), suppress_simfab_trace=True, cmaddr=cmaddr) + # 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) # Store input/output information from metadata self.inputs = self.metadata.inputs diff --git a/spada/syntax/csl/constants.py b/spada/syntax/csl/constants.py index 6c8e7292..2163e68c 100644 --- a/spada/syntax/csl/constants.py +++ b/spada/syntax/csl/constants.py @@ -6,15 +6,28 @@ # 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': 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)), @@ -35,21 +48,31 @@ } 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] +# 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)), +} +MICROTHREAD_IDS = _MICROTHREAD_IDS[ARCH] + _HARDWARE_FABRIC_DIMS = { 'wse2': (757, 996), 'wse3': (762, 1172), @@ -61,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. @@ -85,3 +102,10 @@ # 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] + +# 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/dsd_ops.py b/spada/syntax/csl/dsd_ops.py index f1e76ee0..3cecf599 100644 --- a/spada/syntax/csl/dsd_ops.py +++ b/spada/syntax/csl/dsd_ops.py @@ -23,11 +23,18 @@ 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): - 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});' + # 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_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, statement: spir.Statement, diff --git a/spada/syntax/csl/routing.py b/spada/syntax/csl/routing.py index 1e49a966..94311776 100644 --- a/spada/syntax/csl/routing.py +++ b/spada/syntax/csl/routing.py @@ -25,7 +25,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 analysis, shift_bundles, 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 @@ -65,6 +65,31 @@ def changes_both_sides(self, previous: 'RouteConfig') -> bool: return self.rx != previous.rx and self.tx != previous.tx +@dataclass(frozen=True) +class FilterConfig: + """Hardware counter filter configuration for a router color. + + 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. + + 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 + 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. @@ -99,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. @@ -136,9 +161,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 +195,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: """ @@ -258,7 +288,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. @@ -311,10 +341,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 +379,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: + """Compute the initial counter value for a destination filter as an affine expression. + + 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 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: + 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']]: + """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' + 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 +516,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 +531,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 +556,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 +565,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) @@ -469,16 +686,21 @@ 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, 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 + 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. + 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. @@ -516,6 +738,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: @@ -527,8 +750,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]), @@ -544,11 +770,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 configuration for this color and never routes - # anything again, so wavelets passing through may over-advance it harmlessly. + 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: @@ -564,7 +799,18 @@ 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, 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. 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 planned += 1 return planned @@ -626,10 +872,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 +921,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/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 693c01b2..7a2a97ff 100644 --- a/spada/syntax/csl/structures.py +++ b/spada/syntax/csl/structures.py @@ -60,15 +60,30 @@ class FabricDSD(DataStructureDescriptor): color: str 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 + #: 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) + 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 '' - 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 '' + 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}{advance} .{queue_type} = @get_{queue_type}({self.queue}) }})') def __hash__(self): return hash(("FabricDSD", self.as_csl())) diff --git a/spada/syntax/csl/task_recycling.py b/spada/syntax/csl/task_recycling.py index 96dc6d78..2b278edf 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 @@ -281,12 +363,80 @@ 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; + 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 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): + 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/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/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/spada/syntax/spatial_ir/canonicalization.py b/spada/syntax/spatial_ir/canonicalization.py index b5bd68cf..48e51531 100644 --- a/spada/syntax/spatial_ir/canonicalization.py +++ b/spada/syntax/spatial_ir/canonicalization.py @@ -609,28 +609,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 @@ -666,8 +679,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. """ @@ -755,14 +768,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/irnodes.py b/spada/syntax/spatial_ir/irnodes.py index 11396838..9890e3c4 100644 --- a/spada/syntax/spatial_ir/irnodes.py +++ b/spada/syntax/spatial_ir/irnodes.py @@ -579,6 +579,9 @@ 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. + + 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" @@ -617,7 +620,11 @@ 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}", + ] + return ", \n".join(lines) @dataclass @@ -925,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. 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: assert isinstance(self.stream_name, (Identifier, ArraySlice)) @@ -935,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/spada/syntax/spatial_ir/language.lark b/spada/syntax/spatial_ir/language.lark index 5d41c17b..3eb7724b 100644 --- a/spada/syntax/spatial_ir/language.lark +++ b/spada/syntax/spatial_ir/language.lark @@ -116,7 +116,10 @@ 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_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 "}")? 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 f1d62a51..c8cdefed 100644 --- a/spada/syntax/spatial_ir/lark_to_ir.py +++ b/spada/syntax/spatial_ir/lark_to_ir.py @@ -188,7 +188,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): @@ -326,6 +325,27 @@ 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_field(self, args): + return args[0] + + def routing(self, args): + kwargs = {'hops': 'auto', 'channel': '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']) + 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..9b1b81a3 --- /dev/null +++ b/spada/syntax/spatial_ir/shift_bundles.py @@ -0,0 +1,203 @@ +"""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. + +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 + +from dataclasses import dataclass +from typing import Literal, Optional + +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: + """ + One run of consecutive sources shifted onto an equally long run of destinations. + + 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 + #: 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 + #: 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]]: + """ + Returns the axis and signed distance of a stream that runs straight along one axis. + + ``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. + """ + stream = declaration.stream + if not isinstance(stream, spir.RelativeStreamDeclaration) or stream.routing is None: + return None + try: + dx, dy = int(stream.dx.eval()), int(stream.dy.eval()) + except Exception: # pragma: no cover - defensive: a non-constant offset + return None + if (dx == 0) == (dy == 0): + return None + axis: Literal['x', 'y'] = 'x' if dy == 0 else 'y' + delta = dx if axis == 'x' else dy + + 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 _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: set[int]) -> list[tuple[int, int]]: + """Splits a set of coordinates into ``(start, length)`` runs of consecutive values.""" + if not values: + return [] + runs: list[tuple[int, int]] = [] + ordered = sorted(values) + start = previous = ordered[0] + for value in ordered[1:]: + if value == previous + 1: + previous = value + continue + runs.append((start, previous - start + 1)) + start = previous = value + runs.append((start, previous - start + 1)) + return runs + + +def detect_shift_bundles(rectangles: list[Rectangle[PEBlock]]) -> list[ShiftBundle]: + """Detect interval shifts with overlapping router paths suitable for bundling. + + 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: 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]] = {} + 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 + shift = _straight_shift(declaration) + words = _bound(declaration) + if shift is None or words is None: + continue + axis, delta = shift + if abs(delta) < MIN_BUNDLE_DISTANCE: + continue + channel = declaration.stream.routing.resolved_channel + if channel == 'auto': + continue + + 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])) + + 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 + 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 bundles_by_group(bundles: list[ShiftBundle]) -> dict[str, list[ShiftBundle]]: + """ + Indexes bundles by the routing identity of their stream, for the route collector to look up. + """ + 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..ec39a4a5 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 @@ -126,14 +127,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 +228,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 +260,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 @@ -638,6 +740,110 @@ 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, + exclusive_keys: frozenset[str] | None = None) -> dict[str, int]: + """Assign fabric queues to stream groups based on occupancy spans. + + 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. + + 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: 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 {} + if not queue_ids: + 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 (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: + assigned[key] = queue + break + else: + overlapping = sorted( + other for other, (other_start, other_end) in spans.items() + 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 ' + 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{extra}') + return assigned + + +def assign_microthreads(live: dict[str, list[tuple[int, int]]], microthread_ids: list[int], *, + location: str) -> dict[str, int]: + """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 {} + + 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(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 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 ' + 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/csl_runtime/Makefile b/tests/csl_runtime/Makefile index 34499e77..ee8d0fc3 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 97b89c13..d54f76c6 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() { @@ -189,7 +201,8 @@ vm "if ! python3 -m pip --version >/dev/null 2>&1; then \ python3 -m pip install --quiet -e '$REPO_ROOT[dev]'" # ── 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) @@ -213,6 +226,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 diff --git a/tests/csl_runtime/run_tests.sh b/tests/csl_runtime/run_tests.sh index 98a1948d..88fc0d95 100755 --- a/tests/csl_runtime/run_tests.sh +++ b/tests/csl_runtime/run_tests.sh @@ -1,5 +1,8 @@ #!/bin/bash +# 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' GREEN='\033[0;32m' @@ -17,6 +20,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 "" 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..6ee917ed --- /dev/null +++ b/tests/csl_runtime/samples/data_task_two_epochs.sptl @@ -0,0 +1,68 @@ +/** + * Reusing a channel and its data task across two phases. + * + * 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 two-phase sequence 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/samples/shift_bundle_1D.sptl b/tests/csl_runtime/samples/shift_bundle_1D.sptl new file mode 100644 index 00000000..1d900b6e --- /dev/null +++ b/tests/csl_runtime/samples/shift_bundle_1D.sptl @@ -0,0 +1,47 @@ +/** + * Shift values from PEs [0:M) to PEs [D:D+M) on a single channel. + * + * 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. + **/ +kernel @shift_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 + } + + phase { + compute i16 i, i16 j in [0:D + M, 0] { + await receive(val, inp[i, j]) + } + } + + phase { + // 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 = 0 + } + } + compute i16 i, i16 j in [0:M, 0] { + await send(val, fwd) + } + compute i16 i, i16 j in [D:D + M, 0] { + await receive(val, fwd) + } + } + + phase { + compute i16 i, i16 j in [0:D + M, 0] { + await send(val, out[i, j]) + } + } +} 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..84022f8d --- /dev/null +++ b/tests/csl_runtime/test_batcher_oddeven_wse3_1d.sh @@ -0,0 +1,63 @@ +#!/bin/sh +# E2E test: batcher_oddeven_wse3_1D.sptl params: L, K, R + +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 kernel targets wse3." + exit 0 +fi + +run_batcher() { + l=$1 + k=$2 + r=${3:-1} + echo "--- batcher_oddeven_wse3_1d L=$l K=$k R=$r ---" + + sptlc "$SORT_DIR/batcher_oddeven_wse3_1D.sptl" "$FOLDER" -p L=$l -p K=$k -p R=$r + + python3 - <() { + 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) diff --git a/tests/spatial_ir/test_dsd_ops.py b/tests/spatial_ir/test_dsd_ops.py index 61cf3737..a5ef8eb4 100644 --- a/tests/spatial_ir/test_dsd_ops.py +++ b/tests/spatial_ir/test_dsd_ops.py @@ -1,8 +1,10 @@ +import os +import re import pytest 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(): @@ -153,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) } @@ -207,6 +208,63 @@ def test_sequential_for_does_not_make_element_accesses_into_dsds(): assert '@fadds(a_dsd' not in code, code +@pytest.mark.parametrize('k', [1, 4]) +def test_an_array_of_one_element_still_gets_a_dsd(k: int): + """ + A single-element array is a scalar in all but name, but CSL keeps treating it as an array: the + bare name is rejected as the operand of a transfer or of a move between memory locations. It + therefore needs a DSD like any other array, whatever its extent. + """ + kernel = parser.parse_string(code=""" +kernel @blockcopy (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_wse3_concurrent_transfers_use_distinct_microthreads(): + """Two transfers in flight at once may not share a microthread. + + 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``. + """ + 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_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_routing.py b/tests/spatial_ir/test_routing.py index d8cdff47..0f761f06 100644 --- a/tests/spatial_ir/test_routing.py +++ b/tests/spatial_ir/test_routing.py @@ -216,20 +216,24 @@ 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 (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) - 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'] + + 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 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(): @@ -420,53 +424,149 @@ def test_switch_positions_beyond_capacity_are_rejected(): ### -# bitonic_sort_1D: the heaviest channel reuse in the samples +# odd_even_sort_1D_looped: N rounds as a runtime loop on four static channels ### -_BITONIC = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', 'sorting', - 'bitonic_sort_1D.sptl') +_ODD_EVEN_LOOPED = os.path.join(os.path.dirname(__file__), '..', '..', 'samples', 'spatial', + 'sort', 'odd_even_sort_1D_looped.sptl') -def _lower_bitonic(L: int, K: int = 4) -> dict[str, str]: - kernel = parser.parse_file(_BITONIC) +def _lower_odd_even_looped(L: int, K: int = 4) -> 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)} + return {f.filename: f.code for f in s2c.lower_spatial_ir_to_csl(kernel, disable_benchmarking=True)} -@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(): +def test_odd_even_sort_looped_uses_four_static_channels(): """ - 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. + 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 + + +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 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). """ - files = _lower_bitonic(2) - colors = set(re.findall(r'@get_color\((\d+)\)', files['layout.csl'])) - assert len(colors) <= 2, sorted(colors) + 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 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 - 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(): + +### +# 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', + 'sort', 'shearsort_2D.sptl') + + +def _lower_shearsort_looped(L: int, K: int = 1) -> 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(): """ - 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. + 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. """ - with pytest.raises(SyntaxError, match='switch positions'): - _lower_bitonic(2) + _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 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 + + +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 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) + 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__': diff --git a/tests/spatial_ir/test_shift_bundles.py b/tests/spatial_ir/test_shift_bundles.py new file mode 100644 index 00000000..c24e0059 --- /dev/null +++ b/tests/spatial_ir/test_shift_bundles.py @@ -0,0 +1,376 @@ +""" +Tests for bundling overlapping 1D interval shifts onto one color. +""" +import os +import re + +import pytest + +from spada.lowering.spatial_ir_to_csl import canonicalize_kernel, lower_spatial_ir_to_csl +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 + +_SHIFT = """ +kernel @shift( + stream[D + M, 1] readonly inp, + stream[D + M, 1] writeonly out +) { + place i16 i, i16 j in [0:D + M, 0] { + f32[K] val + } + phase { + compute i16 i, i16 j in [0:D + M, 0] { + await receive(val, inp[i, j]) + } + } + phase { + dataflow i16 i, i16 j in [0:D + M, 0] { + stream fwd = relative_stream(D, 0) { + hops = auto, + channel = 0 + } + } + compute i16 i, i16 j in [0:M, 0] { + await send(val, fwd) + } + compute i16 i, i16 j in [D:D + M, 0] { + await receive(val, fwd) + } + } + phase { + compute i16 i, i16 j in [0:D + M, 0] { + await send(val, out[i, j]) + } + } +} +""" + +_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') + + +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)) + + +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 _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. + + 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_detects_the_overlapping_shift(): + bundles = _bundles(_SHIFT, M=3, D=3, K=1) + assert len(bundles) == 1 + 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_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_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_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_a_single_source_is_not_bundled(): + assert _bundles(_SHIFT, M=1, D=3, K=1) == [] + + +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 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(): + 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'] + # 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'] + 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(): + 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 _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)} + + +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, R=1) + 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_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', 'simple', 'exchange_bundle_1D.sptl' + ) + kernel = parser.parse_file(path) + 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] + assert pe_codes + for code in pe_codes: + # 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 + 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 + + +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_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() diff --git a/tests/spatial_ir/test_stream_lifetime.py b/tests/spatial_ir/test_stream_lifetime.py index 0b363c6e..f2c989c8 100644 --- a/tests/spatial_ir/test_stream_lifetime.py +++ b/tests/spatial_ir/test_stream_lifetime.py @@ -458,5 +458,95 @@ 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)') + + +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'})) + + +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)') + 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_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)') + + +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__]) diff --git a/tests/spatial_ir/test_task_recycling.py b/tests/spatial_ir/test_task_recycling.py index 9b026d8b..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 @@ -72,8 +73,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) @@ -255,3 +257,28 @@ 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. + + 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) + 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']) + + +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 = 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 diff --git a/tests/spatial_ir/test_task_recycling_codegen.py b/tests/spatial_ir/test_task_recycling_codegen.py index dc023cf0..763a1648 100644 --- a/tests/spatial_ir/test_task_recycling_codegen.py +++ b/tests/spatial_ir/test_task_recycling_codegen.py @@ -4,11 +4,93 @@ import pytest from spada.lowering.spatial_ir_to_csl import lower_spatial_ir_to_csl +from spada.syntax.csl import constants, 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( 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 _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) + 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( @@ -17,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] = {} @@ -52,15 +134,114 @@ 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' +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 a phase of its own and the fabric deadlocks behind the send it never made. + """ + csl_files = _scalar_exchange_chain() + + 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_(?: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 + 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, '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(): + """Several receives on one channel at one PE share the data task the channel binds. + + 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 + 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(hardware_ids), ( + f'{file.filename}: binds {len(bound)} data tasks for {len(hardware_ids)} hardware IDs') + + 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 + 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}: 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 when it is done') + + 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(): + """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) 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) @@ -69,6 +250,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)}' )