这份代码开了 O2 之后 T 了第 10 个点, 但不开 O2 就过了,请各位大佬帮忙找错
#include <bits/stdc++.h>
using namespace std;
#define gc() (p1 == p2 && (p2 = (p1 = buf) + fread(buf, 1, MAXSIZE, stdin), p1 == p2) ? EOF : *p1++)
const int MAXSIZE = 1 << 20;
char buf[MAXSIZE], *p1, *p2;
void read() {}
template <class T1, class ...T2>
void read(T1& ret,T2&... rest) {
ret = 0; char c; bool f = false;
while (!isdigit(c = gc() ) ) f = c == '-';
while(isdigit(c) ) {
ret = (ret << 3) + (ret << 1) + (c ^ '0');
c = gc();
}
if(f) ret = -ret;
read(rest...);
}
const int N = 250010;
struct edge {
int to, nxt, w;
}e[N << 1];
int head[N];
int cnt = 1;
inline void add(int x, int y, int z) {
e[++cnt] = {y, head[x], z};
head[x] = cnt;
}
vector<int> v[N];
long long dep[N], fa[N][23], mi[N][23], dfn[N];
int a[N], tmp[N];
int k, tot;
bool vis[N];
void init(int x) {
dfn[x] = ++tot;
for(int i = 1; i <= 22; i++) {
fa[x][i] = fa[fa[x][i - 1]][i - 1];
mi[x][i] = min(mi[x][i - 1], mi[fa[x][i - 1]][i - 1]);
}
for(int i = head[x]; i; i = e[i].nxt) {
int y = e[i].to, z = e[i].w;
if(y == fa[x][0]) continue;
dep[y] = dep[x] + 1;
fa[y][0] = x;
mi[y][0] = z;
init(y);
}
}
int get_lca(int x, int y) {
if(dep[x] < dep[y]) swap(x, y);
for(int i = 22; i >= 0; i--) if(dep[fa[x][i]] >= dep[y]) x = fa[x][i];
if(x == y) return x;
for(int i = 22; i >= 0; i--) if(fa[x][i] != fa[y][i]) x = fa[x][i], y = fa[y][i];
return fa[x][0];
}
long long get_len(int x, int y) {
int lca = get_lca(x, y);
long long ret = 0x3f3f3f3f3f3f3f3fll;
for(int i = 22; i >= 0; i--) if(dep[fa[x][i]] >= dep[lca]) ret = min(ret, mi[x][i]), x = fa[x][i];
for(int i = 22; i >= 0; i--) if(dep[fa[y][i]] >= dep[lca]) ret = min(ret, mi[y][i]), y = fa[y][i];
return ret;
}
inline bool cmp(int x, int y) {
return dfn[x] < dfn[y];
}
long long dfs(int x) {
long long ret = 0;
for(auto y : v[x]) ret += dfs(y);
if(vis[x]) ret = get_len(1, x);
else ret = min(get_len(1, x), ret);
vis[x] = 0;
v[x].clear();
return ret;
}
int main() {
int n;
read(n);
for(int i = 1, x, y, z; i < n; i++) read(x, y, z), add(x, y, z), add(y, x, z);
memset(mi, 0x3f, sizeof(mi));
dep[1] = 1;
init(1);
int m;
read(m);
while(m--) {
read(k);
for(int i = 1; i <= k; i++) read(a[i]), vis[a[i]] = 1;
sort(a + 1, a + 1 + k, cmp);
a[0] = 1;
for(int i = 0, tmp = k; i <= tmp - 1; i++) a[++k] = get_lca(a[i], a[i + 1]);
sort(a + 1, a + 1 + k, cmp);
k = unique(a + 1, a + 1 + k) - a - 1;
for(int i = 1; i < k; i++) {
int lca = get_lca(a[i], a[i + 1]);
v[lca].push_back(a[i + 1]);
//v[a[i + 1]].push_back(lca);
}
cout << dfs(1) << endl;
}
return 0;
}