Skip to content

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

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

vchuravy wants to merge 5 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).

@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
@vchuravy
vchuravy added this pull request to stack #834 October 4, 2026 06:46
@vchuravy
vchuravy force-pushed the vc/groupreduce branch 6 times, most recently from 81e4230 to 06fdf57 Compare October 4, 2026 10:43
@vchuravy
vchuravy force-pushed the vc/groupreduce branch 3 times, most recently from 5f5c328 to a424d5d Compare October 4, 2026 12:31
vchuravy and others added 5 commits October 4, 2026 18:34
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