Sorting is the most reused primitive in data systems: joins, group-bys, deduplication, index builds, ranking and the expert routing step in mixture-of-experts layers all sort or partially sort keys. On one core the problem is solved; a good introsort or radix sort is close to memory bandwidth. Parallel sorting is harder, because the work must be split so that every core gets the same amount, the pieces must be stitched back together without a sequential bottleneck, and the data movement between cores or machines usually costs more than the comparisons.
This article builds the three designs that real libraries use: parallel merge sort with merge-path partitioning, sample sort, and scan-based radix sort. Each comes with tested code and, for sample sort, measured load imbalance. It ends with the library calls to reach for first and the failure modes that show up in production.
Work, span and memory traffic
Use the work-span model. Work W is the total number of operations; span S is the length of the longest chain of dependent operations, the time with infinitely many processors. A greedy scheduler on p processors finishes in at most W/p + S, so a good parallel algorithm keeps W close to the best sequential cost and makes S tiny. For comparison sorting W cannot go below order n log n; radix sort escapes that bound by not comparing, at the cost of work proportional to n times the number of digit passes.
The model leaves out the thing that dominates on real hardware: memory traffic. Sorting a billion 8-byte keys moves 8 GB per pass over the data. At tens of gigabytes per second on a CPU socket, each pass costs a substantial fraction of a second regardless of how many cores you add. Good parallel sorts minimise the number of passes over memory, and on clusters, the number of times each key crosses the network, which ideally is once.
Parallel merge sort and merge path
Parallel merge sort recursively sorts the two halves in parallel and merges. With a sequential merge, the final merge alone takes n steps, so the span is order n and the speed-up is capped at about log n regardless of core count. The fix is a parallel merge.
The textbook parallel merge (CLRS's P-MERGE) takes the median of the larger array, binary searches it in the other, and recurses on both sides in parallel; its span is order log squared n, giving merge sort a span of order log cubed n. Cole's merge sort reaches order log n span but is too intricate for practical use. What practical libraries do instead is partition the output: to give p workers equal shares of the merged output, find for each output position i the split (j, k) with j + k = i such that the first i outputs are exactly the first j elements of a and the first k of b. This is called co-ranking, or merge path, because it is where the merge path crosses a diagonal of the a-by-b grid. Each split is one binary search, and every worker then merges its own slice sequentially.
def co_rank(i, a, b):
"""(j, k) with j + k == i and merge(a, b)[:i] == merge(a[:j], b[:k]); ties favour a."""
lo, hi = max(0, i - len(b)), min(i, len(a))
while lo < hi:
j = (lo + hi) // 2
if a[j] <= b[i - j - 1]: # a[j] precedes b[i-j-1]: take more of a
lo = j + 1
else:
hi = j
return lo, i - lo
def parallel_merge(a, b, p):
n = len(a) + len(b)
cuts = [co_rank(n * t // p, a, b) for t in range(p + 1)]
parts = []
for t in range(p): # independent: one worker per t
(j0, k0), (j1, k1) = cuts[t], cuts[t + 1]
parts.append(merge_seq(a[j0:j1], b[k0:k1]))
return parts # concatenation is the merged output
def merge_seq(x, y):
out, i, j = [], 0, 0
while i < len(x) and j < len(y):
if x[i] <= y[j]: # ties take from x (a): stable
out.append(x[i]); i += 1
else:
out.append(y[j]); j += 1
return out + x[i:] + y[j:]Worked example. Merge a = [64, 120, 137, 261, 460, 483, 507, 582, 779, 782, 821, 867] with b = [29, 96, 214, 388, 499, 667, 807, 914] on 4 workers. The cuts come out as (0,0), (3,2), (6,4), (9,6), (12,8): worker 1, for example, merges a[3:6] = [261, 460, 483] with b[2:4] = [214, 388]. Every worker gets exactly 5 outputs, whatever the distribution of values between the two inputs. The routine was checked against Python's sorted() on 2,000 random cases with heavy duplication and 1 to 7 workers. The tie rule in the comparison is what keeps the merge stable; get it wrong and equal keys cross between workers.
Sample sort and splitter quality
Merge-based sorts touch every key once per level of the merge tree. Sample sort instead moves each key once: choose p - 1 splitters that cut the key space into p ranges, send each key to the worker owning its range, and sort locally. The concatenated buckets are already in global order. It is the standard design for distributed sorting, where the all-to-all exchange is the dominant cost and you can afford only one.
import bisect, random
def sample_sort(data, p, s, rng):
"""p buckets, oversampling factor s (samples per bucket)."""
sample = sorted(rng.sample(data, min(len(data), p * s)))
splitters = [sample[t * len(sample) // p] for t in range(1, p)]
buckets = [[] for _ in range(p)]
for x in data: # in parallel: each worker routes its own chunk
buckets[bisect.bisect_right(splitters, x)].append(x)
return [sorted(b) for b in buckets] # in parallel: one bucket per workerThe whole algorithm lives or dies on splitter quality, because the slowest worker sets the finish time. Measured on 1,000,000 uniform random floats with p = 64 buckets (ideal bucket 15,625 keys), averaged over 5 sampling seeds:
| samples per bucket s | largest bucket / ideal (mean) | worst of 5 |
|---|---|---|
| 1 | 4.64 | 5.90 |
| 4 | 2.61 | 2.99 |
| 16 | 1.74 | 2.22 |
| 64 | 1.30 | 1.41 |
| 256 | 1.15 | 1.19 |
With one sample per bucket, the unluckiest worker does nearly five times its share, so 64 workers deliver the speed of about 14. Oversampling by 64 brings the slowest worker within 30 to 40 percent of ideal for a sample that is only 0.4 percent of the input. Duplicates are the other trap: with keys drawn from just three values, the same code put a third of the input in one bucket and left 61 of 64 buckets empty, because equal keys cannot be split by value. Breaking ties with the original index, sorting (key, index) pairs, brought the largest bucket back to 1.27 times ideal.
Radix sort built on scans
Radix sort processes keys a digit at a time, typically 8 bits. Least-significant-digit (LSD) radix sort needs every pass to be stable, and each pass is three parallel steps: count how many keys fall in each of the 256 digit values, exclusive-scan the counts to get output offsets, and scatter every key to its offset. The scan is the same primitive covered in the parallel prefix sum article.
def radix_pass(keys, shift, bits=8):
r = 1 << bits
hist = [0] * r
for x in keys: # parallel: per-block histograms, then summed
hist[(x >> shift) & (r - 1)] += 1
offs, run = [0] * r, 0
for d in range(r): # exclusive scan of the counts
offs[d], run = run, run + hist[d]
out = [0] * len(keys)
for x in keys: # stable scatter
d = (x >> shift) & (r - 1)
out[offs[d]] = x
offs[d] += 1
return out
# 32-bit unsigned keys: four passes, low byte first.
for shift in (0, 8, 16, 24):
keys = radix_pass(keys, shift)On a GPU, each thread block histograms its tile, a scan over (digit, block) counts gives every block its private output offset for every digit, and a second sweep scatters, so keys stay stable without any locking. That classic design reads the keys twice per pass. Onesweep (Adinets and Merrill, NVIDIA, 2022) cuts this to essentially one read per pass with a single-pass chained scan across blocks, and its authors report about 1.5 times the speed of the previous CUB radix sort on an A100. Radix sort needs keys that order correctly as unsigned integers: for signed integers flip the sign bit; for IEEE floats flip all bits of negative values and only the sign bit of non-negative ones. Variable-length strings suit most-significant-digit variants or comparison sorts better.
Sorting networks such as bitonic sort are the fourth family. Their fixed, data-independent comparisons suit SIMD registers and small GPU tiles, and they often sort the base cases inside the designs above; the sorting networks article covers them in detail. At large n they lose to radix and sample sort because they do order n log squared n work.
Library calls to use first
Reach for a library before writing any of this. The calls below are the ones in common use.
// C++17 parallel algorithms (GCC's libstdc++ needs TBB as the backend; link -ltbb)
#include <algorithm>
#include <execution>
std::sort(std::execution::par_unseq, v.begin(), v.end());
// CUB device radix sort: first call sizes temp storage, second call sorts
void* d_temp = nullptr; size_t temp_bytes = 0;
cub::DeviceRadixSort::SortKeys(d_temp, temp_bytes, d_keys_in, d_keys_out, n);
cudaMalloc(&d_temp, temp_bytes);
cub::DeviceRadixSort::SortKeys(d_temp, temp_bytes, d_keys_in, d_keys_out, n);
// Java: fork/join parallel merge sort for arrays
java.util.Arrays.parallelSort(a);
// Rust with rayon
use rayon::prelude::*;
v.par_sort_unstable();Thrust's thrust::sort dispatches to CUB's radix sort for primitive keys with the default comparator and to a merge sort otherwise, so a custom comparator can silently cost you the fast path. In data frames, Polars and DuckDB already sort in parallel; check their thread settings before reaching lower. For more background on the sequential building blocks, see mergesort and quicksort.
Operational guidance
- Benchmark against the sequential baseline on the same keys. A parallel sort that is 3 times faster on 32 cores may be memory-bound; the fix is fewer passes, not more threads.
- Sort keys and indices, then gather. Moving wide records through every pass multiplies memory traffic; sort (key, row id) pairs and permute the payload once.
- Pick the algorithm by key type. Fixed-width integers and floats: radix. Strings and custom comparators: merge or sample sort. Distributed data: sample sort with oversampling.
- Watch skew in production. Log the largest-to-mean bucket ratio for distributed sorts; a rise usually means a hot key or a new duplicate-heavy column.
- Respect NUMA and huge pages. First-touch allocation on the wrong socket can sharply cut bandwidth; initialise buffers with the same threads that will sort them.
- Decide on stability up front. Unstable sorts are faster, but multi-key orderings built from successive sorts, and reproducible output, need stability.
Failure modes
- Sequential merge at the top level. The final merge serialises n work and caps the speed-up near log n; use merge-path partitioning.
- Too few samples. One sample per bucket gave a slowest bucket 4.6 times ideal on average in the measurement above.
- Duplicate keys collapsing buckets. Break ties with a secondary key or route equal runs round-robin.
- Float keys sorted as raw bits. Negative floats come out reversed and after the positives; apply the bit transform, and decide where NaN belongs.
- Comparator that is not a strict weak ordering. Parallel implementations can crash or return unsorted output where a sequential sort happened to survive.
- Out-of-memory on the scatter. Radix and sample sort need an output buffer the size of the input; budget twice the data, or sort in chunks and merge.
Trade-offs
| design | work | passes over data | best for | weak spot |
|---|---|---|---|---|
| merge sort + merge path | n log n | log(n / block) merges | stable, any comparator | many passes |
| sample sort | n log n | about 2 (route, local sort) | distributed, large n | skew, duplicates |
| LSD radix sort | n x passes | key bits / digit bits | integer and float keys on GPU | wide or variable keys |
| bitonic network | n log squared n | many | small tiles, SIMD | extra work at scale |
What to do next
- Measure your current sort: keys per second, and memory bandwidth used against the machine's peak.
- Replace hand-written sorts with std::sort with an execution policy, rayon, Arrays.parallelSort or CUB, and benchmark on production-sized inputs.
- If records are wide, switch to sorting (key, index) pairs and gathering once.
- For distributed sorts, add oversampling of at least 32 to 64 per bucket and log the bucket skew ratio.
- Test with adversarial inputs: all-equal keys, already sorted, reverse sorted, negative floats and NaN.