Skip to content

KernelInterface: more sub-group communication (shuffles, votes, constant width) - #831

Open
vchuravy wants to merge 18 commits into
mainfrom
vc/ki-subgroup-ops
Open

vchuravy wants to merge 18 commits into
mainfrom
vc/ki-subgroup-ops

Conversation

@vchuravy

@vchuravy vchuravy commented Oct 3, 2026 •

Copy link
Copy Markdown
Member

Adds the sub-group primitives that Molly.jl's CUDA kernels use (ext/MollyCUDAExt.jl), so that they can be written portably on top of KernelInterface. This is a companion to #830 (@groupreduce/@subgroupreduce).

New device functions

Function Requirement Molly.jl use
shfl(val, lane) supports_shuffle(backend, T) rotating j-atom data around the warp in the kernel without a neighbor list
shfl_up(val, offset) supports_shuffle(backend, T) building block for sub-group scans (its neighbor finder builds a prefix count by hand)
shfl_xor(val, mask) supports_shuffle(backend, T) butterfly all-reduce (min/max/|) of tile masks
sub_group_any(pred) / sub_group_all(pred) supports_subgroups uniform skip of pair evaluations (vote_any_sync)
sub_group_ballot(pred)::UInt64 supports_subgroups, width ≤ 64 interaction bitmasks; count_ones/trailing_zeros on the result
  • Struct shuffles: backends implement the shuffles for primitive types only, with signatures that only match those. A generic fallback shuffles isbits structs and tuples field by field (e.g. SVector, Unitful quantities). For such types, supports_shuffle checks their fields. Molly currently gets this by overriding CUDA.jl's internal shfl_recurse.
  • Contract change: supports_shuffle and the backend implementation notes now cover all four shuffles.

Warp size

get_max_sub_group_size() is now required to be a compile-time constant of the generated code. Loops over lanes and shuffle butterflies get specialized for it, which every backend can provide since kernels are already compiled for a fixed width (sub_group_size(backend)). For type-level decisions (e.g. a UInt32 vs UInt64 mask) the docs point to passing sub_group_size(backend) from the host.

POCL

  • Shuffles: implemented with SPIRVIntrinsics' sub_group_shuffle/sub_group_shuffle_xor. Out-of-range lanes give an unspecified value instead of an InexactError.
  • Votes: implemented with SPIRVIntrinsics' sub_group_any/sub_group_all/sub_group_ballot (cl_khr_subgroups, cl_khr_subgroup_ballot).
  • Constant width: finish_module! replaces loads of __spirv_BuiltInSubgroupMaxSize with the width the kernel is compiled for (intel_reqd_sub_group_size), before optimization.

Depends on JuliaGPU/OpenCL.jl#526, which adds the votes and unchecked shuffle lanes to SPIRVIntrinsics. Until it's released, this PR takes SPIRVIntrinsics from that branch:

  • through [sources] in Project.toml;
  • explicitly on Julia 1.10 CI, which ignores [sources];
  • in the Buildkite jobs, where the OpenCL job used to develop SPIRVIntrinsics from OpenCL.jl's ka-0.10 branch.

The branch has the LLVM 10 upgrade that ka-0.10 lacks. Before merging, these should be replaced by a compat bound on the release; all of them are marked TODO.

Built on the above (fallbacks, backends may override)

Added after a survey of packages that use warp operations: KomaMRI, KernelIntrinsics/KernelForge, AcceleratedKernels#93, ParallelStencil, ClimaCore, IntervalMDP, …

Function Fallback
shuffles of other primitive types (Bool, Char, 64-bit types on Metal, …) shuffled as UInt32 words. Backends now have to support UInt32 natively.
shfl(val, lane, width), shfl_down/up/xor(val, x, width) segments of width lanes with CUDA's semantics (reads from outside the segment return the own value), built on shfl
sub_group_match_any(val)::UInt64 the lanes with the same value (===), found group by group with shfl + sub_group_ballot. CUDA could use match.any.sync
sub_group_any/all(pred, width), sub_group_ballot(pred, width), sub_group_match_any(val, width) votes within segments of width lanes (masks have a bit per lane of the segment), built on the ballot/match of the whole sub-group. Together with the shuffles with a width, this lets e.g. 32-lane tiles run on 64-wide sub-groups
sub_group_reduce(op, val) ordered tree with shfl_down, then broadcast with shfl. Backends can dispatch on typeof(op) for native reductions (Metal simd_sum, SPIR-V GroupNonUniformIAdd, CUDA redux.sync)
sub_group_scan(op, val) inclusive Hillis-Steele scan with shfl_up

