세그먼트 트리(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;
}