A wavelet tree stores a sequence of symbols from an alphabet of size sigma in about n log2(sigma) bits. That is roughly the size of the sequence itself. Yet it answers a family of questions in O(log sigma) time each: what is the i-th symbol, how many times does c occur before position i, where is the j-th c, what is the k-th smallest value in positions lo to hi, and how many values in that range fall below x. Other structures answer some of these; the wavelet tree answers all of them from one compressed structure.
This article builds the structure from first principles: rank on bit vectors, the bit-by-bit partition, and the wavelet matrix layout that most production libraries use instead of a pointer tree. It gives tested Python code, traces two queries by hand, and covers select, applications and engineering costs. All ranges in this article are half-open: positions lo to hi means indices lo, lo+1, ..., hi-1.
Rank and select on bit vectors
Everything rests on one primitive. Given a bit vector B of length n, rank1(i) counts the ones in B[0:i] and rank0(i) = i - rank1(i) counts the zeros. A plain prefix-sum array answers both in O(1) but uses a machine word per position, which is 64 times the size of the bits. Succinct bit vectors get O(1) rank with o(n) extra bits by using two levels of counts. Superblocks store the absolute count of ones every 512 bits or so. Blocks store a small relative count every 64 bits. The remainder within a word is a single hardware popcount on a masked word. Overhead depends on block sizes. select1(j), the position of the j-th one, is the inverse of rank. It is usually implemented with sampled positions plus a short scan, and is slower in practice than rank.
Rank lets a query map a position in one level to the matching position in the next, with no pointers. Fenwick trees handle prefix sums that change. Here the bits never change, so rank can be constant-time and tiny.
From bit partitions to the wavelet matrix
Write each value as a fixed number of bits, L = ceil(log2 sigma). Look at the most significant bit of every element and record it as a bit vector. Now split the sequence, stably, into the elements whose bit was 0 and those whose bit was 1. Each part is a sequence over half the alphabet. Recurse on each part with the next bit. After L levels every part holds a single value. That recursion, drawn as a binary tree with one bit vector per node, is the wavelet tree of Grossi, Gupta and Vitter (2003).
A pointer tree over a large alphabet has about sigma nodes, each with its own bit vector and rank directory, and the overhead adds up. The wavelet matrix of Claude and Navarro (2012) fixes this. At every level, instead of keeping each node's part separate, it concatenates the whole level into one bit vector of length n. It does this by stably moving all 0-bit elements to the front and all 1-bit elements to the back. Each level then needs only one bit vector and one integer, Z[level], the number of zeros at that level. If an element at position i has bit 0, its position on the next level is rank0(i). If the bit is 1, it is Z[level] + rank1(i). Elements with the same prefix of bits stay contiguous and in their original order, so a range of the original sequence maps to a contiguous range at every level, which is the property every query uses.
A complete implementation
The implementation below is complete and was tested against brute force on random sequences. The bit vector uses a prefix-count list for clarity. Swap in a succinct rank structure for real sizes.
class BitVector:
# bits plus prefix counts of ones: rank1(i) = ones in bits[0:i]
def __init__(self, bits):
self.bits = bits
self.ones = [0]
for b in bits:
self.ones.append(self.ones[-1] + b)
def rank1(self, i): return self.ones[i]
def rank0(self, i): return i - self.ones[i]
class WaveletMatrix:
def __init__(self, seq, sigma_bits):
self.n, self.L = len(seq), sigma_bits
self.levels, self.zeros = [], []
cur = list(seq)
for lvl in range(self.L):
shift = self.L - 1 - lvl # most significant bit first
bits = [(v >> shift) & 1 for v in cur]
bv = BitVector(bits)
self.levels.append(bv)
self.zeros.append(bv.rank0(self.n))
cur = [v for v, b in zip(cur, bits) if b == 0] + \
[v for v, b in zip(cur, bits) if b == 1] # stable split
def access(self, i):
v = 0
for lvl, bv in enumerate(self.levels):
b = bv.bits[i]
v = (v << 1) | b
i = bv.rank0(i) if b == 0 else self.zeros[lvl] + bv.rank1(i)
return v
def rank(self, c, i): # occurrences of c in seq[0:i]
lo, hi = 0, i
for lvl, bv in enumerate(self.levels):
if (c >> (self.L - 1 - lvl)) & 1 == 0:
lo, hi = bv.rank0(lo), bv.rank0(hi)
else:
lo = self.zeros[lvl] + bv.rank1(lo)
hi = self.zeros[lvl] + bv.rank1(hi)
return hi - lo
def quantile(self, lo, hi, k): # k-th smallest (0-based) in seq[lo:hi]
v = 0
for lvl, bv in enumerate(self.levels):
z = bv.rank0(hi) - bv.rank0(lo) # zeros inside the range
if k < z:
lo, hi, v = bv.rank0(lo), bv.rank0(hi), v << 1
else:
k -= z
lo = self.zeros[lvl] + bv.rank1(lo)
hi = self.zeros[lvl] + bv.rank1(hi)
v = (v << 1) | 1
return v
def count_less(self, lo, hi, x): # values < x in seq[lo:hi]
if x >= (1 << self.L):
return hi - lo
count = 0
for lvl, bv in enumerate(self.levels):
if (x >> (self.L - 1 - lvl)) & 1:
count += bv.rank0(hi) - bv.rank0(lo) # everything with a 0 here is smaller
lo = self.zeros[lvl] + bv.rank1(lo)
hi = self.zeros[lvl] + bv.rank1(hi)
else:
lo, hi = bv.rank0(lo), bv.rank0(hi)
return countEach query is one loop over L levels with at most two rank calls per level, so it runs in O(log sigma) time regardless of n or of the range width. Construction is O(n log sigma). A range frequency query, "how many values in [lo, hi) lie in [a, b)", is count_less(lo, hi, b) - count_less(lo, hi, a).
Worked example: tracing two queries
Take S = [5, 2, 7, 1, 2, 6, 3, 2] with L = 3. Figure 1 shows the three levels, and running the code prints these rows (comments added):
level 0 10100100 zeros 5
level 1 10111011 zeros 2
level 2 11001010 zeros 4
access [5, 2, 7, 1, 2, 6, 3, 2]
rank(2, 6) 2
quantile(1,6,2) 2 # sorted(S[1:6]) = [1, 2, 2, 6, 7]
count_less(2,8,3) 3 # S[2:8] = [7, 1, 2, 6, 3, 2]access(3). At level 0, position 3 holds bit 0. Remap to rank0(3), the number of zeros in "101", which is 1. At level 1, position 1 holds bit 0. Remap to rank0(1) = 0. At level 2, position 0 holds bit 1. The bits read are 0, 0, 1, so the value is 1, which matches S[3].
quantile(1, 6, 2) asks for the third smallest of S[1:6] = [2, 7, 1, 2, 6]. At level 0 the range is positions 1 to 6. It contains rank0(6) - rank0(1) = 4 - 0 = 4 zeros. Since k = 2 is less than 4, the answer has top bit 0, and the range becomes [rank0(1), rank0(6)) = [0, 4). At level 1, bits 0 to 4 are "1011", with one zero. Since k = 2 is not below 1, the answer's middle bit is 1. Subtract to get k = 1, and move to [Z1 + rank1(0), Z1 + rank1(4)) = [2, 5). At level 2, bits 2 to 5 are "001", with two zeros. Since k = 1 is below 2, the last bit is 0. The bits are 0, 1, 0, so the answer is 2, which is correct.
The logic is a binary search over values, guided by counts of zeros in the range.
Select: walking back up
Select, the position of the j-th occurrence of c, runs the rank walk in reverse. First go down the levels following c's bits from position 0, recording the start of c's block at each level. At the bottom, the j-th occurrence is at offset start + j. Then climb back up. At each level, if c's bit there is 0, the position p came from the p-th zero, so p = select0(p). Otherwise it came from the (p - Z)-th one, so p = select1(p - Z[level]).
def select(wm, c, j): # position of the j-th (0-based) occurrence of c
p = 0
for lvl, bv in enumerate(wm.levels): # descend: start of c's block
bit = (c >> (wm.L - 1 - lvl)) & 1
p = bv.rank0(p) if bit == 0 else wm.zeros[lvl] + bv.rank1(p)
p += j
for lvl in reversed(range(wm.L)): # ascend with select on each level
bv, bit = wm.levels[lvl], (c >> (wm.L - 1 - lvl)) & 1
p = bv.select0(p) if bit == 0 else bv.select1(p - wm.zeros[lvl])
return pThe sketch assumes a bit vector with select0 and select1, which the simple class above omits. The caller must check that j is less than rank(c, n).
Applications
Where wavelet trees earn their place:
- FM-index over large alphabets. Backward search in a compressed full-text index needs rank(c, i) over the Burrows-Wheeler transform. For DNA a few bit vectors suffice. For bytes, words or token ids, a wavelet tree over the BWT gives rank in O(log sigma) in compressed space. This is the classic use, and it pairs naturally with suffix arrays.
- Two-dimensional range counting. Sort n points by x and store their y values, after rank reduction, as the sequence. The number of points in [x1, x2) times [y1, y2) is then a range frequency query: two binary searches for the x range and two count_less calls.
- Range k-th smallest and medians over static logs or time-ordered values.
- Compressed sequences. Give the tree a Huffman shape, so frequent symbols get short paths, and back it with compressed bit vectors. Space then approaches the entropy of the sequence.
Engineering for real sizes
Several engineering details decide whether the structure is fast in practice.
- Compress the alphabet first. Map values to their ranks among distinct values, so L is ceil(log2 of the distinct count) rather than 64 bits. Keep the sorted distinct list to map answers back.
- Count the cache misses. A query costs about two cache misses per level, so 40 for 20 levels on large inputs. Interleaving rank counts with the bits in one cache line helps.
- Space. The bits take n L bits, plus the rank and select overhead, plus L integers. For 10^8 values over 2^20 distinct symbols, that is 2 x 10^9 bits, about 250 MB before the rank directory. The toy prefix list in the code stores a Python integer per bit and would need many times that. Use a library: sdsl-lite provides
wt_intandwm_intfor C++, and similar structures exist in Rust and Java succinct libraries. - Construction memory. The stable split needs O(n) words of scratch space.
- The structure is static. Dynamic variants add a log factor; prefer batch rebuilds.
Trade-offs against other range structures
| Structure | Range k-th / count | Space | Updates | Best when |
|---|---|---|---|---|
| Wavelet matrix | O(log sigma) | n log sigma bits + o(n) | Static | Large static data, many query types, memory-bound |
| Persistent segment tree | O(log n) | O(n log n) words | Append versions | Online k-th queries, simpler code |
| Merge sort tree | O(log^2 n) to O(log^3 n) | O(n log n) words | Static | Small n, quick to write |
| Mo's algorithm | Offline, about O((n + q) sqrt n) | O(n) | Offline only | Queries known in advance, odd aggregates |
For contest-sized inputs, a persistent segment tree gives the same k-th smallest answers with code many people find easier to get right. If all queries are known ahead of time, Mo's algorithm handles aggregates neither tree supports. Choose the wavelet matrix when memory matters, the data is static, and you need several query types on the same sequence.
Failure modes
Bugs that show up repeatedly:
- Mixing closed and half-open ranges between callers and the structure. Pick half-open and assert it at the API boundary.
- Forgetting the + Z[level] offset on the 1-branch. Access still works on some inputs, which hides the bug until rank or quantile breaks.
- Not handling x at or above 2^L in count_less, or k at or beyond the range width in quantile. Validate both.
- Querying with raw values after compressing the alphabet. Map query bounds through the same sorted list, and use lower bounds for the half-open value ranges.
What to do next
- Run the code above on S, then add a brute-force random test like the one used to validate this article.
- Implement a two-level rank directory with popcount and compare its speed and memory with the prefix list.
- Add
select0andselect1with sampled positions, then implement select and check it against brute force. - Solve a two-dimensional point counting problem with the matrix, using alphabet compression on y.
- Build a small FM-index over a token-id sequence, using the matrix for rank on the BWT.
- For production sizes, benchmark sdsl-lite's
wm_inton your data before writing your own.