The docs now also say how partial sub-groups behave: lanes without a work-item give unspecified shuffle values, and the votes, match, reduce and scan only take the existing work-items into account.

The new tests are in small functions of their own, subgroup_communication_testsuite and helpers, with concrete loops. An earlier version iterated over heterogeneous tuples of functions and passed the resulting union-of-singletons value through KI.@launch, whose GC.@preserve then hit JuliaLang/julia#63482: on 1.12 and 1.13, codegen emits a null gc_preserve_begin operand, and LLVM segfaults. That is fixed on master by #63483, but the backport to 1.12/1.13 is still pending.

Guarantees added after porting Molly, KomaMRI and ParallelStencil

Follows from a cross-backend investigation with experiments on CUDA, AMD (wave32 and wave64), PoCL, rusticl and Intel's CPU OpenCL runtime.

  • Sub-group layout. If a work-group is 1-D, or its x extent is a multiple of the sub-group width W, sub-groups are formed from consecutive work-items, x fastest: sub-group (lin - 1) ÷ W + 1, lane (lin - 1) % W + 1.
    • This holds on CUDA (documented), AMD, PoCL, rusticl and Intel's CPU runtime. Intel's CPU runtime forms sub-groups per row for other shapes.
    • Still to check on Metal and Intel GPUs.
    • It lets kernels with e.g. (32, 8) work-groups (ParallelStencil's default) rely on the layout.
  • shfl_down/shfl_up past the sub-group width return the work-item's own value, as CUDA's do. That makes them consistent with the shuffles with a width.
    • Cost: free on CUDA and Metal, a select on AMD and SPIR-V.
    • shfl_xor requires mask < W.
  • Faster sub_group_reduce fallback. It is now an ordered shfl_xor butterfly over the constant width, so it unrolls and needs no broadcast in full sub-groups. KomaMRI measured the previous fallback at ~40% of a reduction-heavy kernel on CUDA. All sub-groups run the same shuffles. sub_group_scan also loops over the constant width.
  • POCL limitations:
    • PoCL implements sub-group operations with work-group barriers, so all sub-groups of a work-group have to run the same sub-group operations. This is now documented in POCLBackend.
    • Until pocl_standalone_jll includes the PoCL fix ([pocl] Backport upstream PR #2239 JuliaPackaging/Yggdrasil#15001), POCL overrides the reduce/scan fallbacks with a loop it compiles correctly in bounds-checked kernels.

For backend packages

KernelInterface 0.4 is still unreleased, so this extends its contract. CUDA.jl, AMDGPU.jl, oneAPI.jl, Metal.jl and OpenCL.jl will need to:

  • implement shfl, shfl_up, shfl_xor, sub_group_any, sub_group_all and sub_group_ballot;
  • narrow their shfl_down overrides to the primitive types, so that struct shuffles reach the fallback;
  • return a constant from get_max_sub_group_size (e.g. 32 % T on CUDA).

One open question: sub_group_ballot is required for widths ≤ 64. If a backend can't support ballot, it could get its own capability query instead.

Tests

  • KI testsuite:
    • rotation via shfl, shfl_up lanes and a shfl_xor butterfly all-reduce for every supported type
    • struct shuffles and supports_shuffle for structs
    • any/all/ballot with several predicate patterns
    • get_max_sub_group_size returns the width
  • Stub tests: the shuffle fallback throws for primitive types.
  • POCL test: the width is a constant in the IR.
  • Results: the full test suite passes locally on POCL (KI on POCL: 568/568).

🤖 Generated with Claude Code

@github-actions

github-actions Bot commented Oct 3, 2026 •

Copy link
Copy Markdown
Contributor

Benchmark Results

Show table
main 64ac668... main / 64ac668...
const/@Const/Float32/262144 0.312 ± 0.013 ms 0.322 ± 0.015 ms 0.968 ± 0.061
const/@Const/Float32/65536 0.109 ± 0.0036 ms 0.105 ± 0.0057 ms 1.04 ± 0.066
const/@Const/Float64/262144 0.588 ± 0.012 ms 0.588 ± 0.014 ms 0.999 ± 0.031
const/@Const/Float64/65536 0.181 ± 0.0098 ms 0.182 ± 0.0091 ms 0.996 ± 0.073
const/unmarked/Float32/262144 0.512 ± 0.013 ms 0.602 ± 0.014 ms 0.85 ± 0.029
const/unmarked/Float32/65536 0.151 ± 0.0094 ms 0.152 ± 0.0088 ms 0.993 ± 0.084
const/unmarked/Float64/262144 0.996 ± 0.017 ms 0.988 ± 0.019 ms 1.01 ± 0.026
const/unmarked/Float64/65536 0.266 ± 0.012 ms 0.266 ± 0.012 ms 1 ± 0.063
launch/3D static workgroup, dynamic ndrange 10.9 ± 0.42 μs 9.03 ± 0.3 μs 1.21 ± 0.062
launch/3D static workgroup, static ndrange 8.65 ± 0.32 μs 10.5 ± 2.1 μs 0.824 ± 0.17
launch/dynamic workgroup, dynamic ndrange 7.91 ± 0.3 μs 7.94 ± 0.18 μs 0.996 ± 0.044
launch/dynamic workgroup, dynamic ndrange, workgroupsize given 8.01 ± 1.9 μs 8.26 ± 1.8 μs 0.97 ± 0.31
launch/static workgroup, dynamic ndrange 10.7 ± 0.34 μs 8.86 ± 2.1 μs 1.2 ± 0.28
launch/static workgroup, static ndrange 8.55 ± 0.38 μs 8.58 ± 0.19 μs 0.996 ± 0.049
partition/dynamic workgroup, dynamic ndrange 0.0486 ± 0.0014 μs 0.0463 ± 0.0019 μs 1.05 ± 0.052
partition/static workgroup, dynamic ndrange 0.0525 ± 0.011 μs 0.0456 ± 0.011 μs 1.15 ± 0.38
partition/static workgroup, static ndrange 1.35 ± 0 ns 1.35 ± 0 ns 1 ± 0
saxpy/default/Float16/1024 11.5 ± 0.49 μs 11.6 ± 0.33 μs 0.989 ± 0.051
saxpy/default/Float16/1048576 0.266 ± 0.016 ms 0.268 ± 0.018 ms 0.995 ± 0.09
saxpy/default/Float16/16384 17.1 ± 19 μs 17 ± 18 μs 1.01 ± 1.6
saxpy/default/Float16/2048 9.71 ± 1.9 μs 9.84 ± 1.8 μs 0.986 ± 0.27
saxpy/default/Float16/256 8.82 ± 0.43 μs 9.06 ± 0.31 μs 0.973 ± 0.058
saxpy/default/Float16/262144 0.095 ± 0.0074 ms 0.0948 ± 0.0074 ms 1 ± 0.11
saxpy/default/Float16/32768 0.0377 ± 0.0036 ms 0.0385 ± 0.0026 ms 0.98 ± 0.12
saxpy/default/Float16/4096 10.2 ± 1.7 μs 12.4 ± 1.7 μs 0.824 ± 0.18
saxpy/default/Float16/512 9.29 ± 2.5 μs 9.67 ± 2.5 μs 0.96 ± 0.36
saxpy/default/Float16/64 8.73 ± 0.2 μs 11.2 ± 0.33 μs 0.778 ± 0.029
saxpy/default/Float16/65536 0.0461 ± 0.0042 ms 0.0471 ± 0.0034 ms 0.979 ± 0.11
saxpy/default/Float32/1024 9.41 ± 0.22 μs 11.3 ± 0.39 μs 0.832 ± 0.035
saxpy/default/Float32/1048576 0.239 ± 0.031 ms 0.249 ± 0.033 ms 0.962 ± 0.18
saxpy/default/Float32/16384 15.2 ± 2.8 μs 15.6 ± 1.6 μs 0.979 ± 0.21
saxpy/default/Float32/2048 9.66 ± 0.65 μs 9.57 ± 0.21 μs 1.01 ± 0.072
saxpy/default/Float32/256 8.86 ± 2.2 μs 11.3 ± 2 μs 0.786 ± 0.24
saxpy/default/Float32/262144 0.0848 ± 0.013 ms 0.0854 ± 0.013 ms 0.993 ± 0.21
saxpy/default/Float32/32768 0.038 ± 0.0036 ms 0.0379 ± 0.0034 ms 1 ± 0.13
saxpy/default/Float32/4096 11.4 ± 0.27 μs 10.1 ± 1.8 μs 1.14 ± 0.2
saxpy/default/Float32/512 9.26 ± 0.24 μs 9.1 ± 0.22 μs 1.02 ± 0.035
saxpy/default/Float32/64 8.65 ± 0.2 μs 9.02 ± 2.2 μs 0.959 ± 0.24
saxpy/default/Float32/65536 0.0482 ± 0.006 ms 0.0471 ± 0.0051 ms 1.02 ± 0.17
saxpy/default/Float64/1024 11.3 ± 1.6 μs 9.77 ± 1.9 μs 1.16 ± 0.28
saxpy/default/Float64/1048576 0.563 ± 0.061 ms 0.583 ± 0.054 ms 0.966 ± 0.14
saxpy/default/Float64/16384 0.0374 ± 0.0037 ms 0.0381 ± 0.0026 ms 0.981 ± 0.12
saxpy/default/Float64/2048 9.97 ± 0.27 μs 11.3 ± 1.7 μs 0.886 ± 0.14
saxpy/default/Float64/256 9.01 ± 0.22 μs 11.4 ± 0.53 μs 0.787 ± 0.041
saxpy/default/Float64/262144 0.124 ± 0.015 ms 0.126 ± 0.019 ms 0.989 ± 0.19
saxpy/default/Float64/32768 0.0485 ± 0.0055 ms 0.0479 ± 0.0046 ms 1.01 ± 0.15
saxpy/default/Float64/4096 10.5 ± 0.43 μs 13.8 ± 0.67 μs 0.761 ± 0.048
saxpy/default/Float64/512 9.67 ± 1.9 μs 9.51 ± 1.7 μs 1.02 ± 0.28
saxpy/default/Float64/64 11.1 ± 2.4 μs 11.2 ± 2.6 μs 0.994 ± 0.32
saxpy/default/Float64/65536 0.0617 ± 0.0092 ms 0.0627 ± 0.0099 ms 0.983 ± 0.21
saxpy/static workgroup=(1024,)/Float16/1024 9.48 ± 2.3 μs 9.43 ± 0.24 μs 1.01 ± 0.24
saxpy/static workgroup=(1024,)/Float16/1048576 0.272 ± 0.018 ms 0.271 ± 0.021 ms 1 ± 0.1
saxpy/static workgroup=(1024,)/Float16/16384 17 ± 17 μs 17 ± 21 μs 0.998 ± 1.6
saxpy/static workgroup=(1024,)/Float16/2048 11.2 ± 2 μs 11.3 ± 1.9 μs 0.99 ± 0.24
saxpy/static workgroup=(1024,)/Float16/256 9.17 ± 0.31 μs 9.17 ± 0.3 μs 1 ± 0.047
saxpy/static workgroup=(1024,)/Float16/262144 0.0964 ± 0.0081 ms 0.0954 ± 0.008 ms 1.01 ± 0.12
saxpy/static workgroup=(1024,)/Float16/32768 0.0377 ± 0.0033 ms 0.0388 ± 0.0035 ms 0.971 ± 0.12
saxpy/static workgroup=(1024,)/Float16/4096 9.98 ± 0.24 μs 10.2 ± 0.26 μs 0.982 ± 0.034
saxpy/static workgroup=(1024,)/Float16/512 9.21 ± 0.83 μs 9.3 ± 0.59 μs 0.991 ± 0.11
saxpy/static workgroup=(1024,)/Float16/64 11.3 ± 1.7 μs 9.17 ± 0.7 μs 1.23 ± 0.21
saxpy/static workgroup=(1024,)/Float16/65536 0.0471 ± 0.0034 ms 0.0469 ± 0.0047 ms 1 ± 0.12
saxpy/static workgroup=(1024,)/Float32/1024 9.4 ± 2.1 μs 9.49 ± 1.8 μs 0.991 ± 0.3
saxpy/static workgroup=(1024,)/Float32/1048576 0.232 ± 0.03 ms 0.239 ± 0.025 ms 0.973 ± 0.16
saxpy/static workgroup=(1024,)/Float32/16384 15.3 ± 1.7 μs 15.3 ± 19 μs 0.997 ± 1.2
saxpy/static workgroup=(1024,)/Float32/2048 9.51 ± 0.24 μs 9.69 ± 1.8 μs 0.982 ± 0.18
saxpy/static workgroup=(1024,)/Float32/256 11.3 ± 0.45 μs 11.3 ± 0.37 μs 0.997 ± 0.051
saxpy/static workgroup=(1024,)/Float32/262144 0.0854 ± 0.013 ms 0.0843 ± 0.011 ms 1.01 ± 0.2
saxpy/static workgroup=(1024,)/Float32/32768 0.038 ± 0.0052 ms 0.0399 ± 0.0041 ms 0.952 ± 0.16
saxpy/static workgroup=(1024,)/Float32/4096 11.4 ± 1.7 μs 9.94 ± 0.4 μs 1.15 ± 0.18
saxpy/static workgroup=(1024,)/Float32/512 9.42 ± 1 μs 11 ± 2.3 μs 0.855 ± 0.2
saxpy/static workgroup=(1024,)/Float32/64 8.98 ± 0.83 μs 9.08 ± 1.2 μs 0.99 ± 0.16
saxpy/static workgroup=(1024,)/Float32/65536 0.0475 ± 0.0063 ms 0.0479 ± 0.0047 ms 0.991 ± 0.16
saxpy/static workgroup=(1024,)/Float64/1024 9.59 ± 0.38 μs 9.79 ± 2.1 μs 0.98 ± 0.21
saxpy/static workgroup=(1024,)/Float64/1048576 0.531 ± 0.059 ms 0.585 ± 0.06 ms 0.908 ± 0.14
saxpy/static workgroup=(1024,)/Float64/16384 0.0376 ± 0.0045 ms 0.0385 ± 0.0034 ms 0.976 ± 0.15
saxpy/static workgroup=(1024,)/Float64/2048 10.2 ± 1.5 μs 9.96 ± 1.7 μs 1.02 ± 0.23
saxpy/static workgroup=(1024,)/Float64/256 11.8 ± 10 μs 11.9 ± 9.9 μs 0.998 ± 1.2
saxpy/static workgroup=(1024,)/Float64/262144 0.124 ± 0.016 ms 0.126 ± 0.019 ms 0.985 ± 0.2
saxpy/static workgroup=(1024,)/Float64/32768 0.0489 ± 0.0057 ms 0.0484 ± 0.0054 ms 1.01 ± 0.16
saxpy/static workgroup=(1024,)/Float64/4096 11.9 ± 1.6 μs 10.5 ± 0.41 μs 1.13 ± 0.16
saxpy/static workgroup=(1024,)/Float64/512 11.5 ± 14 μs 11.9 ± 11 μs 0.974 ± 1.5
saxpy/static workgroup=(1024,)/Float64/64 11.2 ± 2.2 μs 11.3 ± 0.43 μs 0.992 ± 0.2
saxpy/static workgroup=(1024,)/Float64/65536 0.0621 ± 0.0071 ms 0.0618 ± 0.0094 ms 1 ± 0.19
time_to_load 0.449 ± 0.015 s 0.458 ± 0.011 s 0.98 ± 0.041
main 64ac668... main / 64ac668...
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 9 allocs: 0.203 kB 5 allocs: 0.0938 kB 2.17
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 8 allocs: 0.141 kB 12 allocs: 0.25 kB 0.562
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 12 allocs: 0.25 kB 12 allocs: 0.25 kB 1
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 8 allocs: 0.141 kB 1.78
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 8 allocs: 0.141 kB 1
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 8 allocs: 0.141 kB 8 allocs: 0.141 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).

Comment thread src/pocl/device/subgroups.jl Outdated
@codecov

codecov Bot commented Oct 3, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 47.69231% with 34 lines in your changes missing coverage. Please review.
✅ Project coverage is 74.15%. Comparing base (39ef02f) to head (c3e0b7a).
⚠️ Report is 4 commits behind head on main.

Files with missing lines Patch % Lines
src/pocl/backend.jl 44.68% 26 Missing ⚠️
src/KernelAbstractions.jl 0.00% 8 Missing ⚠️
Additional details and impacted files
@@            Coverage Diff             @@
##             main     #831      +/-   ##
==========================================
- Coverage   79.06%   74.15%   -4.92%     
==========================================
  Files          24       24              
  Lines        2040     2275     +235     
==========================================
+ Hits         1613     1687      +74     
- Misses        427      588     +161     

☔ 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.

Comment on lines +68 to +80
# the sub-group width is fixed, so make `get_max_sub_group_size` a constant, as
# KernelInterface requires (this runs before optimization)
gvs = LLVM.globals(mod)
if haskey(gvs, "__spirv_BuiltInSubgroupMaxSize")
gv = gvs["__spirv_BuiltInSubgroupMaxSize"]
for use in collect(LLVM.uses(gv))
load = LLVM.user(use)
load isa LLVM.LoadInst || continue
LLVM.replace_uses!(load, ConstantInt(LLVM.value_type(load), sg_size))
LLVM.erase!(load)
end
isempty(LLVM.uses(gv)) && LLVM.erase!(gv)
end

@vchuravy vchuravy Oct 4, 2026 •

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is perhaps a bit sketchy, and would need to be replicated for OpenCL/oneAPI

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Eh, I'm not sure I like this. Isn't the driver responsible for specializing on this? It looks like other drivers do this, so this might be a PoCL deficiency if anything.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Alternative: pocl/pocl#2375
So I'd prefer if we leave this out for now, and I'll work with upstream to get this polished and backported.

@vchuravy vchuravy Oct 6, 2026 •

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude:

PoCL does specialize: its workgroup pass replaces _pocl_sub_group_size with the intel_reqd_sub_group_size value, so the final kernel has a constant either way. It does so late, though, after it has formed the barrier regions (on the CPU device every shuffle is a work-group barrier) and only right before its final O3. So the shuffle loops in the KI reduce/scan fallbacks stay as barrier loops around work-item loops and are never unrolled. Without a constant on the Julia side, the testsuite passes but those fallbacks are ~1.2–1.7x slower. The other backends get their constant on the Julia side too: AMDGPU.jl folds llvm.amdgcn.wavefrontsize in finish_module!, and the CUDA/Metal KI PRs hardcode 32, since NVPTX on LLVM 18 doesn't fold %warpsize in the middle end.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, hence my PR.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

In addition, my PR also folds get_local_size exposing further optimizations.

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)
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)
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)
….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)
- Require the released SPIRVIntrinsics 1.3 instead of the deleted
  `vc/subgroup-votes` branch of OpenCL.jl, in `[sources]` and in CI.
  The OpenCL.jl job installs 1.3 rather than the ka-0.10 copy (1.1.4).
