答案偏小。。
#include <bits/stdc++.h>
const int N = 250000 + 5;
const long long inf = 1ll * 1e18;
int head[N], nxt[N << 1], ver[N << 1], tot = 0, n;
std::vector<int>VT[N];
int anc[20][N], dep[N], m, k;
int dfn[N], dsum = 0, key[N];
long long len[N << 1], mind[20][N];
inline bool cmp(int x, int y){
return dfn[x] < dfn[y];
}
inline int lca(int x, int y){
if(dep[x] > dep[y]) std::swap(x, y);
int D = dep[y] - dep[x];
for(int i = 0; i <= 19; ++i) if(D & (1 << i)) y = anc[i][y];
for(int i = 0; i <= 19; ++i) if(anc[i][x] ^ anc[i][y])
x = anc[i][x], y = anc[i][y];
return x == y ? x : anc[0][x];
}
inline long long getmin(int x, int y){
int D = dep[x] - dep[y];
long long res = inf;
for(int i = 0; i <= 19; ++i) if(D & (1 << i))
res = std::min(res, mind[i][x]),
x = anc[i][x];
return res;
}
void dfs1(int x){
dfn[x] = ++dsum;
for(int i = 1; i <= 19; ++i)
anc[i][x] = anc[i - 1][anc[i - 1][x]],
mind[i][x] = std::min(mind[i - 1][x], mind[i - 1][anc[i - 1][x]]);
for(int i = head[x]; i; i = nxt[i]){
if(ver[i] == anc[0][x]) continue;
dep[ver[i]] = dep[x] + 1;
anc[0][ver[i]] = x;
mind[0][ver[i]] = len[i];
dfs1(ver[i]);
}
return ;
}
int q[250005];
inline void adde(int x, int y, long long z){
nxt[++tot] = head[x];
head[x] = tot;
ver[tot] = y;
len[tot] = z;
return ;
}
int tp, stk[N];
inline void addedge(int x, int y){
VT[x].push_back(y);
return ;
}
void build(){
tp = 0;
std::sort(q + 1, q + k + 1, cmp);
stk[++tp] = 1;
VT[1].clear();
for(int i = 1; i <= k; ++i){
int Lca = lca(q[i], stk[tp]);
if(Lca ^ stk[tp]){
while(dfn[Lca] < dfn[stk[tp - 1]]){
addedge(stk[tp - 1], stk[tp]);
--tp;
}
if(dfn[Lca] > dfn[stk[tp - 1]]){
VT[Lca].clear();
addedge(Lca, stk[tp]);
--tp;
stk[++tp] = Lca;
}
else addedge(Lca, stk[tp]), --tp;
}
VT[q[i]].clear();
stk[++tp] = q[i];
}
for(int i = 1; i <= tp - 1; ++i)
addedge(stk[i], stk[i + 1]);
return ;
}
long long dp[N];
void dfs2(int x){
int S = VT[x].size();
dp[x] = 0ll;
for(int i = 0; i < S; ++i){
int v = VT[x][i];
dfs2(v);
if(key[v]) dp[x] += getmin(v, x);
else dp[x] += std::min(getmin(v, x), dp[v]);
}
return ;
}
int main(){
freopen("4.in", "r", stdin);
scanf("%d", &n);
int x, y;
long long z;
for(int i = 1; i <= n - 1; ++i){
scanf("%d%d%lld", &x, &y, &z);
adde(x, y, z);
adde(y, x, z);
}
for(int i = 0; i <= 19; ++i)
for(int j = 0; j <= n; ++j)
mind[i][j] = inf;
dep[1] = 1;
dfs1(1);
scanf("%d", &m);
for(int i = 1; i <= m; ++i){
scanf("%d", &k);
for(int j = 1; j <= k; ++j){
scanf("%d", &x);
q[j] = x;
key[x] = 1;
}
build();
dfs2(1);
printf("%lld\n", dp[1]);
for(int j = 1; j <= k; ++j) key[q[j]] = 0;
}
return 0;
}