按照 OI-Wiki 的第一种建虚树的方法,大佬看看。
#include<bits/stdc++.h>
using namespace std;
const int N = 250005;
int len , n, x, y, z, fir[N], nxt[N << 1], son[N << 1], tot, w[N << 1], m, sum, a[N << 1], top[N], sz[N], hs[N], fir2[N], nxt2[N << 1], son2[N << 1], w2[N << 1], dis[N], lca, dep[N], b[N];
int f[N], dfsn[N], fa[N][19], L[N][19], LL;
//f[i] ; i 不和 i 的子树中任意一个节点联通的最小值
bool vis[N];
inline void add(int x, int y, int z){
nxt[++tot] = fir[x];
fir[x] = tot;
son[tot] = y;
w[tot] = z;
}
inline void add2(int x, int y, int z){
nxt2[++tot] = fir2[x];
fir2[x] = tot;
son2[tot] = y;
w2[tot] = z;
}
inline void dfs(int x, int ff){
sz[x] = 1, dep[x] = dep[ff] + 1;
fa[x][0] = ff, dfsn[x] = ++tot;
for(int i = 1; i <= 18; i++){
L[x][i] = min(L[x][i - 1], L[fa[x][i - 1]][i - 1]);
fa[x][i] = fa[fa[x][i - 1]][i - 1];
}
for(int i = fir[x]; i; i = nxt[i]){
if(son[i] == ff) continue;
dis[son[i]] = dis[x] + w[i];
L[son[i]][0] = w[i];
dfs(son[i], x);
sz[x] += sz[son[i]];
if(sz[son[i]] > sz[hs[x]]) hs[x] = son[i];
}
}
inline bool cmp(int x, int y){
return dfsn[x] < dfsn[y];
}
inline void LCA(int x, int y){
LL = 9e18;
if(dep[x] > dep[y]) swap(x, y);
for(int i = 18; ~i; i--){
if(dep[fa[y][i]] >= dep[x]) LL = min(LL, L[y][i]), y = fa[y][i];
}
if(x == y){
lca = x;
return ;
}
for(int i = 18; ~i; i--){
if(fa[y][i] != fa[x][i]){
x = fa[x][i];
y = fa[y][i];
}
}
lca = fa[x][0];
}
inline void dp(int x){
for(int i = fir2[x]; i ; i = nxt2[i]){
dp(son2[i]);
if(vis[son2[i]]) f[x] += w2[i];
else f[x] += min(w2[i], f[son2[i]]);
}
}
int main(){
scanf("%lld", &n);
for(int i = 1; i < n; i++){
scanf("%lld%lld%lld", &x, &y, &z);
add(x, y, z), add(y, x, z);
}
tot = 0;
dfs(1, 0);
scanf("%lld", &m);
tot = 0;
while(m--){
scanf("%lld", &sum);
for(int i = 1; i <= sum; i++) scanf("%lld", &b[i]), vis[b[i]] = 1;
sort(b + 1, b + sum + 1, cmp);
for(int i = 1; i < sum; i++){
a[++len] = b[i];
LCA(b[i], b[i + 1]);
a[++len] = lca;
}
a[++len] = b[sum];
a[++len] = 1;
sort(a + 1, a + len + 1, cmp);
len = unique(a + 1, a + len + 1) - a - 1;
for(int i = 1; i < len; i++){
LCA(a[i], a[i + 1]);
LCA(lca, a[i + 1]);
add2(lca, a[i + 1], LL);
}
dp(1);
printf("%lld\n", f[1]);
for(int i = 1; i <= sum; i++) vis[b[i]] = 0;
for(int i = 1; i <= len; i++) f[a[i]] = 0, fir2[a[i]] = 0;
len = 0, tot = 0;
}
return 0;
}