Skip to content

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

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

vchuravy wants to merge 6 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_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.

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

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)
@github-actions

github-actions Bot commented Oct 3, 2026 •

Copy link
Copy Markdown
Contributor

Benchmark Results

Show table
main 5109dae... main / 5109dae...
const/@Const/Float32/262144 0.261 ± 0.023 ms 0.262 ± 0.018 ms 0.996 ± 0.11
const/@Const/Float32/65536 0.0993 ± 0.012 ms 0.0994 ± 0.012 ms 0.999 ± 0.17
const/@Const/Float64/262144 0.462 ± 0.026 ms 0.475 ± 0.023 ms 0.972 ± 0.073
const/@Const/Float64/65536 0.217 ± 0.013 ms 0.221 ± 0.017 ms 0.983 ± 0.093
const/unmarked/Float32/262144 0.471 ± 0.021 ms 0.471 ± 0.018 ms 1 ± 0.059
const/unmarked/Float32/65536 0.155 ± 0.016 ms 0.166 ± 0.019 ms 0.934 ± 0.14
const/unmarked/Float64/262144 0.87 ± 0.11 ms 0.856 ± 0.01 ms 1.02 ± 0.13
const/unmarked/Float64/65536 0.311 ± 0.018 ms 0.313 ± 0.017 ms 0.994 ± 0.079
launch/3D static workgroup, dynamic ndrange 16.5 ± 8.8 μs 13.3 ± 2.6 μs 1.24 ± 0.71
launch/3D static workgroup, static ndrange 15.6 ± 0.96 μs 13.6 ± 1.6 μs 1.15 ± 0.15
launch/dynamic workgroup, dynamic ndrange 15.7 ± 0.98 μs 14.9 ± 3.1 μs 1.05 ± 0.23
launch/dynamic workgroup, dynamic ndrange, workgroupsize given 12.9 ± 2.5 μs 13.3 ± 3.3 μs 0.972 ± 0.31
launch/static workgroup, dynamic ndrange 12.8 ± 2.4 μs 12.6 ± 0.93 μs 1.01 ± 0.21
launch/static workgroup, static ndrange 15.6 ± 1 μs 15.1 ± 2.4 μs 1.04 ± 0.18
partition/dynamic workgroup, dynamic ndrange 0.0432 ± 0.015 μs 0.0409 ± 0.0093 μs 1.06 ± 0.43
partition/static workgroup, dynamic ndrange 0.0387 ± 0.0085 μs 0.0434 ± 0.0088 μs 0.89 ± 0.27
partition/static workgroup, static ndrange 1.03 ± 0.16 ns 1.12 ± 0.088 ns 0.92 ± 0.16
saxpy/default/Float16/1024 13.8 ± 1.6 μs 14.1 ± 2.3 μs 0.982 ± 0.2
saxpy/default/Float16/1048576 0.414 ± 0.037 ms 0.419 ± 0.052 ms 0.99 ± 0.15
saxpy/default/Float16/16384 0.053 ± 0.011 ms 0.0523 ± 0.0064 ms 1.01 ± 0.24
saxpy/default/Float16/2048 15.2 ± 15 μs 0.0466 ± 0.04 ms 0.326 ± 0.43
saxpy/default/Float16/256 14 ± 2.1 μs 15.6 ± 2.1 μs 0.895 ± 0.18
saxpy/default/Float16/262144 0.153 ± 0.039 ms 0.141 ± 0.026 ms 1.08 ± 0.34
saxpy/default/Float16/32768 0.0573 ± 0.0051 ms 0.058 ± 0.0088 ms 0.988 ± 0.17
saxpy/default/Float16/4096 14.1 ± 1.5 μs 17.2 ± 10 μs 0.818 ± 0.5
saxpy/default/Float16/512 14.1 ± 1.4 μs 17 ± 2.7 μs 0.831 ± 0.16
saxpy/default/Float16/64 13.5 ± 1.3 μs 14 ± 1.8 μs 0.966 ± 0.16
saxpy/default/Float16/65536 0.0719 ± 0.018 ms 0.079 ± 0.055 ms 0.91 ± 0.67
saxpy/default/Float32/1024 13.4 ± 2.3 μs 16.3 ± 35 μs 0.817 ± 1.7
saxpy/default/Float32/1048576 0.604 ± 0.092 ms 0.552 ± 0.088 ms 1.09 ± 0.24
saxpy/default/Float32/16384 15.3 ± 38 μs 0.0555 ± 0.044 ms 0.276 ± 0.72
saxpy/default/Float32/2048 14.4 ± 3.3 μs 13.6 ± 2 μs 1.06 ± 0.29
saxpy/default/Float32/256 15.4 ± 2.1 μs 16.6 ± 1.8 μs 0.925 ± 0.16
saxpy/default/Float32/262144 0.2 ± 0.028 ms 0.199 ± 0.03 ms 1 ± 0.21
saxpy/default/Float32/32768 0.0618 ± 0.0067 ms 0.0587 ± 0.012 ms 1.05 ± 0.24
saxpy/default/Float32/4096 15.9 ± 4.2 μs 16.8 ± 4.2 μs 0.951 ± 0.35
saxpy/default/Float32/512 12.6 ± 1.2 μs 15.4 ± 2 μs 0.823 ± 0.13
saxpy/default/Float32/64 14.2 ± 2.1 μs 16.7 ± 33 μs 0.851 ± 1.7
saxpy/default/Float32/65536 0.0769 ± 0.016 ms 0.0805 ± 0.094 ms 0.954 ± 1.1
saxpy/default/Float64/1024 15.2 ± 3 μs 13.8 ± 1.5 μs 1.11 ± 0.25
saxpy/default/Float64/1048576 1.24 ± 0.13 ms 1.17 ± 0.2 ms 1.06 ± 0.21
saxpy/default/Float64/16384 0.0641 ± 0.018 ms 0.0597 ± 0.013 ms 1.07 ± 0.39
saxpy/default/Float64/2048 13.9 ± 3.8 μs 18.7 ± 37 μs 0.744 ± 1.5
saxpy/default/Float64/256 13.1 ± 1.4 μs 15.5 ± 1.6 μs 0.848 ± 0.13
saxpy/default/Float64/262144 0.326 ± 0.046 ms 0.312 ± 0.044 ms 1.05 ± 0.21
saxpy/default/Float64/32768 0.0801 ± 0.07 ms 0.0782 ± 0.013 ms 1.02 ± 0.9
saxpy/default/Float64/4096 0.0367 ± 0.032 ms 14 ± 0.95 μs 2.62 ± 2.3
saxpy/default/Float64/512 15.2 ± 2.2 μs 16.2 ± 3.3 μs 0.936 ± 0.23
saxpy/default/Float64/64 14.4 ± 1.7 μs 14.2 ± 2.2 μs 1.01 ± 0.2
saxpy/default/Float64/65536 0.116 ± 0.011 ms 0.117 ± 0.019 ms 0.995 ± 0.18
saxpy/static workgroup=(1024,)/Float16/1024 15.1 ± 2.9 μs 13.5 ± 1.4 μs 1.11 ± 0.24
saxpy/static workgroup=(1024,)/Float16/1048576 0.416 ± 0.04 ms 0.414 ± 0.041 ms 1 ± 0.14
saxpy/static workgroup=(1024,)/Float16/16384 0.0513 ± 0.0049 ms 0.053 ± 0.0044 ms 0.967 ± 0.12
saxpy/static workgroup=(1024,)/Float16/2048 13.4 ± 0.89 μs 15 ± 2 μs 0.897 ± 0.14
saxpy/static workgroup=(1024,)/Float16/256 14.8 ± 2 μs 15.7 ± 2.5 μs 0.943 ± 0.2
saxpy/static workgroup=(1024,)/Float16/262144 0.149 ± 0.039 ms 0.144 ± 0.029 ms 1.03 ± 0.34
saxpy/static workgroup=(1024,)/Float16/32768 0.059 ± 0.0078 ms 0.0669 ± 0.015 ms 0.882 ± 0.23
saxpy/static workgroup=(1024,)/Float16/4096 16.5 ± 32 μs 16.1 ± 1.7 μs 1.03 ± 2
saxpy/static workgroup=(1024,)/Float16/512 14.7 ± 2.2 μs 13.8 ± 1.4 μs 1.07 ± 0.19
saxpy/static workgroup=(1024,)/Float16/64 16.6 ± 36 μs 15.3 ± 2.4 μs 1.08 ± 2.4
saxpy/static workgroup=(1024,)/Float16/65536 0.0743 ± 0.036 ms 0.0907 ± 0.049 ms 0.819 ± 0.59
saxpy/static workgroup=(1024,)/Float32/1024 13.2 ± 1.9 μs 15.2 ± 1.8 μs 0.864 ± 0.16
saxpy/static workgroup=(1024,)/Float32/1048576 0.543 ± 0.082 ms 0.535 ± 0.08 ms 1.02 ± 0.21
saxpy/static workgroup=(1024,)/Float32/16384 0.051 ± 0.041 ms 0.055 ± 0.04 ms 0.928 ± 1
saxpy/static workgroup=(1024,)/Float32/2048 15.3 ± 3.9 μs 14.9 ± 3.2 μs 1.03 ± 0.34
saxpy/static workgroup=(1024,)/Float32/256 15.6 ± 1.1 μs 16.4 ± 1.6 μs 0.953 ± 0.11
saxpy/static workgroup=(1024,)/Float32/262144 0.199 ± 0.027 ms 0.2 ± 0.037 ms 0.997 ± 0.23
saxpy/static workgroup=(1024,)/Float32/32768 0.0608 ± 0.0093 ms 0.0683 ± 0.045 ms 0.891 ± 0.6
saxpy/static workgroup=(1024,)/Float32/4096 14.4 ± 5.6 μs 14.7 ± 32 μs 0.977 ± 2.2
saxpy/static workgroup=(1024,)/Float32/512 15.7 ± 1.2 μs 13.7 ± 2.6 μs 1.14 ± 0.23
saxpy/static workgroup=(1024,)/Float32/64 15.5 ± 1.7 μs 14 ± 3.9 μs 1.1 ± 0.33
saxpy/static workgroup=(1024,)/Float32/65536 0.0787 ± 0.076 ms 0.078 ± 0.014 ms 1.01 ± 0.99
saxpy/static workgroup=(1024,)/Float64/1024 17.3 ± 34 μs 13.1 ± 1.5 μs 1.32 ± 2.6
saxpy/static workgroup=(1024,)/Float64/1048576 1.09 ± 0.24 ms 1.19 ± 0.085 ms 0.918 ± 0.21
saxpy/static workgroup=(1024,)/Float64/16384 0.0638 ± 0.043 ms 0.0609 ± 0.013 ms 1.05 ± 0.73
saxpy/static workgroup=(1024,)/Float64/2048 14.7 ± 36 μs 18.4 ± 37 μs 0.802 ± 2.5
saxpy/static workgroup=(1024,)/Float64/256 14.4 ± 2.1 μs 13.7 ± 1.5 μs 1.05 ± 0.19
saxpy/static workgroup=(1024,)/Float64/262144 0.315 ± 0.036 ms 0.315 ± 0.034 ms 0.997 ± 0.16
saxpy/static workgroup=(1024,)/Float64/32768 0.095 ± 0.083 ms 0.0781 ± 0.011 ms 1.22 ± 1.1
saxpy/static workgroup=(1024,)/Float64/4096 0.0495 ± 0.039 ms 14.3 ± 32 μs 3.46 ± 8.1
saxpy/static workgroup=(1024,)/Float64/512 15.7 ± 1.6 μs 16.4 ± 2.4 μs 0.956 ± 0.17
saxpy/static workgroup=(1024,)/Float64/64 14 ± 1.5 μs 16.2 ± 1.2 μs 0.863 ± 0.11
saxpy/static workgroup=(1024,)/Float64/65536 0.116 ± 0.018 ms 0.114 ± 0.021 ms 1.02 ± 0.25
time_to_load 0.398 ± 0.0061 s 0.402 ± 0.003 s 0.989 ± 0.017
main 5109dae... main / 5109dae...
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 9 allocs: 0.203 kB 1
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 12 allocs: 0.25 kB 8 allocs: 0.141 kB 1.78
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 12 allocs: 0.25 kB 12 allocs: 0.25 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 12 allocs: 0.25 kB 12 allocs: 0.25 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).

