C++ 세그먼트 트리 구현 및 지연 전파(Lazy Propagation) 완벽 가이드

세그먼트 트리(Segment Tree)는 펜윅 트리(Fenwick Tree)와 유사하게 구간 합을 구하는 데 주로 사용되지만, 이 외에도 구간 최소/최대값 탐색, 구간 색칠 등 다양한 구간 연산을 효율적으로 처리할 수 있는 강력한 자료구조입니다. 본 가이드에서는 C++를 사용하여 세그먼트 트리의 기본 구현부터 지연 전파(Lazy Propagation)를 활용한 고급 기법까지 단계별로 다룹니다.

1. 단일 원소 업데이트 및 구간 쿼리

가장 기본이 되는 세그먼트 트리 형태입니다. 특정 인덱스의 값을 변경(Point Update)하고, 특정 구간의 합을 계산(Range Query)하는 연산을 $O(\log N)$ 시간에 수행합니다.

#include <iostream>
#include <vector>

using namespace std;
typedef long long ll;

struct Node {
    int left, right;
    ll value;
};

const int MAX_SIZE = 500005;
Node segTree[MAX_SIZE * 4];
ll arr[MAX_SIZE];
int N, M;

void buildTree(int node, int start, int end) {
    segTree[node].left = start;
    segTree[node].right = end;
    if (start == end) {
        segTree[node].value = arr[start];
        return;
    }
    int mid = (start + end) / 2;
    buildTree(node * 2, start, mid);
    buildTree(node * 2 + 1, mid + 1, end);
    segTree[node].value = segTree[node * 2].value + segTree[node * 2 + 1].value;
}

void pointUpdate(int node, int idx, ll val) {
    if (segTree[node].left == segTree[node].right) {
        segTree[node].value += val;
        return;
    }
    int mid = (segTree[node].left + segTree[node].right) / 2;
    if (idx <= mid) pointUpdate(node * 2, idx, val);
    else pointUpdate(node * 2 + 1, idx, val);
    segTree[node].value = segTree[node * 2].value + segTree[node * 2 + 1].value;
}

ll rangeQuery(int node, int qLeft, int qRight) {
    if (segTree[node].right < qLeft || segTree[node].left > qRight) return 0;
    if (qLeft <= segTree[node].left && segTree[node].right <= qRight) return segTree[node].value;
    
    ll leftSum = rangeQuery(node * 2, qLeft, qRight);
    ll rightSum = rangeQuery(node * 2 + 1, qLeft, qRight);
    return leftSum + rightSum;
}

int main() {
    ios_base::sync_with_stdio(false);
    cin.tie(NULL);
    
    cin >> N >> M;
    for (int i = 1; i <= N; i++) cin >> arr[i];
    buildTree(1, 1, N);
    
    while (M--) {
        int op, x;
        ll y;
        cin >> op >> x >> y;
        if (op == 1) pointUpdate(1, x, y);
        else cout << rangeQuery(1, x, (int)y) << "\n";
    }
    return 0;
}

2. 구간 업데이트 및 단일 원소 쿼리

특정 구간의 모든 원소에 값을 더하고, 특정 인덱스의 값을 조회하는 형태입니다. 이 경우 리프 노드까지 내려가지 않고 겹치는 구간에만 값을 누적해 두었다가, 쿼리 시 루트에서 리프까지의 경로를 따라 누적된 값을 합산하여 결과를 도출합니다.

#include <iostream>

using namespace std;
typedef long long ll;

struct Node {
    int left, right;
    ll value;
};

const int MAX_SIZE = 500005;
Node segTree[MAX_SIZE * 4];
ll arr[MAX_SIZE];
int N, M;

void buildTree(int node, int start, int end) {
    segTree[node].left = start;
    segTree[node].right = end;
    segTree[node].value = 0;
    if (start == end) return;
    int mid = (start + end) / 2;
    buildTree(node * 2, start, mid);
    buildTree(node * 2 + 1, mid + 1, end);
}

void rangeUpdate(int node, int qLeft, int qRight, ll val) {
    if (qLeft <= segTree[node].left && segTree[node].right <= qRight) {
        segTree[node].value += val;
        return;
    }
    if (segTree[node * 2].right >= qLeft) rangeUpdate(node * 2, qLeft, qRight, val);
    if (segTree[node * 2 + 1].left <= qRight) rangeUpdate(node * 2 + 1, qLeft, qRight, val);
}

ll pointQuery(int node, int idx) {
    ll res = segTree[node].value;
    if (segTree[node].left == segTree[node].right) return res;
    int mid = (segTree[node].left + segTree[node].right) / 2;
    if (idx <= mid) res += pointQuery(node * 2, idx);
    else res += pointQuery(node * 2 + 1, idx);
    return res;
}

