rt,应该是 check 函数炸了,最终答案偏小。
using namespace std;
typedef long long int ll;
const int maxn = 2e5 + 10;
struct edge {
int to, nxt;
ll w;
}node[maxn];
int head[maxn], cnt = 0, dct = 0, n, m;
ll fa[maxn], top[maxn], siz[maxn], dep[maxn], son[maxn], dis[maxn], dfn[maxn], id[maxn], disf[maxn], que[maxn], tree[maxn];
bool vis[maxn];
inline void add(int u, int v, ll w) {
node[cnt].nxt = head[u];
node[cnt].to = v;
node[cnt].w = w;
head[u] = cnt++;
}
void dfs1(int u, int f) {
fa[u] = f; dep[u] = dep[f] + 1; siz[u] = 1;
tree[u] = (dep[u] == 2 ? u : tree[f]);
int mson = 0;
for (int i = head[u]; ~i; i = node[i].nxt) {
int v = node[i].to;
if (v == f)continue;
dis[v] = dis[u] + node[i].w;
disf[v] = node[i].w;
dfs1(v, u);
siz[u] += siz[v];
if (siz[v] > mson)son[u] = v, mson = siz[v];
}
}
void dfs2(int u, int t) {
top[u] = t; dfn[u] = ++dct; id[dfn[u]] = u;
if (son[u])dfs2(son[u], t);
for (int i = head[u]; ~i; i = node[i].nxt) {
int v = node[i].to;
if (v == fa[u] || v == son[u])continue;
dfs2(v, v);
}
}
ll tot[maxn], hcnt = 0, tcnt = 0;
pair<ll, int>hist[maxn];
int jump(int u, ll mdis) {
while (dis[u] - mdis <= dis[top[u]]) {
mdis -= dis[u] - dis[top[u]];
u = top[u];
if (mdis < disf[u])break;
mdis -= disf[u];
u = fa[u];
}
int l = dfn[top[u]], r = dfn[u], mid;
while (l < r) {
mid = l + r >> 1;
if (dis[u] - dis[mid] <= mdis)r = mid;
else l = mid + 1;
}
return id[l];
}
ll need[maxn], neds[maxn], ncnt = 0;
bool isfull(int u) {
bool flag = true;
if (vis[u])return true;
for (int i = head[u]; ~i; i = node[i].nxt) {
int v = node[i].to;
if (fa[u] == v)continue;
flag = false;
if (!isfull(v))return false;
}
if (flag)return false;
return true;
}
bool check(ll MaxDis) {
for (int i = 1; i <= n; i++) {
vis[i] = need[i] = neds[i] = tot[i] = 0;
hist[i].first = hist[i].second = 0;
}
ncnt = tcnt = hcnt = 0;
for (int i = 1; i <= m; i++) {
if (dis[que[i]] <= MaxDis)hist[++hcnt] = make_pair(MaxDis - dis[que[i]], tree[que[i]]);
else vis[jump(que[i], MaxDis)] = true;
}
for (int i = head[1]; ~i; i = node[i].nxt)if (!isfull(node[i].to))need[node[i].to] = true;
sort(hist + 1, hist + hcnt + 1);
for (int i = 1; i <= hcnt; i++) {
if (need[hist[i].second] && hist[i].first < dis[hist[i].second])need[hist[i].second] = false;
else tot[++tcnt] = hist[i].first;
}
for (int i = head[1]; ~i; i = node[i].nxt)if (need[node[i].to])neds[++ncnt] = dis[node[i].to];
if (tcnt < ncnt)return 0;
sort(tot + 1, tot + tcnt + 1); sort(neds + 1, neds + ncnt + 1);
int l = 1, r = 1;
while (l <= ncnt && r <= tcnt) {
if (tot[r] >= neds[l])++l, ++r;
else ++r;
}
if (l > ncnt)return true;
else return false;
}
void init() {
memset(head, -1, sizeof(head));
}
int main() {
ios::sync_with_stdio(false);
cin.tie(); cout.tie(); init();
cin >> n;
int u, v; ll w, ans = 0;
ll l = 0, r = 0, mid;
for (int i = 1; i < n; i++) {
cin >> u >> v >> w;
add(u, v, w);
add(v, u, w);
r += w;
}
dfs1(1, 1); dfs2(1, 1);
cin >> m;
for (int i = 1; i <= m; i++)cin >> que[i];
while (l <= r) {
mid = l + r >> 1;
cout << mid << endl;
if (check(mid))r = mid - 1, ans = mid;
else l = mid + 1;
}
cout << ans << endl;
return 0;
}