A merge sort tree is a segment tree whose nodes store sorted lists instead of single numbers. Each node keeps every value of its index range in sorted order, which is exactly the state merge sort reaches after merging that range. The structure answers a family of static range questions that plain segment trees cannot: how many values in positions l to r are at most x, and what the k-th smallest value in positions l to r is.
This article builds the tree from first principles, works a full example by hand, gives tested Python for counting and k-th queries, and then shows two improvements: fractional cascading, which removes a log factor from counting, and an index-sorted variant that answers k-th queries in O(log² n). It ends with a comparison against wavelet trees and persistent segment trees, so you can choose the right one.
The structure
Start with a segment tree over positions 0 to n-1. Node v covers a contiguous range of positions; its children split that range in half. In a merge sort tree, node v stores the sorted multiset of a[i] for i in its range. A leaf stores one value, and an internal node's list is the merge of its two children's lists, built in linear time, the same step merge sort uses.
Every position appears once per level, and there are about log₂ n levels, so the total stored size is n log₂ n plus n. For n = 1,000,000 that is about 21 million integers: roughly 84 MB as 32-bit integers in C++, and several times more in Python lists. Building costs O(n log n) time with merging. Sorting each node from scratch costs O(n log² n), which is a common slip.
A range query on positions l to r is split into the usual canonical segment tree nodes, at most two per level, so O(log n) of them. Each canonical node can say how many of its values are at most x with a binary search in O(log n). Adding the answers gives the count in O(log² n).
Worked example by hand
Worked example. Take a = [5, 1, 4, 2, 3, 1, 6, 2] and the query: how many values in positions 1 to 6 are at most 3? Directly, the slice is 1, 4, 2, 3, 1, 6, and four of those values are at most 3.
With the tree, positions 1 to 6 decompose into four canonical nodes: leaf 1 holding [1], node 2-3 holding [2, 4], node 4-5 holding [1, 3], and leaf 6 holding [6]. Binary searches for the last value at most 3 give 1, 1, 2 and 0, which sum to 4.
Now ask for the 3rd smallest value in positions 1 to 6. The sorted slice is 1, 1, 2, 3, 4, 6, so the answer is 2. The tree finds it by binary search over candidate answers. Counting values at most 1 gives 2, which is less than 3, so the answer is larger than 1. Counting values at most 2 gives 3, which reaches k, so the answer is 2. This is the classic O(log³ n) k-th query: a binary search over values, with an O(log² n) count inside each step.
Python implementation
The implementation below uses an iterative bottom-up segment tree of size equal to the next power of two, so padding leaves stay empty and need no special case. Queries use inclusive 0-based positions.
from bisect import bisect_left, bisect_right
from heapq import merge
class MergeSortTree:
def __init__(self, a):
self.n = len(a)
self.size = 1
while self.size < self.n:
self.size *= 2
self.t = [[] for _ in range(2 * self.size)]
for i, v in enumerate(a):
self.t[self.size + i] = [v]
for v in range(self.size - 1, 0, -1): # linear merge per node
self.t[v] = list(merge(self.t[2 * v], self.t[2 * v + 1]))
self.values = sorted(set(a)) # candidate answers for kth
def count_leq(self, l, r, x):
# How many a[i] <= x for l <= i <= r. O(log^2 n).
res = 0
l += self.size
r += self.size + 1 # half-open [l, r)
while l < r:
if l & 1:
res += bisect_right(self.t[l], x)
l += 1
if r & 1:
r -= 1
res += bisect_right(self.t[r], x)
l //= 2
r //= 2
return res
def kth(self, l, r, k):
# k-th smallest (1-based) in a[l..r]. O(log^3 n).
if not 1 <= k <= r - l + 1:
raise ValueError("k out of range")
lo, hi = 0, len(self.values) - 1
while lo < hi:
mid = (lo + hi) // 2
if self.count_leq(l, r, self.values[mid]) >= k:
hi = mid
else:
lo = mid + 1
return self.values[lo]
t = MergeSortTree([5, 1, 4, 2, 3, 1, 6, 2])
assert t.count_leq(1, 6, 3) == 4
assert t.kth(1, 6, 3) == 2Two details matter. The binary search runs over the sorted distinct values, not over array positions, so the answer is always a value that exists. And counting uses bisect_right for at most x; for strictly less than x, use bisect_left. Mixing the two is the most common source of wrong answers with duplicates.
Faster queries
Fractional cascading for counting. The O(log² n) count repeats a binary search in every canonical node, although the searches are related: if p values in a node are at most x, then the values at most x are exactly the first p entries of its list. Because the merge is stable, store for each node an array left_upto, where left_upto[i] is how many of the node's first i entries came from the left child. Search once at the root, then walk down: the left child's count is left_upto[p] and the right child's is the rest. Counting becomes O(log n) at the cost of a second array the size of the tree.
def count_leq_fc(v, nl, nr, l, r, p):
# p = number of values <= x stored in node v, which covers positions nl..nr.
if r < nl or nr < l or p == 0:
return 0
if l <= nl and nr <= r:
return p
mid = (nl + nr) // 2
pl = left_upto[v][p]
return (count_leq_fc(2 * v, nl, mid, l, r, pl) +
count_leq_fc(2 * v + 1, mid + 1, nr, l, r, p - pl))
# answer = count_leq_fc(1, 0, size - 1, l, r, bisect_right(root_list, x))Index-sorted tree for k-th. Swap the roles of values and positions. Sort the positions by value (breaking ties by position) and build the tree over that order, so each node stores the original positions of a contiguous band of value ranks, sorted by position. To find the k-th smallest in l..r, start at the root and count how many positions in the left child fall inside l..r, with two binary searches. If that count c is at least k, descend left; otherwise subtract c from k and descend right. The leaf reached is the answer. Each level costs O(log n), so a query costs O(log² n), and adding cascading arrays brings it to O(log n), which is the idea a wavelet tree compresses into bit vectors.
class KthTree:
def __init__(self, a):
self.a = a
order = sorted(range(len(a)), key=lambda i: (a[i], i))
self.size = 1
while self.size < len(a):
self.size *= 2
self.t = [[] for _ in range(2 * self.size)]
for j, i in enumerate(order):
self.t[self.size + j] = [i]
for v in range(self.size - 1, 0, -1):
self.t[v] = list(merge(self.t[2 * v], self.t[2 * v + 1]))
def kth(self, l, r, k):
v = 1
while v < self.size:
left = self.t[2 * v]
c = bisect_right(left, r) - bisect_left(left, l)
if k <= c:
v = 2 * v
else:
k -= c
v = 2 * v + 1
return self.a[self.t[v][0]]
Where it fits and where it does not
A realistic use is percentile queries over a static log. Store one day of request latencies in arrival order. A dashboard that asks for the 95th percentile between any two timestamps maps the timestamps to positions with a binary search, then asks for the k-th smallest with k equal to the ceiling of 0.95 times the slice length. A question such as how many requests in the window exceeded 500 ms is the slice length minus a count of values at most 500. Offline, the same tree answers two-dimensional dominance counts: points with index in a range and value below a threshold.
The tree is a poor fit when values change. A point update must delete and insert a value in O(log n) sorted lists, and list insertion is linear, so updates cost O(n) in the worst case. If updates are rare, rebuild periodically. If they are frequent, use a Fenwick tree of order-statistic trees, square-root decomposition with sorted blocks, or offline methods such as Mo's algorithm.
Choosing between range k-th structures
| Structure | Memory | Count at most x | k-th smallest | Updates |
|---|---|---|---|---|
| Merge sort tree | O(n log n) words | O(log² n) | O(log³ n) | O(n) worst case |
| With fractional cascading | about 2x the above | O(log n) | O(log² n) | rebuild |
| Index-sorted tree | O(n log n) words | not direct | O(log² n) | rebuild |
| Persistent segment tree | O(n log n) nodes | O(log n) | O(log n) | not in place |
| Wavelet tree or matrix | n log σ bits plus rank | O(log σ) | O(log σ) | hard |
The merge sort tree wins on simplicity: about thirty lines, no pointer juggling and no bit-level rank structures, which makes it a good first choice in contests and prototypes. Persistent segment trees and wavelet trees answer k-th in a single descent and use less memory, so prefer them for large n or tight latency targets.
Failure modes
- Mixing inclusive and half-open bounds between the build and the query, which drops or double-counts the last position.
- Using
bisect_leftwhere at most x is meant, which undercounts every duplicate of x. - Binary searching k-th over array positions instead of sorted values, which returns a value that is not in the slice.
- Re-sorting every node instead of merging, which turns an O(n log n) build into O(n log² n) and can time out.
- Python memory blow-up: 20 million list entries cost far more than 80 MB. Use
array('i')or NumPy arrays per level. - Underestimating build cost in Python: merging 20 million entries through interpreted loops takes seconds, so build once, persist the levels, and reuse them across queries.
- Floating-point values containing NaN, which break the ordering that every binary search assumes.
What to do next
- Implement
MergeSortTreeand test it against a brute-force slice-and-sort on random arrays with many duplicates. - Add the k-th query and assert that the answer always appears in the slice.
- Implement the index-sorted variant and compare query times for n = 200,000.
- Add fractional cascading to counting and measure the speed-up against the extra memory.
- Read the segment tree article if canonical decomposition is unfamiliar, then wavelet trees for the compact version.
- Decide per workload: static data and simple code suit the merge sort tree; large n or updates need another structure.