Introduction
A segment tree handles range queries and point updates in O(log n) time. Think of it as a mutable prefix sum.
When to use:
- Prefix sums: O(1) query, but O(n) update
- Segment tree: O(log n) query AND O(log n) update
Structure: A binary tree where:
- Each leaf represents one array element
- Each internal node represents a "segment" (range of elements)
- Each node stores aggregated info (sum, min, max, etc.)
- Parent's segment = union of children's segments
Three questions when designing a segment tree:
1. What operation? (sum, min, max, xor, gcd, etc.)
2. What do indices represent?
3. What do values represent?
The indices don't have to map to array positions! They can represent values for frequency counting.
O(n) build, O(log n) query, O(log n) updateO(n)Sum Segment Tree
Most common segment tree. Supports range sum queries and point updates.
class SegTree:
def __init__(self, nums: List[int], lo: int, hi: int):
self.lo, self.hi = lo, hi
self.left = self.right = None
self.val = nums[lo]
if lo < hi:
mid = (lo + hi) // 2
self.left = SegTree(nums, lo, mid)
self.right = SegTree(nums, mid + 1, hi)
self.val = self.left.val + self.right.val # combine
def update(self, i: int, x: int):
if i < self.lo or i > self.hi:
return
if self.lo == self.hi: # leaf
self.val = x
return
self.left.update(i, x)
self.right.update(i, x)
self.val = self.left.val + self.right.val # recombine
def query(self, l: int, r: int) -> int:
# Case 1: completely inside query range
if l <= self.lo and self.hi <= r:
return self.val
# Case 2: completely outside query range
if self.hi < l or r < self.lo:
return 0 # identity for sum
# Case 3: partial overlap - recurse
return self.left.query(l, r) + self.right.query(l, r)Min/Max Segment Tree
Only changes are: combine operation and identity value for query.
# For MIN segment tree:
# combine: self.val = min(self.left.val, self.right.val)
# query identity: return inf (for completely outside)
# query combine: return min(self.left.query(l,r), self.right.query(l,r))
# For MAX segment tree:
# combine: self.val = max(self.left.val, self.right.val)
# query identity: return -inf (for completely outside)
# query combine: return max(self.left.query(l,r), self.right.query(l,r))Coordinate Compression
When values are large but count is small, map values to their ranks.
Example: [0, 100, 10000, 1000000000] → [0, 1, 2, 3]
We don't care about actual values, only relative ordering. This lets us use values as segment tree indices without MLE.
Coordinate Compression
Map values to consecutive ranks.
def compress(nums):
sorted_unique = sorted(set(nums))
rank = {v: i for i, v in enumerate(sorted_unique)}
return rank
# Usage: rank[nums[i]] gives the compressed indexProblems
Practice problems for segment trees.
01Range Sum Query - MutableMedium
Handle array updates and range sum queries.
class NumArray:
def __init__(self, nums: List[int]):
self.n = len(nums)
self.tree = SegTree(nums, 0, self.n - 1)
def update(self, index: int, val: int) -> None:
self.tree.update(index, val)
def sumRange(self, left: int, right: int) -> int:
return self.tree.query(left, right)O(log n) per operationO(n)Direct application of sum segment tree template. This is the canonical segment tree problem.
02Count of Smaller Numbers After SelfHard
For each element, count how many smaller elements are to its right.
def countSmaller(self, nums: List[int]) -> List[int]:
# Coordinate compression
rank = {v: i for i, v in enumerate(sorted(set(nums)))}
n = len(nums)
# Segment tree for frequency counting
tree = SegTree([0] * len(rank), 0, len(rank) - 1)
res = [0] * n
# Process right to left
for i in range(n - 1, -1, -1):
r = rank[nums[i]]
# Count elements with rank < r (smaller values)
res[i] = tree.query(0, r - 1) if r > 0 else 0
# Add current element to tree (increment frequency)
tree.update(r, tree.query(r, r) + 1)
return resO(n log n)O(n)Key insight: segment tree indices = value ranks, values = frequencies. Query [0, rank-1] gives count of smaller elements. Process right-to-left so tree only contains elements to the right.
03132 PatternHard
Return true if there exist indices i < j < k such that nums[i] < nums[k] < nums[j].
from bisect import bisect_left, bisect_right
def find132pattern(self, nums: List[int]) -> bool:
n = len(nums)
if n < 3:
return False
prefix_min = [0] * n
prefix_min[0] = nums[0]
for i in range(1, n):
prefix_min[i] = min(prefix_min[i - 1], nums[i])
values = sorted(set(nums))
rank = {v: i for i, v in enumerate(values)}
tree = SegTree([0] * len(values), 0, len(values) - 1)
# Tree stores frequencies of candidate nums[k] to the right of j.
for num in nums[2:]:
r = rank[num]
tree.update(r, tree.query(r, r) + 1)
for j in range(1, n - 1):
lo = bisect_right(values, prefix_min[j - 1])
hi = bisect_left(values, nums[j]) - 1
if lo <= hi and tree.query(lo, hi) > 0:
return True
r = rank[nums[j + 1]]
tree.update(r, tree.query(r, r) - 1)
return FalseO(n log n)O(n)Fix the middle index j. The segment tree tracks values to the right of j; query whether any right-side value lies strictly between the prefix minimum on the left and nums[j].