Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
120 changes: 91 additions & 29 deletions ext/AMDGPUExt/AMDGPUExt.jl
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,27 @@ function (blkIter::BlockIndexerSwapped)()
end
# COV_EXCL_STOP

# Coalescing-aware 2D block shape for a column-major (m rows x n cols) launch.
# Consecutive workitems (workitemIdx().x, the wavefront direction) must walk the
# stride-1 (row) dimension for global-memory access to coalesce. The aspect-ratio
# heuristic gives good shapes for square/tall arrays, but starves x (down to 1
# workitem) for wide arrays — so we floor x at a full warp whenever there are
# enough rows, and otherwise use every available row.
@inline function _block_shape_2d(maxthreads, m, n)
maxThreadsX = sqrt(maxthreads)
y_thr = clamp(floor(Int, (n / m) * maxThreadsX), 1, maxthreads)
x_thr = fld(maxthreads, y_thr)
if x_thr < 32 && m >= 32
x_thr = 32
y_thr = fld(maxthreads, x_thr)
elseif m < 32 && x_thr < m
x_thr = m
y_thr = max(1, fld(maxthreads, x_thr))
end
return (x_thr, y_thr)
end

# Generic launcher used for the swapped-axes fallback (uncoalesced, last resort).
function _parallel_for(indexer::TI, f, (m, n), (M, N), x...) where {TI}
kargs = _kernel_args(indexer, (M, N), f, x...)
kernel, shmem_size = _kernel_maxshmem(_parallel_for_amdgpu_MN, kargs)
Expand All @@ -118,10 +139,21 @@ function JACC.parallel_for(
dev = AMDGPU.device()
props = AMDGPU.HIP.properties(dev)
maxBlocks = (x = props.maxGridSize[1], y = props.maxGridSize[2])
if M < N && maxBlocks.x >= maxBlocks.y
# Prefer the coalesced (basic) mapping and only swap axes when it is
# *necessary* — i.e. when the basic grid would exceed the y-dimension limit.
# Swapping moves the large extent onto the x-axis (which has a far larger
# limit) but sacrifices memory coalescing, so it must be a last resort.
kargs = _kernel_args(BlockIndexerBasic(), (M, N), f, x...)
kernel, shmem_size = _kernel_maxshmem(_parallel_for_amdgpu_MN, kargs)
config = AMDGPU.launch_configuration(kernel; shmem = shmem_size)
x_thr, y_thr = _block_shape_2d(config.groupsize, M, N)
if cld(N, y_thr) > maxBlocks.y && maxBlocks.x >= maxBlocks.y
_parallel_for(BlockIndexerSwapped(), f, (N, M), (M, N), x...)
else
_parallel_for(BlockIndexerBasic(), f, (M, N), (M, N), x...)
blocks = (cld(M, x_thr), cld(N, y_thr))
kernel(kargs...; groupsize = (x_thr, y_thr), gridsize = blocks,
shmem = shmem_size)
AMDGPU.synchronize()
end
end

Expand Down Expand Up @@ -158,10 +190,37 @@ function JACC.parallel_for(
dev = AMDGPU.device()
props = AMDGPU.HIP.properties(dev)
maxBlocks = (x = props.maxGridSize[1], y = props.maxGridSize[2])
if M < N && maxBlocks.x >= maxBlocks.y
_parallel_for(BlockIndexerSwapped(), f, spec, (N, M), (M, N), x...)
# Determine the y-threads the basic launch would use (from the user's block
# if given, otherwise the coalescing-aware default), then swap axes only if
# the basic grid would overflow the y-dimension limit.
if spec.threads == 0
kargs = _kernel_args(BlockIndexerBasic(), (M, N), f, x...)
kernel, shmem_size = _kernel_maxshmem(_parallel_for_amdgpu_MN, kargs)
if spec.shmem_size < 0
spec.shmem_size = shmem_size
end
config = AMDGPU.launch_configuration(kernel; shmem = spec.shmem_size)
x_thr, y_thr = _block_shape_2d(config.groupsize, M, N)
if cld(N, y_thr) > maxBlocks.y && maxBlocks.x >= maxBlocks.y
_parallel_for(BlockIndexerSwapped(), f, spec, (N, M), (M, N), x...)
else
spec.threads = (x_thr, y_thr)
if spec.blocks == 0
spec.blocks = (cld(M, x_thr), cld(N, y_thr))
end
kernel(kargs...; groupsize = spec.threads, gridsize = spec.blocks,
shmem = spec.shmem_size, stream = spec.stream)
if spec.sync
AMDGPU.synchronize(spec.stream)
end
end
else
_parallel_for(BlockIndexerBasic(), f, spec, (M, N), (M, N), x...)
y_thr = spec.threads[2]
if cld(N, y_thr) > maxBlocks.y && maxBlocks.x >= maxBlocks.y
_parallel_for(BlockIndexerSwapped(), f, spec, (N, M), (M, N), x...)
else
_parallel_for(BlockIndexerBasic(), f, spec, (M, N), (M, N), x...)
end
end
end

Expand Down Expand Up @@ -224,13 +283,14 @@ function JACC.reduce_workspace(::AMDGPUBackend, tmp::AMDGPU.ROCArray{T},
AMDGPUReduceWorkspace{T, JACC.Unmanaged}(tmp, init)
end

# Ensure the partial-results buffer is sized for this launch. `init` is seeded
# directly inside the kernels (passed as an argument), so no `fill!` is needed
# here — the reduce is a pure sequence of async kernel launches on the stream.
@inline function _init!(
wk::AMDGPUReduceWorkspace{T, JACC.Managed}, blocks, init) where {T}
if length(wk.tmp) != prod(blocks)
wk.tmp = AMDGPU.ROCArray{typeof(init)}(undef, blocks)
end
fill!(wk.tmp, init)
fill!(wk.ret, init)
return nothing
end

Expand All @@ -247,10 +307,10 @@ function JACC._parallel_reduce!(reducer::JACC.ParallelReduce{AMDGPUBackend},
op = reducer.op
init = reducer.init

kargs1 = _kernel_args(N, op, wk.ret, f, x...)
kargs1 = _kernel_args(N, op, wk.ret, init, f, x...)
kernel1, threads1 = _kernel_maxthreads(_parallel_reduce_amdgpu, kargs1)

kargs2 = _kernel_args(1, op, wk.ret, wk.ret)
kargs2 = _kernel_args(1, op, wk.ret, init, wk.ret)
kernel2, threads2 = _kernel_maxthreads(_reduce_kernel_amdgpu, kargs2)

threads = min(threads1, threads2, 512)
Expand All @@ -259,11 +319,11 @@ function JACC._parallel_reduce!(reducer::JACC.ParallelReduce{AMDGPUBackend},

_init!(wk, blocks, init)

kargs1 = _kernel_args(N, op, wk.tmp, f, x...)
kargs1 = _kernel_args(N, op, wk.tmp, init, f, x...)
kernel1(kargs1...; groupsize = threads, gridsize = blocks,
shmem = shmem_size, stream = reducer.stream)

kargs2 = _kernel_args(blocks, op, wk.tmp, wk.ret)
kargs2 = _kernel_args(blocks, op, wk.tmp, init, wk.ret)
kernel2(kargs2...; groupsize = threads, gridsize = 1,
shmem = shmem_size, stream = reducer.stream)

Expand All @@ -277,24 +337,25 @@ end
function JACC.parallel_reduce(f, ::AMDGPUBackend, N::Integer, x...; op, init)
ret_inst = AMDGPU.ROCArray{typeof(init)}(undef, 0)

kargs1 = _kernel_args(N, op, ret_inst, f, x...)
kargs1 = _kernel_args(N, op, ret_inst, init, f, x...)
kernel1, threads1 = _kernel_maxthreads(_parallel_reduce_amdgpu, kargs1)

rret = AMDGPU.ROCArray([init])
kargs2 = _kernel_args(1, op, ret_inst, rret)
kargs2 = _kernel_args(1, op, ret_inst, init, rret)
kernel2, threads2 = _kernel_maxthreads(_reduce_kernel_amdgpu, kargs2)

threads = min(threads1, threads2, 512)
blocks = cld(N, threads)

shmem_size = threads * sizeof(init)

ret = fill!(AMDGPU.ROCArray{typeof(init)}(undef, blocks), init)
kargs1 = _kernel_args(N, op, ret, f, x...)
# no fill! — `init` is seeded inside the kernels; every block writes its slot
ret = AMDGPU.ROCArray{typeof(init)}(undef, blocks)
kargs1 = _kernel_args(N, op, ret, init, f, x...)
kernel1(
kargs1...; groupsize = threads, gridsize = blocks, shmem = shmem_size)

kargs2 = _kernel_args(blocks, op, ret, rret)
kargs2 = _kernel_args(blocks, op, ret, init, rret)
kernel2(kargs2...; groupsize = threads, gridsize = 1, shmem = shmem_size)
AMDGPU.synchronize()

Expand All @@ -317,12 +378,12 @@ function JACC._parallel_reduce!(reducer::JACC.ParallelReduce{AMDGPUBackend},
wk = reducer.workspace
_init!(wk, blocks, init)

kargs1 = _kernel_args((M, N), op, wk.tmp, f, x...)
kargs1 = _kernel_args((M, N), op, wk.tmp, init, f, x...)
kernel1 = _make_kernel(_parallel_reduce_amdgpu_MN, kargs1)
kernel1(kargs1...; groupsize = threads, gridsize = blocks,
shmem = shmem_size, stream = reducer.stream)

kargs2 = _kernel_args(blocks, op, wk.tmp, wk.ret)
kargs2 = _kernel_args(blocks, op, wk.tmp, init, wk.ret)
kernel2 = _make_kernel(_reduce_kernel_amdgpu_MN, kargs2)
kernel2(kargs2...; groupsize = threads, gridsize = (1, 1),
shmem = shmem_size, stream = reducer.stream)
Expand All @@ -344,14 +405,15 @@ function JACC.parallel_reduce(f, ::AMDGPUBackend, (M, N)::NTuple{2, Integer},
Nblocks = cld(N, Nthreads)
blocks = (Mblocks, Nblocks)
shmem_size = 16 * 16 * sizeof(init)
ret = fill!(AMDGPU.ROCArray{typeof(init)}(undef, blocks), init)
# no fill! — `init` is seeded inside the kernels; every block writes its slot
ret = AMDGPU.ROCArray{typeof(init)}(undef, blocks)
rret = AMDGPU.ROCArray([init])

kargs1 = _kernel_args((M, N), op, ret, f, x...)
kargs1 = _kernel_args((M, N), op, ret, init, f, x...)
kernel1 = _make_kernel(_parallel_reduce_amdgpu_MN, kargs1)
kernel1(kargs1...; groupsize = threads, gridsize = blocks, shmem = shmem_size)

kargs2 = _kernel_args(blocks, op, ret, rret)
kargs2 = _kernel_args(blocks, op, ret, init, rret)
kernel2 = _make_kernel(_reduce_kernel_amdgpu_MN, kargs2)
kernel2(kargs2...; groupsize = threads, gridsize = (1, 1), shmem = shmem_size)

Expand Down Expand Up @@ -394,12 +456,12 @@ end
return nothing
end

@inline function _parallel_reduce_amdgpu(N, op, ret, f, x...)
@inline function _parallel_reduce_amdgpu(N, op, ret, init, f, x...)
shmem_length = workgroupDim().x
shared_mem = @ROCDynamicLocalArray(eltype(ret), shmem_length, false)
i = (workgroupIdx().x - 1) * workgroupDim().x + workitemIdx().x
ti = workitemIdx().x
@inbounds shared_mem[ti] = ret[workgroupIdx().x]
@inbounds shared_mem[ti] = init

if i <= N
tmp = @inline f(i, x...)
Expand All @@ -421,12 +483,12 @@ end
return nothing
end

function _reduce_kernel_amdgpu(N, op, red, ret)
function _reduce_kernel_amdgpu(N, op, red, init, ret)
shmem_length = workgroupDim().x
shared_mem = @ROCDynamicLocalArray(eltype(ret), shmem_length, false)
i = workitemIdx().x
ii = i
@inbounds tmp = ret[1]
tmp = init
for ii in i:shmem_length:N
tmp = op(tmp, @inbounds red[ii])
end
Expand All @@ -447,7 +509,7 @@ function _reduce_kernel_amdgpu(N, op, red, ret)
return nothing
end

function _parallel_reduce_amdgpu_MN((M, N), op, ret, f, x...)
function _parallel_reduce_amdgpu_MN((M, N), op, ret, init, f, x...)
shared_mem = @ROCDynamicLocalArray(eltype(ret), (16, 16), false)
i = (workgroupIdx().x - 1) * workgroupDim().x + workitemIdx().x
j = (workgroupIdx().y - 1) * workgroupDim().y + workitemIdx().y
Expand All @@ -456,7 +518,7 @@ function _parallel_reduce_amdgpu_MN((M, N), op, ret, f, x...)
bi = workgroupIdx().x
bj = workgroupIdx().y

@inbounds shared_mem[ti, tj] = ret[bi, bj]
@inbounds shared_mem[ti, tj] = init

if (i <= M && j <= N)
tmp = @inline f(i, j, x...)
Expand All @@ -481,12 +543,12 @@ function _parallel_reduce_amdgpu_MN((M, N), op, ret, f, x...)
return nothing
end

function _reduce_kernel_amdgpu_MN((M, N), op, red, ret)
function _reduce_kernel_amdgpu_MN((M, N), op, red, init, ret)
shared_mem = @ROCDynamicLocalArray(eltype(ret), (16, 16), false)
i = workitemIdx().x
j = workitemIdx().y

@inbounds tmp = ret[1]
tmp = init
for ci in CartesianIndices((i:16:M, j:16:N))
tmp = op(tmp, @inbounds red[ci])
end
Expand Down
Loading
Loading