Skip to content

Add @groupreduce, @subgroupreduce, @groupscan and @subgroupscan - #830

Open
vchuravy wants to merge 20 commits into
vc/ki-subgroup-opsfrom
vc/groupreduce
Open

vchuravy wants to merge 20 commits into
vc/ki-subgroup-opsfrom
vc/groupreduce

Conversation

@vchuravy

@vchuravy vchuravy commented Oct 3, 2026 •

Copy link
Copy Markdown
Member

Supersedes #559 by @pxl-th (credited as co-author), reworked on top of KernelInterface.

Stacked on #831: @subgroupscan needs KI.shfl_up (and the struct shuffles) from there, so this PR targets that branch for now. Merge #831 first; this one then retargets to main.

API

  • @groupreduce(op, val, neutral[, groupsize]; subgroups = false) reduces val over the workgroup and returns the result on every work-item.
    • Default: tree reduction in local memory.
    • subgroups = true: each sub-group reduces with KI.shfl_down, then the sub-group results are combined (constant number of barriers). Gate it on the host with KI.supports_shuffle(backend, T) and pass it in as a constant (e.g. ::Val{S}).
    • Local memory is sized by the static workgroup size, or by groupsize as a compile-time upper bound for dynamic workgroup sizes.
  • @subgroupreduce(op, val, neutral) reduces over the sub-group with KI.sub_group_reduce (backends may use native reductions); the result is defined on every lane. It only needs op to be associative.
  • @groupscan(op, val, neutral[, groupsize]; inclusive = true) scans in the order of @index(Local, Linear), inclusive or (with inclusive = false) exclusive. It uses a Hillis-Steele scan in double-buffered local memory, sized like @groupreduce. There is no sub-group variant: how work-items form sub-groups is unspecified in KernelInterface, so a scan built from sub-group scans wouldn't follow the local index order.
  • @subgroupscan(op, val, neutral; inclusive = true) scans over the lanes of a sub-group with KI.sub_group_scan.
  • Both scans only need op to be associative. They are collectives like the reductions, so padding work-items contribute neutral. A use case is stream compaction: offsets from an exclusive @groupscan of the predicates (cf. Molly.jl's neighbor finder).

Differences to #559

  • shfl_down/supports_warp_reduction are gone: they're KI.shfl_down/KI.supports_shuffle now, and the sub-group width comes from KI instead of a hardcoded 32.
  • Partial workgroups: both macros are collectives in the @kernel split, like @synchronize. Padding work-items take part, contribute neutral and don't evaluate val. This fixes the uninitialized-local-memory issue raised in Implement groupreduce API #559, and is why neutral is required.
  • Using a collective inside a larger expression (y[i] = @groupreduce(...)) is a macro-expansion error, since it would run on padding work-items.
  • Non-power-of-two and multi-dimensional workgroups are supported.

Values of different types

The macros convert val to the type of neutral at the call site. Otherwise, a val whose type differs between work-items makes Julia union-split the call, and the work-items run different copies of its barriers and shuffles. Two cases trigger this:

  • the padding work-items' neutral having a different type than val (e.g. @groupreduce(+, x[i]::Float32, 0.0));
  • an accumulator that only some work-items promote.

On POCL this silently returned wrong results. It was found through a NaN in the Molly.jl port.

Tests

test/groupreduce.jl runs as part of the backend testsuite and covers:

  • both algorithms
  • +/max on Int32, Int64 and Float32
  • partial workgroups, dynamic workgroup size with a bound, non-power-of-two sizes
  • reuse in a loop, Cartesian workgroups, unsafe_indices
  • @groupscan (inclusive/exclusive): partial and non-power-of-two workgroups, a dynamic workgroup size with a bound, Cartesian workgroups with padding in the middle, reuse in a loop
  • @subgroupscan, checked against the lane order the kernel reports
  • the scans use the composition of affine maps, which is associative but not commutative, so a wrong combining order fails
  • @groupreduce of (value, index) pairs (argmin)
  • @subgroupreduce on every lane, without assuming that sub-groups are consecutive work-items, and the macro errors

The full test suite passes locally on POCL.

🤖 Generated with Claude Code

@github-actions

github-actions Bot commented Oct 3, 2026 •

Copy link
Copy Markdown
Contributor

Benchmark Results

Show table
main 9b20320... main / 9b20320...
const/@Const/Float32/262144 0.31 ± 0.0099 ms 0.311 ± 0.0096 ms 0.996 ± 0.044
const/@Const/Float32/65536 0.107 ± 0.003 ms 0.106 ± 0.0038 ms 1.02 ± 0.047
const/@Const/Float64/262144 0.589 ± 0.013 ms 0.589 ± 0.012 ms 1 ± 0.03
const/@Const/Float64/65536 0.182 ± 0.0046 ms 0.182 ± 0.0046 ms 0.999 ± 0.036
const/unmarked/Float32/262144 0.509 ± 0.014 ms 0.591 ± 0.012 ms 0.861 ± 0.029
const/unmarked/Float32/65536 0.15 ± 0.0057 ms 0.154 ± 0.0059 ms 0.976 ± 0.053
const/unmarked/Float64/262144 0.981 ± 0.0051 ms 0.981 ± 0.008 ms 1 ± 0.0097
const/unmarked/Float64/65536 0.263 ± 0.0091 ms 0.269 ± 0.01 ms 0.977 ± 0.05
launch/3D static workgroup, dynamic ndrange 11 ± 2.4 μs 9.11 ± 0.39 μs 1.21 ± 0.27
launch/3D static workgroup, static ndrange 8.74 ± 0.92 μs 8.82 ± 0.26 μs 0.991 ± 0.11
launch/dynamic workgroup, dynamic ndrange 9.89 ± 0.13 μs 10 ± 1.9 μs 0.99 ± 0.19
launch/dynamic workgroup, dynamic ndrange, workgroupsize given 9.78 ± 2 μs 8.14 ± 0.2 μs 1.2 ± 0.25
launch/static workgroup, dynamic ndrange 8.78 ± 2.1 μs 10.9 ± 0.2 μs 0.803 ± 0.2
launch/static workgroup, static ndrange 8.5 ± 0.2 μs 8.96 ± 2.3 μs 0.949 ± 0.24
partition/dynamic workgroup, dynamic ndrange 0.0471 ± 0.0036 μs 0.0458 ± 0.0021 μs 1.03 ± 0.091
partition/static workgroup, dynamic ndrange 0.0474 ± 0.011 μs 0.0473 ± 0.011 μs 1 ± 0.33
partition/static workgroup, static ndrange 1.35 ± 0 ns 2.15 ± 0.01 ns 0.628 ± 0.0029
saxpy/default/Float16/1024 9.49 ± 0.21 μs 10.6 ± 2.2 μs 0.895 ± 0.19
saxpy/default/Float16/1048576 0.266 ± 0.016 ms 0.264 ± 0.015 ms 1.01 ± 0.083
saxpy/default/Float16/16384 0.034 ± 0.021 ms 0.0339 ± 0.019 ms 1 ± 0.84
saxpy/default/Float16/2048 9.73 ± 0.22 μs 9.84 ± 0.19 μs 0.989 ± 0.03
saxpy/default/Float16/256 11 ± 2.2 μs 9.28 ± 2.3 μs 1.18 ± 0.38
saxpy/default/Float16/262144 0.0941 ± 0.005 ms 0.0942 ± 0.005 ms 0.999 ± 0.075
saxpy/default/Float16/32768 0.0387 ± 0.0029 ms 0.0382 ± 0.0029 ms 1.01 ± 0.11
saxpy/default/Float16/4096 10.2 ± 0.25 μs 10.4 ± 0.34 μs 0.979 ± 0.04
saxpy/default/Float16/512 9.17 ± 0.31 μs 9.27 ± 0.73 μs 0.99 ± 0.085
saxpy/default/Float16/64 8.9 ± 0.82 μs 9.03 ± 2.3 μs 0.985 ± 0.26
saxpy/default/Float16/65536 0.0457 ± 0.003 ms 0.0474 ± 0.0028 ms 0.963 ± 0.085
saxpy/default/Float32/1024 11.5 ± 0.53 μs 11.6 ± 0.56 μs 0.99 ± 0.066
saxpy/default/Float32/1048576 0.228 ± 0.02 ms 0.219 ± 0.017 ms 1.04 ± 0.12
saxpy/default/Float32/16384 15.9 ± 1.3 μs 15.8 ± 0.66 μs 1.01 ± 0.095
saxpy/default/Float32/2048 9.62 ± 0.25 μs 10.2 ± 2 μs 0.947 ± 0.19
saxpy/default/Float32/256 11.2 ± 0.27 μs 9.12 ± 0.25 μs 1.23 ± 0.045
saxpy/default/Float32/262144 0.0843 ± 0.0066 ms 0.0832 ± 0.0073 ms 1.01 ± 0.12
saxpy/default/Float32/32768 0.0382 ± 0.0031 ms 0.0382 ± 0.0033 ms 0.999 ± 0.12
saxpy/default/Float32/4096 9.96 ± 1.4 μs 10.1 ± 1.6 μs 0.981 ± 0.21
saxpy/default/Float32/512 9.35 ± 0.27 μs 11.6 ± 1.8 μs 0.804 ± 0.13
saxpy/default/Float32/64 11.1 ± 0.25 μs 11.4 ± 0.2 μs 0.971 ± 0.028
saxpy/default/Float32/65536 0.0465 ± 0.0069 ms 0.0468 ± 0.0053 ms 0.992 ± 0.18
saxpy/default/Float64/1024 11.2 ± 1.9 μs 9.81 ± 0.52 μs 1.14 ± 0.21
saxpy/default/Float64/1048576 0.505 ± 0.038 ms 0.498 ± 0.042 ms 1.01 ± 0.11
saxpy/default/Float64/16384 0.0389 ± 0.0036 ms 0.0383 ± 0.0027 ms 1.02 ± 0.12
saxpy/default/Float64/2048 10.3 ± 1.8 μs 10.1 ± 0.24 μs 1.01 ± 0.18
saxpy/default/Float64/256 9.15 ± 2.1 μs 11.6 ± 2 μs 0.791 ± 0.23
saxpy/default/Float64/262144 0.123 ± 0.0098 ms 0.123 ± 0.01 ms 0.995 ± 0.11
saxpy/default/Float64/32768 0.0486 ± 0.0043 ms 0.0468 ± 0.0059 ms 1.04 ± 0.16
saxpy/default/Float64/4096 14 ± 1.2 μs 13.6 ± 1.3 μs 1.03 ± 0.13
saxpy/default/Float64/512 9.54 ± 2 μs 11.7 ± 0.25 μs 0.816 ± 0.17
saxpy/default/Float64/64 10.9 ± 2.3 μs 9.74 ± 2.4 μs 1.12 ± 0.36
saxpy/default/Float64/65536 0.0603 ± 0.0047 ms 0.0617 ± 0.0067 ms 0.977 ± 0.13
saxpy/static workgroup=(1024,)/Float16/1024 9.51 ± 0.35 μs 9.64 ± 0.33 μs 0.986 ± 0.05
saxpy/static workgroup=(1024,)/Float16/1048576 0.268 ± 0.016 ms 0.266 ± 0.014 ms 1.01 ± 0.081
saxpy/static workgroup=(1024,)/Float16/16384 0.0341 ± 0.018 ms 0.0344 ± 0.019 ms 0.993 ± 0.76
saxpy/static workgroup=(1024,)/Float16/2048 9.82 ± 2.3 μs 11.9 ± 0.53 μs 0.823 ± 0.19
saxpy/static workgroup=(1024,)/Float16/256 9.27 ± 0.34 μs 9.45 ± 2.3 μs 0.982 ± 0.24
saxpy/static workgroup=(1024,)/Float16/262144 0.0951 ± 0.005 ms 0.0952 ± 0.0053 ms 1 ± 0.076
saxpy/static workgroup=(1024,)/Float16/32768 0.0389 ± 0.0024 ms 0.0383 ± 0.0032 ms 1.02 ± 0.1
saxpy/static workgroup=(1024,)/Float16/4096 12 ± 0.28 μs 10.7 ± 2 μs 1.13 ± 0.21
saxpy/static workgroup=(1024,)/Float16/512 9.25 ± 0.29 μs 9.34 ± 0.26 μs 0.991 ± 0.041
saxpy/static workgroup=(1024,)/Float16/64 11.5 ± 0.18 μs 9.45 ± 3.4 μs 1.22 ± 0.44
saxpy/static workgroup=(1024,)/Float16/65536 0.0473 ± 0.0034 ms 0.0468 ± 0.0038 ms 1.01 ± 0.11
saxpy/static workgroup=(1024,)/Float32/1024 9.55 ± 1.9 μs 9.55 ± 0.21 μs 1 ± 0.2
saxpy/static workgroup=(1024,)/Float32/1048576 0.219 ± 0.017 ms 0.216 ± 0.017 ms 1.01 ± 0.11
saxpy/static workgroup=(1024,)/Float32/16384 15.8 ± 2.1 μs 15.8 ± 3.9 μs 0.998 ± 0.28
saxpy/static workgroup=(1024,)/Float32/2048 9.62 ± 0.2 μs 10.1 ± 2 μs 0.956 ± 0.19
saxpy/static workgroup=(1024,)/Float32/256 11.4 ± 1.6 μs 9.24 ± 0.37 μs 1.24 ± 0.18
saxpy/static workgroup=(1024,)/Float32/262144 0.0832 ± 0.0067 ms 0.0831 ± 0.0064 ms 1 ± 0.11
saxpy/static workgroup=(1024,)/Float32/32768 0.0382 ± 0.0029 ms 0.0383 ± 0.0049 ms 1 ± 0.15
saxpy/static workgroup=(1024,)/Float32/4096 9.75 ± 0.36 μs 9.97 ± 0.23 μs 0.977 ± 0.043
saxpy/static workgroup=(1024,)/Float32/512 9.64 ± 2.3 μs 11.8 ± 0.28 μs 0.82 ± 0.19
saxpy/static workgroup=(1024,)/Float32/64 9.18 ± 1.1 μs 8.31 ± 1.2 μs 1.1 ± 0.21
saxpy/static workgroup=(1024,)/Float32/65536 0.048 ± 0.0047 ms 0.0463 ± 0.0059 ms 1.04 ± 0.17
saxpy/static workgroup=(1024,)/Float64/1024 9.62 ± 0.22 μs 9.79 ± 0.19 μs 0.983 ± 0.029
saxpy/static workgroup=(1024,)/Float64/1048576 0.516 ± 0.05 ms 0.477 ± 0.037 ms 1.08 ± 0.14
saxpy/static workgroup=(1024,)/Float64/16384 0.0378 ± 0.0028 ms 0.0381 ± 0.0027 ms 0.994 ± 0.1
saxpy/static workgroup=(1024,)/Float64/2048 10 ± 1.8 μs 10.9 ± 1.8 μs 0.914 ± 0.23
saxpy/static workgroup=(1024,)/Float64/256 11.9 ± 0.31 μs 11.8 ± 2.2 μs 1.01 ± 0.19
saxpy/static workgroup=(1024,)/Float64/262144 0.123 ± 0.011 ms 0.123 ± 0.0095 ms 1 ± 0.12
saxpy/static workgroup=(1024,)/Float64/32768 0.0466 ± 0.0059 ms 0.0463 ± 0.0061 ms 1.01 ± 0.18
saxpy/static workgroup=(1024,)/Float64/4096 13.8 ± 1.2 μs 10.8 ± 1.5 μs 1.28 ± 0.21
saxpy/static workgroup=(1024,)/Float64/512 9.63 ± 0.45 μs 11.9 ± 0.33 μs 0.806 ± 0.044
saxpy/static workgroup=(1024,)/Float64/64 9.32 ± 1 μs 9.26 ± 1.2 μs 1.01 ± 0.17
saxpy/static workgroup=(1024,)/Float64/65536 0.0612 ± 0.0055 ms 0.0612 ± 0.0062 ms 1 ± 0.14
time_to_load 0.411 ± 0.0021 s 0.403 ± 0.01 s 1.02 ± 0.027
main 9b20320... main / 9b20320...
const/@Const/Float32/262144 9 allocs: 0.203 kB 9 allocs: 0.203 kB 1
const/@Const/Float32/65536 9 allocs: 0.203 kB 9 allocs: 0.203 kB 1
const/@Const/Float64/262144 9 allocs: 0.203 kB 9 allocs: 0.203 kB 1
const/@Const/Float64/65536 9 allocs: 0.203 kB 9 allocs: 0.203 kB 1
const/unmarked/Float32/262144 9 allocs: 0.203 kB 9 allocs: 0.203 kB 1
const/unmarked/Float32/65536 5 allocs: 0.0938 kB 9 allocs: 0.203 kB 0.462
const/unmarked/Float64/262144 9 allocs: 0.203 kB 9 allocs: 0.203 kB 1
const/unmarked/Float64/65536 9 allocs: 0.203 kB 9 allocs: 0.203 kB 1
launch/3D static workgroup, dynamic ndrange 9 allocs: 0.219 kB 9 allocs: 0.219 kB 1
launch/3D static workgroup, static ndrange 9 allocs: 0.219 kB 9 allocs: 0.219 kB 1
launch/dynamic workgroup, dynamic ndrange 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
launch/dynamic workgroup, dynamic ndrange, workgroupsize given 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
launch/static workgroup, dynamic ndrange 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
launch/static workgroup, static ndrange 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
partition/dynamic workgroup, dynamic ndrange 2 allocs: 0.0625 kB 2 allocs: 0.0625 kB 1
partition/static workgroup, dynamic ndrange 2 allocs: 32 B 2 allocs: 32 B 1
partition/static workgroup, static ndrange 0 allocs: 0 B 0 allocs: 0 B
saxpy/default/Float16/1024 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float16/1048576 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/default/Float16/16384 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float16/2048 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float16/256 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/default/Float16/262144 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/default/Float16/32768 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float16/4096 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float16/512 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float16/64 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/default/Float16/65536 8 allocs: 0.141 kB 12 allocs: 0.25 kB 0.562
saxpy/default/Float32/1024 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float32/1048576 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/default/Float32/16384 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float32/2048 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float32/256 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/default/Float32/262144 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/default/Float32/32768 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float32/4096 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float32/512 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float32/64 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/default/Float32/65536 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float64/1024 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float64/1048576 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/default/Float64/16384 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float64/2048 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float64/256 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/default/Float64/262144 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/default/Float64/32768 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float64/4096 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float64/512 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/default/Float64/64 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/default/Float64/65536 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/static workgroup=(1024,)/Float16/1024 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float16/1048576 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/static workgroup=(1024,)/Float16/16384 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float16/2048 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float16/256 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/static workgroup=(1024,)/Float16/262144 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/static workgroup=(1024,)/Float16/32768 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float16/4096 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float16/512 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float16/64 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/static workgroup=(1024,)/Float16/65536 8 allocs: 0.141 kB 12 allocs: 0.25 kB 0.562
saxpy/static workgroup=(1024,)/Float32/1024 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float32/1048576 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/static workgroup=(1024,)/Float32/16384 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float32/2048 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float32/256 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/static workgroup=(1024,)/Float32/262144 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/static workgroup=(1024,)/Float32/32768 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float32/4096 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float32/512 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float32/64 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/static workgroup=(1024,)/Float32/65536 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float64/1024 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float64/1048576 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/static workgroup=(1024,)/Float64/16384 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float64/2048 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float64/256 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/static workgroup=(1024,)/Float64/262144 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
saxpy/static workgroup=(1024,)/Float64/32768 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float64/4096 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float64/512 8 allocs: 0.141 kB 8 allocs: 0.141 kB 1
saxpy/static workgroup=(1024,)/Float64/64 5 allocs: 0.0938 kB 5 allocs: 0.0938 kB 1
saxpy/static workgroup=(1024,)/Float64/65536 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
time_to_load 0.2 k allocs: 11.8 kB 0.2 k allocs: 11.8 kB 1

Benchmark Plots

A plot of the benchmark results have been uploaded as an artifact to the workflow run for this PR.
Go to "Actions"->"Benchmark a pull request"->[the most recent run]->"Artifacts" (at the bottom).

Add the sub-group operations that e.g. Molly.jl's CUDA kernels use, so that they
can be written portably:

- shuffles `shfl` (from a given lane), `shfl_up` and `shfl_xor`, next to
  `shfl_down`. Backends implement them for primitive types; a fallback shuffles
  `isbits` structs and tuples field by field, and `supports_shuffle` checks
  their fields.
- votes `sub_group_any`, `sub_group_all` and `sub_group_ballot` (a `UInt64`
  mask, for sub-groups of at most 64 work-items), required with sub-group
  support.
- `get_max_sub_group_size` is now required to be a constant of the generated
  code.

Implement them for POCL; its sub-group width is folded into the IR before
optimization.

Assisted-by: Claude Code (Opus 5.5)
@codecov

codecov Bot commented Oct 3, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 96.89441% with 5 lines in your changes missing coverage. Please review.
✅ Project coverage is 75.70%. Comparing base (c3e0b7a) to head (9b20320).

Files with missing lines Patch % Lines
src/macros.jl 69.23% 4 Missing ⚠️
src/groupreduction.jl 99.32% 1 Missing ⚠️
Additional details and impacted files
@@                  Coverage Diff                   @@
##           vc/ki-subgroup-ops     #830      +/-   ##
======================================================
+ Coverage               74.15%   75.70%   +1.55%     
======================================================
  Files                      24       25       +1     
  Lines                    2275     2437     +162     
======================================================
+ Hits                     1687     1845     +158     
- Misses                    588      592       +4     

☔ View full report in Codecov by Harness.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@vchuravy vchuravy changed the title Add @groupreduce and @subgroupreduce Add @groupreduce, @subgroupreduce, @groupscan and @subgroupscan Oct 3, 2026
@vchuravy
vchuravy changed the base branch from main to vc/ki-subgroup-ops October 3, 2026 22:49
On Julia 1.10, inference gives up on the recursive call of `shfl_fields` through
the shuffle of a nested field (e.g. a tuple in a struct), leaving a dynamic
invocation in the kernel. Generate the shuffles of all primitive fields
directly instead.

Assisted-by: Claude Code (Opus 5.5)
@vchuravy
vchuravy added this pull request to stack #834 October 4, 2026 06:46
Replace the local workarounds with the votes and unchecked shuffle lanes from
JuliaGPU/OpenCL.jl#526, taken from its branch until it is released: through
`[sources]`, and explicitly where that doesn't apply (Julia 1.10 on CI, and the
Buildkite jobs, whose OpenCL job developed SPIRVIntrinsics from OpenCL.jl's
ka-0.10 branch).

