Cây đoạn (Segment Tree) là một cấu trúc dữ liệu phổ biến trong lập trình thi đấu, đặc biệt hiệu quả cho các bài toán yêu cầu thực hiện xen kẽ các thao tác cập nhật đoạn và truy vấn đoạn trên dãy số.
Cấu trúc này có thể xử lý hiệu quả các phép toán trên đoạn với độ phức tạp O(log n) cho mỗi thao tác cập nhật hoặc truy vấn.
Nguyên lý hoạt động rất đơn giản: ánh xạ một dãy số lên một cây nhị phân cân bằng, mỗi nút trong cây đại diện cho một đoạn của dãy gốc. Bất kỳ đoạn nào trong dãy cũng có thể được biểu diễn bằng sự kết hợp của một số lượng nhỏ các đoạn tương ứng với các nút trong cây, từ đó đảm bảo độ phức tạp logarit.
Chúng ta sẽ minh họa bằng bài toán mẫu: thực hiện các thao tác cộng đoạn và tính tổng đoạn.
Ký hiệu segment_tree[i] là giá trị lưu trữ tại nút i, segment_tree[2*i] là con trái, segment_tree[2*i+1] là con phải.
Giả sử nút i biểu diễn đoạn [left_bound, right_bound], với mid = ⌊(left_bound + right_bound)/2⌋, thì hai con của nó biểu diễn các đoạn [left_bound, mid] và [mid+1, right_bound].
1. Xây dựng cây đoạn
Sử dụng phương pháp đệ quy để xây dựng cây từng cấp một.
Bắt đầu từ nút gốc, đệ quy xuống các con trái và phải cho đến khi đoạn chỉ còn một phần tử (nút lá), sau đó trả về và cập nhật giá trị cho các nút cha dựa trên tổng của hai con.
void propagate_up(int node_id)
{
segment_tree[node_id] = segment_tree[node_id * 2] + segment_tree[node_id * 2 + 1];
}
void construct_tree(int node_id, int left_bound, int right_bound)
{
if (left_bound == right_bound)
{
segment_tree[node_id] = source_array[left_bound];
return;
}
int middle = (left_bound + right_bound) / 2;
construct_tree(node_id * 2, left_bound, middle);
construct_tree(node_id * 2 + 1, middle + 1, right_bound);
propagate_up(node_id);
}
2. Truy vấn đoạn
Cũng sử dụng đệ quy, nếu đoạn của nút hiện tại nằm hoàn toàn trong đoạn cần truy vấn, trả về giá trị đã lưu.
long long range_query(int node_id, int left_bound, int right_bound, int query_left, int query_right)
{
if (query_left > right_bound || query_right < left_bound)
return 0;
if (query_left <= left_bound && query_right >= right_bound)
return segment_tree[node_id];
int middle = (left_bound + right_bound) / 2;
return range_query(node_id * 2, left_bound, middle, query_left, query_right) +
range_query(node_id * 2 + 1, middle + 1, right_bound, query_left, query_right);
}
3. Cập nhật đoạn cơ bản
Đệ quy xuống từng nút, nếu gặp nút lá thì cập nhật, sau đó cập nhật ngược lên các nút cha.
Tuy nhiên, cách tiếp cận này có nhược điểm nghiêm trọng: khi cần cập nhật cả dãy, độ phức tạp có thể lên tới O(n), làm mất đi ưu điểm của cây đoạn.
4. Lazy Propagation (Truyền tải lười biếng)
Để tối ưu, chúng ta áp dụng kỹ thuật lazy propagation. Thay vì cập nhật ngay toàn bộ đoạn, ta lưu lại thông tin cập nhật tại các nút cha và chỉ áp dụng khi thật sự cần thiết.
Ý tưởng giống như việc giao bài tập về nhà: nếu chưa kiểm tra thì học sinh không làm, chỉ làm khi sắp đến lúc kiểm tra.
4.1 Truyền tải đánh dấu
void apply_mark(int node_id, int left_bound, int right_bound, long long value)
{
lazy_mark[node_id] += value;
segment_tree[node_id] += (right_bound - left_bound + 1) * value;
}
void push_down(int node_id, int left_bound, int right_bound)
{
if (lazy_mark[node_id] == 0)
return;
int middle = (left_bound + right_bound) / 2;
apply_mark(node_id * 2, left_bound, middle, lazy_mark[node_id]);
apply_mark(node_id * 2 + 1, middle + 1, right_bound, lazy_mark[node_id]);
lazy_mark[node_id] = 0;
}
4.2 Cập nhật đoạn với lazy propagation
void range_update(int node_id, int left_bound, int right_bound, int update_left, int update_right, long long value)
{
if (update_left > right_bound || update_right < left_bound)
return;
if (update_left <= left_bound && update_right >= right_bound)
{
apply_mark(node_id, left_bound, right_bound, value);
return;
}
push_down(node_id, left_bound, right_bound);
int middle = (left_bound + right_bound) / 2;
range_update(node_id * 2, left_bound, middle, update_left, update_right, value);
range_update(node_id * 2 + 1, middle + 1, right_bound, update_left, update_right, value);
propagate_up(node_id);
}
4.3 Truy vấn đoạn với lazy propagation
long long optimized_query(int node_id, int left_bound, int right_bound, int query_left, int query_right)
{
if (query_left > right_bound || query_right < left_bound)
return 0;
if (query_left <= left_bound && query_right >= right_bound)
return segment_tree[node_id];
push_down(node_id, left_bound, right_bound);
int middle = (left_bound + right_bound) / 2;
return optimized_query(node_id * 2, left_bound, middle, query_left, query_right) +
optimized_query(node_id * 2 + 1, middle + 1, right_bound, query_left, query_right);
}
Mã nguồn hoàn chỉnh
#include <bits/stdc++.h>
#define LL long long
using namespace std;
const int MAX_N = 1e5 + 10;
LL n, q, input_data[MAX_N], segment_tree[4 * MAX_N], lazy_tag[4 * MAX_N];
void propagate_up(int pos) {
segment_tree[pos] = segment_tree[pos * 2] + segment_tree[pos * 2 + 1];
}
void build_tree(int pos, int l, int r) {
if (l == r) {
segment_tree[pos] = input_data[l];
return;
}
int mid = (l + r) / 2;
build_tree(pos * 2, l, mid);
build_tree(pos * 2 + 1, mid + 1, r);
propagate_up(pos);
}
void apply_lazy(int pos, int l, int r, LL val) {
lazy_tag[pos] += val;
segment_tree[pos] += (r - l + 1) * val;
}
void push_down(int pos, int l, int r) {
if (lazy_tag[pos] == 0) return;
int mid = (l + r) / 2;
apply_lazy(pos * 2, l, mid, lazy_tag[pos]);
apply_lazy(pos * 2 + 1, mid + 1, r, lazy_tag[pos]);
lazy_tag[pos] = 0;
}
void update_range(int pos, int l, int r, int ul, int ur, LL val) {
if (ul > r || ur < l) return;
if (ul <= l && ur >= r) {
apply_lazy(pos, l, r, val);
return;
}
push_down(pos, l, r);
int mid = (l + r) / 2;
update_range(pos * 2, l, mid, ul, ur, val);
update_range(pos * 2 + 1, mid + 1, r, ul, ur, val);
propagate_up(pos);
}
LL query_range(int pos, int l, int r, int ql, int qr) {
if (ql > r || qr < l) return 0;
if (ql <= l && qr >= r) return segment_tree[pos];
push_down(pos, l, r);
int mid = (l + r) / 2;
return query_range(pos * 2, l, mid, ql, qr) +
query_range(pos * 2 + 1, mid + 1, r, ql, qr);
}
int main() {
ios_base::sync_with_stdio(false);
cin.tie(NULL);
cin >> n >> q;
for (int i = 1; i <= n; i++) cin >> input_data[i];
build_tree(1, 1, n);
while (q--) {
int type, x, y;
cin >> type >> x >> y;
if (type == 1) {
LL k; cin >> k;
update_range(1, 1, n, x, y, k);
} else {
cout << query_range(1, 1, n, x, y) << "\n";
}
}
return 0;
}