数据结构优化prim,boruvka
引入
k r u s k a l kruskal kruskal是最常见的最小生成树算法,但他有一个严峻的问题,就是总要遍历所有边,因此在完全图上复杂度至少是 O ( n 2 ) O(n^2) O(n2)的,对于完全图 M S T MST MST束手无策。
因此往往会使用数据结构优化 p r i m prim prim,或者 b o r u v k a boruvka boruvka来解决完全图 M S T MST MST问题。
数据结构优化prim
先说 p r i m prim prim,是类似 d i j k s t r a dijkstra dijkstra的思想,维护一个已选中的点集合,每次找到一个从点集合伸向剩余点的最短边,加入 M S T MST MST,这里我们只要能快速求出点集到剩余点的最短边即可,这可以用数据结构维护,不需要枚举全部边,因此可以取得较好的复杂度。
在稀疏图上,流程就是用堆对边排序,每次取出最小边,把端点加入集合,然后把新加入的这个点的出边加入堆。这样需要枚举全部边,对于完全图,还是会达到 O ( n 2 ) O(n^2) O(n2)以上的复杂度。
但实际上,对于一个点,他虽然可能有 O ( n ) O(n) O(n)条出边,但最小的出边其实是可以快速确定的,不必枚举这 O ( n ) O(n) O(n)个边全部放入堆中。比如下面这题
(CCPC Online 2025)C. 造桥与砍树
完全图上边权定义为 ( a i + a j ) m o d k (a_i+a_j)\mod k (ai+aj)modk
每次更新一个点时,确定了一个点 a i a_i ai,如果没有取模, a i + a j a_i+a_j ai+aj要取到最小值, a j a_j aj显然就是 a a a排序后的最小元素,加上取模,也就是把这个单增的函数变成了两段单增的分段函数,最小值对应的 a j a_j aj就是 a a a中最小的,或者 a i + a j a_i+a_j ai+aj第一个不小于 k k k的,或者说第一个不小于 k − a i k-a_i k−ai的。我们可以维护一个 s e t set set,保存还未被加入 M S T MST MST的所有点,那么第一种情况就取 s e t . b e g i n set.begin set.begin,第二种情况就 s e t set set上二分,只需 O ( log n ) O(\log n) O(logn)时间即可完成一个点的加边操作,整体复杂度只有 O ( n log n ) O(n\log n) O(nlogn)
更进一步,由于边权是取模,我们可以在开始就把所有 a i a_i ai对 k k k取模,那么实际上还有一个有趣的性质:我们上面说的两种情况的边权分别是, a i + a j , a i + a k − k a_i+a_j,a_i+a_k-k ai+aj,ai+ak−k,由于取模了 a k < k a_k<k ak<k,所以 a k − k < a j a_k-k<a_j ak−k<aj,所以 a i + a k − k < a i + a j a_i+a_k-k<a_i+a_j ai+ak−k<ai+aj恒成立,实际上只需要加入一条边:如果存在不小于 k − a i k-a_i k−ai的 a j a_j aj,则取第二种情况,否则取第一种情况
需要注意的是,如果我们把一条边 ( u , v ) (u,v) (u,v)加入生成树,那么 u , v u,v u,v长出的最小边都要更新,因此堆里对边排序时需要保存边权,和两个端点,这是与朴素 p r i m prim prim不同的地方。这是因为朴素 p r i m prim prim在 u u u加入生成树时,就把 u u u为端点的所有边都加入堆了,但这里只把最短边加入堆了,如果把这个最短边用了,需要再把新的最短边加入堆。
void solve() {
int n, k;
cin >> n >> k;
vi a(n + 1);
rep(i, 1, n) {
cin >> a[i];
}
multiset<int>s;
rep(i, 1, n) {
s.insert(a[i] % k);
}
auto get = [&](int x)->int{
auto it = s.lower_bound(k - x);
if (it != s.end()) {
return *it;
}
return *s.begin();
};
int ans = 0;
priority_queue<tri, vector<tri>, greater<tri>>q;
int u = *s.begin();
s.extract(u);
int v = get(u);
q.push({(u + v) % k, u, v});
while (q.size() && s.size()) {
auto t = q.top();
auto [w, u, v] = t;
q.pop();
if (!s.contains(v))continue;
s.extract(v);
ans += w;
if (s.size()) {
int t = get(u);
q.push({(u + t) % k, u, t});
t = get(v);
q.push({(v + t) % k, v, t});
}
}
cout << ans << '\n';
}
另一种数据结构优化是线段树,我们可以让线段树上的第 i i i个叶子保存, min w ( i , j ) \min w(i,j) minw(i,j), j j j是已在生成树中的点,这样我们查询线段树整个区间 [ 1 , n ] [1,n] [1,n],就能得到此时生成树点集,和其余点间的最短边,把这条边,不在生成树中的端点 i i i加入生成树;然后这个点 i i i也成为生成树中的点了,可以去更新所有点的 min w ( i , j ) \min w(i,j) minw(i,j),这体现为一个区间更新操作。
这个做法有极大的灵活性,因为我们在线段树上,更新和查询都可以只取一段区间,具体看下面这道题
2026牛客寒假算法基础集训营1 J-MST Problem
边权定义
w
(
i
,
j
)
=
a
i
+
a
j
w(i,j)=a_i+a_j
w(i,j)=ai+aj,完全图上删去
m
m
m个边。
难点在于删去 m m m个边,注意到对于每个点 u u u,和他有关的删除边 ( u , v ) (u,v) (u,v),这些 v v v把整个区间 [ 1 , n ] [1,n] [1,n]划分成了几个区间,可以分别对这些区间进行更新。由于删除的边一共只有 O ( m ) O(m) O(m)个,划分出的区间总数也是 O ( m ) O(m) O(m)的,只不过分散在不同的 u u u更新时,总之线段树更新的复杂度只有 O ( m log n ) O(m\log n) O(mlogn)。更新时,用当前加入 M S T MST MST的点权 a i a_i ai,去更新所有点的 min w ( i , j ) \min w(i,j) minw(i,j)
线段树的定义,肯定要包含每个点 i i i的最短边 w ( i , j ) w(i,j) w(i,j)的权值,以及这条边的另一个端点 j j j,因为我们更新时要把这个 w w w加入答案, j j j加入 M S T MST MST集合。除此之外,由于区间更新,还需要懒标记和下传,懒标记就是一个区间的最小 a i a_i ai, i i i是 M S T MST MST内的点。每次下传时,要利用懒标记 O ( 1 ) O(1) O(1)地确定一个区间的最短边,那么还需要这个区间内的,不在 M S T MST MST内的点的最小点权 a j a_j aj,以及点编号 j j j,这就是一个区间最小值和最小值下标,很好维护,然后把 a i + a j a_i+a_j ai+aj拼起来就是这个区间的最短边。
一个点加入 M S T MST MST后,就不能作为不在 M S T MST MST内的点,形成 w ( i , j ) w(i,j) w(i,j)了,所以把加入 M S T MST MST的点的点权, w ( i , j ) w(i,j) w(i,j)都置为 i n f inf inf,使其不影响后续的最小值查询。如果某一步查出来最小边权是 i n f inf inf,说明没有可以更新的边了, p r i m prim prim结束。
注意这个图可能是不连通的,所以 p r i m prim prim退出后需要检查是否连通,可以通过 a n s > = i n f ans>=inf ans>=inf来判断。
int a[N];
struct Tree {
#define ls u<<1
#define rs u<<1|1
struct Node {
int l, r;
pii mnv, mne;
int todo;
} tr[N << 2];
void pushup(int u) {
tr[u].mnv = min(tr[ls].mnv, tr[rs].mnv);
tr[u].mne = min(tr[ls].mne, tr[rs].mne);
}
void pushdown(int u) {
tr[ls].todo = min(tr[ls].todo, tr[u].todo);
tr[rs].todo = min(tr[rs].todo, tr[u].todo);
tr[ls].mne = min(tr[ls].mne, {tr[u].todo + tr[ls].mnv.fi, tr[ls].mnv.se});
tr[rs].mne = min(tr[rs].mne, {tr[u].todo + tr[rs].mnv.fi, tr[rs].mnv.se});
}
void build(int u, int l, int r) {
tr[u] = {l, r, {a[l], l}, {a[l] + inf, l}, inf};
if (l == r) {
return;
}
int mid = (l + r) >> 1;
build(ls, l, mid);
build(rs, mid + 1, r);
pushup(u);
}
void modify(int u, int l, int r, int val) {
if (tr[u].l >= l && tr[u].r <= r) {
tr[u].todo = min(tr[u].todo, val);
tr[u].mne = min(tr[u].mne, {tr[u].mnv.fi + val, tr[u].mnv.se});
return ;
} else {
int mid = (tr[u].l + tr[u].r) >> 1;
pushdown(u);
if (mid >= l) modify(ls, l, r, val);
if (r > mid) modify(rs, l, r, val);
pushup(u);
}
}
// ll query(int u, int l, int r) {
// if (l <= tr[u].l && tr[u].r <= r) return tr[u].mx;
// pushdown(u);
// int mid = (tr[u].l + tr[u].r) >> 1;
// if (r <= mid)return query(ls, l, r);
// if (l > mid)return query(rs, l, r);
// return query(ls, l, r) + query(rs, l, r);
// }
void del(int u, int idx) {
if (tr[u].l == tr[u].r) {
tr[u].mnv = {inf, idx};
tr[u].mne = {inf + a[idx], idx};
return ;
} else {
int mid = (tr[u].l + tr[u].r) >> 1;
pushdown(u);
if (idx <= mid) del(ls, idx);
else del(rs, idx);
pushup(u);
}
}
} t;
void solve() {
int n, m;
cin >> n >> m;
rep(i, 1, n) {
cin >> a[i];
}
vector<set<int>>s(n + 1);
vvi g(n + 1);
rep(i, 1, m) {
int u, v;
cin >> u >> v;
s[u].insert(v);
s[v].insert(u);
}
rep(i, 1, n) {
for (int x : s[i]) {
g[i].push_back(x);
}
}
t.build(1, 1, n);
int idx = 1;
int ans = 0;
rep(i, 1, n - 1) {
t.del(1, idx);
int l = 0;
for (int r : g[idx]) {
if (l + 1 <= r - 1) {
t.modify(1, l + 1, r - 1, a[idx]);
}
l = r;
}
if (l + 1 <= n) {
t.modify(1, l + 1, n, a[idx]);
}
ans += t.tr[1].mne.fi;
idx = t.tr[1].mne.se;
// cout<<t.tr[1].mne.fi<<' '<<t.tr[1].mne.se<<'\n';
if (ans >= inf)break;
}
if (ans >= inf) {
ans = -1;
}
cout << ans << '\n';
}
boruvka最小生成树
这是一种不太常见的最小生成树算法,结合了 p r i m prim prim和 k r u s k a l kruskal kruskal的优点,专为完全图 M S T MST MST设计。处理完全图的手段和数据结构优化 p r i m prim prim类似,也需要快速求出一个点集和其他点的最小边,如果理解了前面的数据结构优化 p r i m prim prim,理解起来估计不会有太大难度。
基本流程是:开始让每个点都成为一个连通块,接下来每一轮,对于每个联通快,维护块内块和块外点的最小边,然后连接这些边,合并连通块,直到只剩一个连通块。这样每一轮,最坏情况是连通块两两配对,连通块数量折半,实际还可能出现多个连通块在一轮里都连上,连通块减少的速度比折半还要快,所以至多运行 O ( log n ) O(\log n) O(logn)轮
由于运行轮数不太多,每一轮内的操作就可以暴力点,可以去枚举所有点,然后计算这个点和不在同一个点内点的最小边,然后更新这个点所在连通块的最小边,通常要保证每个点的这个查询复杂度在 O ( log n ) O(\log n) O(logn),每轮复杂度不超过 O ( n log n ) O(n\log n) O(nlogn),整体复杂度 O ( n log 2 n ) O(n\log ^2n) O(nlog2n)。
说这个算法是结合了 p r i m , k r u s k a l prim,kruskal prim,kruskal,是因为既有 p r i m prim prim的,对于连通块找最短边,也有 k r u s k a l kruskal kruskal的合并连通块。
还是对于上面的两道题,对于第一个,可以维护一个包含所有点的 s e t set set,然后枚举每个连通块,把这个连通块内的点临时从 s e t set set里删掉,然后就可以枚举连通块内所有点,通过和前面 p r i m prim prim类似的方式,在集合里查最小元素/二分大于 a i − k a_i-k ai−k的最小元素,来找到这个连通块的最小边。完成一个连通块的操作后,再把刚刚删除的点加入 s e t set set
这里在实现上有几个注意点:
- 连通块合并,并不需要启发式合并之类的,而是可以在每一轮开始时,根据每个点在并查集里的祖先,临时计算每个联通块内的点
- 退出条件也可以是:计算这一轮加入 M S T MST MST的边数,如果为0则退出,这是更好的方式,因为可能原图有一些限制导致不连通,用连通块个数是否为1可能出问题
- 一个连通块可能找不到边,所以要看边 i d id id是否合法,合法才加入,这里 m n . s e ( i d ) = 0 mn.se(id)=0 mn.se(id)=0说明没有可以加入的边
- 这里是对于每个联通块,现场计算它的最小边,现场加入合并连通块,还有一种实现方式是先计算每个联通快的最小边,再逐个合并。现在这个实现是更好的,因为先计算,再合并,可能出现一个边合并后,使得其它的连通块最小边,两个端点在一个连通块内了,也就失效了,导致合并变慢。
void solve() {
int n, k;
cin >> n >> k;
vi a(n + 1);
rep(i, 1, n) {
cin >> a[i];
a[i] %= k;
}
int ans = 0;
set<pii>all;
vvi s(n + 1);
rep(i, 1, n) {
f[i] = i;
sz[i] = 1;
all.insert({a[i], i});
}
auto get = [&](int x)->pii{
auto it = all.lower_bound({x, -1});
if (it == all.end()) {
return *all.begin();
} else {
return *it;
}
};
while (1) {
vvi s(n + 1);
rep(i, 1, n) {
s[find(i)].push_back(i);
}
int cnt = 0;
rep(i, 1, n) {
if (find(i) == i) {
for (int x : s[i]) {
all.erase({a[x], x});
}
pii mn = {inf, 0};
for (int x : s[i]) {
if (all.size()) {
auto [v, id] = get(k - a[x]);
mn = min(mn, {(a[x] + v) % k, id});
}
}
for (int x : s[i]) {
all.insert({a[x], x});
}
if (mn.se && find(i) != find(mn.se)) {
++cnt;
ans += mn.fi;
merge(mn.se, i);
}
}
}
if (!cnt)break;
}
cout << ans << '\n';
}
对于前面的第二题,也可以 b o r u v k a boruvka boruvka。完全图上删除一些边,可以和第一题类似的处理,在枚举一个连通块时,不止在全局 s e t set set里把这个连通块内的点删除,还把和这个连通块内点有关删除边的端点也删除,然后查询 s e t set set最小即可,整体框架和第一题类似,复杂度也是 O ( n log 2 n ) O(n\log ^2n) O(nlog2n)
需要注意的是,删除的边,可能两个端点都在当前连通块内,所以临时删掉边的端点时,和恢复端点时,都需要检查点是否在当前连通块内。否则会 w a wa wa,可能是因为把一些不该加的点加入 s e t set set了,导致错误的查询结果
int f[N], sz[N];
int find(int x) {
if (f[x] == x)return x;
return f[x] = find(f[x]);
}
void merge(int x, int y) {
x = find(x), y = find(y);
if (sz[x] > sz[y])swap(x, y);
f[x] = y;
sz[y] += sz[x];
}
void solve() {
int n, m;
cin >> n >> m;
vi a(n + 1);
rep(i, 1, n) {
cin >> a[i];
}
vvi g(n + 1);
rep(i, 1, m) {
int u, v;
cin >> u >> v;
g[u].push_back(v);
g[v].push_back(u);
}
int ans = 0;
set<pii>all;
vvi s(n + 1);
rep(i, 1, n) {
f[i] = i;
sz[i] = 1;
all.insert({a[i], i});
}
while (1) {
vvi s(n + 1);
rep(i, 1, n) {
s[find(i)].push_back(i);
}
int cnt = 0;
rep(i, 1, n) {
if (find(i) == i) {
for (int x : s[i]) {
all.erase({a[x], x});
}
pii mn = {inf, 0};
for (int x : s[i]) {
for (int y : g[x]) {
if (find(x) == find(y))continue;
all.erase({a[y], y});
}
if (all.size()) {
auto [v, id] = *all.begin();
mn = min(mn, {a[x] + v, id});
}
for (int y : g[x]) {
if (find(x) == find(y))continue;
all.insert({a[y], y});
}
}
for (int x : s[i]) {
all.insert({a[x], x});
}
if (mn.se && find(i) != find(mn.se)) {
++cnt;
ans += mn.fi;
merge(mn.se, i);
}
}
}
if (!cnt)break;
}
rep(i, 1, n) {
if (find(i) != find(1)) {
cout << -1 << '\n';
return;
}
}
cout << ans << '\n';
}
更多推荐
所有评论(0)