- The fallback of `KI.sub_group_match_any` takes a step per distinct
  value, which differs between the sub-groups of a work-group. PoCL needs
  them to take the same steps, so override it with a loop to a bound that
  is the same for the whole work-group. Test it with several sub-groups.
- Merge `NATIVE_COLLECTIVES` into `POCL_REPLICA_FIX`.
- `supports_shuffle` of a type without fields follows `UInt32` instead of
  being `true` on every backend.
- Fix the sizes of primitive types in the `shfl` docstring and the
  `supports_shuffle` entry of the backend table.

Assisted-by: Claude Code
pocl_standalone_jll 7.2.2, which main requires now, includes the fix
for PoCL 7.2's peeling of the first work-item (pocl/pocl#2239), so the
native reductions and scans are always used, and the loops to a uniform
bound that replaced them without the fix can't be reached anymore.

Assisted-by: Claude Code
Primitive types that a backend doesn't shuffle natively were always
split into UInt32 words, so a `Ptr` or another 64-bit primitive type
took two shuffles even on backends with native 64-bit shuffles. Shuffle
them as a `UInt64` instead, which a backend without native 64-bit
shuffles splits into two words as before, and 16-byte values as two
`UInt64` words. Check the generated code on PoCL.

Assisted-by: Claude Code
@vchuravy
vchuravy force-pushed the vc/ki-subgroup-ops branch from c3e0b7a to 64ac668 Compare October 6, 2026 12:29
@vchuravy

