세그먼트 트리와 바이너리 인덱스 트리 (템플릿)

세그먼트 트리 1 - 구간 연산 및 합계

이 템플릿은 구간 더하기 연산과 구간 합을 구하는 세그먼트 트리를 구현합니다.

#include <iostream>
#include <cstdio>
#include <cstring>
#include <cmath>
#include <cstdlib>
#include <algorithm>
using namespace std;
typedef long long ll;

int arrSize, queryCount;
const int MAXN = 1e5 + 5;
ll treeSum[MAXN << 2], lazyTag[MAXN << 2];

void refreshParent(int node) {
    treeSum[node] = treeSum[node << 1] + treeSum[node << 1 | 1];
}

void construct(int node, int left, int right) {
    if (left == right) {
        cin >> treeSum[node];
        return;
    }
    int middle = left + right >> 1;
    construct(node << 1, left, middle);
    construct(node << 1 | 1, middle + 1, right);
    refreshParent(node);
}

void apply(int node, int left, int right, ll value) {
    lazyTag[node] += value;
    treeSum[node] += value * (right - left + 1);
}

void distribute(int node, int left, int right) {
    if (lazyTag[node] != 0) {
        int middle = left + right >> 1;
        apply(node << 1, left, middle, lazyTag[node]);
        apply(node << 1 | 1, middle + 1, right, lazyTag[node]);
        lazyTag[node] = 0;
    }
}

void updateRange(int node, int left, int right, int ql, int qr, ll value) {
    if (ql <= left && right <= qr) {
        apply(node, left, right, value);
        return;
    }
    distribute(node, left, right);
    int middle = left + right >> 1;
    if (ql <= middle)
        updateRange(node << 1, left, middle, ql, qr, value);
    if (middle + 1 <= qr)
        updateRange(node << 1 | 1, middle + 1, right, ql, qr, value);
    refreshParent(node);
}

ll queryRange(int node, int left, int right, int ql, int qr) {
    if (ql <= left && right <= qr) {
        return treeSum[node];
    }
    distribute(node, left, right);
    int middle = left + right >> 1;
    ll result = 0;
    if (ql <= middle)
        result += queryRange(node << 1, left, middle, ql, qr);
    if (qr >= middle + 1)
        result += queryRange(node << 1 | 1, middle + 1, right, ql, qr);
    return result;
}

int main() {
    cin >> arrSize >> queryCount;
    construct(1, 1, arrSize);
    
    while (queryCount--) {
        int operation, x, y;
        ll z;
        cin >> operation >> x >> y;
        if (operation == 1) {
            cin >> z;
            updateRange(1, 1, arrSize, x, y, z);
        } else {
            cout << queryRange(1, 1, arrSize, x, y) << endl;
        }
    }
    return 0;
}

바이너리 인덱스 트리 1 - 점 추가 및 구간 합

점 업데이트와 구간 합산을 위한 바이너리 인덱스 트리 구현입니다.

#include <iostream>
#include <cstdio>
#include <cstring>
#include <cmath>
#include <cstdlib>
#include <algorithm>
using namespace std;

const int MAXN = 5e5 + 5;
int arrSize, queryCount;
ll bit[MAXN];

typedef long long ll;

int getLowbit(int x) {
    return x & -x;
}

void addPoint(int index, ll value) {
    while (index <= arrSize) {
        bit[index] += value;
        index += getLowbit(index);
    }
}

ll prefixSum(int index) {
    ll total = 0;
    while (index > 0) {
        total += bit[index];
        index -= getLowbit(index);
    }
    return total;
}

int main() {
    cin >> arrSize >> queryCount;
    
    for (int i = 1; i <= arrSize; i++) {
        ll val;
        cin >> val;
        addPoint(i, val);
    }
    
    for (int i = 1; i <= queryCount; i++) {
        int operation, x, y;
        cin >> operation >> x >> y;
        if (operation == 1) {
            addPoint(x, y);
        } else {
            cout << prefixSum(y) - prefixSum(x - 1) << endl;
        }
    }
    return 0;
}

세그먼트 트리 2 - 구간 연산 및 모듈로 합계

구간 더하기와 곱하기 연산 및 모듈로 합산을 지원하는 세그먼트 트리입니다.

#include <iostream>
#include <cstdio>
#include <cstring>
#include <cmath>
#include <cstdlib>
#include <algorithm>
using namespace std;

const int MAXN = 1e5 + 1;
typedef long long ll;

int arrSize, modValue, queryCount;
ll treeSum[MAXN << 2], addTag[MAXN << 2], mulTag[MAXN << 2];

void refreshNode(int node) {
    treeSum[node] = (treeSum[node << 1] + treeSum[node << 1 | 1]) % modValue;
}

void buildTree(int node, int left, int right) {
    mulTag[node] = 1;
    if (left == right) {
        cin >> treeSum[node];
        treeSum[node] %= modValue;
        return;
    }
    int middle = left + right >> 1;
    buildTree(node << 1, left, middle);
    buildTree(node << 1 | 1, middle + 1, right);
    refreshNode(node);
}

