目录

信竞学习笔记:最小生成树

发表于
更新于
4 1.8~2.3 分钟 796

注:这是本人的课堂笔记。

模板题: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;
}