vchuravy commented Oct 6, 2026

Copy link
Copy Markdown
Member Author

Rebased onto main and addressed the review:

  • SPIRVIntrinsics: require the released 1.3 instead of the deleted vc/subgroup-votes branch, in [sources] and in CI. The OpenCL.jl Buildkite job installs 1.3 rather than the ka-0.10 copy (1.1.4), which lacks the new intrinsics.
  • sub_group_match_any on PoCL: the fallback takes a step per distinct value, which differs between the sub-groups of a work-group, and PoCL needs them to take the same steps. It gave wrong masks with several sub-groups (4 of 20 random trials). Added a PoCL override that loops to a bound that is the same for the whole work-group, and a test with several sub-groups that fails without it.
  • PoCL fallbacks: dropped the reduce/scan loops for PoCL without the replica fix. Main requires pocl_standalone_jll 7.2.2, whose Yggdrasil build carries the same backport of WorkitemLoops: share local memory across region replicas pocl/pocl#2239, so the native collectives are always used. NATIVE_COLLECTIVES and POCL_REPLICA_FIX are gone.
  • 64-bit shuffles: a primitive type the backend doesn't shuffle natively and that has 8 bytes (e.g. Ptr) is now shuffled as a UInt64, so a backend with native 64-bit shuffles does one shuffle instead of two UInt32 ones. Backends without them still split it into two words. 16-byte values become two UInt64 words. A codegen check covers the 64-bit case on PoCL (128-bit integers don't compile with PoCL's SPIR-V back-end at all).
  • Small fixes: supports_shuffle of a type without fields follows UInt32, and the docs on which primitive types are shuffled as words are corrected.