Assisted-by: Claude Code (Opus 5.5)
…ce and scan

Fill the gaps that a survey of the packages using warp operations (KomaMRI,
KernelIntrinsics/KernelForge, AcceleratedKernels#93, ParallelStencil, ClimaCore,
...) showed:

- Primitive types that a backend doesn't support natively (e.g. `Bool`,
  `Char`, or 64-bit types on Metal) are shuffled as `UInt32` words, which
  backends now have to support. Structs keep being shuffled field by field.
- Shuffles within segments of `width` lanes (`shfl(val, lane, width)` etc.),
  with CUDA's semantics, built on `shfl`.
- `sub_group_match_any(val)`, the mask of the lanes with the same value, with a
  fallback built on `shfl` and `sub_group_ballot`.
- `sub_group_reduce(op, val)` and `sub_group_scan(op, val)` with fallbacks
  built on the shuffles, which backends can implement with native operations.
- Document how partial sub-groups behave.

The new tests are in a function of their own: as part of `interface_testsuite`,
compiling the host code crashed LLVM.

Assisted-by: Claude Code (Opus 5.5)
Implement `KI.sub_group_reduce` and `KI.sub_group_scan` with the collectives of
`cl_khr_subgroups` from SPIRVIntrinsics (JuliaGPU/OpenCL.jl#526): for `+` on
32- and 64-bit integers and floats, and `min`/`max` on integers. Floats keep the
fallback for `min` and `max`, as OpenCL treats NaN and the sign of zero
differently.

Test the operators and types that backends may implement natively, including a
NaN, and that POCL uses the native reduction.

Assisted-by: Claude Code (Opus 5.5)
PoCL's `cl_khr_subgroups` reductions and scans lose the values of work-items
that computed them in a divergent branch (PoCL 7.2; Intel's OpenCL runtime is
fine), so `@groupreduce` in a `@kernel`, whose padding work-items are masked,
returned garbage. Use KernelInterface's fallbacks again, and test reductions
and scans of values from a divergent branch.

Assisted-by: Claude Code (Opus 5.5)
Like the shuffles, the votes, `sub_group_match_any`, `sub_group_reduce` and
`sub_group_scan` exchange values, not memory; communicating through memory
within a sub-group needs `sub_group_barrier`.

Assisted-by: Claude Code (Opus 5.5)
The wrong results of the native `cl_khr_subgroups` collectives had two causes:
SPIRVIntrinsics declared them without `convergent`, so LLVM duplicated the
calls into divergent branches (fixed in JuliaGPU/OpenCL.jl#526), and PoCL 7.2
gives a peeled work-item its own copy of a collective's scratch memory after a
branch with an early exit, as bounds checks emit (fixed on PoCL's main branch,
backport to 7.2 in pocl/pocl#2373).

With both fixed, i.e. with `POCL_WORK_GROUP_METHOD=cbs` for now, the native
reductions and scans pass the tests. Keep them behind `NATIVE_COLLECTIVES`
until `pocl_standalone_jll` includes the PoCL fix.

Assisted-by: Claude Code (Opus 5.5)
Assisted-by: Claude Code (Opus 5.5)
`sub_group_any`, `sub_group_all`, `sub_group_ballot` and `sub_group_match_any`
with a `width`, like the shuffles: the votes of segments of `width` lanes, with
masks that have a bit per lane of the segment. Fallbacks use the ballot and
match of the whole sub-group. With the shuffles with a width, this allows e.g.
tiles of 32 lanes on sub-groups of 64.

Assisted-by: Claude Code (Opus 5.5)
… the width

- If a work-group is 1-D or its x extent is a multiple of the sub-group width,
  sub-groups are formed from consecutive work-items, x fastest. This holds on
  CUDA, AMD, PoCL, rusticl and Intel's CPU OpenCL runtime (which forms
  sub-groups per row otherwise), and lets kernels with e.g. (32, 8)
  work-groups rely on the layout.
- `shfl_down` and `shfl_up` return the work-item's own value where the source
  lane is past the sub-group width, like CUDA's shuffles. That's free on CUDA
  and Metal, and a select on AMD and SPIR-V (POCL), and makes them consistent
  with the shuffles with a `width`. `shfl_xor` requires a mask below the width.

Test both.

Assisted-by: Claude Code (Opus 5.5)
Porting KomaMRI.jl showed the fallback of `sub_group_reduce` to cost ~40% of a
reduction-heavy kernel on CUDA: its loop ran to the run-time sub-group size, so it
didn't unroll, and every lane was masked.

`sub_group_reduce` is now an ordered butterfly with `shfl_xor` over the constant
sub-group width (combining the lower block first, so it still only needs
associativity), which gives every work-item of a full sub-group the result
without a broadcast. Partial sub-groups skip the blocks without work-items and
broadcast the result of the first lane. All sub-groups run the same shuffles,
which PoCL needs; document that requirement of the POCL backend. `sub_group_scan`
loops to the constant width as well.

Test reductions and scans in a work-group of several sub-groups, the last one
partial.

Assisted-by: Claude Code (Opus 5.5)
The new fallbacks of `sub_group_reduce` and `sub_group_scan` loop over the
constant sub-group width. PoCL 7.2 miscompiles the unrolled shuffles after a
branch with an early exit, as bounds-checked `@kernel`s have (the bug fixed by
pocl/pocl#2239), so `@groupreduce` with sub-groups gave wrong results. Override
them for POCL with a loop bounded by the smaller of the width and the work-group
size, which doesn't unroll and is the same for all sub-groups of a work-group.

Assisted-by: Claude Code (Opus 5.5)
…ork-items

Assisted-by: Claude Code (Opus 5.5)
vchuravy and others added 6 commits October 4, 2026 18:33
….2.1+1

pocl_standalone_jll 7.2.1+1 includes the WorkitemLoops fix (pocl/pocl#2239,
JuliaPackaging/Yggdrasil#15001), so implement `KI.sub_group_reduce` and
`KI.sub_group_scan` with the native `cl_khr_subgroups` collectives where they
have Julia's semantics, and only keep the workaround for the fallbacks with
older builds. The version is checked at precompile time, since the compat bound
can't distinguish builds.

Note: PoCL's kernel cache doesn't distinguish 7.2.1+0 and +1, so binaries
miscompiled by the former can be reused by the latter until the cache is
cleared (pocl/pocl#2374, JuliaPackaging/Yggdrasil#15002).

Assisted-by: Claude Code (Opus 5.5)
Rework of #559 on top of KernelInterface:

- `@groupreduce(op, val, neutral[, groupsize]; subgroups=false)` reduces over the
  workgroup and returns the result on every work-item. It uses a local-memory tree
  by default, or a two-level reduction based on `KI.shfl_down` with
  `subgroups=true` (gated on the host by `KI.supports_shuffle`). The local memory
  is sized by the static workgroup size or an explicit upper bound.
- `@subgroupreduce(op, val, neutral)` reduces over the sub-group with shuffles;
  the result is defined on the first lane.
- Both are collectives in `@kernel`: the split treats them like `@synchronize`,
  and padding work-items contribute `neutral` without evaluating `val`, so
  ndranges that are not a multiple of the workgroup size work.

Co-authored-by: Anton Smirnov <tonysmn97@gmail.com>
Assisted-by: Claude Code (Opus 5.5)
- `@groupscan(op, val, neutral[, groupsize]; inclusive = true)` scans over the
  workgroup in the order of the local linear index, with a Hillis-Steele scan in
  double-buffered local memory, sized like `@groupreduce`.
- `@subgroupscan(op, val, neutral; inclusive = true)` scans over the lanes of a
  sub-group with `KI.shfl_up`.

Both only need `op` to be associative, and are collectives in `@kernel` like the
reductions: padding work-items contribute `neutral`.

Assisted-by: Claude Code (Opus 5.5)
`@subgroupreduce` and `@subgroupscan`, and the sub-group stage of
`@groupreduce`, now use `KI.sub_group_reduce` and `KI.sub_group_scan`, which
backends can implement with native operations. `@subgroupreduce` returns the
result on every work-item of the sub-group.

Test reductions of (value, index) pairs, and don't assume that sub-groups are
formed from consecutive work-items.

Assisted-by: Claude Code (Opus 5.5)
KernelAbstractions' collectives are executed by the padding work-items of a
partial workgroup, but direct calls of KernelInterface's sub-group functions
aren't: kernels that use them need `unsafe_indices=true`.

Assisted-by: Claude Code (Opus 5.5)
The type of the value passed to `@groupreduce`, `@subgroupreduce`,
`@groupscan` or `@subgroupscan` may differ between the work-items: in a
`@kernel` the padding work-items contribute `neutral` instead of `val`, so
`@groupreduce(+, x[i]::Float32, 0.0)` reduces a `Union{Float32, Float64}`,
and an accumulator that only some work-items add a `Float64` to is a `Union`
as well. Julia union-splits the call of the collective with such an argument
into one call per type, so the work-items of a workgroup executed different
copies of its barriers and shuffles. On PoCL this silently gave wrong
results (0 for the reduction of a padded workgroup, NaN for the energy of a
Float32 system with a Float64 Coulomb constant in Molly).

Convert the value to the type of `neutral` at the call site instead, so that
only the conversion is union-split, and test mixed types for all four
collectives.

Assisted-by: Claude Code (Opus 5.5)

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant