Skip to content
Merged
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
12 changes: 12 additions & 0 deletions docs/src/api/sort.md
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,18 @@ v = ROCArray(rand(Int32, 100_000))
AK.sort!(v)
```

Multidimensional arrays are sorted as one flat vector by default; pass `dims` to sort each 1D slice
along that dimension independently, like `Base.sort!(A; dims)`. `sortperm` along `dims` returns
linear indices into the array, so `A[ix]` is sorted along `dims`:
```julia
A = ROCArray(rand(Float32, 1000, 1000))
AK.sort!(A; dims=1) # each column sorted
ix = AK.sortperm(A; dims=2) # A[ix] has each row sorted
```

On GPU backends `dims` uses merge sort (`RadixSort()` does not support it); on CPU backends each
slice is sorted with `Base.sort!`.

As GPU memory is more expensive, all functions in AcceleratedKernels.jl expose any temporary arrays they will use (the `temp` argument); you can supply your own buffers to make the algorithms not allocate additional GPU storage, e.g.:
```julia
v = ROCArray(rand(Float32, 100_000))
Expand Down
89 changes: 52 additions & 37 deletions src/sort/merge_sort.jl
Original file line number Diff line number Diff line change
@@ -1,26 +1,29 @@
@kernel inbounds=true cpu=false unsafe_indices=true function _merge_sort_block!(vec, comp)
@kernel inbounds=true cpu=false unsafe_indices=true function _merge_sort_block!(
vec, comp, layout, blocks_per_slice,
)

@uniform N = @groupsize()[1]
s_buf = @localmem eltype(vec) (N * 0x2,)

T = eltype(vec)
I = typeof(N)
len = length(vec)

# NOTE: for many index calculations in this library, computation using zero-indexing leads to
# fewer operations (also code is transpiled to CUDA / ROCm / oneAPI / Metal code which do zero
# indexing). Internal calculations will be done using zero indexing except when actually
# accessing memory. As with C, the lower bound is inclusive, the upper bound exclusive.
# Use zero-based indices internally and half-open search bounds.

# Group (block) and local (thread) indices
iblock = @index(Group, Linear) - 0x1
ithread = @index(Local, Linear) - 0x1

# Each block sorts a tile of one slice
islice, iblock = slice_block(layout, iblock, blocks_per_slice)
elems = slice(vec, layout, islice)
len = layout.len

i = ithread + iblock * N * 0x2
i < len && (s_buf[ithread + 0x1] = vec[i + 0x1])
i < len && (s_buf[ithread + 0x1] = elems[i + 0x1])

i = ithread + N + iblock * N * 0x2
i < len && (s_buf[ithread + N + 0x1] = vec[i + 0x1])
i < len && (s_buf[ithread + N + 0x1] = elems[i + 0x1])

@synchronize()

Expand Down Expand Up @@ -68,28 +71,30 @@
end

i = ithread + iblock * N * 0x2
i < len && (vec[i + 0x1] = s_buf[ithread + 0x1])
i < len && (elems[i + 0x1] = s_buf[ithread + 0x1])

i = ithread + N + iblock * N * 0x2
i < len && (vec[i + 0x1] = s_buf[ithread + N + 0x1])
i < len && (elems[i + 0x1] = s_buf[ithread + N + 0x1])
end


@kernel inbounds=true cpu=false unsafe_indices=true function _merge_sort_global!(
@Const(vec_in), vec_out, comp, half_size_group,
@Const(vec_in), vec_out, comp, half_size_group, layout, blocks_per_slice,
)
len = length(vec_in)
N = @groupsize()[1]

# NOTE: for many index calculations in this library, computation using zero-indexing leads to
# fewer operations (also code is transpiled to CUDA / ROCm / oneAPI / Metal code which do zero
# indexing). Internal calculations will be done using zero indexing except when actually
# accessing memory. As with C, the lower bound is inclusive, the upper bound exclusive.
# Use zero-based indices internally and half-open search bounds.

# Group (block) and local (thread) indices
iblock = @index(Group, Linear) - 0x1
ithread = @index(Local, Linear) - 0x1