void pushDown(int node, int segmentSize) {
    if (mulTag[node] == 1 && addTag[node] == 0)
        return;
    
    ll leftSize = segmentSize - (segmentSize >> 1);
    ll rightSize = segmentSize >> 1;
    
    treeSum[node << 1] = (treeSum[node << 1] * mulTag[node] + leftSize * addTag[node]) % modValue;
    treeSum[node << 1 | 1] = (treeSum[node << 1 | 1] * mulTag[node] + rightSize * addTag[node]) % modValue;
    
    mulTag[node << 1] = (mulTag[node << 1] * mulTag[node]) % modValue;
    mulTag[node << 1 | 1] = (mulTag[node << 1 | 1] * mulTag[node]) % modValue;
    
    addTag[node << 1] = (addTag[node << 1] * mulTag[node] + addTag[node]) % modValue;
    addTag[node << 1 | 1] = (addTag[node << 1 | 1] * mulTag[node] + addTag[node]) % modValue;
    
    mulTag[node] = 1;
    addTag[node] = 0;
}

void applyAdd(int node, int left, int right, ll value) {
    addTag[node] = (addTag[node] + value) % modValue;
    treeSum[node] = (treeSum[node] + value * (right - left + 1)) % modValue;
}

void applyMul(int node, int left, int right, ll value) {
    treeSum[node] = (treeSum[node] * value) % modValue;
    mulTag[node] = (mulTag[node] * value) % modValue;
    addTag[node] = (addTag[node] * value) % modValue;
}

void updateAdd(int node, int left, int right, int ql, int qr, ll value) {
    if (ql <= left && right <= qr) {
        applyAdd(node, left, right, value);
        return;
    }
    pushDown(node, right - left + 1);
    int middle = left + right >> 1;
    if (ql <= middle)
        updateAdd(node << 1, left, middle, ql, qr, value);
    if (middle + 1 <= qr)
        updateAdd(node << 1 | 1, middle + 1, right, ql, qr, value);
    refreshNode(node);
}

void updateMul(int node, int left, int right, int ql, int qr, ll value) {
    if (ql <= left && right <= qr) {
        applyMul(node, left, right, value);
        return;
    }
    pushDown(node, right - left + 1);
    int middle = left + right >> 1;
    if (ql <= middle)
        updateMul(node << 1, left, middle, ql, qr, value);
    if (middle + 1 <= qr)
        updateMul(node << 1 | 1, middle + 1, right, ql, qr, value);
    refreshNode(node);
}

ll querySum(int node, int left, int right, int ql, int qr) {
    if (ql <= left && right <= qr) {
        return treeSum[node];
    }
    pushDown(node, right - left + 1);
    int middle = left + right >> 1;
    ll result = 0;
    if (ql <= middle) {
        result += querySum(node << 1, left, middle, ql, qr);
        result %= modValue;
    }
    if (qr >= middle + 1) {
        result += querySum(node << 1 | 1, middle + 1, right, ql, qr);
        result %= modValue;
    }
    return result;
}

int main() {
    cin >> arrSize >> queryCount >> modValue;
    buildTree(1, 1, arrSize);
    
    while (queryCount--) {
        int operation, a, b;
        ll c;
        cin >> operation >> a >> b;
        if (operation == 1) {
            cin >> c;
            updateMul(1, 1, arrSize, a, b, c);
        } else if (operation == 2) {
            cin >> c;
            updateAdd(1, 1, arrSize, a, b, c);
        } else {
            cout << querySum(1, 1, arrSize, a, b) << endl;
        }
    }
    return 0;
}

바이너리 인덱스 트리 2 - 구간 추가 및 점 조회

구간 업데이트와 점 조회를 위한 바이너리 인덱스 트리 구현입니다.

#include <iostream>
#include <cstdio>
#include <cstring>
#include <cmath>
#include <cstdlib>
#include <algorithm>
using namespace std;

typedef long long ll;
const int MAXN = 5e5 + 9;

int arrSize, queryCount;
ll original[MAXN], diffTree[MAXN], weightedTree[MAXN];

int getLowbit(int x) {
    return x & (-x);
}

void treeAdd(ll* tree, int index, ll value) {
    while (index <= arrSize) {
        tree[index] += value;
        index += getLowbit(index);
    }
}

ll treeQuery(ll* tree, int index) {
    ll result = 0;
    while (index > 0) {
        result += tree[index];
        index -= getLowbit(index);
    }
    return result;
}

ll getValueAt(int position) {
    return (ll)treeQuery(diffTree, position) * (position + 1) - treeQuery(weightedTree, position);
}

int main() {
    cin >> arrSize >> queryCount;
    
    for (int i = 1; i <= arrSize; i++) {
        cin >> original[i];
        ll diff = original[i] - original[i - 1];
        treeAdd(diffTree, i, diff);
        treeAdd(weightedTree, i, i * diff);
    }
    
    while (queryCount--) {
        ll operation, x, y, z;
        cin >> operation >> x;
        if (operation == 1) {
            cin >> y >> z;
            treeAdd(diffTree, x, z);
            treeAdd(weightedTree, x, x * z);
            treeAdd(diffTree, y + 1, -z);
            treeAdd(weightedTree, y + 1, (y + 1) * (-z));
        } else {
            cout << getValueAt(x) - getValueAt(x - 1) << endl;
        }
    }
    return 0;
}

태그: segment-tree binary-indexed-tree data-structures cpp algorithm

7월 28일 18:20에 게시됨