int main() {
    ios_base::sync_with_stdio(false);
    cin.tie(NULL);
    
    cin >> N >> M;
    for (int i = 1; i <= N; i++) cin >> arr[i];
    buildTree(1, 1, N);
    
    while (M--) {
        int op, x, y;
        ll k;
        cin >> op;
        if (op == 1) {
            cin >> x >> y >> k;
            rangeUpdate(1, x, y, k);
        } else {
            cin >> x;
            cout << pointQuery(1, x) + arr[x] << "\n";
        }
    }
    return 0;
}

3. 고급 세그먼트 트리: 지연 전파(Lazy Propagation)

구간 업데이트와 구간 쿼리를 모두 $O(\log N)$에 처리하기 위해서는 지연 전파 기법이 필수적입니다. 업데이트 시 하위 노드들의 값을 즉시 갱신하는 대신, 'Lazy Tag'에_pending_ 상태를 기록해 두었다가 실제 쿼리나 추가 업데이트가 발생할 때 자식 노드로 전파(propagate)하는 방식입니다.

지연 전파 핵심 로직

void propagate(int idx) {
    if (tree[idx].lazy != 0) {
        int mid = (tree[idx].left + tree[idx].right) / 2;
        
        // 좌측 자식 노드 전파
        tree[idx * 2].lazy += tree[idx].lazy;
        tree[idx * 2].sum += tree[idx].lazy * (mid - tree[idx * 2].left + 1);
        
        // 우측 자식 노드 전파
        tree[idx * 2 + 1].lazy += tree[idx].lazy;
        tree[idx * 2 + 1].sum += tree[idx].lazy * (tree[idx * 2 + 1].right - mid);
        
        // 현재 노드의 Lazy Tag 초기화
        tree[idx].lazy = 0;
    }
}

전체 구현 (구간 합 및 구간 덧셈)

#include <iostream>

using namespace std;
typedef long long ll;

struct Node {
    int left, right;
    ll sum;
    ll lazy;
};

const int MAX_SIZE = 500005;
Node tree[MAX_SIZE * 4];
ll data[MAX_SIZE];
int N, M;

void build(int idx, int l, int r) {
    tree[idx].left = l;
    tree[idx].right = r;
    tree[idx].lazy = 0;
    if (l == r) {
        tree[idx].sum = data[l];
        return;
    }
    int mid = (l + r) / 2;
    build(idx * 2, l, mid);
    build(idx * 2 + 1, mid + 1, r);
    tree[idx].sum = tree[idx * 2].sum + tree[idx * 2 + 1].sum;
}

void propagate(int idx) {
    if (tree[idx].lazy != 0) {
        int mid = (tree[idx].left + tree[idx].right) / 2;
        
        tree[idx * 2].lazy += tree[idx].lazy;
        tree[idx * 2].sum += tree[idx].lazy * (mid - tree[idx * 2].left + 1);
        
        tree[idx * 2 + 1].lazy += tree[idx].lazy;
        tree[idx * 2 + 1].sum += tree[idx].lazy * (tree[idx * 2 + 1].right - mid);
        
        tree[idx].lazy = 0;
    }
}

void updateRange(int idx, int l, int r, ll val) {
    if (r < tree[idx].left || tree[idx].right < l) return;
    if (l <= tree[idx].left && tree[idx].right <= r) {
        tree[idx].sum += val * (tree[idx].right - tree[idx].left + 1);
        tree[idx].lazy += val;
        return;
    }
    propagate(idx);
    updateRange(idx * 2, l, r, val);
    updateRange(idx * 2 + 1, l, r, val);
    tree[idx].sum = tree[idx * 2].sum + tree[idx * 2 + 1].sum;
}

ll queryRange(int idx, int l, int r) {
    if (r < tree[idx].left || tree[idx].right < l) return 0;
    if (l <= tree[idx].left && tree[idx].right <= r) return tree[idx].sum;
    
    propagate(idx);
    return queryRange(idx * 2, l, r) + queryRange(idx * 2 + 1, l, r);
}

int main() {
    ios_base::sync_with_stdio(false);
    cin.tie(NULL);
    
    cin >> N >> M;
    for (int i = 1; i <= N; i++) cin >> data[i];
    build(1, 1, N);
    
    while (M--) {
        int type, x, y;
        ll k;
        cin >> type;
        if (type == 1) {
            cin >> x >> y >> k;
            updateRange(1, x, y, k);
        } else {
            cin >> x >> y;
            cout << queryRange(1, x, y) << "\n";
        }
    }
    return 0;
}

4. 다중 지연 태그: 구간 곱셈과 덧셈 동시 처리

