← All patterns
Advanced · ST

Segment Tree

Answer changing range queries with a recursive index.

3lessons
3worked problems
Freefull access
01 / 03

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

KEY INSIGHT

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.

TimeO(n) build, O(log n) query, O(log n) update
SpaceO(n)

Sum Segment Tree

Most common segment tree. Supports range sum queries and point updates.

pythonREFERENCE
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.

pythonREFERENCE
# 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))
02 / 03

Coordinate Compression

When values are large but count is small, map values to their ranks.

KEY INSIGHT

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.

pythonREFERENCE
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 index
03 / 03

Problems

Practice problems for segment trees.

WORKED PROBLEMS3
01Range Sum Query - MutableMedium

Handle array updates and range sum queries.

pythonREFERENCE
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)
TimeO(log n) per operation
SpaceO(n)
WHY IT WORKS

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.

pythonREFERENCE
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 res
TimeO(n log n)
SpaceO(n)
WHY IT WORKS

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].

pythonREFERENCE
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 False
TimeO(n log n)
SpaceO(n)
WHY IT WORKS

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].

NEXT PATTERNParenthesis