diff --git a/Project.toml b/Project.toml index 2a992247..63ac73c4 100644 --- a/Project.toml +++ b/Project.toml @@ -34,7 +34,7 @@ oneAPI_Support_jll = "b049733a-a71d-5ed3-8eba-7d323ac00b36" [compat] AbstractFFTs = "1.5.0" -AcceleratedKernels = "0.3.1, 0.4" +AcceleratedKernels = "0.5" Adapt = "4" CEnum = "0.4, 0.5" ExprTools = "0.1" diff --git a/src/accumulate.jl b/src/accumulate.jl index d3647faa..193a0cb0 100644 --- a/src/accumulate.jl +++ b/src/accumulate.jl @@ -3,19 +3,23 @@ import oneAPI: oneArray, oneAPIBackend import AcceleratedKernels as AK # Use a smaller block size on Intel GPUs to work around a scan correctness issue -# with the Blelloch parallel prefix sum at larger block sizes (>=128). +# with the parallel prefix sum at larger block sizes (>=128). const _ACCUMULATE_BLOCK_SIZE = 64 +# The scan algorithm for the given `dims`: whole-array scans and scans along a dimension +# use different algorithms, so pick the one AcceleratedKernels' `Auto()` would, with our +# block size. +_scan_alg(dims) = dims === nothing ? AK.ScanPrefixes(block_size = _ACCUMULATE_BLOCK_SIZE) : + AK.SliceScan(block_size = _ACCUMULATE_BLOCK_SIZE) + # Accumulate operations using AcceleratedKernels -Base.accumulate!(op, B::oneArray, A::oneArray; init = zero(eltype(A)), - block_size = _ACCUMULATE_BLOCK_SIZE, kwargs...) = - AK.accumulate!(op, B, A, oneAPIBackend(); init, block_size, kwargs...) +Base.accumulate!(op, B::oneArray, A::oneArray; dims = nothing, alg = _scan_alg(dims), kwargs...) = + AK.accumulate!(op, B, A; dims, alg, kwargs...) -Base.accumulate(op, A::oneArray; init = zero(eltype(A)), - block_size = _ACCUMULATE_BLOCK_SIZE, kwargs...) = - AK.accumulate(op, A, oneAPIBackend(); init, block_size, kwargs...) +Base.accumulate(op, A::oneArray; dims = nothing, alg = _scan_alg(dims), kwargs...) = + AK.accumulate(op, A; dims, alg, kwargs...) -Base.cumsum(src::oneArray; block_size = _ACCUMULATE_BLOCK_SIZE, kwargs...) = - AK.cumsum(src, oneAPIBackend(); block_size, kwargs...) -Base.cumprod(src::oneArray; block_size = _ACCUMULATE_BLOCK_SIZE, kwargs...) = - AK.cumprod(src, oneAPIBackend(); block_size, kwargs...) +Base.cumsum(src::oneArray; dims = nothing, alg = _scan_alg(dims), kwargs...) = + AK.cumsum(src; dims, alg, kwargs...) +Base.cumprod(src::oneArray; dims = nothing, alg = _scan_alg(dims), kwargs...) = + AK.cumprod(src; dims, alg, kwargs...)