Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Wavelet tree: subarray counts and order statistics

Last updated: 5 Oct 20269 min read
tutorial
IntermediateBy AITrove Editorial

A wavelet tree indexes a fixed sequence by recursively partitioning its value range. Each node stores prefix counts of values routed left. A query over positions [start, stop) maps that interval into child coordinates with two prefix reads. Count-below returns how many values in the subarray are strictly less than a threshold. Select returns the value at a zero-based ordinal within the subarray's sorted multiset, so duplicate values occupy distinct ordinals. This is a static index; editing a latency after construction would invalidate stored partitions and prefixes. A plain sorted copy cannot answer arbitrary original-position subarray questions with the same information.

Operational case

Index latency readings 61, 19, 47, 19, 83, 26, and 52. In positions [1, 6), three readings are below 47: 19, 19, and 26. The value at sorted ordinal 2 for that subarray is 26. Across the full sequence, ordinal 6 selects 83. Repeated readings remain separate occurrences even though they share one leaf value. A query for an empty interval can count zero, but it cannot select any ordinal. Invalid position or ordinal bounds raise rather than returning a misleading sentinel reading.

Working Python program

python
"""Static wavelet tree for half-open subarray count and selection queries."""

from dataclasses import dataclass


@dataclass
class WaveletNode:
    low: int
    high: int
    left_prefix: list[int]
    left: "WaveletNode | None" = None
    right: "WaveletNode | None" = None


class IncidentQuantiles:
    def __init__(self, latencies: list[int]):
        self.length = len(latencies)
        self.root = self._build(latencies, min(latencies), max(latencies)) if latencies else None

    @classmethod
    def _build(cls, values: list[int], low: int, high: int) -> WaveletNode:
        node = WaveletNode(low, high, [0])
        if low == high:
            node.left_prefix.extend(range(1, len(values) + 1))
            return node
        middle = (low + high) // 2
        smaller, larger = [], []
        for latency in values:
            is_left = latency <= middle
            node.left_prefix.append(node.left_prefix[-1] + is_left)
            (smaller if is_left else larger).append(latency)
        if smaller:
            node.left = cls._build(smaller, low, middle)
        if larger:
            node.right = cls._build(larger, middle + 1, high)
        return node

    def _check_range(self, start: int, stop: int) -> None:
        if not 0 <= start <= stop <= self.length:
            raise IndexError((start, stop))

    def count_below(self, start: int, stop: int, threshold: int) -> int:
        self._check_range(start, stop)

        def visit(node: WaveletNode | None, left: int, right: int) -> int:
            if node is None or left == right or threshold <= node.low:
                return 0
            if node.high < threshold:
                return right - left
            left_start, left_stop = node.left_prefix[left], node.left_prefix[right]
            return (visit(node.left, left_start, left_stop)
                    + visit(node.right, left - left_start, right - left_stop))

        return visit(self.root, start, stop)

    def select(self, start: int, stop: int, ordinal: int) -> int:
        self._check_range(start, stop)
        if not 0 <= ordinal < stop - start:
            raise IndexError(ordinal)
        node = self.root
        while node.low != node.high:
            left_start, left_stop = node.left_prefix[start], node.left_prefix[stop]
            left_count = left_stop - left_start
            if ordinal < left_count:
                node, start, stop = node.left, left_start, left_stop
            else:
                ordinal -= left_count
                node, start, stop = node.right, start - left_start, stop - left_stop
        return node.low


quantiles = IncidentQuantiles([61, 19, 47, 19, 83, 26, 52])
print(quantiles.count_below(1, 6, 47))
print(quantiles.select(1, 6, 2))
print(quantiles.select(0, 7, 6))

Output

Output
3
26
83

Time, space, and tradeoff

Let N be the number of readings and sigma the covered integer value span. Building partitions and prefix arrays takes O(N log sigma) time and O(N log sigma) integer entries in this straightforward Python representation. Each count or select query follows at most O(log sigma) levels and uses O(log sigma) stack space for count or O(1) working space for iterative select. Wide sparse integer domains can make this uncompressed value-depth larger than needed; coordinate compression reduces height to O(log D) for D distinct values. A scan costs O(subarray length) per query and is simpler when query volume is low.

Common Mistakes

  • Do not drop duplicate values when answering an ordinal query.
  • Do not mix global positions with child-local positions after routing through prefix counts.
  • Do not promise point updates on this immutable index.
  • Do not quote compressed-alphabet bounds for this uncompressed implementation.

Connected lessons

Apply the invariant in the sparse readings and task constraints project, then check the operations quiz.

Merge-sort trees: count readings below a threshold in one interval extends this range-query decision.

Fenwick frequency index: select the kth stored key extends this operation choice.

Wavelet matrices: count frequencies and find subarray quantiles adds a distinct structure contract to compare.

data structures
range-query-structures
Storage details