P4180 严格次小生成树 - 洛谷

题意

要找一棵“严格大于最小生成树、且权值和最小”的生成树

思路

核心策略是“非树边替换树边”:

  1. 先用 Kruskal 算法求出最小生成树,记录其权值总和为 sum
  2. 遍历每一条不在最小生成树中的非树边 $e(u, v, w)$。如果把这条边强行加到树中,树上就会形成一个环(即 $u$ 到 $v$ 的树上路径加上这条新边)。
  3. 为了重新变成一棵树,我们需要在这个环里删掉一条原本就在树上的边
  4. 为了让新的树权值总和增加得最少(且严格大于 sum),我们需要在 $u$ 到 $v$ 的树上路径中,找到权值最大、且严格小于 $w$ 的那条树边替换掉。

我使用了Kruskal来求最小生成树,倍增lca来查找 $u$ 到 $v$ 的树上路径中的替换边

分块代码

看完思路后,我们一步一步来分块拆解代码

1.Kruskal 求最小生成树

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
struct Edge{
int u,v,w;
bool vis;//记录是不是树边
bool operator<(const Edge& other) const {
return w < other.w;
}
};
void kru()
{
int cnt=0;
//edg存储所有边,按照权值 w 从小到大排序
sort(edg,edg+m);
//并查集初始化
for(int i=1;i<=n;++i) fa[i]=i;
for(int i=0;i<m;++i)
{
auto& [u,v,w,vis]=edg[i];
int ru=find(u),rv=find(v);
if(ru!=rv)
{
fa[ru]=rv;
cnt++;
sum+=w;
// 标记这条边是树边
vis=1;
// 建立树的邻接表 adj
adj[v].push_back({u,w});
adj[u].push_back({v,w});
if(cnt==n-1) break;
}
else vis=0;
}
}

2.树上倍增预处理

为了高效查询树上 $u$ 到 $v$ 路径间的最大边权次大边权,使用 LCA(最近公共祖先)的倍增算法

  • up[u][i]:表示节点 u 向上跳 $2^i$ 步到达的祖先。
  • w1[u][i]:表示从 u 向上到 up[u][i] 的路径上的最大边权
  • w2[u][i]:表示该路径上的次大边权

其中dfs中它们被动态维护

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
void dfs(int u,int p,int ww)
{
dep[u]=dep[p]+1;
up[u][0]=p;
w1[u][0]=ww,w2[u][0]=-inf;
for(int i=1;i<=19;++i)
{
up[u][i]=up[up[u][i-1]][i-1];
// 路径最大值,显然是两段各自最大值中的较大者
w1[u][i]=max(w1[u][i-1],w1[up[u][i-1]][i-1]);
//路径次大值,情况稍微复杂一点
// 先从两段的次大值里挑个大的
w2[u][i]=max(w2[u][i-1],w2[up[u][i-1]][i-1]);
// 如果两段的最大值不一样,那么其中较小的那个最大值,是整段的“次大值”
if(w1[u][i-1]!=w1[up[u][i-1]][i-1])
{
w2[u][i]=max(w2[u][i],min(w1[u][i-1],w1[up[u][i-1]][i-1]));
}
}
for(auto [vv,ww]:adj[u])
{
if(vv==p) continue;
dfs(vv,u,ww);
}
}

3.查询与替换函数

当要尝试加入一条非树边 $(u, v, w)$ 时,我们要去捞出 $u$ 到 $v$ 路径上符合要求的最大替换边
也就是$u->l$ 和$l->v$这个路径上面的最大替换边

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
int F(int u,int v,int w)
{
int res=-inf;
int l=lca(u,v);
//u->l
for(int i=19;i>=0;--i)
{
//向上倍增跳跃
if(up[u][i]&&dep[up[u][i]]>=dep[l])
{
// 如果最大树边小于 w,直接尝试用它
if(w1[u][i]<w)
{
res=max(res,w1[u][i]);
}
else res=max(res,w2[u][i]);// 如果最大树边等于 w,只能退而求其次用次大树边
u=up[u][i];
}
}
//l->v就是v->l
for(int i=19;i>=0;--i)
{
if(up[v][i]&&dep[up[v][i]]>=dep[l])
{
if(w1[v][i]<w)
{
res=max(res,w1[v][i]);
}
else res=max(res,w2[v][i]);
v=up[v][i];
}
}
return res;
}

