WA30pts 求助
查看原帖
WA30pts 求助
416521
NATURAL6楼主2022/10/17 16:27

rt

#include<bits/stdc++.h>
using namespace std;
inline int qread()
{
	int a=0,f=1;char ch=getchar();
	while(ch>'9'||ch<'0'){if(ch=='-')f=-1;ch=getchar();}
	while(ch>='0'&&ch<='9'){(a*=10)+=(ch^48);ch=getchar();}
	return a*f;
}
int n,m,a[100010],b[100010],pos[100010],op,x,y;
int head[100010],tot,dx[200010],p,siz[100010],rot=1,Log[100010],dep[100010],fa[100010][20];
int cnt,al,ar,bl,br,sq,l,r,sl[100010],sr[100010];
int cl,c[200010];
long long ans[500010],an;
struct xw 
{
	int l,r,id,op;
}q[2000010];
inline bool cmp(xw x,xw y){return c[x.l]==c[y.l]?((c[x.l]&1)?x.r<y.r:x.r>y.r):x.l<y.l;}
struct qxx
{
	int nex,t;
}e[200010];
inline void adde(register int x,register int y)
{
	e[++tot].t=y;
	e[tot].nex=head[x];
	head[x]=tot;
	return ;
}
void dfs(register int rt,register int da)
{
	dx[++p]=rt;dx[p+n]=rt;siz[rt]=1;pos[rt]=p;dep[rt]=dep[da]+1;fa[rt][0]=da;
	for(register int i=1;i<=Log[n];++i)fa[rt][i]=fa[fa[rt][i-1]][i-1];
	for(register int i=head[rt];i;i=e[i].nex)
	{
		if(e[i].t==da)continue;
		dfs(e[i].t,rt);
		siz[rt]+=siz[e[i].t];
	}
	return ;
}
inline void U(register int x,register int op)
{
	if(op==1)
	{
		an+=sl[a[x]];
		++sr[a[x]];
	}
	else 
	{
		an+=sr[a[x]];
		++sl[a[x]];
	}
	return ;
}
inline void M(register int x,register int op)
{
	if(op==1)
	{
		an-=sl[a[x]];
		--sr[a[x]];
	}
	else 
	{
		an-=sr[a[x]];
		--sl[a[x]];
	}
	return ;
}
int main()
{
	n=qread(),m=qread();
	Log[0]=-1;
	for(register int i=1;i<=n;++i)b[i]=a[i]=qread(),Log[i]=Log[i>>1]+1;
	b[0]=n;
	sort(b+1,b+1+n);
	b[0]=unique(b+1,b+1+b[0])-b-1;
	for(register int i=1;i<=n;++i)a[i]=lower_bound(b+1,b+1+b[0],a[i])-b;
	for(register int i=1,u,v;i<n;++i)
	{
		u=qread();
		v=qread();
		adde(u,v);
		adde(v,u);
	}
	dfs(1,0);
	cl=sqrt(n)+1;
	for(register int i=(n<<1);i;--i)c[i]=(i-1)/cl+1;
	for(register int i=1;i<=m;++i)
	{
		op=qread();
		x=qread();
		if(op==1)rot=x;
		else
		{
			++cnt;
			y=qread();
			if(rot==x)
			{
				al=pos[x];
				ar=pos[x]+n-1;
			}
			else if(pos[rot]>pos[x]&&pos[rot]<pos[x]+siz[x])
			{
				al=pos[rot]+1;
				l=rot;
				for(register int i=Log[n];i>=0;--i)
					while(dep[fa[l][i]]>dep[x])l=fa[l][i];
				ar=pos[l]+n-1;
			}
			else 
			{
				al=pos[x];
				ar=pos[x]+siz[x]-1;
			}
			if(rot==y)
			{
				bl=pos[y];
				br=pos[y]+n-1;
			}
			else if(pos[rot]>pos[y]&&pos[rot]<pos[y]+siz[y])
			{
				bl=pos[rot]+1;
				r=rot;
				for(register int i=Log[n];i>=0;--i)
					while(dep[fa[r][i]]>dep[y])r=fa[r][i];
				br=pos[r]+n-1;
			}
			else 
			{
				bl=pos[y];
				br=pos[y]+siz[y]-1;
			}
			--al,--bl;
			q[++sq].l=al,q[sq].r=bl,q[sq].id=cnt,q[sq].op=1;
			q[++sq].l=al,q[sq].r=br,q[sq].id=cnt,q[sq].op=-1;
			q[++sq].l=ar,q[sq].r=br,q[sq].id=cnt,q[sq].op=1;
			q[++sq].l=ar,q[sq].r=bl,q[sq].id=cnt,q[sq].op=-1;
		}
	}
	sort(q+1,q+1+sq,cmp);
	l=r=0;
	for(register int i=1;i<=sq;++i)
	{
		while(r<q[i].r)U(dx[++r],1);
		while(l<q[i].l)U(dx[++l],0);
		while(r>q[i].r)M(dx[r--],1);
		while(l>q[i].l)M(dx[l--],0);
		ans[q[i].id]+=an*q[i].op;
	}
	for(register int i=1;i<=cnt;++i)printf("%lld\n",ans[i]);
	return 0;
}
2022/10/17 16:27
加载中...