The full test suite passes locally on Julia 1.10 and 1.12.

Still open from the review: Float16/Float64 on PoCL devices without fp16/fp64 are reported as unsupported for shuffles, even though the UInt32-word fallback could handle them.

#830 is based on this branch and needs a rebase.

Comment thread .buildkite/pipeline.yml
# not the ka-0.10 copy of SPIRVIntrinsics: KernelAbstractions needs the
# sub-group votes and collectives of SPIRVIntrinsics 1.3
Pkg.add(name="SPIRVIntrinsics", version="1.3")' || exit 3

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

SPIRVIntrinsics 1.3 has been released

Comment on lines +272 to +275
# Shuffles exchange values between the work-items of a sub-group. Backends implement them for
# the primitive types for which `supports_shuffle` returns `true`. The fallbacks below shuffle
# other primitive types as unsigned words, and other `isbits` types field by field.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is in the docstring

Suggested change
# Shuffles exchange values between the work-items of a sub-group. Backends implement them for
# the primitive types for which `supports_shuffle` returns `true`. The fallbacks below shuffle
# other primitive types as unsigned words, and other `isbits` types field by field.

[`get_sub_group_local_id`](@ref) equal to `get_sub_group_local_id() + offset`, for an `offset`
of at least 0. When that lane is past the sub-group width, i.e.
`get_sub_group_local_id() + offset > get_max_sub_group_size()`, the result is `val` of the
work-item itself, like CUDA's `shfl_down_sync`. When the lane is within the width but has no

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is there a performance impact to ensuring this on backends that don't make that promise, and if so, is it worth ensuring?