4.最后枚举非树边寻找答案

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
ll ans=1e18;
for(int i=0;i<m;++i)
{
auto [u,v,w,vis]=edg[i];
if(!vis&&u!=v)
{
//新树的总权值=原来的sum-被删掉的树边val+新加入的非树边w
int val=F(u,v,w);
if(val>-inf)
{
ans=min(ans,sum-val+w);
}
}
}
cout<<ans<<'\n';

最终代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int N=1e5+5,M=3e5+5,inf=2e9;
struct Edge{
int u,v,w;
bool vis;
bool operator<(const Edge& other) const {
return w < other.w;
}
};
struct Edge1{
int v,w;
};
int fa[N];
int find(int x)
{
if(fa[x]==x) return x;
return fa[x]=find(fa[x]);
}
vector<Edge1> adj[N];
int up[N][20],dep[N],w1[N][20],w2[N][20];
Edge edg[M];
ll sum=0;
int n,m;
void kru()
{
int cnt=0;
sort(edg,edg+m);
for(int i=1;i<=n;++i) fa[i]=i;
for(int i=0;i<m;++i)
{
auto& [u,v,w,vis]=edg[i];
int ru=find(u),rv=find(v);
if(ru!=rv)
{
fa[ru]=rv;
cnt++;
sum+=w;
vis=1;
adj[v].push_back({u,w});
adj[u].push_back({v,w});
if(cnt==n-1) break;
}
else vis=0;
}
}
void dfs(int u,int p,int ww)
{
dep[u]=dep[p]+1;
up[u][0]=p;
w1[u][0]=ww,w2[u][0]=-inf;
for(int i=1;i<=19;++i)
{
up[u][i]=up[up[u][i-1]][i-1];
w1[u][i]=max(w1[u][i-1],w1[up[u][i-1]][i-1]);
w2[u][i]=max(w2[u][i-1],w2[up[u][i-1]][i-1]);
if(w1[u][i-1]!=w1[up[u][i-1]][i-1])
{
w2[u][i]=max(w2[u][i],min(w1[u][i-1],w1[up[u][i-1]][i-1]));
}
}
for(auto [vv,ww]:adj[u])
{
if(vv==p) continue;
dfs(vv,u,ww);
}
}
int lca(int u,int v)
{
if(dep[v]>dep[u]) swap(u,v);
for(int i=19;i>=0;--i)
{
if(dep[up[u][i]]>=dep[v])
{
u=up[u][i];
}
}
if(u==v)
{
return u;
}
for(int i=19;i>=0;--i)
{
if(up[u][i]!=up[v][i])
{
u=up[u][i];
v=up[v][i];
}
}
return up[u][0];
}
int F(int u,int v,int w)
{
int res=-inf;
int l=lca(u,v);
for(int i=19;i>=0;--i)
{
if(up[u][i]&&dep[up[u][i]]>=dep[l])
{
if(w1[u][i]<w)
{
res=max(res,w1[u][i]);
}
else res=max(res,w2[u][i]);
u=up[u][i];
}
}
for(int i=19;i>=0;--i)
{
if(up[v][i]&&dep[up[v][i]]>=dep[l])
{
if(w1[v][i]<w)
{
res=max(res,w1[v][i]);
}
else res=max(res,w2[v][i]);
v=up[v][i];
}
}
return res;
}
void solve()
{
cin>>n>>m;
for(int i=0;i<m;++i)
{
int x,y,z;
cin>>x>>y>>z;
edg[i]={x,y,z};
}
kru();
dfs(1,0,-inf);
ll ans=1e18;
for(int i=0;i<m;++i)
{
auto [u,v,w,vis]=edg[i];
if(!vis&&u!=v)
{
int val=F(u,v,w);
if(val>-inf)
{
ans=min(ans,sum-val+w);
}
}
}
cout<<ans<<'\n';
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int T=1;
//cin>>T;
while(T--)
{
solve();
}
return 0;
}

后记

有人问既然是在$u$到$v$的路径,那么能不能直接在倍增lca的过程中直接求res呢
当然可以!!

