引入

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';
}
Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