본문 바로가기
카테고리 없음

Segment Tree(세그먼트 트리)

by SuldenLion 2026. 2. 21.
반응형

Segment Tree: 구간 쿼리의 마스터

 

들어가며

Segment Tree(세그먼트 트리)는 구간 쿼리(Range Query)를 효율적으로 처리하기 위한 트리 자료구조입니다. 배열의 특정 구간에 대한 합, 최소값, 최대값 등을 O(log n)에 구하고, 원소 업데이트도 O(log n)에 수행합니다. 1977년 Jon Louis Bentley가 제안한 이후, 알고리즘 대회, 데이터베이스 인덱싱, 계산 기하학 등에서 필수적으로 사용됩니다. Segment Tree의 원리부터 구현, 변형, 실전 문제까지 모든 것을 깊이 있게 탐구해봅시다.

 

1. Segment Tree의 본질

1.1 세그먼트 트리란?

문제: 배열의 구간 합 구하기 + 원소 업데이트

배열: [1, 3, 5, 7, 9, 11]
쿼리 1: sum(0, 2) = 1 + 3 + 5 = 9
쿼리 2: sum(1, 4) = 3 + 5 + 7 + 9 = 24
업데이트: arr[2] = 6
쿼리 3: sum(0, 2) = 1 + 3 + 6 = 10

Naive 방법:
- 구간 합: O(n) - 루프로 합산
- 업데이트: O(1) - 직접 수정
→ Q개 쿼리 시 O(Q×n)

누적 합(Prefix Sum):
- 전처리: O(n)
- 구간 합: O(1) - prefix[r] - prefix[l-1]
- 업데이트: O(n) - 모든 prefix 재계산
→ 업데이트 많으면 느림

Segment Tree:
- 전처리: O(n)
- 구간 합: O(log n)
- 업데이트: O(log n)
→ 균형 잡힌 성능!

시간 복잡도 비교 (n=100만, Q=100만):
Naive:      100만 × 100만 = 1조 (불가능)
Prefix Sum: 100만 × 100만 = 1조 (업데이트 많으면)
Segment Tree: 100만 × log(100만) ≈ 2천만 (가능!)

 

1.2 세그먼트 트리 구조

배열: [1, 3, 5, 7, 9, 11]
구간 합 세그먼트 트리:

                  [0-5: 36]
                 /          \
          [0-2: 9]          [3-5: 27]
          /      \          /        \
     [0-1: 4]  [2: 5]  [3-4: 16]  [5: 11]
      /    \            /      \
  [0: 1] [1: 3]    [3: 7]  [4: 9]

노드 표기: [구간: 값]

특징:
1. 리프 노드: 배열의 개별 원소
2. 내부 노드: 자식 구간의 합
3. 루트: 전체 배열의 합
4. 높이: O(log n)
5. 노드 수: ~4n

이진 트리:
- 왼쪽 자식: 구간의 왼쪽 절반
- 오른쪽 자식: 구간의 오른쪽 절반

예: [0-5] → [0-2], [3-5]
    [0-2] → [0-1], [2-2]
    [0-1] → [0-0], [1-1]

 

1.3 배열 표현

python
# 1-based 인덱싱 (구현 간단)

배열: [1, 3, 5, 7, 9, 11]
트리: [0, 36, 9, 27, 4, 5, 16, 11, 1, 3, 0, 0, 7, 9, 0, 0]
      ↑   ↑   ↑   ↑   ↑  ↑   ↑   ↑  ↑  ↑
      0   1   2   3   4  5   6   7  8  9 ...

인덱스 1: 루트 [0-5: 36]
인덱스 2: [0-2: 9]
인덱스 3: [3-5: 27]
인덱스 4: [0-1: 4]
...

관계:
- 부모: i // 2
- 왼쪽 자식: i * 2
- 오른쪽 자식: i * 2 + 1

트리 크기:
- n개 원소 → 최대 4n 크기 배열 필요
- 2의 거듭제곱이 아니어도 OK
- 여유 공간 포함

예: n=6 → 배열 크기 24

 

2. 기본 구현

2.1 세그먼트 트리 클래스

