Segment Tree Range Sum DSA Solution - Study Chapter | QuizMaker

Segment Tree explained with repeated range-sum brute force, optimized tree build/update/query, dry run, edge cases, complexity, and Python, C++, Java code.

Read
12m
Type
Chapter
Access
Free

Course

DSA Course: Interview Patterns and Problem Solving

Topic

Module 16: Trie & Advanced Data Structures

Learning Outcome

After this lesson, you should be able to split ranges into tree segments and update only affected ancestors.

Problem Statement

Design a NumArray with update(index,val) and sumRange(left,right).

InputOutputWhy
nums = [1,3,5], sumRange(0,2), update(1,2), sumRange(0,2)9, then 8Initial sum is 1+3+5 = 9. After changing index 1 to 2, sum is 1+2+5 = 8.

Brute Force Approach

Compute every requested range by scanning from left to right. This costs O(n) per query.

Optimized Approach

Build a segment tree where each node stores a range sum. Point update changes one leaf and refreshes ancestors; query combines only relevant segments.

Exact Pseudocode

build(node, left, right):
  if left == right:
    tree[node] = nums[left]
  else:
    build both halves
    tree[node] = sum of children

update(index, value):
  walk to the leaf for index
  update leaf and refresh ancestors

query(ql, qr):
  return 0 for disjoint range
  return node sum for fully covered range
  otherwise query both halves

Reference Code

class NumArray:
    def __init__(self, nums):
        self.n = len(nums)
        self.tree = [0] * (4 * self.n)
        self._build(nums, 1, 0, self.n - 1)

    def _build(self, nums, node, left, right):
        if left == right:
            self.tree[node] = nums[left]
            return
        mid = (left + right) // 2
        self._build(nums, node * 2, left, mid)
        self._build(nums, node * 2 + 1, mid + 1, right)
        self.tree[node] = self.tree[node * 2] + self.tree[node * 2 + 1]

    def update(self, index, val):
        self._update(1, 0, self.n - 1, index, val)

    def _update(self, node, left, right, index, val):
        if left == right:
            self.tree[node] = val
            return
        mid = (left + right) // 2
        if index <= mid:
            self._update(node * 2, left, mid, index, val)
        else:
            self._update(node * 2 + 1, mid + 1, right, index, val)
        self.tree[node] = self.tree[node * 2] + self.tree[node * 2 + 1]

    def sumRange(self, left, right):
        return self._query(1, 0, self.n - 1, left, right)

    def _query(self, node, left, right, ql, qr):
        if qr < left or right < ql:
            return 0
        if ql <= left and right <= qr:
            return self.tree[node]
        mid = (left + right) // 2
        return self._query(node * 2, left, mid, ql, qr) + self._query(node * 2 + 1, mid + 1, right, ql, qr)
class NumArray {
    int n;
    vector<int> tree;

    void build(vector<int>& nums, int node, int left, int right) {
        if (left == right) {
            tree[node] = nums[left];
            return;
        }
        int mid = (left + right) / 2;
        build(nums, node * 2, left, mid);
        build(nums, node * 2 + 1, mid + 1, right);
        tree[node] = tree[node * 2] + tree[node * 2 + 1];
    }

    void update(int node, int left, int right, int index, int val) {
        if (left == right) {
            tree[node] = val;
            return;
        }
        int mid = (left + right) / 2;
        if (index <= mid) update(node * 2, left, mid, index, val);
        else update(node * 2 + 1, mid + 1, right, index, val);
        tree[node] = tree[node * 2] + tree[node * 2 + 1];
    }

    int query(int node, int left, int right, int ql, int qr) {
        if (qr < left || right < ql) return 0;
        if (ql <= left && right <= qr) return tree[node];
        int mid = (left + right) / 2;
        return query(node * 2, left, mid, ql, qr)
             + query(node * 2 + 1, mid + 1, right, ql, qr);
    }

public:
    NumArray(vector<int>& nums) {
        n = nums.size();
        tree.assign(4 * n, 0);
        build(nums, 1, 0, n - 1);
    }

    void update(int index, int val) {
        update(1, 0, n - 1, index, val);
    }

    int sumRange(int left, int right) {
        return query(1, 0, n - 1, left, right);
    }
};
class NumArray {
    private int n;
    private int[] tree;

    public NumArray(int[] nums) {
        n = nums.length;
        tree = new int[4 * n];
        build(nums, 1, 0, n - 1);
    }

    private void build(int[] nums, int node, int left, int right) {
        if (left == right) {
            tree[node] = nums[left];
            return;
        }
        int mid = (left + right) / 2;
        build(nums, node * 2, left, mid);
        build(nums, node * 2 + 1, mid + 1, right);
        tree[node] = tree[node * 2] + tree[node * 2 + 1];
    }

    public void update(int index, int val) {
        update(1, 0, n - 1, index, val);
    }

    private void update(int node, int left, int right, int index, int val) {
        if (left == right) {
            tree[node] = val;
            return;
        }
        int mid = (left + right) / 2;
        if (index <= mid) update(node * 2, left, mid, index, val);
        else update(node * 2 + 1, mid + 1, right, index, val);
        tree[node] = tree[node * 2] + tree[node * 2 + 1];
    }

    public int sumRange(int left, int right) {
        return query(1, 0, n - 1, left, right);
    }

    private int query(int node, int left, int right, int ql, int qr) {
        if (qr < left || right < ql) return 0;
        if (ql <= left && right <= qr) return tree[node];
        int mid = (left + right) / 2;
        return query(node * 2, left, mid, ql, qr)
             + query(node * 2 + 1, mid + 1, right, ql, qr);
    }
}

Sample Dry Run

StepStateResult
Buildroot stores sum 1+3+5tree root = 9
sumRange 0 to 2query fully covers rootreturn 9
update index 1 to 2change leaf 3 to 2 and refresh ancestorsroot becomes 8
sumRange 0 to 2query root againreturn 8

Complexity

MeasureValueReason
TimeO(log n) per update or queryEach operation follows or combines a logarithmic number of tree segments.
SpaceO(n)The segment tree stores range sums for the array.

Edge Cases

Interview Checklist

FAQs

Why does query need three overlap cases?

It must ignore disjoint ranges, return stored sums for fully covered ranges, and split partial ranges.

Why does update refresh ancestors?

Every parent sum depends on the changed leaf, so ancestor sums must be recomputed.

What is the core pattern?

Range decomposition tree.

Tags

Open on QuizMaker