这道题想到用LCA了,又是最大值最小化问题,也想到用二分,但还是不会写啊55555...不知道怎么二分,怎么判断条件。。我好菜啊55555...
后来看了大佬博客Orz:https://blog.csdn.net/qq_40993793/article/details/82762766
明白了二分答案,小于答案的节点就是true(mid偏大),所有大于答案的点求它们的公共祖先,涂红之后计算它们的新权值是否小于答案。
另外就是在LCA在线算法中,RMQ的查询部分,用mm数组比直接求要快,不然会TLE。。。
附上AC代码:
#include<cstdio>
#include<cstring>
#include<cstdlib>
#include<cmath>
#include<iostream>
#include<algorithm>
using namespace std;
#define ll long long
typedef pair<ll,int>pp;
#define mkp make_pair
#define pb push_back
const double pi=acos(-1.0);
const double eps=1e-9;
const int INF=0x3f3f3f3f;
const ll MOD=1e9+(ll)7;
const int MAX=1e5+5;
int n,m,q;
bool red[MAX];
int redroot[MAX];
ll cost[MAX];//离最近的红色结点的距离
ll dis[MAX];//离根节点的距离
struct Edge
{
int to,next;
ll w;
}edge[MAX*2];
int head[MAX],tot;
void init()
{
tot=0;
memset(head,-1,sizeof(head));
}
void addedge(int u,int v,ll w)
{
edge[tot].to=v;edge[tot].w=w;
edge[tot].next=head[u];head[u]=tot++;
}
int rmq[MAX*2];
struct ST
{
int dp[MAX*2][20];
int mm[MAX*2];
void init(int n)
{
mm[0]=-1;
for(int i=1;i<=n;i++)
{
mm[i]=((i&(i-1))==0)?mm[i-1]+1:mm[i-1];
dp[i][0]=i;
}
for(int j=1;j<=mm[n];j++)
for(int i=1;i+(1<<j)-1<=n;i++)
dp[i][j]=rmq[dp[i][j-1]]<rmq[dp[i+(1<<(j-1))][j-1]]?dp[i][j-1]:dp[i+(1<<(j-1))][j-1];
}
int query(int x,int y)
{
if(x>y)
swap(x,y);
int k=mm[y-x+1];//这样快一些
return rmq[dp[x][k]]<=rmq[dp[y-(1<<k)+1][k]]?dp[x][k]:dp[y-(1<<k)+1][k];
}
};
ST st;
int f[MAX*2];
int p[MAX*2];
int cnt;
void lca_dfs(int u,int pre,int dep)
{
f[++cnt]=u;
rmq[cnt]=dep;
p[u]=cnt;
for(int i=head[u];i!=-1;i=edge[i].next)
{
int v=edge[i].to;
if(v==pre)
continue;
if(red[v])
{
cost[v]=0;
redroot[v]=v;
}
else
{
cost[v]=cost[u]+edge[i].w;
redroot[v]=redroot[u];
}
dis[v]=dis[u]+edge[i].w;
lca_dfs(v,u,dep+1);
f[++cnt]=u;
rmq[cnt]=dep;
}
}
void lca_init(int root,int n)
{
cnt=0;
memset(cost,0,sizeof(cost));
memset(dis,0,sizeof(dis));
for(int i=1;i<=n;i++)
redroot[i]=1;
lca_dfs(root,root,0);
st.init(2*n-1);
/*for(int i=1;i<=n;i++)
cout<<"i="<<i<<" dis="<<dis[i]<<" cost="<<cost[i]<<" redroot="<<redroot[i]<<endl;
cout<<endl;*/
}
int lca_query(int u,int v)
{
return f[st.query(p[u],p[v])];
}
int tmp[MAX];
bool cmp(int a,int b)
{
return cost[a]>cost[b];
}
bool check(ll x,int n)
{
if(cost[tmp[1]]<=x)//x偏大
return true;
//所有大于答案结点的LCA
int rt=tmp[1];
for(int i=2;i<=n;i++)
{
if(cost[tmp[i]]<=x)
break;
rt=lca_query(rt,tmp[i]);
}
//cout<<"rt="<<rt<<endl;
for(int i=1;i<=n;i++)
{
if(cost[tmp[i]]<=x)
break;
if(dis[tmp[i]]-dis[rt]>x)
return false;
}
return true;
}
int main()
{
int t,num;
scanf("%d",&t);
int u,v;ll w;
ll l,r;
while (t--)
{
scanf("%d%d%d",&n,&m,&q);
memset(red,false,sizeof(red));
for(int i=1;i<=m;i++)
{
scanf("%d",&num);
red[num]=true;
}
red[1]=true;
init();//注意!!!
for(int i=1;i<n;i++)
{
scanf("%d%d%lld",&u,&v,&w);
addedge(u,v,w);
addedge(v,u,w);
}
lca_init(1,n);
//memset(tmp,0,sizeof(tmp));
while(q--)
{
scanf("%d",&num);
for(int i=1;i<=num;i++)
scanf("%d",&tmp[i]);
sort(tmp+1,tmp+1+num,cmp);//cost从大到小
l=0,r=cost[tmp[1]];
while(l<r)
{
ll mid=(l+r)/2;
//cout<<"l="<<l<<" r="<<r<<" mid="<<mid<<endl;
if(check(mid,num))
r=mid;
else
l=mid+1;
}
printf("%lld\n",l);
}
}
return 0;
}
果然对二分的理解还是不够啊55555...