A prefix sum, or scan, turns [x0, x1, x2, ...] into [x0, x0+x1, x0+x1+x2, ...]. Written as a loop it is the most sequential thing imaginable, because every output depends on the one before. Yet scan is one of the most important parallel primitives there is. GPU radix sort, stream compaction, sparse matrix formats, histogram equalisation, and the selective state-space layers in models like Mamba all reduce to it.
This page builds parallel scan from first principles. It explains why the loop is the wrong shape, walks through the two classic algorithms with real traces, writes a CUDA warp and block scan, explains how libraries scan a whole array in one pass, and covers what scan is used to build. It ends with the floating-point behaviour that surprises people the first time a parallel cumsum disagrees with a sequential one.
Inclusive, exclusive and associative
Given a binary operator ⊕ and an input array x, the inclusive scan produces y[i] = x[0] ⊕ ... ⊕ x[i]. The exclusive scan shifts by one and starts from the identity: y[0] = identity and y[i] = x[0] ⊕ ... ⊕ x[i-1]. Exclusive scan is the one you want for computing output offsets, because element i's offset is the count of everything before it.
The operator does not have to be addition. Any associative operator works: max, min, multiplication, bitwise OR, matrix product, or the pair operator for linear recurrences shown later. Associativity, (a ⊕ b) ⊕ c = a ⊕ (b ⊕ c), is the whole trick: it lets you group the work into a tree instead of a chain. Commutativity is not required, which matters for matrix products and recurrences.
Two measures describe a parallel algorithm. Work is the total number of operations. Span is the length of the longest chain of dependent operations, the time with unlimited processors. The sequential loop has work n and span n. A good parallel scan keeps work at O(n) and cuts span to O(log n). The work-span view is developed further in the BSP model in depth.
Hillis-Steele: short span, too much work
The simplest parallel scan, from Hillis and Steele, runs log2 n steps. In the step with distance d, every element i ≥ d adds the value d places to its left, all at once. After the step with d = 1 each element holds the sum of 2 inputs; after d = 2, of 4; after d = 4, of 8. On [3, 1, 7, 0, 4, 1, 6, 3] the three steps produce [3, 4, 8, 7, 4, 5, 7, 9], [3, 4, 11, 11, 12, 12, 11, 14] and [3, 4, 11, 11, 15, 16, 22, 25], the inclusive scan.
Span is log2 n, which is optimal, but work is about n log2 n, because almost every element does an add at every step. For a million elements that is twenty times the adds of the sequential loop. Each step also needs double buffering, since element i must read the old value of i - d. Hillis-Steele is still the right choice at small sizes where the hardware runs in lockstep anyway, which is exactly what happens inside a warp.
Blelloch: work-efficient up-sweep and down-sweep
Blelloch's scan reaches O(n) work and O(log n) span with two passes over an implicit binary tree. The up-sweep is a reduction: at distance d, every element at index 2d-1, 4d-1, ... adds the element d to its left, so the last element ends up holding the total. The down-sweep clears the root to the identity and walks back down; at each node, the left child receives the parent's value and the right child receives the parent's value plus the old left child. The result is the exclusive scan, in place.
def blelloch_exclusive(x): # len(x) must be a power of two
x = list(x); n = len(x); d = 1
while d < n: # up-sweep: reduce in a tree
for i in range(2*d - 1, n, 2*d): # every iteration is independent
x[i] += x[i - d]
d *= 2
total, x[n-1] = x[n-1], 0 # clear the root
d = n // 2
while d >= 1: # down-sweep: distribute prefixes
for i in range(2*d - 1, n, 2*d):
t = x[i - d]; x[i - d] = x[i]; x[i] += t
d //= 2
return x, totalEach inner loop is a parallel step; its iterations touch disjoint pairs. Both sweeps do n - 1 adds, so total work is about 2n, with span 2 log2 n. Running it on the example returns [0, 3, 4, 11, 11, 15, 16, 22] with total 25. Figure 1 shows every intermediate array.
A CUDA scan: warp, block, grid
A GPU scan is built in layers that match the hardware. A warp of 32 threads executes together and can exchange registers with shuffle instructions, so the bottom layer is a Hillis-Steele scan across a warp. Its extra work is free, because idle lanes would do nothing anyway. The next layer combines warps within a block through shared memory. The top layer combines blocks across the grid.
__device__ float warp_inclusive_scan(float x) {
unsigned lane = threadIdx.x & 31;
for (int d = 1; d < 32; d <<= 1) {
float y = __shfl_up_sync(0xffffffff, x, d); // value from lane - d
if (lane >= d) x += y;
}
return x;
}
__global__ void block_scan(const float* in, float* out, float* block_sums, int n) {
__shared__ float warp_totals[32];
int i = blockIdx.x * blockDim.x + threadIdx.x;
int lane = threadIdx.x & 31, warp = threadIdx.x >> 5;
float x = (i < n) ? in[i] : 0.0f;
x = warp_inclusive_scan(x); // layer 1: within a warp
if (lane == 31) warp_totals[warp] = x;
__syncthreads();
if (warp == 0) { // layer 2: scan the warp totals
float t = (lane < (blockDim.x >> 5)) ? warp_totals[lane] : 0.0f;
warp_totals[lane] = warp_inclusive_scan(t);
}
__syncthreads();
if (warp > 0) x += warp_totals[warp - 1];
if (i < n) out[i] = x;
if (threadIdx.x == blockDim.x - 1) block_sums[blockIdx.x] = x; // for layer 3
}Production kernels have each thread first scan several elements serially in registers, then run the warp and block layers on per-thread totals. That raises the work done per memory access, which matters because scan does almost no arithmetic. Shared-memory banking and synchronisation are explained in GPU shared memory architecture.
Scanning a whole array: three grid strategies
Scan is memory-bound. The floor is reading n inputs and writing n outputs, about 2n memory operations, and the grid-level strategy decides how close you get to it.
| Strategy | How blocks combine | Memory traffic | Notes |
|---|---|---|---|
| Scan-then-propagate | Scan each block, scan the block sums, add each block's offset | about 4n | Three kernels; easy to write |
| Reduce-then-scan | Reduce each block, scan the sums, rescan each block with its offset | about 3n | Reads input twice |
| Single pass with decoupled look-back | Each block finds its offset from predecessors' published results | about 2n | One kernel; used by CUB |
Decoupled look-back, described by Merrill and Garland in a 2016 NVIDIA technical report, works like this. Each block takes a tile index from an atomic counter, so tiles start in order. It reduces its tile and publishes the aggregate with a flag. Then it looks back at its predecessors: it adds their aggregates until it reaches one that has published an inclusive prefix, at which point it knows its own offset, publishes its own inclusive prefix and writes the output. Most blocks find a finished predecessor within a few steps, so the serial chain between blocks stays short. The atomic tile counter matters: the hardware does not promise to schedule blocks in blockIdx order, so a block that used blockIdx as its tile number could spin forever on a predecessor that has not started.
Use the library
Write your own scan to learn it, then call a library. CUB follows a two-call pattern: the first call with a null temporary buffer only reports the scratch size.
// CUDA C++ with CUB
void* d_temp = nullptr; size_t temp_bytes = 0;
cub::DeviceScan::InclusiveSum(d_temp, temp_bytes, d_in, d_out, n); // size query
cudaMalloc(&d_temp, temp_bytes);
cub::DeviceScan::InclusiveSum(d_temp, temp_bytes, d_in, d_out, n); // the scan
# PyTorch
y = torch.cumsum(x, dim=0)
# JAX: any associative function over pytrees
y = jax.lax.associative_scan(jnp.add, x)Thrust exposes the same operation as thrust::inclusive_scan and thrust::exclusive_scan. JAX's associative_scan is the one to reach for when the operator is not addition, because it takes any associative function.
What scan builds
Stream compaction keeps the elements that pass a predicate, in order, without gaps. Compute a 0/1 flag per element, take the exclusive scan of the flags, and each kept element writes to its scanned offset. Keeping positives from [5, -2, 9, 0, -7, 3, 8, -1] gives flags [1, 0, 1, 0, 0, 1, 1, 0], offsets [0, 1, 1, 2, 2, 2, 3, 4] and output [5, 9, 3, 8]. GPU filtering, culling and sparse output all use this.
Radix sort is compaction repeated per digit: each pass scans digit counts to find where each key goes. Segmented scan restarts at segment boundaries by carrying a flag with each value, which lets one kernel scan many variable-length rows. That is how CSR sparse-matrix products and ragged batches are handled.
Linear recurrences h[t] = a[t] h[t-1] + b[t] look sequential, but pairs (a, b) combine associatively: (a1, b1) then (a2, b2) equals (a1 a2, a2 b1 + b2). Scanning the pairs gives every h[t] in O(log n) span. With a = [0.5, 0.9, 0.1, 1.0] and b = [1, 2, 3, 4], both the loop and the scan give [1.0, 2.9, 3.29, 7.29]. Selective state-space models train in parallel using this idea; see selective SSMs.
Numerics and determinism
Floating-point addition is not associative, so a parallel scan rearranges the rounding. Summing 220 standard normals in float32 gave 934.5408 sequentially and 934.5487 when NumPy summed it as 1,024 block sums; the float64 answer is 934.54884. The blocked order was closer, because its partial sums stay smaller, but neither is bit-identical to the other. Tests that compare a GPU cumsum with a CPU loop must use a tolerance.
Low-precision accumulators fail harder. A float16 running sum of 4,096 copies of 0.1 stops at 256.0: from there the gap between representable values is 0.25, so adding 0.1 rounds back to the same number. Accumulating in float32 gives 409.5. Keep the accumulator in float32 or wider whatever the input type.
Determinism is a separate question. A fixed tree is deterministic. Schemes where the combination order depends on timing, such as atomics, can change results between runs. If you need bitwise reproducibility, check what your library documents.
Failure modes
- Inclusive or exclusive confusion. Using an inclusive scan for offsets shifts every write by one element, overwriting a neighbour and leaving slot 0 empty.
- Non-associative operator. Subtraction, an average of two values, or a floating-point operator with clamping all give order-dependent results. Check associativity on random triples before parallelising.
- Assumed commutativity. Matrix products and recurrence pairs must keep left and right in order. A kernel that swaps operands passes tests with additions and fails on the real operator.
- Padding with the wrong identity. Pad a max-scan with zeros and negative inputs come out wrong. The pad value must be the operator's identity.
- Overflow. Offsets from a scan over billions of flags overflow 32-bit integers. Use 64-bit counters for large inputs.
- Missing synchronisation. Leaving out a barrier between block layers reads warp totals before they are written, which fails rarely and only on some hardware.
Trade-offs
Hillis-Steele trades work for simplicity and is ideal inside a warp. Blelloch is work-efficient but has two passes and more synchronisation. Across a grid, the single-pass look-back scan gets closest to the 2n memory floor, at the cost of a subtle protocol that is best left to a library. Fewer, larger tiles per block raise efficiency but need more registers and shared memory. In every case scan is bandwidth-bound, so the right goal is a kernel that runs at memory speed, not fewer adds. Matrix multiply is the opposite case, compute-bound with heavy reuse; compare how a matrix multiply runs on a GPU.
What to do next
- Implement Blelloch in Python, check it against
itertools.accumulateon random arrays, and reproduce the trace in Figure 1. - Write the warp and block CUDA kernels above and check them against
torch.cumsumwith a float tolerance. - Benchmark
cub::DeviceScan::InclusiveSumagainst your kernel and report GB/s against your GPU's memory bandwidth. - Rewrite a filter in your codebase as flag, exclusive scan and scatter.
- Try
jax.lax.associative_scanwith the (a, b) pair operator on a linear recurrence and compare it with a Python loop. - Audit every scan for float32 accumulators, 64-bit offsets and tolerance-based tests.