트리 기반 구간 연산 최적화 기법

트리 구조를 활용한 구간 갱신과 합계 계산은 알고리즘에서 핵심적인 기술입니다. 기본적으로 차분 배열을 사용해 구간 갱신을 단일 원소 갱신으로 변환하고, 두 개의 트리 구조를 유지해 효율적으로 계산합니다.

구간 합계 계산 공식: $$\sum_{i=1}^{x}\sum_{j=1}^{i}c_j = (x+1)\sum_{i=1}^{x}c_i - \sum_{i=1}^{x}(i \cdot c_i)$$ 따라서 $\sum c_i$와 $\sum (i \cdot c_i)$를 별도로 관리하면 됩니다.

코드 예시

#include <iostream>
#define lowbit(x) (x & -x)
using namespace std;
typedef long long ll;
const int MAXN = 200010;

int n;
ll arr[MAXN];
ll tree1[MAXN]; // Σc_i
ll tree2[MAXN]; // Σ(i * c_i)

void update(ll tree[], int idx, ll val) {
    for (int i = idx; i <= n; i += lowbit(i))
        tree[i] += val;
}

void range_update(int l, int r, ll val) {
    update(tree1, l, val);
    update(tree1, r + 1, -val);
    update(tree2, l, val * l);
    update(tree2, r + 1, -val * (r + 1));
}

ll prefix_sum(ll tree[], int idx) {
    ll res = 0;
    for (int i = idx; i; i -= lowbit(i))
        res += tree[i];
    return res;
}

ll query(int l, int r) {
    ll sum1 = prefix_sum(tree1, r) * (r + 1);
    ll sum2 = prefix_sum(tree2, r);
    if (l > 1) {
        sum1 -= prefix_sum(tree1, l - 1) * l;
        sum2 -= prefix_sum(tree2, l - 1);
    }
    return sum1 - sum2;
}

int main() {
    cin >> n;
    for (int i = 1; i <= n; i++) {
        cin >> arr[i];
        range_update(i, i, arr[i] - arr[i - 1]);
    }
    int m;
    cin >> m;
    while (m--) {
        int op, l, r;
        ll k;
        cin >> op;
        if (op == 1) {
            cin >> l >> r >> k;
            range_update(l, r, k);
        } else {
            cin >> l >> r;
            cout << query(l, r) << '\n';
        }
    }
    return 0;
}

구간 쿼리의 기대 시간 복잡도는 다음과 같이 분석됩니다: $$E_{query} = \frac{1}{2^k} + \frac{k}{2}$$ 단일 원소 갱신의 경우 $E_{modify} = E_{query} + 1$로 계산됩니다. 트리 구조의 효율성은 이로 인해 높은 성능을 보입니다.

선형 세그먼트 트리 구현에서는 두 가지 태그를 관리합니다:

  • delta: 상수 추가 값
  • slope: 선형 계수

구간 갱신 시 다음 식을 사용합니다: $$\text{sum} = \text{delta} \times \text{length} + \text{slope} \times \frac{\text{length} \times (\text{start} + \text{end})}{2}$$

선형 세그먼트 트리 코드

#include <iostream>
using namespace std;
typedef long long ll;
const int MAXN = 500010;

struct SegmentNode {
    ll delta;
    ll slope;
    ll sum;
    SegmentNode() : delta(0), slope(0), sum(0) {}
};

SegmentNode tree[MAXN << 2];
ll arr[MAXN];

void build(int node, int start, int end) {
    if (start == end) {
        tree[node].sum = arr[start];
        return;
    }
    int mid = (start + end) >> 1;
    build(node << 1, start, mid);
    build(node << 1 | 1, mid + 1, end);
    tree[node].sum = tree[node << 1].sum + tree[node << 1 | 1].sum;
}

void push_down(int node, int start, int end) {
    if (tree[node].delta) {
        int mid = (start + end) >> 1;
        tree[node << 1].delta += tree[node].delta;
        tree[node << 1 | 1].delta += tree[node].delta;
        tree[node << 1].sum += tree[node].delta * (mid - start + 1);
        tree[node << 1 | 1].sum += tree[node].delta * (end - mid);
        tree[node].delta = 0;
    }
    if (tree[node].slope) {
        int mid = (start + end) >> 1;
        tree[node << 1].slope += tree[node].slope;
        tree[node << 1 | 1].slope += tree[node].slope;
        tree[node << 1].sum += tree[node].slope * (mid - start + 1) * (start + mid) / 2;
        tree[node << 1 | 1].sum += tree[node].slope * (end - mid) * (mid + 1 + end) / 2;
        tree[node].slope = 0;
    }
}

