트리 구조를 활용한 구간 갱신과 합계 계산은 알고리즘에서 핵심적인 기술입니다. 기본적으로 차분 배열을 사용해 구간 갱신을 단일 원소 갱신으로 변환하고, 두 개의 트리 구조를 유지해 효율적으로 계산합니다.
구간 합계 계산 공식: $$\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;
}