Comment thread src/pocl/device/subgroups.jl Outdated
Comment on lines +2 to +25

# `sub_group_shuffle`, with the lane passed modulo `UInt32`, so that an out-of-range lane
# gives an unspecified value rather than an `InexactError`
for T in SPIRVIntrinsics.gentypes
@eval @device_function shuffle(x::$T, lane::Integer) =
@builtin_ccall(
"__spirv_GroupNonUniformShuffle", $T, (UInt32, $T, UInt32),
UInt32(Scope.Subgroup), x, (lane - 1) % UInt32
)
end

# Votes, from `cl_khr_subgroups` and `cl_khr_subgroup_ballot`. The SPIR-V back-end lowers
# these OpenCL built-ins, which have to be listed in `subgroup_intrinsics`.
const subgroup_intrinsics = ["_Z13sub_group_anyi", "_Z13sub_group_alli", "_Z16sub_group_balloti"]

@device_function sub_group_any(pred::Bool) =
ccall("extern _Z13sub_group_anyi", llvmcall, Int32, (Int32,), pred) != Int32(0)

@device_function sub_group_all(pred::Bool) =
ccall("extern _Z13sub_group_alli", llvmcall, Int32, (Int32,), pred) != Int32(0)

# bit `i` of the result is set for the lane with (0-based) id `i`
@device_function sub_group_ballot(pred::Bool) =
ccall("extern _Z16sub_group_balloti", llvmcall, NTuple{4, VecElement{UInt32}}, (Int32,), 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.

Should this be added to SPIRVIntrinsics.jl instad or are they hacks?

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.

Upstream PR: JuliaGPU/OpenCL.jl#526

@codecov

codecov Bot commented Oct 3, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 0% with 20 lines in your changes missing coverage. Please review.
✅ Project coverage is 0.00%. Comparing base (39ef02f) to head (02990c9).
⚠️ Report is 2 commits behind head on main.

Files with missing lines Patch % Lines
src/pocl/backend.jl 0.00% 10 Missing ⚠️
src/pocl/compiler/compilation.jl 0.00% 10 Missing ⚠️

❗ There is a different number of reports uploaded between BASE (39ef02f) and HEAD (02990c9). Click for more details.

HEAD has 24 uploads less than BASE
Flag BASE (39ef02f) HEAD (02990c9)
44 20
Additional details and impacted files
@@            Coverage Diff             @@
##             main    #831       +/-   ##
==========================================
- Coverage   79.06%   0.00%   -79.07%     
==========================================
  Files          24      22        -2     
  Lines        2040    1857      -183     
==========================================
- Hits         1613       0     -1613     
- Misses        427    1857     +1430     

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

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)
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

…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)
vchuravy added a commit to JuliaGPU/Metal.jl that referenced this pull request Oct 4, 2026
KernelInterface (JuliaGPU/KernelAbstractions.jl#831) now shuffles primitive
types that a back-end doesn't support natively as `UInt32` words, so the
Metal-specific split into halves isn't needed anymore.

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)

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.

2 participants