查看原帖
256970
xie_lzh楼主2022/11/6 22:52

写了一下午加一晚上,但0pt,挂一下

#include<bits/stdc++.h>
using namespace std;
const int N=5e5+5,M=1e6+5;
int n,m,u,v;
int to[N<<1],nxt[N<<1],head[N],cnt;
vector<int> key[N],pos[N],To[N];
void add(int u,int v)
{
	to[++cnt]=v;
	nxt[cnt]=head[u];
	head[u]=cnt;
}
int dfn[N],dep[N],fa[N][22],idx,siz[N];
void dfs(int now,int f)
{
	fa[now][0]=f;
	for(int i=1;i<=20;i++)
		fa[now][i]=fa[fa[now][i-1]][i-1];
	dfn[now]=++idx; dep[now]=dep[f]+1;
	siz[now]=1;
	for(int i=head[now];i;i=nxt[i])
	{
		int v=to[i];
		if(v==f) continue;
		dfs(v,now);
		siz[now]+=siz[v];
	}
}

bool cmp(int x,int y)
{
	return dfn[x]<dfn[y];
}
int LCA(int x,int y)
{
	if(dep[x]<dep[y]) swap(x,y);
	for(int i=20;i>=0;i--)
		if(dep[fa[x][i]]>=dep[y])
			x=fa[x][i];
	if(x==y) return x;
	for(int i=20;i>=0;i--)
		if(fa[x][i]!=fa[y][i])
			x=fa[x][i],y=fa[y][i];
	return fa[x][0];
}

int stk[N],top,tot;
bool vis[N];
struct node
{
	int a,b,c,d;
}Matrix[N*10];
int gett(int x,int y)
{ 
	for(int i=20;i>=0;i--)
		if(dep[fa[x][i]]>dep[y])
			x=fa[x][i];
	return x;
}
void insert(int x,int y)
{
	int lca=LCA(x,y);
	if(lca==x)
	{
		int tmp=gett(y,x);
		Matrix[++tot]=node{1,dfn[tmp]-1,dfn[y],dfn[y]+siz[y]-1};
		Matrix[++tot]=node{dfn[tmp]+siz[tmp],n,dfn[y],dfn[y]+siz[y]-1};
	}
	else if(lca==y) 
	{
		int tmp=gett(x,y);
		Matrix[++tot]=node{dfn[x],dfn[x]+siz[x]-1,1,dfn[tmp]-1};
		Matrix[++tot]=node{dfn[x],dfn[x]+siz[x]-1,dfn[tmp]+siz[tmp],n};
	}
	else
	{
		Matrix[++tot]=node{dfn[x],dfn[x]+siz[x]-1,dfn[y],dfn[y]+siz[y]-1};
	}
}
void dfs2(int now,int f,int cnt,int id)
{
	for(auto v:To[now])
	{
		if(v==f) continue;
		if(vis[v]) dfs2(v,now,cnt+1,id);
		else
		{
			if(cnt==0)
				insert(id,v);
			else dfs2(v,now,cnt-1,id);
		}
	}
}
void build(int col)
{
	sort(pos[col].begin(),pos[col].end(),cmp);
	top=0;
//	if(pos[col][0]!=1) stk[++top]=1;
	for(int i=0;i<pos[col].size();i++)
	{
		int now=pos[col][i];
		if(!top)
		{
			stk[++top]=now;
			continue;
		}
		int lca=LCA(now,stk[top]);
		if(lca==stk[top])
		{
			stk[++top]=now;
			continue;
		}
		while(top>1&&dep[stk[top-1]]>dep[lca])
		{
			To[stk[top-1]].push_back(stk[top]);
			To[stk[top]].push_back(stk[top-1]);
			top--;
		}
		if(top==1)
		{
			To[lca].push_back(stk[top]);
			To[stk[top]].push_back(lca);
			top--;
			stk[++top]=lca;
			stk[++top]=now;
		}
		else if(stk[top-1]==lca)
		{
			To[stk[top-1]].push_back(stk[top]);
			To[stk[top]].push_back(stk[top-1]);
			top--;
			stk[++top]=now;
		}
		else
		{
			To[stk[top-1]].push_back(stk[top]);
			To[stk[top]].push_back(stk[top-1]);
			top--;
			To[lca].push_back(stk[top]);
			To[stk[top]].push_back(lca);
			top--;
			stk[++top]=lca;
			stk[++top]=now;
		}
	}
	while(top>1)
	{
		To[stk[top-1]].push_back(stk[top]);
		To[stk[top]].push_back(stk[top-1]);
		top--;
	}
	for(auto i:key[col])
	{
		dfs2(i,0,0,i);
	}
	for(int i=0;i<pos[col].size();i++)
		To[pos[col][i]].clear();
}
struct cc
{
	int x,l,r,id,ad;
}upd[N*10];
bool cmp2(cc a,cc b)
{
	if(a.x==b.x)
	{
		return (a.id!=0);
	}
	return a.x<b.x;
}
int ans[M];
int c[N];
int lowbit(int x){return x&(-x);}
void ADD(int x,int val)
{
	for(;x<=n;x+=lowbit(x))
		c[x]+=val;
}
int ask(int x)
{
	int sum=0;
	for(;x;x-=lowbit(x))
		sum+=c[x];
	return sum;
}
int main()
{
	ios::sync_with_stdio(false);
	cin.tie(0); cout.tie(0);
	cin>>n>>m;
	for(int i=1;i<=n;i++)
	{
		cin>>u>>v;
		if(u==1)
		{
			key[v].push_back(i);
			vis[i]=1;
		} 
		pos[v].push_back(i);
	}
	for(int i=1;i<n;i++)
	{
		cin>>u>>v;
		add(u,v); add(v,u);
	}
	dfs(1,0);
	for(int i=1;i<=m;i++)
	{
		cin>>upd[i].x>>upd[i].l;
		upd[i].ad=i;
		upd[i].x=dfn[upd[i].x];
		upd[i].l=dfn[upd[i].l];
	}
	for(int i=1;i<=n;i++)
		build(i);
	// puts("?");
	int cnttt=m;
	for(int i=1;i<=tot;i++)
	{
		upd[++m]=cc{Matrix[i].a,Matrix[i].c,Matrix[i].d,1,0};
		upd[++m]=cc{Matrix[i].b+1,Matrix[i].c,Matrix[i].d,-1,0};
	}
	sort(upd+1,upd+1+m,cmp2);
	for(int i=1;i<=m;i++)
	{
		if(upd[i].id==0)
		{
			ans[upd[i].ad]=ask(upd[i].l);
		}
		else
		{
			ADD(upd[i].l,upd[i].id);
			ADD(upd[i].r+1,-upd[i].id);
		}
	}
	for(int i=1;i<=cnttt;i++)
		printf("%d\n",ans[i]);
}


2022/11/6 22:52
加载中...