线段树模板
单点修线段树:
#include "bits/stdc++.h"
using namespace std;
// 1. 定义节点信息
struct Info {
long long sum = 0; // 区间求和建议开 long long 防止溢出
};
// 2. 重载 + 运算符用于合并节点信息
Info operator+(const Info &a, const Info &b) {
return {a.sum + b.sum};
}
struct SegmentTree {
vector<Info> info;
int n;
// 默认构造(建空树)
SegmentTree(int size) {
n = size;
info.assign(4 * (n + 5), Info{});
build(1, 0, n);
}
// 用初始数组建树
SegmentTree(const vector<long long>& arr) {
n = arr.size();
info.assign(4 * (n + 5), Info{});
build(1, 0, n, arr);
}
void pull(int p) {
info[p] = info[2 * p] + info[2 * p + 1];
}
// 辅助函数:空建树
void build(int p, int l, int r) {
if (r - l == 1) {
info[p] = Info();
return;
}
int m = l + (r - l) / 2;
build(2 * p, l, m);
build(2 * p + 1, m, r);
pull(p);
}
// 辅助函数:带初始数组建树
void build(int p, int l, int r, const vector<long long> &arr) {
if (r - l == 1) {
info[p] = {arr[l]};
return;
}
int m = l + (r - l) / 2;
build(2 * p, l, m, arr);
build(2 * p + 1, m, r, arr);
pull(p);
}
// 单点修改:将 pos 位置的值修改为 v
// (如果是 "在原基础上增加 v",把 info[p].sum = v 改为 += 即可)
void modify(int p, int l, int r, int pos, long long v) {
if (r - l == 1) {
info[p].sum = v;
return;
}
int m = l + (r - l) / 2;
if (pos < m) {
modify(2 * p, l, m, pos, v);
} else {
modify(2 * p + 1, m, r, pos, v);
}
pull(p);
}
// 区间查询:返回 Info 结构体
Info rangeQuery(int p, int l, int r, int x, int y) {
// 越界返回单位元(sum 默认为 0)
if (l >= y || r <= x) {
return Info();
}
// 完全包含直接返回
if (l >= x && r <= y) {
return info[p];
}
int m = l + (r - l) / 2;
return rangeQuery(2 * p, l, m, x, y) + rangeQuery(2 * p + 1, m, r, x, y);
}
// --- 外部调用接口 ---
// 单点修改
void modify(int pos, long long v) {
modify(1, 0, n, pos, v);
}
// 区间查询(返回区间 [l, r) 的总和)
long long rangeQuery(int l, int r) {
return rangeQuery(1, 0, n, l, r).sum;
}
};
#include "bits/stdc++.h"
using namespace std;
struct Info {
int x = 0;
int sum = 1;
};
struct SegmentTree {
vector<Info> info;
int n;
SegmentTree(int size) {
n = size;
info.assign(4 * (n + 5), Info{});
build(1, 0, n);
}
void pull(int p) {
if (info[p].x == 0)
info[p].sum = 1;
else
info[p].sum = info[2 * p].sum + info[2 * p + 1].sum;
}
void build(int p, int l, int r) {
if (r - l == 1) {
info[p] = Info();
return;
}
int m = (l + r + 1) / 2;
build(2 * p, l, m);
build(2 * p + 1, m, r);
pull(p);
}
void modify(int p, int l, int r, int x, int y) {
if (l >= y || r <= x) {
return;
}
if (l >= x && r <= y) {
info[p].x = 1;
pull(p);
return;
}
int m = (l + r + 1) / 2;
modify(2 * p, l, m, x, y);
modify(2 * p + 1, m, r, x, y);
pull(p);
}
int rangeQuery(int p, int l, int r, int x, int y) {
if (l >= y || r <= x) {
return 0;
}
if (l >= x && r <= y) {
return info[p].sum;
}
int m = (l + r + 1) / 2;
return rangeQuery(2 * p, l, m, x, y) + rangeQuery(2 * p + 1, m, r, x, y) +
!info[p].x;
}
void modify(int x, int y) { modify(1, 0, n, x, y); }
int rangeQuery(int l, int r) { return rangeQuery(1, 0, n, l, r); }
};
懒标记线段树:
#include "bits/stdc++.h"
using namespace std;
struct Tag {
int tag = 0;
void apply(const Tag &t) { tag += t.tag; }
};
struct Info {
int x = -2e9;
void apply(const Tag &t) { x += t.tag; }
};
Info operator+(const Info &a, const Info &b) { return {max(a.x, b.x)}; };
struct SegmentTree {
vector<Info> info;
vector<Tag> tag;
int n;
// 选一个写
SegmentTree(int size) {
n = size;
info.assign(4 * (n + 5), Info{});
tag.assign(4 * (n + 5), Tag{});
build(1, 0, n);
}
void build(int p, int l, int r) {
if (r - l == 1) {
info[p] = Info();
return;
}
int m = (l + r) / 2;
build(2 * p, l, m);
build(2 * p + 1, m, r);
pull(p);
}
SegmentTree(const vector<Info> &arr) {
n = arr.size();
info.assign(4 * (n + 5), Info{});
tag.assign(4 * (n + 5), Tag{});
build(1, 0, n, arr);
}
void build(int p, int l, int r, const vector<Info> &arr) {
if (r - l == 1) {
info[p] = arr[l];
return;
}
int m = (l + r) / 2;
build(2 * p, l, m, arr);
build(2 * p + 1, m, r, arr);
pull(p);
}
void pull(int p) { info[p] = info[2 * p] + info[2 * p + 1]; }
void applyNode(int p, const Tag &v) {
info[p].apply(v);
tag[p].apply(v);
}
void push(int p) {
applyNode(2 * p, tag[p]);
applyNode(2 * p + 1, tag[p]);
tag[p] = Tag();
}
void modify(int p, int l, int r, int x, const Info &v) {
if (r - l == 1) {
info[p] = v;
return;
}
int m = (l + r) / 2;
push(p);
if (x < m) {
modify(2 * p, l, m, x, v);
} else {
modify(2 * p + 1, m, r, x, v);
}
pull(p);
}
Info rangeQuery(int p, int l, int r, int x, int y) {
if (l >= y || r <= x) {
return Info();
}
if (l >= x && r <= y) {
return info[p];
}
int m = (l + r) / 2;
push(p);
return rangeQuery(2 * p, l, m, x, y) + rangeQuery(2 * p + 1, m, r, x, y);
}
void rangeApply(int p, int l, int r, int x, int y, const Tag &v) {
if (l >= y || r <= x) {
return;
}
if (l >= x && r <= y) {
applyNode(p, v);
return;
}
int m = (l + r) / 2;
push(p);
rangeApply(2 * p, l, m, x, y, v);
rangeApply(2 * p + 1, m, r, x, y, v);
pull(p);
}
template <class F>
int findFirst(int p, int l, int r, int x, int y, F &&pred) {
if (l >= y || r <= x) {
return -1;
}
if (l >= x && r <= y && !pred(info[p])) {
return -1;
}
if (r - l == 1) {
return l;
}
int m = (l + r) / 2;
push(p);
int res = findFirst(2 * p, l, m, x, y, pred);
if (res == -1) {
res = findFirst(2 * p + 1, m, r, x, y, pred);
}
return res;
}
template <class F> int findLast(int p, int l, int r, int x, int y, F &&pred) {
if (l >= y || r <= x) {
return -1;
}
if (l >= x && r <= y && !pred(info[p])) {
return -1;
}
if (r - l == 1) {
return l;
}
int m = (l + r) / 2;
push(p);
int res = findLast(2 * p + 1, m, r, x, y, pred);
if (res == -1) {
res = findLast(2 * p, l, m, x, y, pred);
}
return res;
}
template <class F> int findLast(int l, int r, F &&pred) {
return findLast(1, 0, n, l, r, pred);
}
template <class F> int findFirst(int l, int r, F &&pred) {
return findFirst(1, 0, n, l, r, pred);
}
void modify(int x, const Info &v) { modify(1, 0, n, x, v); }
Info rangeQuery(int x, int y) { return rangeQuery(1, 0, n, x, y); }
void rangeApply(int x, int y, const Tag &v) { rangeApply(1, 0, n, x, y, v); }
};
int main(){
}
这里是极简版的线段树二分(findFirst / findLast)使用总结:
核心原理: 利用线段树节点维护的 Info 进行剪枝,时间内找到区间 [l, r) 中满足条件的点。
谓词函数 pred 怎么写:
传入 Lambda 表达式,判断当前区间有没有可能包含答案:
-
返回
true:这区间可能有答案,放行向下搜。 -
返回
false:这区间绝对没答案,直接剪枝(提速关键)。
极简代码模板(以寻找区间 [l, r) 第一个 >= k 的位置为例):
C++
// 1. 写条件:当前区间最大值 >= k,说明里面可能有我们要的数
auto pred = [&](const Info& info) {
return info.x >= k;
};
// 2. 调接口:找第一个用 findFirst,找最后一个用 findLast
int pos = seg.findFirst(l, r, pred);
// 3. 判无解:找不到一定会返回 -1,千万记得特判!
if (pos != -1) {
// 找到了,pos 是下标
}
三大避坑指南:
-
前提条件:你的
Info必须维护了能用于判断的属性(比如找 >= k,Info里必须有区间最大值)。 -
pred的逻辑:是判断“区间是否有潜力”,而不是“区间是否全满足”。 -
防 RE:一定要处理返回的
-1,别直接拿去当数组下标。
开点线段树:
#include "bits/stdc++.h"
using namespace std;
struct node{
int v;
int l;
int c;
};
const int N = 5e5 + 10;
vector<node> g[N];
template<class Info>
struct DynamicSegmentTree {
struct Node {
int ls = 0; // 左孩子下标
int rs = 0; // 右孩子下标
Info info;
};
int n;
vector<Node> nodes;
DynamicSegmentTree(int n_) : n(n_) {
nodes.push_back(Node{});
}
int newNode() {
nodes.emplace_back();
return (int)nodes.size() - 1;
}
// 上拉更新
void pull(int p) {
const Info &lhs = nodes[p].ls ? nodes[nodes[p].ls].info : Info();
const Info &rhs = nodes[p].rs ? nodes[nodes[p].rs].info : Info();
nodes[p].info = lhs + rhs;
}
//如果是新的树开一个点
//如果是旧树不变
int modify(int p, int l, int r, int x, const Info &v) {
if (!p) p = newNode(); // 动态开点
if (r - l == 1) {
nodes[p].info = v;
return p;
}
int m = l + (r - l) / 2;
// 移除 push(p)
if (x < m) {
int new_ls = modify(nodes[p].ls, l, m, x, v);
nodes[p].ls = new_ls;
} else {
int new_rs = modify(nodes[p].rs, m, r, x, v);
nodes[p].rs = new_rs;
}
pull(p); // 更新当前节点信息
return p;
}
// 外部接口:单点修改
void modify(int &root, int x, const Info &v) {
root = modify(root, 0, n, x, v);
}
// 区间查询
Info rangeQuery(int p, int l, int r, int x, int y) {
// 如果节点不存在,或者区间无交集,返回单位元
if (!p || l >= y || r <= x) {
return Info();
}
// 完全包含
if (l >= x && r <= y) {
return nodes[p].info;
}
int m = l + (r - l) / 2;
return rangeQuery(nodes[p].ls, l, m, x, y) + rangeQuery(nodes[p].rs, m, r, x, y);
}
// 外部接口:区间查询
Info rangeQuery(int root, int l, int r) {
return rangeQuery(root, 0, n, l, r);
}
// 线段树二分:查找区间 [x, y) 内第一个满足 pred 的位置
template<class F>
int findFirst(int p, int l, int r, int x, int y, F &&pred) {
if (!p || l >= y || r <= x) return -1;
// 剪枝:如果当前区间的信息都不满足条件,直接返回
if (l >= x && r <= y && !pred(nodes[p].info)) return -1;
if (r - l == 1) return l;
int m = l + (r - l) / 2;
// 移除 push(p)
int res = findFirst(nodes[p].ls, l, m, x, y, pred);
if (res == -1) res = findFirst(nodes[p].rs, m, r, x, y, pred);
return res;
}
template<class F>
int findFirst(int root, int l, int r, F &&pred) {
return findFirst(root, 0, n, l, r, pred);
}
// --- 你要求的 findLast ---
template<class F>
int findLast(int p, int l, int r, int x, int y, F &&pred) {
if (!p || l >= y || r <= x) return -1;
if (l >= x && r <= y && !pred(nodes[p].info)) return -1;
if (r - l == 1) return l;
int m = l + (r - l) / 2;
// 优先找右边
int res = findLast(nodes[p].rs, m, r, x, y, pred);
if (res == -1) res = findLast(nodes[p].ls, l, m, x, y, pred);
return res;
}
template<class F>
int findLast(int root, int l, int r, F &&pred) {
return findLast(root, 0, n, l, r, pred);
}
};
struct Info {
int mx = -1; // 默认为 0 或 -INF,视具体题目而定
};
// 重载 + 运算符用于合并
Info operator+(const Info &a, const Info &b) {
Info res;
res.mx = max(a.mx, b.mx);
return res;
}
开点懒标记线段树:
#include "bits/stdc++.h"
using namespace std;
struct node{
int v;
int l;
int c;
};
const int N = 5e5 + 10;
vector<node> g[N];
template<class Info>
struct DynamicSegmentTree {
struct Node {
int ls = 0; // 左孩子下标
int rs = 0; // 右孩子下标
Info info;
};
int n;
vector<Node> nodes;
DynamicSegmentTree(int n_) : n(n_) {
nodes.push_back(Node{});
}
int newNode() {
nodes.emplace_back();
return (int)nodes.size() - 1;
}
// 上拉更新
void pull(int p) {
const Info &lhs = nodes[p].ls ? nodes[nodes[p].ls].info : Info();
const Info &rhs = nodes[p].rs ? nodes[nodes[p].rs].info : Info();
nodes[p].info = lhs + rhs;
}
//如果是新的树开一个点
//如果是旧树不变
int modify(int p, int l, int r, int x, const Info &v) {
if (!p) p = newNode(); // 动态开点
if (r - l == 1) {
nodes[p].info = v;
return p;
}
int m = l + (r - l) / 2;
// 移除 push(p)
if (x < m) {
int new_ls = modify(nodes[p].ls, l, m, x, v);
nodes[p].ls = new_ls;
} else {
int new_rs = modify(nodes[p].rs, m, r, x, v);
nodes[p].rs = new_rs;
}
pull(p); // 更新当前节点信息
return p;
}
// 外部接口:单点修改
void modify(int &root, int x, const Info &v) {
root = modify(root, 0, n, x, v);
}
// 区间查询
Info rangeQuery(int p, int l, int r, int x, int y) {
// 如果节点不存在,或者区间无交集,返回单位元
if (!p || l >= y || r <= x) {
return Info();
}
// 完全包含
if (l >= x && r <= y) {
return nodes[p].info;
}
int m = l + (r - l) / 2;
return rangeQuery(nodes[p].ls, l, m, x, y) + rangeQuery(nodes[p].rs, m, r, x, y);
}
// 外部接口:区间查询
Info rangeQuery(int root, int l, int r) {
return rangeQuery(root, 0, n, l, r);
}
// 线段树二分:查找区间 [x, y) 内第一个满足 pred 的位置
template<class F>
int findFirst(int p, int l, int r, int x, int y, F &&pred) {
if (!p || l >= y || r <= x) return -1;
// 剪枝:如果当前区间的信息都不满足条件,直接返回
if (l >= x && r <= y && !pred(nodes[p].info)) return -1;
if (r - l == 1) return l;
int m = l + (r - l) / 2;
// 移除 push(p)
int res = findFirst(nodes[p].ls, l, m, x, y, pred);
if (res == -1) res = findFirst(nodes[p].rs, m, r, x, y, pred);
return res;
}
template<class F>
int findFirst(int root, int l, int r, F &&pred) {
return findFirst(root, 0, n, l, r, pred);
}
// --- 你要求的 findLast ---
template<class F>
int findLast(int p, int l, int r, int x, int y, F &&pred) {
if (!p || l >= y || r <= x) return -1;
if (l >= x && r <= y && !pred(nodes[p].info)) return -1;
if (r - l == 1) return l;
int m = l + (r - l) / 2;
// 优先找右边
int res = findLast(nodes[p].rs, m, r, x, y, pred);
if (res == -1) res = findLast(nodes[p].ls, l, m, x, y, pred);
return res;
}
template<class F>
int findLast(int root, int l, int r, F &&pred) {
return findLast(root, 0, n, l, r, pred);
}
};
struct Info {
int mx = -1; // 默认为 0 或 -INF,视具体题目而定
};
// 重载 + 运算符用于合并
Info operator+(const Info &a, const Info &b) {
Info res;
res.mx = max(a.mx, b.mx);
return res;
}

浙公网安备 33010602011771号