출시·고도화 중
구간 트리 안내서 · 3/6
이 장에서는 가장 널리 쓰는 반복문(아래에서 위로) 구간 트리를 Python으로 작성하고 한 줄씩 설명합니다. 이어서 같은 루틴을 C++, Java, TypeScript로 옮기고, 합 전용의 가벼운 대안인 펜윅 트리도 구현합니다. 구간은 모두 반열린 구간 [left, right)입니다.
class SegmentTree:
def __init__(self, values):
self.n = len(values)
self.tree = [0] * self.n + list(values)
for i in range(self.n - 1, 0, -1):
self.tree[i] = self.tree[2 * i] + self.tree[2 * i + 1]
def update(self, index, value):
i = index + self.n
self.tree[i] = value
while i > 1:
i //= 2
self.tree[i] = self.tree[2 * i] + self.tree[2 * i + 1]
def query(self, left, right):
total = 0
lo, hi = left + self.n, right + self.n
while lo < hi:
if lo & 1:
total += self.tree[lo]
lo += 1
if hi & 1:
hi -= 1
total += self.tree[hi]
lo //= 2
hi //= 2
return total
st = SegmentTree([5, 3, 7, 9, 6, 4, 1, 2])
print(st.query(2, 6)) # 26
st.update(3, 0)
print(st.query(2, 6)) # 17한 줄씩 보면 다음과 같습니다.
self.tree = [0] * self.n + list(values): 크기 2n 배열을 만들고 뒤쪽 절반 tree[n..2n)에 원소를 놓습니다. 인덱스 0은 쓰지 않습니다.for 문: n-1부터 1까지 거꾸로 돌며 두 자식 2i, 2i+1을 합칩니다. 자식이 먼저 계산되므로 O(n)에 끝납니다.update: 리프 index + n을 바꾼 뒤 i //= 2로 부모로 올라가며 다시 계산합니다. 높이만큼, 곧 O(log n)번 돕니다.lo, hi: 질의 구간의 양 끝을 리프 위치로 옮긴 값입니다.lo & 1: lo가 오른쪽 자식(홀수)이면 그 부모는 구간 밖까지 덮으므로, 이 노드만 결과에 넣고 lo를 한 칸 오른쪽으로 옮깁니다.hi & 1: 반열린 오른쪽 끝이 홀수이면 바로 왼쪽 노드 hi - 1이 구간 안에 있으므로 그 노드를 넣습니다.lo //= 2, hi //= 2: 한 층 위로 올라갑니다. 두 끝이 만나면 멈춥니다.최솟값 트리로 바꾸려면 +를 min으로, 초깃값 0을 float("inf")로 바꾸면 됩니다. 순서가 중요한 연산이라면 왼쪽과 오른쪽에서 모은 값을 따로 두었다가 마지막에 순서대로 합칩니다.
합이 int 범위를 넘기 쉬우므로 long long을 씁니다.
#include <vector>
struct SegmentTree {
int n;
std::vector<long long> tree;
explicit SegmentTree(const std::vector<long long>& values)
: n(static_cast<int>(values.size())), tree(2 * values.size()) {
for (int i = 0; i < n; ++i) tree[n + i] = values[i];
for (int i = n - 1; i > 0; --i) tree[i] = tree[2 * i] + tree[2 * i + 1];
}
void update(int index, long long value) {
int i = index + n;
tree[i] = value;
while (i > 1) { i /= 2; tree[i] = tree[2 * i] + tree[2 * i + 1]; }
}
long long query(int left, int right) const { // [left, right)
long long total = 0;
for (int lo = left + n, hi = right + n; lo < hi; lo /= 2, hi /= 2) {
if (lo & 1) total += tree[lo++];
if (hi & 1) total += tree[--hi];
}
return total;
}
};public final class SegmentTree {
private final int n;
private final long[] tree;
public SegmentTree(long[] values) {
n = values.length;
tree = new long[2 * n];
System.arraycopy(values, 0, tree, n, n);
for (int i = n - 1; i > 0; i--) tree[i] = tree[2 * i] + tree[2 * i + 1];
}
public void update(int index, long value) {
int i = index + n;
tree[i] = value;
while (i > 1) { i /= 2; tree[i] = tree[2 * i] + tree[2 * i + 1]; }
}
public long query(int left, int right) { // [left, right)
long total = 0;
for (int lo = left + n, hi = right + n; lo < hi; lo /= 2, hi /= 2) {
if ((lo & 1) == 1) total += tree[lo++];
if ((hi & 1) == 1) total += tree[--hi];
}
return total;
}
}class SegmentTree {
private readonly n: number;
private readonly tree: number[];
constructor(values: number[]) {
this.n = values.length;
this.tree = new Array<number>(this.n).fill(0).concat(values);
for (let i = this.n - 1; i > 0; i--) this.tree[i] = this.tree[2 * i] + this.tree[2 * i + 1];
}
update(index: number, value: number): void {
let i = index + this.n;
this.tree[i] = value;
while (i > 1) { i >>= 1; this.tree[i] = this.tree[2 * i] + this.tree[2 * i + 1]; }
}
query(left: number, right: number): number {
let total = 0;
for (let lo = left + this.n, hi = right + this.n; lo < hi; lo >>= 1, hi >>= 1) {
if (lo & 1) total += this.tree[lo++];
if (hi & 1) total += this.tree[--hi];
}
return total;
}
}펜윅 트리(BIT)는 i & -i, 곧 i의 가장 낮은 1비트만큼씩 이동하며 누적 합을 나누어 저장합니다. 인덱스를 1부터 쓰는 것이 핵심입니다.
class FenwickTree:
def __init__(self, n):
self.n = n
self.tree = [0] * (n + 1) # tree[0]은 쓰지 않는다
def add(self, index, delta): # a[index] += delta
i = index + 1
while i <= self.n:
self.tree[i] += delta
i += i & -i
def prefix_sum(self, count): # a[0] + ... + a[count-1]
total = 0
while count > 0:
total += self.tree[count]
count -= count & -count
return total
def range_sum(self, left, right): # [left, right)
return self.prefix_sum(right) - self.prefix_sum(left)
bit = FenwickTree(8)
for i, v in enumerate([5, 3, 7, 9, 6, 4, 1, 2]):
bit.add(i, v)
print(bit.range_sum(2, 6)) # 26
bit.add(3, -9) # a[3] = 0 과 같다
print(bit.range_sum(2, 6)) # 17펜윅 트리는 값을 "더하는" 연산만 받으므로, 값을 덮어쓰려면 차이(새 값 - 옛 값)를 더합니다. 구간 합이 뺄셈으로 구해지므로 최솟값처럼 되돌릴 수 없는 연산에는 그대로 쓸 수 없습니다.
반복문 구간 트리는 크기 2n 배열, 생성 반복문 하나, 갱신과 질의 반복문 각각 하나로 끝나며 언어가 바뀌어도 구조가 같습니다. 합에는 더 짧은 펜윅 트리를 쓸 수 있고, 최솟값 · 최댓값처럼 빼기가 안 되는 연산에는 구간 트리를 씁니다.
댓글 0개
로그인 · 로그인하면 댓글을 남길 수 있습니다.
첫 댓글을 남겨 보세요.