구간에 대한 곱셈과 덧셈 연산이 혼재될 경우, 연산의 우선순위를 올바르게 처리하는 것이 핵심입니다. 일반적으로 곱셈을 먼저 수행한 후 덧셈을 수행하는 규칙을 적용하며, 이를 위해 두 가지 Lazy Tag(mulLazy, addLazy)를 관리해야 합니다.

자식 노드로 전파할 때의 수식은 다음과 같이 적용됩니다:
새로운 값 = (기존 값 * 곱셈 태그) + (덧셈 태그 * 구간 길이)
새로운 덧셈 태그 = (기존 덧셈 태그 * 부모 곱셈 태그) + 부모 덧셈 태그
새로운 곱셈 태그 = 기존 곱셈 태그 * 부모 곱셈 태그

#include <iostream>

using namespace std;
typedef long long ll;

struct Node {
    int left, right;
    ll sum;
    ll addLazy;
    ll mulLazy;
};

const int MAX_SIZE = 500005;
Node tree[MAX_SIZE * 4];
ll data[MAX_SIZE];
int N, M;
ll MOD;

void build(int idx, int l, int r) {
    tree[idx].left = l;
    tree[idx].right = r;
    tree[idx].addLazy = 0;
    tree[idx].mulLazy = 1;
    if (l == r) {
        tree[idx].sum = data[l] % MOD;
        return;
    }
    int mid = (l + r) / 2;
    build(idx * 2, l, mid);
    build(idx * 2 + 1, mid + 1, r);
    tree[idx].sum = (tree[idx * 2].sum + tree[idx * 2 + 1].sum) % MOD;
}

void propagate(int idx) {
    ll m = tree[idx].mulLazy;
    ll a = tree[idx].addLazy;
    
    if (m != 1 || a != 0) {
        for (int child : {idx * 2, idx * 2 + 1}) {
            int len = tree[child].right - tree[child].left + 1;
            tree[child].sum = (tree[child].sum * m % MOD + a * len % MOD) % MOD;
            tree[child].mulLazy = (tree[child].mulLazy * m) % MOD;
            tree[child].addLazy = (tree[child].addLazy * m % MOD + a) % MOD;
        }
        tree[idx].mulLazy = 1;
        tree[idx].addLazy = 0;
    }
}

void updateMul(int idx, int l, int r, ll val) {
    if (r < tree[idx].left || tree[idx].right < l) return;
    if (l <= tree[idx].left && tree[idx].right <= r) {
        tree[idx].sum = (tree[idx].sum * val) % MOD;
        tree[idx].mulLazy = (tree[idx].mulLazy * val) % MOD;
        tree[idx].addLazy = (tree[idx].addLazy * val) % MOD;
        return;
    }
    propagate(idx);
    updateMul(idx * 2, l, r, val);
    updateMul(idx * 2 + 1, l, r, val);
    tree[idx].sum = (tree[idx * 2].sum + tree[idx * 2 + 1].sum) % MOD;
}

void updateAdd(int idx, int l, int r, ll val) {
    if (r < tree[idx].left || tree[idx].right < l) return;
    if (l <= tree[idx].left && tree[idx].right <= r) {
        int len = tree[idx].right - tree[idx].left + 1;
        tree[idx].sum = (tree[idx].sum + val * len) % MOD;
        tree[idx].addLazy = (tree[idx].addLazy + val) % MOD;
        return;
    }
    propagate(idx);
    updateAdd(idx * 2, l, r, val);
    updateAdd(idx * 2 + 1, l, r, val);
    tree[idx].sum = (tree[idx * 2].sum + tree[idx * 2 + 1].sum) % MOD;
}

ll query(int idx, int l, int r) {
    if (r < tree[idx].left || tree[idx].right < l) return 0;
    if (l <= tree[idx].left && tree[idx].right <= r) return tree[idx].sum;
    
    propagate(idx);
    return (query(idx * 2, l, r) + query(idx * 2 + 1, l, r)) % MOD;
}

int main() {
    ios_base::sync_with_stdio(false);
    cin.tie(NULL);
    
    cin >> N >> M >> MOD;
    for (int i = 1; i <= N; i++) cin >> data[i];
    build(1, 1, N);
    
    while (M--) {
        int type, x, y;
        ll k;
        cin >> type;
        if (type == 1) {
            cin >> x >> y >> k;
            updateMul(1, x, y, k % MOD);
        } else if (type == 2) {
            cin >> x >> y >> k;
            updateAdd(1, x, y, k % MOD);
        } else {
            cin >> x >> y;
            cout << query(1, x, y) << "\n";
        }
    }
    return 0;
}

태그: SegmentTree LazyPropagation C++ DataStructure algorithm

7월 23일 20:44에 게시됨