信竞学习笔记:最小生成树
注:这是本人的课堂笔记。
模板题:P3366
Prim 算法
Prim 算法的过程和代码类似于求最短路的 Dijkstra 算法。具体地说,Prim 算法维护一个 dis 数组,代表每个节点到当前已经求得的部分生成树的距离。将 1 号点放入生成树中打上标记,然后更新 dis 数组。随后,每次选择所有未被标记的节点中 dis 最小的点并打上标记,并更新数组 dis。重复以上过程,直至所有点都被标记(求得最小生成树)或所有未被标记的点都不存在与已标记的点相连的边(图不连通)。可以通过统计打上标记的点的数量来判断是否有解。
和 Dijkstra 算法类似,Prim 算法同样存在堆优化版本。将 dis 数组及其对应的节点编号放入一个小根堆中,依据 dis 排序。每次取出 dis 最小的点,随后进行和上述基本相同的过程即可。注意,当一个点的 dis 数组被修改时,就需要将这个节点重新压入。时间复杂度 O(m \log n)。完整代码如下:
#include <bits/stdc++.h>
using namespace std;
const int N = 5010;
int n, m, ans;
int tot, dis[N];
bool vis[N];
struct node {
int pos, dis;
bool operator<(const node &other) const {
return dis > other.dis;
}
};
priority_queue<node> q;
struct edge {
int v, w;
};
vector<edge> g[N];
void prim() {
memset(dis, 0x3f, sizeof dis);
dis[1] = 0;
q.push({1, 0});
while (!q.empty()) {
node p = q.top();
q.pop();
int u = p.pos;
if (vis[u]) continue;
tot++;
vis[u] = 1;
ans += dis[u];
for (const edge &e : g[u]) {
int v = e.v, w = e.w;
if (dis[v] > w) {
dis[v] = w;
q.push({v, w});
}
}
}
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr); cout.tie(nullptr);
cin >> n >> m;
for (int i=1; i<=m; i++) {
int u, v, w;
cin >> u >> v >> w;
g[u].push_back({v, w});
g[v].push_back({u, w});
}
prim();
if (tot == n) cout << ans;
else cout << "orz";
return 0;
}
Kruskal 算法
Kruskal 算法使用并查集实现,具体地说,我们将所有边的起点、终点和边权打包存储起来,并以边权为依据从小到大的排序。我们依次加入需要加入的边,即加入两端没有连通的边,使得图连通。用并查集维护连通性集合,如果一条边的两个端点在一个集合中,则不用加入这条边;否则要加入这条边。如果加入的边数(用一个变量统计)等于 n-1,代表求得了最小生成树;否则代表图不连通。时间复杂度 O(m \log m)。完整代码如下:
#include <bits/stdc++.h>
using namespace std;
const int N = 5010, M = 2e5 + 100;
struct edge {
int u, v, w;
} e[M];
int n, m, ans, cnt, p[N];
int find(int x) {
if (x != p[x]) p[x] = find(p[x]);
return p[x];
}
void kruskal() {
sort(e+1, e+m+1, [](const edge &A, const edge &B) {
return A.w < B.w;
});
for (int i=1; i<=m; i++) {
int u = find(e[i].u);
int v = find(e[i].v);
if (u == v) continue;
ans += e[i].w;
p[v] = u;
cnt++;
}
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr); cout.tie(nullptr);
cin >> n >> m;
for (int i=1; i<=n; i++) p[i] = i;
for (int i=1; i<=m; i++) cin >> e[i].u >> e[i].v >> e[i].w;
kruskal();
if (cnt == n - 1) cout << ans;
else cout << "orz";
return 0;
}