void update(int node, int start, int end, int l, int r, ll delta, ll slope) {
    if (r < start || end < l) return;
    if (l <= start && end <= r) {
        tree[node].delta += delta;
        tree[node].slope += slope;
        tree[node].sum += delta * (end - start + 1) + slope * (end - start + 1) * (start + end) / 2;
        return;
    }
    push_down(node, start, end);
    int mid = (start + end) >> 1;
    update(node << 1, start, mid, l, r, delta, slope);
    update(node << 1 | 1, mid + 1, end, l, r, delta, slope);
    tree[node].sum = tree[node << 1].sum + tree[node << 1 | 1].sum;
}

ll query(int node, int start, int end, int l, int r) {
    if (r < start || end < l) return 0;
    if (l <= start && end <= r) return tree[node].sum;
    push_down(node, start, end);
    int mid = (start + end) >> 1;
    return query(node << 1, start, mid, l, r) + query(node << 1 | 1, mid + 1, end, l, r);
}

고차 분할 차이를 활용한 구현에서는 세 개의 트리 구조를 관리합니다:

  • $sum_d$: $\sum dd_i$
  • $sum_id$: $\sum (i \cdot dd_i)$
  • $sum_i2d$: $\sum (i^2 \cdot dd_i)$

이를 통해 구간 합계를 다음과 같이 계산합니다: $$\text{sum} = (r+1) \cdot \text{sum_id} - \text{sum_i2d} + \text{sum_d} \cdot (r^2 + 3r + 2)$$

고차 차분 구현 코드

#include <iostream>
using namespace std;
typedef long long ll;
const int MAXN = 500010;

int n;
ll arr[MAXN];
ll tree1[MAXN]; // Σdd_i
ll tree2[MAXN]; // Σ(i * dd_i)
ll tree3[MAXN]; // Σ(i^2 * dd_i)

void update_tree(ll tree[], int idx, ll val) {
    for (int i = idx; i <= n; i += (i & -i))
        tree[i] += val;
}

void range_update(int l, int r, ll k, ll d) {
    update_tree(tree1, l, k);
    update_tree(tree1, r + 1, -k);
    update_tree(tree2, l, k * l);
    update_tree(tree2, r + 1, -k * (r + 1));
    update_tree(tree3, l, k * l * l);
    update_tree(tree3, r + 1, -k * (r + 1) * (r + 1));
    
    update_tree(tree1, l + 1, d);
    update_tree(tree1, r + 2, -d);
    update_tree(tree2, l + 1, d * (l + 1));
    update_tree(tree2, r + 2, -d * (r + 2));
    update_tree(tree3, l + 1, d * (l + 1) * (l + 1));
    update_tree(tree3, r + 2, -d * (r + 2) * (r + 2));
}

ll query_tree(ll tree[], int idx) {
    ll res = 0;
    for (int i = idx; i; i -= (i & -i))
        res += tree[i];
    return res;
}

ll query(int l, int r) {
    ll sum1 = query_tree(tree1, r) * (r + 1) - query_tree(tree2, r);
    ll sum2 = query_tree(tree3, r) + query_tree(tree2, r) * (-3 - 2 * r) + query_tree(tree1, r) * (r * r + 3 * r + 2);
    if (l > 1) {
        sum1 -= query_tree(tree1, l - 1) * l - query_tree(tree2, l - 1);
        sum2 -= query_tree(tree3, l - 1) + query_tree(tree2, l - 1) * (-3 - 2 * (l - 1)) + query_tree(tree1, l - 1) * ((l - 1) * (l - 1) + 3 * (l - 1) + 2);
    }
    return sum1 + sum2 / 2;
}

나눗셈 연산 최적화를 위한 방법으로, 수의 값이 절반 이하로 감소할 때까지 반복적으로 모듈러 연산을 수행합니다. 이는 최대 $O(\log a_i)$ 번의 연산으로 해결 가능합니다.

모듈러 최적화 코드

#include <iostream>
#include <algorithm>
using namespace std;
typedef long long ll;
const int MAXN = 100010;

struct SegmentNode {
    ll sum;
    ll max_val;
    SegmentNode() : sum(0), max_val(0) {}
};

SegmentNode tree[MAXN << 2];
ll arr[MAXN];

void build(int node, int start, int end) {
    if (start == end) {
        tree[node].sum = arr[start];
        tree[node].max_val = arr[start];
        return;
    }
    int mid = (start + end) >> 1;
    build(node << 1, start, mid);
    build(node << 1 | 1, mid + 1, end);
    tree[node].sum = tree[node << 1].sum + tree[node << 1 | 1].sum;
    tree[node].max_val = max(tree[node << 1].max_val, tree[node << 1 | 1].max_val);
}

void mod_update(int node, int start, int end, int l, int r, int mod) {
    if (tree[node].max_val < mod) return;
    if (start == end) {
        tree[node].sum %= mod;
        tree[node].max_val %= mod;
        return;
    }
    int mid = (start + end) >> 1;
    if (l <= mid) mod_update(node << 1, start, mid, l, r, mod);
    if (r > mid) mod_update(node << 1 | 1, mid + 1, end, l, r, mod);
    tree[node].sum = tree[node << 1].sum + tree[node << 1 | 1].sum;
    tree[node].max_val = max(tree[node << 1].max_val, tree[node << 1 | 1].max_val);
}