# Merges never cross slice boundaries
islice, iblock = slice_block(layout, iblock, blocks_per_slice)
slice_in = slice(vec_in, layout, islice)
slice_out = slice(vec_out, layout, islice)
len = layout.len

idx = ithread + iblock * N
size_group = half_size_group * 0x2
gid = idx ÷ half_size_group
Expand All @@ -101,23 +106,23 @@ end
if lo >= len
# Incomplete left half, nothing to swap on the right, simply copy elements to be sorted
# in next iteration
pos_in < len && (vec_out[pos_in + 0x1] = vec_in[pos_in + 0x1])
pos_in < len && (slice_out[pos_in + 0x1] = slice_in[pos_in + 0x1])
else

hi = (gid + 0x1) * size_group
hi > len && (hi = len)

pos_out = pos_in + _lower_bound_s0(vec_in, vec_in[pos_in + 0x1], lo, hi, comp) - lo
vec_out[pos_out + 0x1] = vec_in[pos_in + 0x1]
pos_out = pos_in + _lower_bound_s0(slice_in, slice_in[pos_in + 0x1], lo, hi, comp) - lo
slice_out[pos_out + 0x1] = slice_in[pos_in + 0x1]

# Right half
pos_in = gid * size_group + half_size_group + idx % half_size_group

if pos_in < len
lo = gid * size_group
hi = lo + half_size_group
pos_out = pos_in - half_size_group + _upper_bound_s0(vec_in, vec_in[pos_in + 0x1], lo, hi, comp) - lo
vec_out[pos_out + 0x1] = vec_in[pos_in + 0x1]
pos_out = pos_in - half_size_group + _upper_bound_s0(slice_in, slice_in[pos_in + 0x1], lo, hi, comp) - lo
slice_out[pos_out + 0x1] = slice_in[pos_in + 0x1]
end
end
end
Expand All @@ -134,6 +139,9 @@ end

block_size::Int=256,
temp::Union{Nothing, AbstractArray}=nothing,