LCA代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
int lca(int u,int v,int w)
{
int res=-inf;
if(dep[v]>dep[u]) swap(u,v);
for(int i=19;i>=0;--i)
{
if(dep[up[u][i]]>=dep[v])
{
// 如果最大树边小于 w,直接尝试用它
if(w1[u][i]<w) res=max(res,w1[u][i]);
else res=max(res,w2[u][i]);// 如果最大树边等于 w,只能退而求其次用次大树边
u=up[u][i];
}
}
if(u==v)
{
return res;
}
for(int i=19;i>=0;--i)
{
if(up[u][i]!=up[v][i])
{
if(w1[u][i]<w)
res=max(res,w1[u][i]);
else res=max(res,w2[u][i]);
if(w1[v][i]<w)
res=max(res,w1[v][i]);
else res=max(res,w2[v][i]);
u=up[u][i];
v=up[v][i];
}
}
//最后需要判断一下up[u][0]和up[v][0]
if(w1[u][0]<w)
res=max(res,w1[u][0]);
else res=max(res,w2[u][0]);
if(w1[v][0]<w)
res=max(res,w1[v][0]);
else res=max(res,w2[v][0]);
return res;
}

总代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int N=1e5+5,M=3e5+5,inf=2e9;
struct Edge{
int u,v,w;
bool vis;
bool operator<(const Edge& other) const {
return w < other.w;
}
};
struct Edge1{
int v,w;
};
int fa[N];
int find(int x)
{
if(fa[x]==x) return x;
return fa[x]=find(fa[x]);
}
vector<Edge1> adj[N];
int up[N][20],dep[N],w1[N][20],w2[N][20];
Edge edg[M];
ll sum=0;
int n,m;
void kru()
{
int cnt=0;
sort(edg,edg+m);
for(int i=1;i<=n;++i) fa[i]=i;
for(int i=0;i<m;++i)
{
auto& [u,v,w,vis]=edg[i];
int ru=find(u),rv=find(v);
if(ru!=rv)
{
fa[ru]=rv;
cnt++;
sum+=w;
vis=1;
adj[v].push_back({u,w});
adj[u].push_back({v,w});
if(cnt==n-1) break;
}
else vis=0;
}
}
void dfs(int u,int p,int ww)
{
dep[u]=dep[p]+1;
up[u][0]=p;
w1[u][0]=ww,w2[u][0]=-inf;
for(int i=1;i<=19;++i)
{
up[u][i]=up[up[u][i-1]][i-1];
w1[u][i]=max(w1[u][i-1],w1[up[u][i-1]][i-1]);
w2[u][i]=max(w2[u][i-1],w2[up[u][i-1]][i-1]);
if(w1[u][i-1]!=w1[up[u][i-1]][i-1])
{
w2[u][i]=max(w2[u][i],min(w1[u][i-1],w1[up[u][i-1]][i-1]));
}
}
for(auto [vv,ww]:adj[u])
{
if(vv==p) continue;
dfs(vv,u,ww);
}
}
int lca(int u,int v,int w)
{
int res=-inf;
if(dep[v]>dep[u]) swap(u,v);
for(int i=19;i>=0;--i)
{
if(dep[up[u][i]]>=dep[v])
{
if(w1[u][i]<w) res=max(res,w1[u][i]);
else res=max(res,w2[u][i]);
u=up[u][i];
}
}
if(u==v)
{
return res;
}
for(int i=19;i>=0;--i)
{
if(up[u][i]!=up[v][i])
{
if(w1[u][i]<w)
res=max(res,w1[u][i]);
else res=max(res,w2[u][i]);
if(w1[v][i]<w)
res=max(res,w1[v][i]);
else res=max(res,w2[v][i]);
u=up[u][i];
v=up[v][i];
}
}
if(w1[u][0]<w)
res=max(res,w1[u][0]);
else res=max(res,w2[u][0]);
if(w1[v][0]<w)
res=max(res,w1[v][0]);
else res=max(res,w2[v][0]);
return res;
}
void solve()
{
cin>>n>>m;
for(int i=0;i<m;++i)
{
int x,y,z;
cin>>x>>y>>z;
edg[i]={x,y,z};
}
kru();
dfs(1,0,-inf);
ll ans=1e18;
for(int i=0;i<m;++i)
{
auto [u,v,w,vis]=edg[i];
if(!vis&&u!=v)
{
int val=lca(u,v,w);
if(val>-inf)
{
ans=min(ans,sum-val+w);
}
}
}
cout<<ans<<'\n';
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int T=1;
while(T--)
{
solve();
}
return 0;
}