ll query_sum(int node, int start, int end, int l, int r) {
    if (r < start || end < l) return 0;
    if (l <= start && end <= r) return tree[node].sum;
    int mid = (start + end) >> 1;
    return query_sum(node << 1, start, mid, l, r) + query_sum(node << 1 | 1, mid + 1, end, l, r);
}

물리적 충돌 시뮬레이션 문제에서는 인접한 개체 간 충돌 시간을 계산하고, 최소 시간을 갖는 충돌을 처리합니다. 충돌 시뮬레이션은 다음과 같은 조건을 만족합니다:

  • $v_x \geq d$: 1단계 내 충돌
  • $v_x > v_y$: 충돌 시간 = $\lceil \frac{d}{v_x - v_y} \rceil$
  • $v_x \leq v_y$: 충돌 불가

충돌 시뮬레이션 코드

#include <iostream>
#include <set>
#include <algorithm>
using namespace std;
typedef long long ll;
const int MAXN = 100000;

struct Frog {
    ll pos;
    ll speed;
    int id;
};

bool cmp(const Frog &a, const Frog &b) { return a.pos < b.pos; }

int n, m;
Frog frogs[MAXN + 10];
int prev_node[MAXN + 10];
int next_node[MAXN + 10];
ll collision_time[MAXN + 10];
ll current_time[MAXN + 10];
int ans[MAXN + 10];
int ans_size;

ll distance(int a, int b) {
    ll d = frogs[b].pos - frogs[a].pos;
    return d < 0 ? d + m : d;
}

ll calculate_time(int a, int b, ll t) {
    ll d = distance(a, b);
    ll pos_a = frogs[a].pos + frogs[a].speed * (t - current_time[a]);
    ll pos_b = frogs[b].pos + frogs[b].speed * (t - current_time[b]);
    if (frogs[a].id < frogs[b].id && d <= frogs[a].speed) return 1;
    if (frogs[a].id < frogs[b].id && frogs[a].speed > frogs[b].speed) {
        return (d - (pos_a - pos_b) + frogs[a].speed - 1) / (frogs[a].speed - frogs[b].speed);
    }
    return frogs[a].speed <= frogs[b].speed ? LLONG_MAX : (d - (pos_a - pos_b) + frogs[a].speed - 1) / (frogs[a].speed - frogs[b].speed);
}

int main() {
    cin >> n >> m;
    for (int i = 0; i < n; i++) {
        cin >> frogs[i].pos >> frogs[i].speed;
        frogs[i].id = i;
    }
    sort(frogs, frogs + n, cmp);
    
    for (int i = 0; i < n; i++) {
        prev_node[i] = (i - 1 + n) % n;
        next_node[i] = (i + 1) % n;
    }
    
    set<tuple<ll, int, int>> events;
    for (int i = 0; i < n; i++) {
        ll t = calculate_time(i, next_node[i], 0);
        events.insert({t, i, next_node[i]});
    }
    
    while (!events.empty() && events.size() > 1) {
        auto [t, a, b] = *events.begin();
        events.erase(events.begin());
        
        if (next_node[a] != b || prev_node[b] != a) continue;
        
        if (t >= LLONG_MAX) break;
        
        current_time[a] = t;
        frogs[a].speed--;
        prev_node[next_node[a]] = prev_node[a];
        next_node[prev_node[a]] = next_node[a];
        
        while (next_node[a] != a && 
               (frogs[a].pos + frogs[a].speed * (t - current_time[a])) - 
               (frogs[next_node[a]].pos + frogs[next_node[a]].speed * (t - current_time[next_node[a]])) >= distance(a, next_node[a])) {
            frogs[a].speed--;
            prev_node[next_node[a]] = prev_node[a];
            next_node[prev_node[a]] = next_node[a];
        }
        
        if (next_node[a] != a) {
            ll new_time = calculate_time(prev_node[a], a, t);
            events.insert({new_time, prev_node[a], a});
        }
        if (next_node[a] != a) {
            ll new_time = calculate_time(a, next_node[a], t);
            events.insert({new_time, a, next_node[a]});
        }
    }
    
    ans_size = 0;
    for (auto [t, a, b] : events) {
        ans[ans_size++] = frogs[a].id;
    }
    sort(ans, ans + ans_size);
    cout << ans_size << '\n';
    for (int i = 0; i < ans_size; i++) cout << ans[i] + 1 << ' ';
    return 0;
}

태그: fenwick_tree segment_tree difference_array modulo_optimization collision_simulation

10월 5일 00:35에 게시됨