# Sort each 1D slice along this dimension; `:` sorts the whole array as one vector
dims::Union{Colon, Integer}=Colon(),
)
"""
function merge_sort!(
Expand All @@ -146,40 +154,43 @@ function merge_sort!(

block_size::Int=256,
temp::Union{Nothing, AbstractArray}=nothing,
dims::Union{Colon, Integer}=Colon(),
)
# Simple sanity checks
@argcheck block_size > 0
layout = slice_layout(v, dims)
ord = Base.Order.ord(lt, by, rev, order)
if !isnothing(temp)
@argcheck length(temp) == length(v)
@argcheck eltype(temp) === eltype(v)
end
(isempty(v) || layout.len <= 1) && return v

# Hoist `by` transform: broadcast it once to produce a key array, then sort
# (key, value) pairs. Without hoisting, by(elem) fires inside every binary-search
# comparison in the O(n log²n) merge hot-path — once per element per merge step.
# Compute keys once instead of evaluating `by` in every comparison.
if by !== identity
keys = by.(v)
merge_sort_by_key!(
keys, v, backend;
lt, rev, order, block_size,
lt, rev, order, block_size, dims,
temp_values=temp, # temp was for v swap buffer; maps to temp_values here
)
return v
end

if !isnothing(temp)
@argcheck length(temp) == length(v)
@argcheck eltype(temp) === eltype(v)
end

# Construct comparator
ord = Base.Order.ord(lt, by, rev, order)
comp = (x, y) -> Base.Order.lt(ord, x, y)

# Block level
blocks = (length(v) + block_size * 2 - 1) ÷ (block_size * 2)
_merge_sort_block!(backend, block_size)(v, comp, ndrange=(block_size * blocks,))
# Block level: each block sorts a tile of one slice in local memory
len = layout.len
blocks = (len + block_size * 2 - 1) ÷ (block_size * 2)
_merge_sort_block!(backend, block_size)(
v, comp, layout, blocks,
ndrange=(block_size * blocks * slice_count(layout),),
)

# Global level
# Global level: merge the sorted tiles of each slice, doubling the run length every pass
half_size_group = Int32(block_size * 2)
size_group = half_size_group * 2
len = length(v)
if len > half_size_group
p1 = v
p2 = isnothing(temp) ? similar(v) : temp
Expand All @@ -189,7 +200,10 @@ function merge_sort!(
niter = 0
while len > half_size_group
blocks = ((len + half_size_group - 1) ÷ half_size_group + 1) ÷ 2 * (half_size_group ÷ block_size)
kernel!(p1, p2, comp, half_size_group, ndrange=(block_size * blocks,))
kernel!(
p1, p2, comp, half_size_group, layout, blocks,
ndrange=(block_size * blocks * slice_count(layout),),
)

half_size_group = half_size_group << 1;
size_group = size_group << 1;
Expand Down Expand Up @@ -218,6 +232,7 @@ end

block_size::Int=256,
temp::Union{Nothing, AbstractArray}=nothing,
dims::Union{Colon, Integer}=Colon(),
)
"""
function merge_sort(
Expand Down
76 changes: 50 additions & 26 deletions src/sort/merge_sort_by_key.jl
Original file line number Diff line number Diff line change
@@ -1,31 +1,35 @@
@kernel inbounds=true cpu=false unsafe_indices=true function _merge_sort_by_key_block!(keys, values, comp)
@kernel inbounds=true cpu=false unsafe_indices=true function _merge_sort_by_key_block!(
keys, values, comp, layout, blocks_per_slice,
)

@uniform N = @groupsize()[1]
s_keys = @localmem eltype(keys) (N * 0x2,)
s_values = @localmem eltype(values) (N * 0x2,)

I = typeof(N)
len = length(keys)

# NOTE: for many index calculations in this library, computation using zero-indexing leads to
# fewer operations (also code is transpiled to CUDA / ROCm / oneAPI / Metal code which do zero
# indexing). Internal calculations will be done using zero indexing except when actually
# accessing memory. As with C, the lower bound is inclusive, the upper bound exclusive.
# Use zero-based indices internally and half-open search bounds.

# Group (block) and local (thread) indices
iblock = @index(Group, Linear) - 0x1
ithread = @index(Local, Linear) - 0x1

# Each block sorts a tile of one slice
islice, iblock = slice_block(layout, iblock, blocks_per_slice)
slice_keys = slice(keys, layout, islice)
slice_values = slice(values, layout, islice)
len = layout.len

i = ithread + iblock * N * 0x2
if i < len
s_keys[ithread + 0x1] = keys[i + 0x1]
s_values[ithread + 0x1] = values[i + 0x1]
s_keys[ithread + 0x1] = slice_keys[i + 0x1]
s_values[ithread + 0x1] = slice_values[i + 0x1]
end

i = ithread + N + iblock * N * 0x2
if i < len
s_keys[ithread + N + 0x1] = keys[i + 0x1]
s_values[ithread + N + 0x1] = values[i + 0x1]
s_keys[ithread + N + 0x1] = slice_keys[i + 0x1]
s_values[ithread + N + 0x1] = slice_values[i + 0x1]
end

@synchronize()
Expand Down Expand Up @@ -85,36 +89,40 @@

i = ithread + iblock * N * 0x2
if i < len
keys[i + 0x1] = s_keys[ithread + 0x1]
values[i + 0x1] = s_values[ithread + 0x1]
slice_keys[i + 0x1] = s_keys[ithread + 0x1]
slice_values[i + 0x1] = s_values[ithread + 0x1]
end

i = ithread + N + iblock * N * 0x2
if i < len
keys[i + 0x1] = s_keys[ithread + N + 0x1]
values[i + 0x1] = s_values[ithread + N + 0x1]
slice_keys[i + 0x1] = s_keys[ithread + N + 0x1]
slice_values[i + 0x1] = s_values[ithread + N + 0x1]
end
end


@kernel inbounds=true cpu=false unsafe_indices=true function _merge_sort_by_key_global!(
@Const(keys_in), keys_out,
@Const(values_in), values_out,
comp, half_size_group,
comp, half_size_group, layout, blocks_per_slice,
)

len = length(keys_in)
N = @groupsize()[1]

# NOTE: for many index calculations in this library, computation using zero-indexing leads to
# fewer operations (also code is transpiled to CUDA / ROCm / oneAPI / Metal code which do zero
# indexing). Internal calculations will be done using zero indexing except when actually
# accessing memory. As with C, the lower bound is inclusive, the upper bound exclusive.
# Use zero-based indices internally and half-open search bounds.

# Group (block) and local (thread) indices
iblock = @index(Group, Linear) - 0x1
ithread = @index(Local, Linear) - 0x1

# Merges never cross slice boundaries
islice, iblock = slice_block(layout, iblock, blocks_per_slice)
keys_in = slice(keys_in, layout, islice)
keys_out = slice(keys_out, layout, islice)
values_in = slice(values_in, layout, islice)
values_out = slice(values_out, layout, islice)
len = layout.len

idx = ithread + iblock * N
size_group = half_size_group * 0x2
gid = idx ÷ half_size_group
Expand Down Expand Up @@ -167,6 +175,9 @@ end
block_size::Int=256,
temp_keys::Union{Nothing, AbstractArray}=nothing,
temp_values::Union{Nothing, AbstractArray}=nothing,

# Sort each 1D slice along this dimension; `:` sorts the whole array as one vector
dims::Union{Colon, Integer}=Colon(),
)
"""
function merge_sort_by_key!(
Expand All @@ -182,10 +193,15 @@ function merge_sort_by_key!(
block_size::Int=256,
temp_keys::Union{Nothing, AbstractArray}=nothing,
temp_values::Union{Nothing, AbstractArray}=nothing,
dims::Union{Colon, Integer}=Colon(),
)
# Simple sanity checks
@argcheck block_size > 0
@argcheck length(keys) == length(values)
layout = slice_layout(keys, dims)
if !(dims isa Colon)
@argcheck axes(keys) == axes(values)
end
if !isnothing(temp_keys)
@argcheck length(temp_keys) == length(keys)
@argcheck eltype(temp_keys) === eltype(keys)
Expand All @@ -197,16 +213,20 @@ function merge_sort_by_key!(

# Construct comparator
ord = Base.Order.ord(lt, by, rev, order)
(isempty(keys) || layout.len <= 1) && return keys, values
comp = (x, y) -> Base.Order.lt(ord, x, y)

# Block level
blocks = (length(keys) + block_size * 2 - 1) ÷ (block_size * 2)
_merge_sort_by_key_block!(backend, block_size)(keys, values, comp, ndrange=(block_size * blocks,))
# Block level: each block sorts a tile of one slice in local memory
len = layout.len
blocks = (len + block_size * 2 - 1) ÷ (block_size * 2)
_merge_sort_by_key_block!(backend, block_size)(
keys, values, comp, layout, blocks,
ndrange=(block_size * blocks * slice_count(layout),),
)

# Global level
# Global level: merge the sorted tiles of each slice, doubling the run length every pass
half_size_group = Int32(block_size * 2)
size_group = half_size_group * 2
len = length(keys)
if len > half_size_group
pk1 = keys
pk2 = isnothing(temp_keys) ? similar(keys) : temp_keys
Expand All @@ -219,7 +239,10 @@ function merge_sort_by_key!(
niter = 0
while len > half_size_group
blocks = ((len + half_size_group - 1) ÷ half_size_group + 1) ÷ 2 * (half_size_group ÷ block_size)
kernel!(pk1, pk2, pv1, pv2, comp, half_size_group, ndrange=(block_size * blocks,))
kernel!(
pk1, pk2, pv1, pv2, comp, half_size_group, layout, blocks,
ndrange=(block_size * blocks * slice_count(layout),),
)

half_size_group = half_size_group << 1;
size_group = size_group << 1;
Expand Down Expand Up @@ -253,6 +276,7 @@ end
block_size::Int=256,
temp_keys::Union{Nothing, AbstractArray}=nothing,
temp_values::Union{Nothing, AbstractArray}=nothing,
dims::Union{Colon, Integer}=Colon(),
)
"""
function merge_sort_by_key(
Expand Down
Loading
Loading