python
class SegmentTree:
    def __init__(self, arr):
        self.n = len(arr)
        self.tree = [0] * (4 * self.n)
        self.arr = arr
        self._build(0, self.n - 1, 1)
    
    def _build(self, start, end, node):
        """트리 구축 - O(n)"""
        if start == end:
            # 리프 노드
            self.tree[node] = self.arr[start]
            return
        
        mid = (start + end) // 2
        # 왼쪽 서브트리
        self._build(start, mid, node * 2)
        # 오른쪽 서브트리
        self._build(mid + 1, end, node * 2 + 1)
        
        # 현재 노드 = 자식들의 합
        self.tree[node] = (self.tree[node * 2] + 
                          self.tree[node * 2 + 1])
    
    def query(self, left, right):
        """구간 합 쿼리 - O(log n)"""
        return self._query(0, self.n - 1, 1, left, right)
    
    def _query(self, start, end, node, 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(start, mid, node * 2, left, right)
        right_sum = self._query(mid + 1, end, node * 2 + 1, left, right)
        
        return left_sum + right_sum
    
    def update(self, index, value):
        """원소 업데이트 - O(log n)"""
        self._update(0, self.n - 1, 1, index, value)
    
    def _update(self, start, end, node, index, value):
        if start == end:
            # 리프 노드 도달
            self.tree[node] = value
            self.arr[index] = value
            return
        
        mid = (start + end) // 2
        
        if index <= mid:
            # 왼쪽 서브트리
            self._update(start, mid, node * 2, index, value)
        else:
            # 오른쪽 서브트리
            self._update(mid + 1, end, node * 2 + 1, index, value)
        
        # 현재 노드 갱신
        self.tree[node] = (self.tree[node * 2] + 
                          self.tree[node * 2 + 1])

# 사용 예시
arr = [1, 3, 5, 7, 9, 11]
seg_tree = SegmentTree(arr)

print(seg_tree.query(0, 2))  # sum(0, 2) = 1+3+5 = 9
print(seg_tree.query(1, 4))  # sum(1, 4) = 3+5+7+9 = 24

seg_tree.update(2, 6)         # arr[2] = 6
print(seg_tree.query(0, 2))  # sum(0, 2) = 1+3+6 = 10

 

2.2 구간 쿼리 과정 상세

배열: [1, 3, 5, 7, 9, 11]
쿼리: sum(1, 4) = ?

트리 구조:
                [0-5: 36]
               /          \
        [0-2: 9]          [3-5: 27]
        /      \          /        \
   [0-1: 4]  [2: 5]  [3-4: 16]  [5: 11]
    /    \            /      \
[0: 1] [1: 3]    [3: 7]  [4: 9]

과정:
1. 루트 [0-5]: 부분 겹침 → 분할
2. 왼쪽 [0-2]: 부분 겹침 → 분할
   - [0-1]: 부분 겹침 → 분할
     - [0]: 벗어남 → 0
     - [1]: 포함됨 → 3 ✓
   - [2]: 포함됨 → 5 ✓
3. 오른쪽 [3-5]: 부분 겹침 → 분할
   - [3-4]: 포함됨 → 16 ✓
   - [5]: 벗어남 → 0

결과: 3 + 5 + 16 = 24

방문 노드: O(log n)
최대 4개 레벨, 각 레벨에서 최대 4개 노드

 

2.3 업데이트 과정 상세

배열: [1, 3, 5, 7, 9, 11]
업데이트: arr[2] = 6

Before:
                [0-5: 36]
               /          \
        [0-2: 9]          [3-5: 27]
        /      \
   [0-1: 4]  [2: 5]

After:
                [0-5: 37]  ← 갱신 (36+1=37)
               /          \
        [0-2: 10] ← 갱신  [3-5: 27]
        /      \
   [0-1: 4]  [2: 6]  ← 갱신 (5→6)

과정:
1. 루트 [0-5]: 2 ∈ [0, 5] → 왼쪽 서브트리
2. [0-2]: 2 ∈ [0, 2] → 오른쪽 서브트리
3. [2]: 도달 → 값 변경 (5 → 6)
4. 역순으로 부모 갱신:
   - [0-2]: 4 + 6 = 10
   - [0-5]: 10 + 27 = 37

방문 노드: O(log n) - 루트에서 리프까지의 경로

 

3. 다양한 쿼리 유형

3.1 구간 최소값 (Range Minimum Query)

python
class SegmentTreeMin:
    def __init__(self, arr):
        self.n = len(arr)
        self.tree = [float('inf')] * (4 * self.n)
        self.arr = arr
        self._build(0, self.n - 1, 1)
    
    def _build(self, start, end, node):
        if start == end:
            self.tree[node] = self.arr[start]
            return
        
        mid = (start + end) // 2
        self._build(start, mid, node * 2)
        self._build(mid + 1, end, node * 2 + 1)
        
        # 최소값
        self.tree[node] = min(self.tree[node * 2],
                             self.tree[node * 2 + 1])
    
    def query(self, left, right):
        return self._query(0, self.n - 1, 1, left, right)
    
    def _query(self, start, end, node, left, right):
        if right < start or end < left:
            return float('inf')
        
        if left <= start and end <= right:
            return self.tree[node]
        
        mid = (start + end) // 2
        left_min = self._query(start, mid, node * 2, left, right)
        right_min = self._query(mid + 1, end, node * 2 + 1, left, right)
        
        return min(left_min, right_min)
    
    def update(self, index, value):
        self._update(0, self.n - 1, 1, index, value)
    
    def _update(self, start, end, node, index, value):
        if start == end:
            self.tree[node] = value
            self.arr[index] = value
            return
        
        mid = (start + end) // 2
        
        if index <= mid:
            self._update(start, mid, node * 2, index, value)
        else:
            self._update(mid + 1, end, node * 2 + 1, index, value)
        
        self.tree[node] = min(self.tree[node * 2],
                             self.tree[node * 2 + 1])

# 사용
arr = [5, 2, 9, 1, 7, 3]
seg_min = SegmentTreeMin(arr)

print(seg_min.query(0, 3))  # min(5,2,9,1) = 1
print(seg_min.query(2, 5))  # min(9,1,7,3) = 1

seg_min.update(3, 8)         # arr[3] = 8
print(seg_min.query(0, 3))  # min(5,2,9,8) = 2

 

3.2 구간 최대값 (Range Maximum Query)

python
class SegmentTreeMax:
    def __init__(self, arr):
        self.n = len(arr)
        self.tree = [float('-inf')] * (4 * self.n)
        self.arr = arr
        self._build(0, self.n - 1, 1)
    
    def _build(self, start, end, node):
        if start == end:
            self.tree[node] = self.arr[start]
            return
        
        mid = (start + end) // 2
        self._build(start, mid, node * 2)
        self._build(mid + 1, end, node * 2 + 1)
        
        # 최대값
        self.tree[node] = max(self.tree[node * 2],
                             self.tree[node * 2 + 1])
    
    def query(self, left, right):
        return self._query(0, self.n - 1, 1, left, right)
    
    def _query(self, start, end, node, left, right):
        if right < start or end < left:
            return float('-inf')
        
        if left <= start and end <= right:
            return self.tree[node]
        
        mid = (start + end) // 2
        left_max = self._query(start, mid, node * 2, left, right)
        right_max = self._query(mid + 1, end, node * 2 + 1, left, right)
        
        return max(left_max, right_max)

 

3.3 구간 GCD

python
import math

class SegmentTreeGCD:
    def __init__(self, arr):
        self.n = len(arr)
        self.tree = [0] * (4 * self.n)
        self.arr = arr
        self._build(0, self.n - 1, 1)
    
    def _build(self, start, end, node):
        if start == end:
            self.tree[node] = self.arr[start]
            return
        
        mid = (start + end) // 2
        self._build(start, mid, node * 2)
        self._build(mid + 1, end, node * 2 + 1)
        
        # GCD
        self.tree[node] = math.gcd(self.tree[node * 2],
                                    self.tree[node * 2 + 1])
    
    def query(self, left, right):
        return self._query(0, self.n - 1, 1, left, right)
    
    def _query(self, start, end, node, 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_gcd = self._query(start, mid, node * 2, left, right)
        right_gcd = self._query(mid + 1, end, node * 2 + 1, left, right)
        
        return math.gcd(left_gcd, right_gcd)

# 사용
arr = [12, 18, 24, 30, 36]
seg_gcd = SegmentTreeGCD(arr)

print(seg_gcd.query(0, 2))  # gcd(12,18,24) = 6
print(seg_gcd.query(1, 4))  # gcd(18,24,30,36) = 6

 

4. Lazy Propagation

4.1 구간 업데이트 문제

문제: 구간 [l, r]의 모든 원소에 값 추가

Naive:
for i in range(l, r+1):
    arr[i] += value
→ O(n) per update

구간 업데이트가 많으면 느림!

해결책: Lazy Propagation
- 업데이트를 "지연"
- 실제로 필요할 때만 적용
- O(log n) per update

 

4.2 Lazy Propagation 구현

python
class SegmentTreeLazy:
    def __init__(self, arr):
        self.n = len(arr)
        self.tree = [0] * (4 * self.n)
        self.lazy = [0] * (4 * self.n)  # Lazy 배열
        self.arr = arr
        self._build(0, self.n - 1, 1)
    
    def _build(self, start, end, node):
        if start == end:
            self.tree[node] = self.arr[start]
            return
        
        mid = (start + end) // 2
        self._build(start, mid, node * 2)
        self._build(mid + 1, end, node * 2 + 1)
        
        self.tree[node] = (self.tree[node * 2] + 
                          self.tree[node * 2 + 1])
    
    def _push(self, start, end, node):
        """Lazy 값을 자식에게 전파"""
        if self.lazy[node] == 0:
            return
        
        # 현재 노드에 lazy 값 적용
        self.tree[node] += (end - start + 1) * self.lazy[node]
        
        if start != end:
            # 자식에게 전파
            self.lazy[node * 2] += self.lazy[node]
            self.lazy[node * 2 + 1] += self.lazy[node]
        
        self.lazy[node] = 0
    
    def update_range(self, left, right, value):
        """구간 업데이트 - O(log n)"""
        self._update_range(0, self.n - 1, 1, left, right, value)
    
    def _update_range(self, start, end, node, left, right, value):
        # Lazy 값 적용
        self._push(start, end, node)
        
        # 구간이 완전히 벗어남
        if right < start or end < left:
            return
        
        # 구간이 완전히 포함됨
        if left <= start and end <= right:
            # Lazy 값 설정
            self.lazy[node] += value
            self._push(start, end, node)
            return
        
        # 부분적으로 겹침
        mid = (start + end) // 2
        self._update_range(start, mid, node * 2, left, right, value)
        self._update_range(mid + 1, end, node * 2 + 1, left, right, value)
        
        # 현재 노드 갱신
        self._push(start, mid, node * 2)
        self._push(mid + 1, end, node * 2 + 1)
        self.tree[node] = (self.tree[node * 2] + 
                          self.tree[node * 2 + 1])
    
    def query(self, left, right):
        """구간 합 쿼리 - O(log n)"""
        return self._query(0, self.n - 1, 1, left, right)
    
    def _query(self, start, end, node, left, right):
        # Lazy 값 적용
        self._push(start, end, node)
        
        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(start, mid, node * 2, left, right)
        right_sum = self._query(mid + 1, end, node * 2 + 1, left, right)
        
        return left_sum + right_sum

# 사용 예시
arr = [1, 3, 5, 7, 9, 11]
seg_lazy = SegmentTreeLazy(arr)

print(seg_lazy.query(0, 2))  # sum(0,2) = 1+3+5 = 9

seg_lazy.update_range(0, 2, 10)  # arr[0..2] += 10
# arr = [11, 13, 15, 7, 9, 11]

print(seg_lazy.query(0, 2))  # sum(0,2) = 11+13+15 = 39

 

4.3 Lazy Propagation 과정

배열: [1, 3, 5, 7, 9, 11]
업데이트: range(1, 4) += 10

트리:
                [0-5: 36]
               /          \
        [0-2: 9]          [3-5: 27]
        /      \          /        \
   [0-1: 4]  [2: 5]  [3-4: 16]  [5: 11]

과정:
1. 루트 [0-5]: 부분 겹침 → 분할
2. 왼쪽 [0-2]: 부분 겹침 → 분할
   - [0-1]: 부분 겹침 → 분할
     - [1]: 포함 → lazy[1] = 10 ✓
   - [2]: 포함 → lazy[2] = 10 ✓
3. 오른쪽 [3-5]: 부분 겹침 → 분할
   - [3-4]: 포함 → lazy[3-4] = 10 ✓
   - [5]: 벗어남

Lazy 배열:
lazy[1] = 10  (노드 [1])
lazy[2] = 10  (노드 [2])
lazy[3-4] = 10  (노드 [3-4])

다음 쿼리 시 lazy 값 적용!

 

5. 고급 응용

5.1 2D 세그먼트 트리

python
class SegmentTree2D:
    """2차원 배열의 구간 합"""
    
    def __init__(self, matrix):
        self.rows = len(matrix)
        self.cols = len(matrix[0])
        self.matrix = matrix
        
        # 행 방향 세그먼트 트리 배열
        self.tree = [[0] * (4 * self.cols) 
                     for _ in range(4 * self.rows)]
        
        self._build_rows(0, self.rows - 1, 1)
    
    def _build_rows(self, start, end, node):
        if start == end:
            self._build_cols(start, 0, self.cols - 1, node, 1)
            return
        
        mid = (start + end) // 2
        self._build_rows(start, mid, node * 2)
        self._build_rows(mid + 1, end, node * 2 + 1)
        
        # 행 병합
        for i in range(4 * self.cols):
            self.tree[node][i] = (self.tree[node * 2][i] + 
                                 self.tree[node * 2 + 1][i])
    
    def _build_cols(self, row, start, end, row_node, col_node):
        if start == end:
            self.tree[row_node][col_node] = self.matrix[row][start]
            return
        
        mid = (start + end) // 2
        self._build_cols(row, start, mid, row_node, col_node * 2)
        self._build_cols(row, mid + 1, end, row_node, col_node * 2 + 1)
        
        self.tree[row_node][col_node] = (
            self.tree[row_node][col_node * 2] + 
            self.tree[row_node][col_node * 2 + 1]
        )
    
    def query(self, r1, c1, r2, c2):
        """2D 구간 합 쿼리"""
        return self._query_rows(0, self.rows - 1, 1, r1, r2, c1, c2)
    
    # 구현 생략 (복잡도: O(log²n))

# 사용
matrix = [
    [1, 2, 3],
    [4, 5, 6],
    [7, 8, 9]
]
seg2d = SegmentTree2D(matrix)
# query(0, 0, 1, 1) = 1+2+4+5 = 12

 

5.2 Persistent Segment Tree

python
class PersistentSegmentTree:
    """
    영속 세그먼트 트리
    - 각 버전 유지
    - 업데이트 시 새 버전 생성
    - 과거 버전 쿼리 가능
    """
    
    class Node:
        def __init__(self, value=0):
            self.value = value
            self.left = None
            self.right = None
    
    def __init__(self, arr):
        self.n = len(arr)
        self.arr = arr
        self.versions = []
        
        # 초기 버전 빌드
        root = self._build(0, self.n - 1)
        self.versions.append(root)
    
    def _build(self, start, end):
        node = self.Node()
        
        if start == end:
            node.value = self.arr[start]
            return node
        
        mid = (start + end) // 2
        node.left = self._build(start, mid)
        node.right = self._build(mid + 1, end)
        node.value = node.left.value + node.right.value
        
        return node
    
    def update(self, version, index, value):
        """새 버전 생성"""
        old_root = self.versions[version]
        new_root = self._update(old_root, 0, self.n - 1, index, value)
        self.versions.append(new_root)
        return len(self.versions) - 1
    
    def _update(self, node, start, end, index, value):
        # 새 노드 생성 (복사)
        new_node = self.Node(node.value)
        
        if start == end:
            new_node.value = value
            return new_node
        
        mid = (start + end) // 2
        
        if index <= mid:
            new_node.left = self._update(node.left, start, mid, index, value)
            new_node.right = node.right  # 공유
        else:
            new_node.left = node.left  # 공유
            new_node.right = self._update(node.right, mid + 1, end, index, value)
        
        new_node.value = new_node.left.value + new_node.right.value
        return new_node
    
    def query(self, version, left, right):
        """특정 버전의 구간 쿼리"""
        root = self.versions[version]
        return self._query(root, 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 node.value
        
        mid = (start + end) // 2
        left_sum = self._query(node.left, start, mid, left, right)
        right_sum = self._query(node.right, mid + 1, end, left, right)
        
        return left_sum + right_sum

# 사용
arr = [1, 3, 5, 7, 9]
pst = PersistentSegmentTree(arr)

# 버전 0: [1, 3, 5, 7, 9]
print(pst.query(0, 0, 2))  # 1+3+5 = 9

# 버전 1: [1, 3, 6, 7, 9]  (arr[2] = 6)
v1 = pst.update(0, 2, 6)
print(pst.query(v1, 0, 2))  # 1+3+6 = 10

# 버전 0은 여전히 유지!
print(pst.query(0, 0, 2))  # 1+3+5 = 9 (변경 안 됨)

# 공간 복잡도: O(n + q log n)
# 각 업데이트마다 O(log n)개 노드만 새로 생성

 

5.3 동적 세그먼트 트리

python
class DynamicSegmentTree:
    """
    동적 세그먼트 트리
    - 범위가 매우 큰 경우 (10^9)
    - 노드를 필요할 때만 생성
    """
    
    class Node:
        def __init__(self):
            self.value = 0
            self.left = None
            self.right = None
    
    def __init__(self, min_val, max_val):
        self.min_val = min_val
        self.max_val = max_val
        self.root = self.Node()
    
    def update(self, index, value):
        self._update(self.root, self.min_val, self.max_val, index, value)
    
    def _update(self, node, start, end, index, value):
        if start == end:
            node.value = value
            return
        
        mid = (start + end) // 2
        
        if index <= mid:
            if node.left is None:
                node.left = self.Node()
            self._update(node.left, start, mid, index, value)
        else:
            if node.right is None:
                node.right = self.Node()
            self._update(node.right, mid + 1, end, index, value)
        
        left_val = node.left.value if node.left else 0
        right_val = node.right.value if node.right else 0
        node.value = left_val + right_val
    
    def query(self, left, right):
        return self._query(self.root, self.min_val, self.max_val, left, right)
    
    def _query(self, node, start, end, left, right):
        if node is None or right < start or end < left:
            return 0
        
        if left <= start and end <= right:
            return node.value
        
        mid = (start + end) // 2
        left_sum = self._query(node.left, start, mid, left, right)
        right_sum = self._query(node.right, mid + 1, end, left, right)
        
        return left_sum + right_sum

# 사용 (범위: 0 ~ 10^9)
dyn_seg = DynamicSegmentTree(0, 10**9)

dyn_seg.update(100, 5)
dyn_seg.update(10000000, 10)
dyn_seg.update(999999999, 3)

print(dyn_seg.query(100, 10000000))  # 5 + 10 = 15

# 장점: 메모리 O(q log n) (q = 쿼리 수)
# vs 일반 세그먼트 트리: O(4 × 10^9) = 불가능!

 

6. 실전 문제

6.1 구간 최소값과 개수

python
class SegmentTreeMinCount:
    """구간 최소값과 그 개수"""
    
    def __init__(self, arr):
        self.n = len(arr)
        self.tree = [(float('inf'), 0)] * (4 * self.n)
        self.arr = arr
        self._build(0, self.n - 1, 1)
    
    def _build(self, start, end, node):
        if start == end:
            self.tree[node] = (self.arr[start], 1)
            return
        
        mid = (start + end) // 2
        self._build(start, mid, node * 2)
        self._build(mid + 1, end, node * 2 + 1)
        
        # 병합
        left_min, left_count = self.tree[node * 2]
        right_min, right_count = self.tree[node * 2 + 1]
        
        if left_min < right_min:
            self.tree[node] = (left_min, left_count)
        elif left_min > right_min:
            self.tree[node] = (right_min, right_count)
        else:
            self.tree[node] = (left_min, left_count + right_count)
    
    def query(self, left, right):
        return self._query(0, self.n - 1, 1, left, right)
    
    def _query(self, start, end, node, left, right):
        if right < start or end < left:
            return (float('inf'), 0)
        
        if left <= start and end <= right:
            return self.tree[node]
        
        mid = (start + end) // 2
        left_result = self._query(start, mid, node * 2, left, right)
        right_result = self._query(mid + 1, end, node * 2 + 1, left, right)
        
        # 병합
        left_min, left_count = left_result
        right_min, right_count = right_result
        
        if left_min < right_min:
            return (left_min, left_count)
        elif left_min > right_min:
            return (right_min, right_count)
        else:
            return (left_min, left_count + right_count)

# 사용
arr = [3, 1, 2, 1, 4, 1]
seg = SegmentTreeMinCount(arr)

min_val, count = seg.query(0, 5)
print(f"최소값: {min_val}, 개수: {count}")  # 최소값: 1, 개수: 3

 

6.2 구간에서 K번째 수

python
class MergeSortTree:
    """Merge Sort Tree - K번째 수 찾기"""
    
    def __init__(self, arr):
        self.n = len(arr)
        self.tree = [[] for _ in range(4 * self.n)]
        self.arr = arr
        self._build(0, self.n - 1, 1)
    
    def _build(self, start, end, node):
        if start == end:
            self.tree[node] = [self.arr[start]]
            return
        
        mid = (start + end) // 2
        self._build(start, mid, node * 2)
        self._build(mid + 1, end, node * 2 + 1)
        
        # 병합 정렬
        self.tree[node] = self._merge(
            self.tree[node * 2],
            self.tree[node * 2 + 1]
        )
    
    def _merge(self, left, right):
        result = []
        i = j = 0
        
        while i < len(left) and j < len(right):
            if left[i] <= right[j]:
                result.append(left[i])
                i += 1
            else:
                result.append(right[j])
                j += 1
        
        result.extend(left[i:])
        result.extend(right[j:])
        return result
    
    def count_less_than(self, left, right, value):
        """구간 [left, right]에서 value보다 작은 수의 개수"""
        return self._count(0, self.n - 1, 1, left, right, value)
    
    def _count(self, start, end, node, left, right, value):
        if right < start or end < left:
            return 0
        
        if left <= start and end <= right:
            # 이진 탐색
            import bisect
            return bisect.bisect_left(self.tree[node], value)
        
        mid = (start + end) // 2
        left_count = self._count(start, mid, node * 2, left, right, value)
        right_count = self._count(mid + 1, end, node * 2 + 1, left, right, value)
        
        return left_count + right_count

# K번째 수 찾기 (이진 탐색)
def kth_smallest(mst, left, right, k):
    """구간 [left, right]에서 K번째로 작은 수"""
    lo, hi = min(mst.arr), max(mst.arr)
    
    while lo < hi:
        mid = (lo + hi) // 2
        count = mst.count_less_than(left, right, mid + 1)
        
        if count < k:
            lo = mid + 1
        else:
            hi = mid
    
    return lo

# 사용
arr = [5, 2, 8, 1, 9, 3, 7]
mst = MergeSortTree(arr)

# [2, 8, 1, 9, 3]에서 3번째로 작은 수
kth = kth_smallest(mst, 1, 5, 3)
print(kth)  # 3 (1, 2, 3, 8, 9 중 3번째)

 

6.3 구간 XOR

python
class SegmentTreeXOR:
    """구간 XOR"""
    
    def __init__(self, arr):
        self.n = len(arr)
        self.tree = [0] * (4 * self.n)
        self.arr = arr
        self._build(0, self.n - 1, 1)
    
    def _build(self, start, end, node):
        if start == end:
            self.tree[node] = self.arr[start]
            return
        
        mid = (start + end) // 2
        self._build(start, mid, node * 2)
        self._build(mid + 1, end, node * 2 + 1)
        
        # XOR
        self.tree[node] = (self.tree[node * 2] ^ 
                          self.tree[node * 2 + 1])
    
    def query(self, left, right):
        return self._query(0, self.n - 1, 1, left, right)
    
    def _query(self, start, end, node, 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_xor = self._query(start, mid, node * 2, left, right)
        right_xor = self._query(mid + 1, end, node * 2 + 1, left, right)
        
        return left_xor ^ right_xor

# 사용
arr = [1, 3, 5, 7, 9]
seg_xor = SegmentTreeXOR(arr)

print(seg_xor.query(0, 2))  # 1^3^5 = 7
print(seg_xor.query(1, 4))  # 3^5^7^9 = 12

 

7. 성능 분석

7.1 시간 복잡도

연산               복잡도        설명
─────────────────────────────────────────
빌드               O(n)         모든 노드 한 번씩
단일 쿼리          O(log n)     높이만큼
단일 업데이트      O(log n)     높이만큼
구간 업데이트      O(log n)     Lazy Propagation
  (Lazy)
Q개 쿼리           O(Q log n)   각 O(log n)

vs 다른 방법:

Naive:
- 쿼리: O(n)
- 업데이트: O(1)
- Q개 쿼리: O(Q×n)

Prefix Sum:
- 전처리: O(n)
- 쿼리: O(1)
- 업데이트: O(n)
- Q개 쿼리 + U개 업데이트: O(n + Q + U×n)

Segment Tree:
- 전처리: O(n)
- 쿼리: O(log n)
- 업데이트: O(log n)
- Q개 쿼리 + U개 업데이트: O(n + (Q+U) log n)

벤치마크 (n=100만, Q=100만, U=100만):
Naive: ~10^12 (불가능)
Prefix Sum: ~10^12 (업데이트 많으면)
Segment Tree: ~4×10^7 (가능!)

 

7.2 공간 복잡도

python
# 배열 크기
n = 1000000
tree_size = 4 * n

# 메모리 계산
int_size = 4  # bytes
memory = tree_size * int_size
print(f"메모리: {memory / (1024**2):.2f} MB")
# 약 15.26 MB

# Lazy Propagation
lazy_memory = tree_size * int_size
total = memory + lazy_memory
print(f"Lazy 포함: {total / (1024**2):.2f} MB")
# 약 30.52 MB

# 최적화
# 1. 2의 거듭제곱이면 2n 크기로 충분
# 2. 동적 할당 (Dynamic Segment Tree)
# 3. 압축 (좌표 압축)

공간 복잡도:
- 기본: O(4n) = O(n)
- Lazy: O(8n) = O(n)
- 2D: O(16n²) = O(n²)
- Dynamic: O(q log n) (q=쿼리 수)
- Persistent: O(n + q log n)

 

7.3 벤치마크

python
import time
import random

n = 100000
arr = [random.randint(1, 100) for _ in range(n)]

# 세그먼트 트리 빌드
start = time.time()
seg_tree = SegmentTree(arr)
build_time = time.time() - start

# 쿼리 10000개
queries = [(random.randint(0, n-1), random.randint(0, n-1)) 
           for _ in range(10000)]
queries = [(min(l, r), max(l, r)) for l, r in queries]

start = time.time()
for l, r in queries:
    seg_tree.query(l, r)
query_time = time.time() - start

# 업데이트 10000개
updates = [(random.randint(0, n-1), random.randint(1, 100)) 
           for _ in range(10000)]

start = time.time()
for idx, val in updates:
    seg_tree.update(idx, val)
update_time = time.time() - start

print(f"빌드: {build_time:.4f}초")
print(f"쿼리 10000개: {query_time:.4f}초")
print(f"업데이트 10000개: {update_time:.4f}초")

# 결과 (예시):
# 빌드: 0.0523초
# 쿼리 10000개: 0.0089초
# 업데이트 10000개: 0.0102초

# vs Naive (쿼리)
start = time.time()
for l, r in queries:
    sum(arr[l:r+1])
naive_time = time.time() - start

print(f"\nNaive 쿼리: {naive_time:.4f}초")
print(f"세그먼트 트리가 {naive_time/query_time:.1f}배 빠름")

# 결과:
# Naive 쿼리: 0.4521초
# 세그먼트 트리가 50.8배 빠름

 

8. 실무 팁

8.1 언제 세그먼트 트리를 사용할까?

세그먼트 트리 사용 적합:
✓ 구간 쿼리 + 업데이트 동시에 필요
✓ 쿼리와 업데이트가 모두 빈번
✓ 범위가 크지 않음 (≤ 10^6)
✓ 구간 합/최소/최대/GCD 등

세그먼트 트리 부적합:
✗ 쿼리만 있고 업데이트 없음 → Prefix Sum
✗ 업데이트만 있고 쿼리 없음 → 배열
✗ 범위가 너무 큼 (> 10^9) → Dynamic Segment Tree
✗ 단순 검색 → Hash Table

대안:
- Fenwick Tree (BIT): 구간 합만, 더 간단
- Sqrt Decomposition: 구현 간단, 약간 느림
- Sparse Table: 쿼리만, 전처리 빠름

 

8.2 구현 체크리스트

python
세그먼트 트리 구현 시:

□ 배열 크기 4n으로 설정
□ 1-based 인덱싱 사용 (구현 간단)
□ 재귀 vs 반복 선택
  - 재귀: 구현 간단
  - 반복: 약간 빠름
□ 쿼리 타입 결정 (합/최소/최대/GCD/XOR)
□ Lazy Propagation 필요 여부
□ 범위 체크 (완전 포함/부분 겹침/완전 벗어남)
□ 초기값 설정 (합:0, 최소:INF, 최대:-INF)
□ 오버플로우 주의 (long long 사용)

최적화:
□ 반복 구현 (스택 오버플로우 방지)
□ 좌표 압축 (값의 범위가 클 때)
□ 메모리 풀 (노드 재사용)
□ 비트 연산 (2로 곱하기/나누기)

 

8.3 디버깅

python
def visualize_tree(seg_tree):
    """세그먼트 트리 시각화"""
    def print_tree(node, start, end, depth=0):
        if node >= len(seg_tree.tree):
            return
        
        indent = "  " * depth
        print(f"{indent}[{start}-{end}]: {seg_tree.tree[node]}")
        
        if start != end:
            mid = (start + end) // 2
            print_tree(node * 2, start, mid, depth + 1)
            print_tree(node * 2 + 1, mid + 1, end, depth + 1)
    
    print_tree(1, 0, seg_tree.n - 1)

# 사용
arr = [1, 3, 5, 7]
seg = SegmentTree(arr)
visualize_tree(seg)

# 출력:
# [0-3]: 16
#   [0-1]: 4
#     [0-0]: 1
#     [1-1]: 3
#   [2-3]: 12
#     [2-2]: 5
#     [3-3]: 7

def validate_tree(seg_tree):
    """트리 무결성 검증"""
    def check(node, start, end):
        if node >= len(seg_tree.tree):
            return True
        
        if start == end:
            # 리프 노드
            return seg_tree.tree[node] == seg_tree.arr[start]
        
        mid = (start + end) // 2
        left_ok = check(node * 2, start, mid)
        right_ok = check(node * 2 + 1, mid + 1, end)
        
        expected = (seg_tree.tree[node * 2] + 
                   seg_tree.tree[node * 2 + 1])
        current_ok = seg_tree.tree[node] == expected
        
        return left_ok and right_ok and current_ok
    
    return check(1, 0, seg_tree.n - 1)

# 사용
assert validate_tree(seg), "트리 무결성 위반!"

 

9. 실무 체크리스트

세그먼트 트리 사용 시:

선택 기준

  • 구간 쿼리 필요한가?
  • 업데이트도 필요한가?
  • 쿼리와 업데이트 모두 빈번한가?
  • 범위가 적절한가? (≤ 10^6)

구현 시

  • 배열 크기 4n 할당
  • 쿼리 타입 결정 (합/최소/최대)
  • 초기값 설정 (0/INF/-INF)
  • Lazy 필요 여부 확인
  • 재귀 깊이 고려

최적화

  • 좌표 압축 고려
  • Dynamic Segment Tree 고려
  • Persistent 필요 여부
  • 메모리 제약 확인

테스트

  • 경계 조건 (0, n-1)
  • 단일 원소 쿼리
  • 전체 구간 쿼리
  • 업데이트 후 쿼리
  • 성능 벤치마크

 

10. 결론

Segment Tree는 구간 쿼리의 마스터입니다.

핵심 교훈:

  1. O(log n) 쿼리 - 빠른 구간 쿼리
  2. O(log n) 업데이트 - 빠른 갱신
  3. 분할 정복 - 구간 분할로 효율성
  4. Lazy Propagation - 구간 업데이트 최적화
  5. 다양한 쿼리 - 합/최소/최대/GCD/XOR
  6. 공간 O(n) - 트레이드오프
  7. 높이 log n - 균형 보장
  8. 응용 다양 - 2D, Persistent, Dynamic

Segment Tree는 구간 쿼리와 업데이트를 모두 효율적으로 처리합니다. Prefix Sum은 쿼리만 빠르고, Naive는 업데이트만 빠르지만, Segment Tree는 둘 다 O(log n)으로 균형 잡힌 성능을 제공합니다. 알고리즘 대회에서 필수이며, 실무에서도 데이터베이스 인덱싱 등에 활용됩니다!

"구간 쿼리 + 업데이트? Segment Tree!"

반응형

댓글