Comment on lines +709 to +710
for n in unique((sg_size, max(sg_size - 3, 1)))
@testset "sub_group_reduce and sub_group_scan, $n work-items" begin

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I find this makes the output much easier to parse when things fail

Suggested change
for n in unique((sg_size, max(sg_size - 3, 1)))
@testset "sub_group_reduce and sub_group_scan, $n work-items" begin
@testset "sub_group_reduce and sub_group_scan" begin
@testset "$n work-items" for n in unique((sg_size, max(sg_size - 3, 1)))

return
end

function subgroup_communication_testsuite(backend::KI.Backend, AT, sg_size)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Needs to actually be added to the testsuite in testsuite.jl for backends to test them.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Interface.jl should probably get split up I made this when the interface was much smaller but that's out of scope for this PR

Comment thread src/pocl/backend.jl
Comment on lines +268 to +273
@device_override KI.shfl_xor(val::T, mask::Integer) where {T <: ShuffleTypes} =
sub_group_shuffle_xor(val, mask)

@device_override KI.sub_group_any(pred::Bool) = SPIRVIntrinsics.sub_group_any(pred)

@device_override KI.sub_group_all(pred::Bool) = SPIRVIntrinsics.sub_group_all(pred)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Any reason why only sub_group_shuffle_xor isn't qualified?

Comment thread src/pocl/backend.jl
Comment on lines +309 to +312
# Native reductions and scans of `cl_khr_subgroups`, for `+` on 32- and 64-bit integers and
# floats, and `min`/`max` on integers (OpenCL's `min` and `max` treat NaN and the sign of zero
# differently from Julia's). They need the fix of PoCL 7.2's peeling of the first work-item
# (pocl/pocl#2239), which `pocl_standalone_jll` includes since 7.2.1+1.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Do we have thorough tests for the reductions and scans behaviour for the commonly implemented operators? Julia sometimes treats edge cases differently and I assume we want to keep Julia's semantics so we might have to stick to the fallback even when backends havenative implementations

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.

3 participants