Skip to main content

KERNEL_RADIX_SORT

Constant KERNEL_RADIX_SORT 

Source
pub const KERNEL_RADIX_SORT: &str = r#"
// Radix sort kernel (LSB first, 4 bits per pass)
// PASS define: 0 = count, 1 = scatter
// Uses prefix sums (from prefix_sum kernel) between passes.

layout(local_size_x = 256) in;

#ifndef PASS
#define PASS 0
#endif

#ifndef RADIX_BITS
#define RADIX_BITS 4
#endif

#define RADIX (1 << RADIX_BITS)

layout(std430, binding = 0) buffer KeysIn {
    uint keys_in[];
};

layout(std430, binding = 1) buffer KeysOut {
    uint keys_out[];
};

layout(std430, binding = 2) buffer ValuesIn {
    uint values_in[];
};

layout(std430, binding = 3) buffer ValuesOut {
    uint values_out[];
};

layout(std430, binding = 4) buffer Offsets {
    uint offsets[];      // RADIX * num_blocks
};

layout(std430, binding = 5) buffer GlobalOffsets {
    uint global_offsets[];  // RADIX prefix sums
};

uniform uint u_n;
uniform uint u_bit_offset;  // which 4-bit nibble (0, 4, 8, 12, ...)

shared uint local_counts[RADIX];

uint extract_digit(uint key, uint bit_offset) {
    return (key >> bit_offset) & uint(RADIX - 1);
}

void main() {
    uint lid = gl_LocalInvocationID.x;
    uint gid = gl_WorkGroupID.x;
    uint idx = gl_GlobalInvocationID.x;

#if PASS == 0
    // Count pass: count occurrences of each digit in this block
    if (lid < uint(RADIX)) {
        local_counts[lid] = 0u;
    }
    barrier();

    if (idx < u_n) {
        uint digit = extract_digit(keys_in[idx], u_bit_offset);
        atomicAdd(local_counts[digit], 1u);
    }
    barrier();

    // Write local counts to global offset table
    if (lid < uint(RADIX)) {
        offsets[lid * gl_NumWorkGroups.x + gid] = local_counts[lid];
    }

#elif PASS == 1
    // Scatter pass: place each element at its globally sorted position
    if (idx < u_n) {
        uint key = keys_in[idx];
        uint digit = extract_digit(key, u_bit_offset);

        // global_offsets[digit] = prefix sum of all counts for this digit
        // We need the exact output position: global prefix + local prefix
        // This is a simplified scatter; a production sort uses local prefix sums
        uint dest = atomicAdd(global_offsets[digit], 1u);
        if (dest < u_n) {
            keys_out[dest] = key;
            values_out[dest] = values_in[idx];
        }
    }

#endif
}
"#;
Expand description

GPU radix sort kernel.