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).