Skip to main content

KERNEL_PREFIX_SUM

Constant KERNEL_PREFIX_SUM 

Source
pub const KERNEL_PREFIX_SUM: &str = r#"
// Blelloch parallel prefix sum (exclusive scan)
// Two-phase: up-sweep (reduce) then down-sweep.
// Works on a single workgroup; for larger arrays, use multi-block with auxiliary sums.
//
// PHASE define: 0 = up-sweep, 1 = down-sweep, 2 = add block offsets

layout(local_size_x = 256) in;

#ifndef PHASE
#define PHASE 0
#endif

layout(std430, binding = 0) buffer Data {
    uint data[];
};

layout(std430, binding = 1) buffer BlockSums {
    uint block_sums[];
};

uniform uint u_n;           // number of elements
uniform uint u_block_size;  // elements per block (2 * local_size)

shared uint temp[512]; // 2 * local_size_x

void main() {
    uint lid = gl_LocalInvocationID.x;
    uint gid = gl_WorkGroupID.x;
    uint block_offset = gid * u_block_size;

#if PHASE == 0
    // Load into shared memory
    uint ai = lid;
    uint bi = lid + 256u;
    uint a_idx = block_offset + ai;
    uint b_idx = block_offset + bi;

    temp[ai] = (a_idx < u_n) ? data[a_idx] : 0u;
    temp[bi] = (b_idx < u_n) ? data[b_idx] : 0u;
    barrier();

    // Up-sweep (reduce)
    uint offset = 1u;
    for (uint d = 512u >> 1u; d > 0u; d >>= 1u) {
        barrier();
        if (lid < d) {
            uint ai2 = offset * (2u * lid + 1u) - 1u;
            uint bi2 = offset * (2u * lid + 2u) - 1u;
            temp[bi2] += temp[ai2];
        }
        offset <<= 1u;
    }
    barrier();

    // Store block sum and clear last element
    if (lid == 0u) {
        block_sums[gid] = temp[511u];
        temp[511u] = 0u;
    }
    barrier();

    // Down-sweep
    for (uint d = 1u; d < 512u; d <<= 1u) {
        offset >>= 1u;
        barrier();
        if (lid < d) {
            uint ai2 = offset * (2u * lid + 1u) - 1u;
            uint bi2 = offset * (2u * lid + 2u) - 1u;
            uint t = temp[ai2];
            temp[ai2] = temp[bi2];
            temp[bi2] += t;
        }
    }
    barrier();

    // Write back
    if (a_idx < u_n) data[a_idx] = temp[ai];
    if (b_idx < u_n) data[b_idx] = temp[bi];

#elif PHASE == 2
    // Add block offsets for multi-block scan
    if (gid > 0u) {
        uint a_idx2 = block_offset + lid;
        uint b_idx2 = block_offset + lid + 256u;
        uint block_sum = block_sums[gid];

        if (a_idx2 < u_n) data[a_idx2] += block_sum;
        if (b_idx2 < u_n) data[b_idx2] += block_sum;
    }

#endif
}
"#;
Expand description

Blelloch parallel prefix sum (exclusive scan).