Skip to content

Advanced Data Structures

Trees that remember ranges: Segment trees are like a librarian who knows the sum of every shelf range instantly — instead of counting books one by one, they pre-compute summaries at every level, so any range query is answered by combining a few pre-computed values.

Why it matters: Segment trees enable O(log n) range queries and updates, which is crucial for competitive programming, database indexing, and computational geometry where you need to query and modify intervals efficiently.

The key insight: The power of segment trees comes from pre-computation — by storing summaries at every level of the tree, you trade O(n) space for O(log n) query time, which is a massive improvement for repeated range queries.

A segment tree is a binary tree data structure for storing information about intervals or segments. It allows efficient range queries and point updates.

class SegmentTree:
"""
Segment tree for range sum queries with point updates.
Build: O(n)
Point update: O(log n)
Range query: O(log n)
Space: O(4n)
"""
def __init__(self, data):
self.n = len(data)
self.tree = [0] * (4 * self.n)
self._build(data, 0, 0, self.n - 1)
def _build(self, data, node, start, end):
if start == end:
self.tree[node] = data[start]
else:
mid = (start + end) // 2
self._build(data, 2 * node + 1, start, mid)
self._build(data, 2 * node + 2, mid + 1, end)
self.tree[node] = self.tree[2 * node + 1] + self.tree[2 * node + 2]
def update(self, idx, value):
self._update(0, 0, self.n - 1, idx, value)
def _update(self, node, start, end, idx, value):
if start == end:
self.tree[node] = value
else:
mid = (start + end) // 2
if idx <= mid:
self._update(2 * node + 1, start, mid, idx, value)
else:
self._update(2 * node + 2, mid + 1, end, idx, value)
self.tree[node] = self.tree[2 * node + 1] + self.tree[2 * node + 2]
def query(self, left, right):
return self._query(0, 0, self.n - 1, left, right)
def _query(self, node, start, end, left, right):
if right < start or end < left:
return 0
if left <= start and end <= right:
return self.tree[node]
mid = (start + end) // 2
left_sum = self._query(2 * node + 1, start, mid, left, right)
right_sum = self._query(2 * node + 2, mid + 1, end, left, right)
return left_sum + right_sum
class SegmentTreeMin:
"""
Segment tree for range minimum queries with point updates.
Build: O(n)
Point update: O(log n)
Range query: O(log n)
Space: O(4n)
"""
def __init__(self, data, neutral=float("inf')):
self.n = len(data)
self.neutral = neutral
self.tree = [neutral] * (4 * self.n)
self._build(data, 0, 0, self.n - 1)
def _build(self, data, node, start, end):
if start == end:
self.tree[node] = data[start]
else:
mid = (start + end) // 2
self._build(data, 2 * node + 1, start, mid)
self._build(data, 2 * node + 2, mid + 1, end)
self.tree[node] = min(self.tree[2 * node + 1], self.tree[2 * node + 2])
def query(self, left, right):
return self._query(0, 0, self.n - 1, left, right)
def _query(self, node, start, end, left, right):
if right < start or end < left:
return self.neutral
if left <= start and end <= right:
return self.tree[node]
mid = (start + end) // 2
return min(
self._query(2 * node + 1, start, mid, left, right),
self._query(2 * node + 2, mid + 1, end, left, right)
)
def update(self, idx, value):
self._update(0, 0, self.n - 1, idx, value)
def _update(self, node, start, end, idx, value):
if start == end:
self.tree[node] = value
else:
mid = (start + end) // 2
if idx <= mid:
self._update(2 * node + 1, start, mid, idx, value)
else:
self._update(2 * node + 2, mid + 1, end, idx, value)
self.tree[node] = min(self.tree[2 * node + 1], self.tree[2 * node + 2])

Lazy propagation allows efficient range updates by deferring updates to child nodes until they are Needed.

class LazySegmentTree:
"""
Segment tree with lazy propagation for range sum queries
with range add updates.
Build: O(n)
Point update: O(log n)
Range update: O(log n)
Range query: O(log n)
Space: O(4n)
"""
def __init__(self, data):
self.n = len(data)
self.tree = [0] * (4 * self.n)
self.lazy = [0] * (4 * self.n)
self._build(data, 0, 0, self.n - 1)
def _build(self, data, node, start, end):
if start == end:
self.tree[node] = data[start]
else:
mid = (start + end) // 2
self._build(data, 2 * node + 1, start, mid)
self._build(data, 2 * node + 2, mid + 1, end)
self.tree[node] = self.tree[2 * node + 1] + self.tree[2 * node + 2]
def _push_down(self, node, start, end):
if self.lazy[node] != 0:
mid = (start + end) // 2
left_len = mid - start + 1
right_len = end - mid
self.tree[2 * node + 1] += self.lazy[node] * left_len
self.lazy[2 * node + 1] += self.lazy[node]
self.tree[2 * node + 2] += self.lazy[node] * right_len
self.lazy[2 * node + 2] += self.lazy[node]
self.lazy[node] = 0
def range_update(self, left, right, value):
self._range_update(0, 0, self.n - 1, left, right, value)
def _range_update(self, node, start, end, left, right, value):
if right < start or end < left:
return
if left <= start and end <= right:
self.tree[node] += value * (end - start + 1)
self.lazy[node] += value
return
self._push_down(node, start, end)
mid = (start + end) // 2
self._range_update(2 * node + 1, start, mid, left, right, value)
self._range_update(2 * node + 2, mid + 1, end, left, right, value)
self.tree[node] = self.tree[2 * node + 1] + self.tree[2 * node + 2]
def range_query(self, left, right):
return self._range_query(0, 0, self.n - 1, left, right)
def _range_query(self, node, start, end, left, right):
if right < start or end < left:
return 0
if left <= start and end <= right:
return self.tree[node]
self._push_down(node, start, end)
mid = (start + end) // 2
return (self._range_query(2 * node + 1, start, mid, left, right) +
self._range_query(2 * node + 2, mid + 1, end, left, right))
class IterativeSegmentTree:
"""
Iterative segment tree — faster in practice due to no recursion overhead.
Build: O(n)
Point update: O(log n)
Range query: O(log n)
Space: O(2n)
"""
def __init__(self, data):
self.n = len(data)
self.size = 1
while self.size < self.n:
self.size <<= 1
self.tree = [0] * (2 * self.size)
for i in range(self.n):
self.tree[self.size + i] = data[i]
for i in range(self.size - 1, 0, -1):
self.tree[i] = self.tree[2 * i] + self.tree[2 * i + 1]
def update(self, idx, value):
idx += self.size
self.tree[idx] = value
idx >>= 1
while idx >= 1:
self.tree[idx] = self.tree[2 * idx] + self.tree[2 * idx + 1]
idx >>= 1
def query(self, left, right):
left += self.size
right += self.size
result = 0
while left <= right:
if left % 2 == 1:
result += self.tree[left]
left += 1
if right % 2 == 0:
result += self.tree[right]
right -= 1
left >>= 1
right >>= 1
return result
graph TD
    ROOT["[0-7]: sum=28"] --> L["[0-3]: sum=10"]
    ROOT --> R["[4-7]: sum=18"]
    L --> LL["[0-1]: sum=3"]
    L --> LR["[2-3]: sum=7"]
    R --> RL["[4-5]: sum=9"]
    R --> RR["[6-7]: sum=9"]
    LL --> LLL["[0]: 1"]
    LL --> LLR["[1]: 2"]
    LR --> LRL["[2]: 3"]
    LR --> LRR["[3]: 4"]
    RL --> RLL["[4]: 5"]
    RL --> RLR["[5]: 4"]
    RR --> RRL["[6]: 4"]
    RR --> RRR["[7]: 5"]

A Fenwick tree (BIT) supports point updates and prefix sum queries in O(logn)O(\log n) time with less Memory and simpler code than a segment tree.

Each node at index ii stores the sum of a range of length i & (-i) (the lowest set bit of ii). This allows prefix sum queries by “climbing up” the tree and point updates by “adding” to all Relevant ranges.

class FenwickTree:
"""
Fenwick tree (Binary Indexed Tree) for prefix sums.
Build: O(n log n) or O(n) with bulk construction
Point update: O(log n)
Prefix sum: O(log n)
Range sum: O(log n) (two prefix sums)
Space: O(n)
"""
def __init__(self, n):
self.n = n
self.tree = [0] * (n + 1)
def update(self, idx, delta):
"""Add delta to element at index idx (1-based). O(log n)."""
idx += 1
while idx <= self.n:
self.tree[idx] += delta
idx += idx & (-idx)
def prefix_sum(self, idx):
"""Sum of elements [0, idx] (0-based). O(log n)."""
idx += 1
result = 0
while idx > 0:
result += self.tree[idx]
idx -= idx & (-idx)
return result
def range_sum(self, left, right):
"""Sum of elements [left, right] (0-based). O(log n)."""
return self.prefix_sum(right) - (self.prefix_sum(left - 1) if left > 0 else 0)
def build(self, data):
"""Build from array in O(n)."""
for i, val in enumerate(data):
self.tree[i + 1] = val
for i in range(1, self.n + 1):
parent = i + (i & (-i))
if parent <= self.n:
self.tree[parent] += self.tree[i]
class FenwickTreeRangeUpdate:
"""
Fenwick tree for range updates and point queries.
Range update: O(log n)
Point query: O(log n)
Space: O(n)
"""
def __init__(self, n):
self.n = n
self.tree = [0] * (n + 1)
def range_add(self, left, right, delta):
"""Add delta to all elements in [left, right] (0-based)."""
self._update(left, delta)
self._update(right + 1, -delta)
def _update(self, idx, delta):
idx += 1
while idx <= self.n:
self.tree[idx] += delta
idx += idx & (-idx)
def point_query(self, idx):
"""Get value at index idx (0-based)."""
idx += 1
result = 0
while idx > 0:
result += self.tree[idx]
idx -= idx & (-idx)
return result
class FenwickTree2D:
"""
2D Fenwick tree for prefix sums on a grid.
Point update: O(log^2 n)
Prefix sum: O(log^2 n)
Space: O(n^2)
"""
def __init__(self, rows, cols):
self.rows = rows
self.cols = cols
self.tree = [[0] * (cols + 1) for _ in range(rows + 1)]
def update(self, row, col, delta):
r = row + 1
while r <= self.rows:
c = col + 1
while c <= self.cols:
self.tree[r][c] += delta
c += c & (-c)
r += r & (-r)
def prefix_sum(self, row, col):
result = 0
r = row + 1
while r > 0:
c = col + 1
while c > 0:
result += self.tree[r][c]
c -= c & (-c)
r -= r & (-r)
return result
def range_sum(self, r1, c1, r2, c2):
return (self.prefix_sum(r2, c2) - self.prefix_sum(r1 - 1, c2) -
self.prefix_sum(r2, c1 - 1) + self.prefix_sum(r1 - 1, c1 - 1))
  • Binary Search Trees — BSTs provide the foundation for understanding tree-based data structures like segment trees.
  • Trie and Pattern Matching — Tries are specialised string data structures that complement the general-purpose structures here.
  • Dynamic Programming — Segment trees are often combined with DP for efficient range query solutions.
  • Deques and Priority Queues — Priority queues underpin Dijkstra’s algorithm when combined with segment trees.