这是100pts代码 但是疑惑点在于insert函数中if ( lca == sta[tp] ) return;这一句话为什么不用加上一句sta[++tp] = u;(也就是将当前处理的节点加进最右链中)就可以return了呢
题解中对此讲道 "首先,每插入一个点,如果栈顶元素是他所在链上的点,那么就可以不加这个点,这个画画图就知道了"
那加上这个点为啥只有20pts了呢/kk/kk/kk
#include <bits/stdc++.h>
using namespace std;
#define int long long
const int N = 5e5 + 5;
const int inf = 0x3f3f3f3f;
inline int read ()
{
int x = 0 , f = 1;
char ch = getchar();
while ( !isdigit ( ch ) ) { if ( ch == '-' ) f = -1; ch = getchar (); }
while ( isdigit ( ch ) ) { x = ( x << 1 ) + ( x << 3 ) + ( ch ^ 48 ); ch = getchar (); }
return x * f;
}
int n , m , a[N] , mn[N];//岛屿数量 敌方机器使用的次数 每一次的关键点编号
//最后一个是辅助数组 表示i到根路径上的最小边权
struct node
{
int head[N] , cnt;
struct edge { int to , nxt , w; } e[N<<1];
void add ( int u , int v , int w = 0 ) { e[++cnt] = { v , head[u] , w }; head[u] = cnt; }
}e1,e2;//e1是原来树的边 e2是现在虚树的边
int sz[N] , dep[N] , fa[N] , top[N] , son[N] , pos[N] , rev[N] , timer;
int tp , sta[N];
void dfs1 ( int u , int f )
{
sz[u] = 1 , fa[u] = f , dep[u] = dep[f] + 1;
for ( int i = e1.head[u] ; i ; i = e1.e[i].nxt )
{
int v = e1.e[i].to;
if ( v == f ) continue;
mn[v] = min ( mn[u] , e1.e[i].w );//这个v节点到根节点路径上的最小边权(也就是将v这个子树整个切掉所需要的最小代价)
dfs1 ( v , u );
sz[u] += sz[v];
if ( sz[v] > sz[son[u]] ) son[u] = v;
}
}
void dfs2 ( int u , int tp )
{
pos[u] = ++timer , rev[timer] = u , top[u] = tp;
if ( son[u] ) dfs2 ( son[u] , tp );
for ( int i = e1.head[u] ; i ; i = e1.e[i].nxt )
{
int v = e1.e[i].to;
if ( v == fa[u] || v == son[u] ) continue;
dfs2 ( v , v );
}
}
int LCA ( int u , int v )
{
while ( top[u] != top[v] )
{
if ( dep[top[u]] < dep[top[v]] ) swap ( u , v );//优先跳链顶深度大的节点
u = fa[top[u]];
}
return dep[u] > dep[v] ? v : u;
}
void insert ( int u )
{
if ( tp == 1 ) { if ( u != 1 ) sta[++tp] = u; return; }
int lca = LCA ( u , sta[tp] );
if ( lca == sta[tp] ) return;
while ( tp > 1 && pos[sta[tp-1]] >= pos[lca] ) e2.add ( sta[tp-1] , sta[tp] ) , tp --;
if ( lca != sta[tp] ) e2.add ( lca , sta[tp] ) , sta[tp] = lca;
sta[++tp] = u;
}
void init()
{
sta[tp=1] = 1;
e2.cnt = 0;
}
int dfs ( int u )
{
if ( !e2.head[u] ) return mn[u];
int res = 0;
for ( int i = e2.head[u] ; i ; i = e2.e[i].nxt )
{
int v = e2.e[i].to;
res += dfs(v);
}
e2.head[u] = 0;
return min ( res , mn[u] );
}
signed main ()
{
n = read();
memset ( mn , inf , sizeof ( mn ) );
for ( int i = 1 , u , v , w ; i < n ; i ++ ) u = read() , v = read() , w = read() , e1.add ( u , v , w ) , e1.add ( v , u , w );
dfs1 ( 1 , 0 ) , dfs2 ( 1 , 1 );
m = read();
for ( int i = 1 , k ; i <= m ; i ++ )
{
k = read();
for ( int i = 1 ; i <= k ; i ++ ) a[i] = read();
sort ( a + 1 , a + k + 1 , [](const int &a , const int &b) { return pos[a] < pos[b]; } );
//按照dfs序排序
init();
for ( int i = 1 ; i <= k ; i ++ ) insert ( a[i] );
for ( ; tp > 1 ; tp -- ) e2.add ( sta[tp-1] , sta[tp] );
printf ( "%lld\n" , dfs(1) );
}
return 0;
}