树剖TLE求调#8#9
查看原帖
树剖TLE求调#8#9
287217
tanyanling楼主2022/11/16 12:57
#include<bits/stdc++.h>
using namespace std;
inline int read()
{
	int x=0;
	char ch=getchar();
	while(ch<'0'||ch>'9')
		ch=getchar();
	while(ch>='0'&&ch<='9')
	{
		x=(x<<1)+(x<<3)+(ch^'0');
		ch=getchar();
	}
	return x;
}
int n,m;
vector<int>g[500001];
int siz[500001],depth[500001],son[500001],fa[500001];
inline void dfs1(int x,int dep,int fath)
{
	depth[x]=dep;
	int len=g[x].size();
	siz[x]=1;
	fa[x]=fath;
	if(len==1&&fath!=x)
		return;
	int maxn=-1;
	for(register int i=0;i<len;i++)
	{
		int y=g[x][i];
		if(y==fath)
			continue;
		dfs1(y,dep+1,x);
		siz[x]+=siz[y];
		if(siz[y]>maxn)
		{
			maxn=siz[y];
			son[x]=y;
		}
	}
}
int seg[500001],id[500001],top[500001],sum;
inline void dfs2(int x,int fath,int topp)
{
	int len=g[x].size();
	id[x]=++sum;
	seg[sum]=1;
	top[x]=topp;
	if(len==1&&fath!=x)
		return;
	dfs2(son[x],x,topp);
	for(register int i=0;i<len;i++)
	{
		int y=g[x][i];
		if(y==fath||y==son[x])
			continue;
		dfs2(y,x,y);
	}
}
int segtree[2000001];
inline void build(int l,int r,int num)
{
	if(l==r)
	{
		segtree[num]=seg[l];
		return;
	}
	int mid=(l+r)>>1;
	build(l,mid,(num<<1));
	build(mid+1,r,(num<<1)+1);
	segtree[num]=segtree[(num<<1)]+segtree[(num<<1)+1];
}
inline int query(int l,int r,int s,int e,int num)
{
	if(l>=s&&r<=e)
		return segtree[num];
	int mid=(l+r)>>1;
	int res=0;
	if(mid>=s)
		res+=query(l,mid,s,e,(num<<1));
	if(mid<e)
		res+=query(mid+1,r,s,e,(num<<1)+1);
	return res;
}
int lca;
inline int waysolve(int x,int y)
{
	int res=0;
	while(top[x]!=top[y])
	{
		if(depth[top[x]]<depth[top[y]])
			swap(x,y);
		res+=query(1,sum,id[top[x]],id[x],1);
		x=fa[top[x]];
	}
	if(depth[x]>depth[y])
		swap(x,y);
	lca=x;
	res+=query(1,sum,id[x],id[y],1);
	res-=query(1,sum,id[x],id[x],1);
	return res;
}
int main()
{
	n=read(),m=read();
	for(register int i=1;i<=n-1;i++)
	{
		int from=read(),to=read();
		g[from].push_back(to);
		g[to].push_back(from);
	}
	dfs1(1,1,1);
	dfs2(1,1,1);
	build(1,sum,1);
	while(m--)
	{
		int x=read(),y=read(),z=read();
		int maxn=0,maxlca;
		int axy=waysolve(x,y),lcaxy=lca;
		if(depth[lcaxy]>maxn)
		{
			maxn=depth[lcaxy];
			maxlca=lcaxy;
		}
		int axz=waysolve(x,z),lcaxz=lca;
		if(depth[lcaxz]>maxn)
		{
			maxn=depth[lcaxz];
			maxlca=lcaxz;
		}
		int ayz=waysolve(y,z),lcayz=lca;
		if(depth[lcayz]>maxn)
		{
			maxn=depth[lcayz];
			maxlca=lcayz;
		}
		int ans=(axy+axz+ayz)/2;
		printf("%d %d\n",maxlca,ans);
	}
	return 0;
}
2022/11/16 12:57
加载中...