A persistent data structure keeps every version of itself. After an update you can still query the version before it, and the version before that, at full speed. In a persistent segment tree each update creates a new version in O(log n) time and space by copying only the nodes it touches and sharing everything else.
The classic payoff is a query that looks impossible online: the k-th smallest value in any subarray a[l..r], answered in O(log n) after O(n log n) preprocessing. All code below was checked against brute force on thousands of random arrays.
Path copying
Start from an ordinary segment tree, covered in depth in the segment tree guide. A point update walks one path from the root to a leaf and recomputes the nodes on that path. Nothing off the path changes.
That observation is the whole trick. To make the update persistent, do not modify the path; build a fresh copy of it. Each new node points to its new child on the path and to the old, unchanged sibling. The new root therefore describes the updated array, the old root still describes the old array, and the two trees share every node except the O(log n) on the path. Copying the whole tree per version would cost O(n) per update; path copying costs one node per level.
The rule that makes this safe is immutability: once a node is created, its fields never change. Any code that writes into an existing node corrupts every version that shares it, which is the bug behind most broken implementations.
A node pool with a shared empty node
Nodes are stored in parallel arrays rather than as objects. A node is an index; left[i], right[i] and cnt[i] are its fields. That avoids per-object overhead across millions of nodes and allows one preallocation.
Index 0 is a sentinel empty node whose children are itself and whose count is 0. An empty tree of any size is then just root 0: no build step is needed, because every missing subtree is the sentinel. Inserting into an empty tree creates exactly the nodes on one path.
Building versions over compressed values
For the k-th smallest query, the tree is built over values, not positions. Compress the distinct values to indices 0..m-1 and let each leaf count how many times its value has been inserted. Version i is the tree after inserting a[0], a[1], ..., a[i-1]; version 0 is empty. So version i answers questions about the prefix of length i.
import bisect
class PersistentCountTree:
"""Counts of compressed values; version i holds a[0..i-1]. Node 0 is the shared empty node."""
def __init__(self, m, capacity):
self.m = m # number of distinct values (leaves)
self.left = [0] * capacity
self.right = [0] * capacity
self.cnt = [0] * capacity
self.size = 1 # node 0 already exists: left = right = 0, cnt = 0
def _new(self, l, r, c):
i = self.size
self.left[i], self.right[i], self.cnt[i] = l, r, c
self.size += 1
return i
def insert(self, prev, pos):
"""Return a new root equal to prev plus one occurrence of value pos."""
lo, hi = 0, self.m - 1
path = []
node = prev
while lo < hi: # walk down, remembering which way we went
mid = (lo + hi) // 2
if pos <= mid:
path.append((node, 0)); node = self.left[node]; hi = mid
else:
path.append((node, 1)); node = self.right[node]; lo = mid + 1
child = self._new(0, 0, self.cnt[node] + 1) # the new leaf
for old, went_right in reversed(path): # copy the path bottom-up
if went_right:
child = self._new(self.left[old], child, self.cnt[old] + 1)
else:
child = self._new(child, self.right[old], self.cnt[old] + 1)
return child
def build(a):
vals = sorted(set(a))
m = len(vals)
depth = max(1, (m - 1).bit_length())
tree = PersistentCountTree(m, 1 + len(a) * (depth + 1))
roots = [0]
for x in a:
roots.append(tree.insert(roots[-1], bisect.bisect_left(vals, x)))
return vals, tree, rootsThe insert walks down recording the path, then creates nodes bottom-up, pairing each fresh child with the untouched sibling. Capacity is one sentinel plus at most depth + 1 nodes per insert.
Why two prefix versions give a range
Why do prefix versions answer range questions? Every version splits the value range in exactly the same way, because the tree shape depends only on m. So for any node position, cnt in version r minus cnt in version l-1 is the number of elements of a[l..r] whose values fall in that node's range. Two prefix trees subtract to a virtual tree for the subarray, which is never built: we just walk both trees in lockstep and subtract counts on the fly.
With 1-based positions l and r, the subarray a[l..r] is version r minus version l-1. It is the prefix-difference idea behind a Fenwick tree, applied to whole trees. It requires the quantity to be subtractable; counts and sums are, minimum and maximum are not.
The k-th smallest and counting queries
To find the k-th smallest, walk down from both roots. At each node, count how many subarray elements lie in the left half. If k is at most that count, the answer is in the left half; otherwise subtract the count from k and go right. At a leaf, the leaf's value is the answer.
# add this method to PersistentCountTree
def kth(self, u, v, k):
"""k-th smallest (1-based) among values in version u minus version v."""
lo, hi = 0, self.m - 1
while lo < hi:
mid = (lo + hi) // 2
in_left = self.cnt[self.left[u]] - self.cnt[self.left[v]]
if k <= in_left:
u, v, hi = self.left[u], self.left[v], mid
else:
k -= in_left
u, v, lo = self.right[u], self.right[v], mid + 1
return lo
def kth_smallest(vals, tree, roots, l, r, k):
"""k-th smallest in a[l..r], 1-based inclusive positions."""
return vals[tree.kth(roots[r], roots[l - 1], k)]
def count_le(vals, tree, roots, l, r, x):
"""How many of a[l..r] are at most x."""
pos = bisect.bisect_right(vals, x) - 1 # last compressed index with value at most x
if pos < 0:
return 0
u, v = roots[r], roots[l - 1]
lo, hi, total = 0, tree.m - 1, 0
while lo < hi:
mid = (lo + hi) // 2
if pos <= mid:
u, v, hi = tree.left[u], tree.left[v], mid
else:
total += tree.cnt[tree.left[u]] - tree.cnt[tree.left[v]]
u, v, lo = tree.right[u], tree.right[v], mid + 1
return total + tree.cnt[u] - tree.cnt[v]Both queries touch one node per level in each of two trees, so they run in O(log m). The median of a subarray is kth with k = (r - l + 2) // 2. Compare this with a k-th smallest query on a BST, which descends by subtree sizes in the same way but over one set rather than a difference of two.
Worked example
Take a = [5, 2, 6, 3, 2, 7, 1]. The distinct values sorted are [1, 2, 3, 5, 6, 7], so m = 6 and the compressed index of 1 is 0, of 2 is 1, of 3 is 2, of 5 is 3, of 6 is 4 and of 7 is 5. Building versions 1 to 7 creates 26 nodes in addition to the sentinel: the root splits [0..5] at mid 2, and paths have three or four nodes depending on which side they go.
Query: the 3rd smallest of a[2..6], positions 2 to 6, which is [2, 6, 3, 2, 7]. Sorted, that is [2, 2, 3, 6, 7], so the answer should be 3. We walk version 6 (prefix 5, 2, 6, 3, 2, 7) and version 1 (prefix 5).
- Range [0..5], mid 2. The left half holds values 1, 2 and 3. Version 6 has three of them (2, 3, 2); version 1 has none. in_left = 3. Since k = 3 is at most 3, go left.
- Range [0..2], mid 1. The left half holds values 1 and 2. Version 6 has two (2, 2); version 1 has none. in_left = 2. Since k = 3 is more than 2, set k = 1 and go right.
- Range [2..2] is a leaf: index 2, which is value 3. The answer is 3.
The 5 at position 1 is in version 6 but never counted: version 1 cancels it in the right half.
Cost and memory
Each insert creates at most ceil(log2 m) + 1 nodes. In an instrumented run of the code above on 100,000 random values (99,995 distinct), the tree held 1,768,924 nodes, or 17.69 per insert, against a bound of 18.
Memory is the real constraint. Each node is three integers; in C++ or Java with 32-bit int arrays that is 12 bytes, about 21 MB for the run above. Python lists cost several times more; use the array module or numpy for large inputs.
Range updates with permanent tags
Range updates on persistent trees need a twist. Ordinary lazy propagation pushes a pending tag down to the children during a query, which means writing to existing nodes, which is forbidden. The clean alternative is permanent tags: a range add leaves its tag on the covering nodes forever, and queries accumulate the tags of the ancestors they pass through instead of pushing them.
class PersistentRangeAdd:
"""Range add / range sum with permanent (never pushed) tags. Node 0 is unused."""
def __init__(self, a):
self.n = len(a)
self.L, self.R, self.S, self.T = [0], [0], [0], [0]
self.roots = [self._build(a, 0, self.n - 1)]
def _node(self, l, r, s, t):
self.L.append(l); self.R.append(r); self.S.append(s); self.T.append(t)
return len(self.S) - 1
def _build(self, a, lo, hi):
if lo == hi:
return self._node(0, 0, a[lo], 0)
mid = (lo + hi) // 2
l, r = self._build(a, lo, mid), self._build(a, mid + 1, hi)
return self._node(l, r, self.S[l] + self.S[r], 0)
def _add(self, node, lo, hi, ql, qr, d):
# S[node] stores the sum of its segment including every tag at or below node.
s = self.S[node] + d * (min(hi, qr) - max(lo, ql) + 1)
if ql <= lo and hi <= qr:
return self._node(self.L[node], self.R[node], s, self.T[node] + d)
mid = (lo + hi) // 2
l, r = self.L[node], self.R[node]
if ql <= mid:
l = self._add(l, lo, mid, ql, qr, d)
if qr > mid:
r = self._add(r, mid + 1, hi, ql, qr, d)
return self._node(l, r, s, self.T[node])
def add(self, version, ql, qr, d):
"""New version = version with d added to a[ql..qr] (0-based inclusive)."""
self.roots.append(self._add(self.roots[version], 0, self.n - 1, ql, qr, d))
return len(self.roots) - 1
def _sum(self, node, lo, hi, ql, qr, carried):
if ql <= lo and hi <= qr:
return self.S[node] + carried * (hi - lo + 1)
carried += self.T[node] # tags above apply to every element below
mid = (lo + hi) // 2
total = 0
if ql <= mid:
total += self._sum(self.L[node], lo, mid, ql, qr, carried)
if qr > mid:
total += self._sum(self.R[node], mid + 1, hi, ql, qr, carried)
return total
def range_sum(self, version, ql, qr):
return self._sum(self.roots[version], 0, self.n - 1, ql, qr, 0)Any version can be the base of an add, so versions form a tree of histories, not a line; the tests applied adds to randomly chosen old versions and checked sums on all of them. An add copies the O(log n) nodes a normal range update visits. On n = 65,536, 2,000 random range adds created 40.7 nodes on average and 54 at most. Permanent tags suit commuting operations such as add; assignment needs pushdown that copies children before writing.
Failure modes
- Off-by-one in the version pair. Using roots[l] instead of roots[l-1] silently drops a[l] from every answer. Keep positions 1-based and version 0 empty, and test l = 1.
- Mutating a shared node. Any in-place write, including a lazy push, corrupts older versions. Old-version queries in tests catch it; current-version tests do not.
- Pool overflow. An undersized preallocation fails mid-run or, in C++, writes past the array. Compute capacity from the bound and assert on it.
- Compressing with the wrong bisect. Positions come from bisect_left on the sorted distinct values; count_le needs bisect_right minus one. Mixing them shifts answers at duplicates.
- k out of range. k must be between 1 and r - l + 1, or the walk returns a meaningless leaf.
Trade-offs
| Approach | Query | Preprocess / memory | When to prefer |
|---|---|---|---|
| Persistent segment tree | O(log n), online | O(n log n) nodes | Online range k-th or counting; access to old versions |
| Merge sort tree | O(log^2 n) or O(log^3 n) for k-th | O(n log n) values | Simpler code, counting queries |
| Wavelet tree | O(log sigma) | O(n log sigma) bits | Very large static arrays where memory matters |
| Mo's algorithm | O(n sqrt n) total, offline | O(n) | All queries known in advance, awkward aggregates |
| Sort per query | O(len log len) | none | A handful of queries, or as a test oracle |
Persistence trades memory for online answers and history; choose it when you need undo, branching or snapshot reads.
What to do next
- Implement PersistentCountTree from memory, then rerun the worked example and confirm the 3rd smallest of a[2..6] is 3.
- Write a brute-force oracle (sort the slice) and a random tester that also queries l = 1 and l = r.
- Instrument node creation and check it against the ceil(log2 m) + 1 bound for your input sizes.
- Implement PersistentRangeAdd and test sums on old versions, not just the newest one.
- Revisit segment trees and Fenwick trees to see which ideas carry over.
- Compare against Mo's algorithm on the same query set to see when offline wins.