萌新袜子求调树分治
查看原帖
萌新袜子求调树分治
708103
封禁用户楼主2022/7/21 20:31

TLE(RE on #4)

#include <bits/stdc++.h>
using namespace std;
const int N=1e5+5;
int a[N],d[N],dfn[N],sz[N],g[N],f[N],lg[N<<1],st[N<<1][20],rt,all,tot=0,ans=0;
bool vis[N];
vector<int> nodes[N];
struct BIT
{
	vector<int> c;
	int lowbit(int x)
	{
		return x&(-x);
	}
	void add(int x,int v)
	{
		for(int i=x;i<(int)c.size();i+=lowbit(x))c[i]+=v;
		return;
	}
	int sum(int x)
	{
		int res=0;
		for(int i=x;i;i-=lowbit(x))res+=c[i];
		return res;
	}
}bit0[N],bit1[N];
void dfs(int u,int fa)
{
	dfn[u]=++tot;
	d[u]=d[fa]+1;
	st[tot][0]=d[u];
	for(int v:nodes[u])
	{
		if(v==fa)continue;
		dfs(v,u);
		st[++tot][0]=d[u];
	}
	return;
}
void init()
{
	for(int i=2;i<=tot;i++)lg[i]=lg[i>>1]+1;
	for(int j=1;j<=lg[tot];j++)for(int i=1;i+(1<<j)-1<=tot;i++)st[i][j]=min(st[i][j-1],st[i+(1<<(j-1))][j-1]);
	return;
}
int d_lca(int u,int v)
{
	int x=dfn[u],y=dfn[v];
	if(x>y)swap(x,y);
	int k=lg[y-x+1];
	return min(st[x][k],st[y-(1<<k)+1][k]);
}
int DIS(int u,int v)
{
	return d[u]+d[v]-2*d_lca(u,v);
}
void get_rt(int u,int fa)
{
	sz[u]=1,g[u]=0;
	for(int v:nodes[u])
	{
		if(v==fa||vis[v])continue;
		get_rt(v,u);
		sz[u]+=sz[v];
		g[u]=max(g[u],sz[v]);
	}
	g[u]=max(g[u],all-sz[u]);
	if(!rt||g[u]<g[rt])rt=u;
	return;
}
void build(int u)
{
	vis[u]=true;
	bit0[u].c.resize(sz[u]+5);
	bit1[u].c.resize(sz[u]+5);
	for(int v:nodes[u])
	{
		if(vis[v])continue;
		rt=0,all=sz[v];
		get_rt(v,0);
		get_rt(rt,0);
		f[rt]=u;
		build(rt);
	}
	return;
}
void QUERY(int x,int k)
{
	ans=bit0[x].sum(k+1);
	for(int i=x;f[i];i=f[i])
	{
		int z=f[i];
		int dis=DIS(x,z);
		if(dis<=k)ans+=bit0[z].sum(k-dis+1)-bit1[i].sum(k-dis+1);
	}
	printf("%d\n",ans);
	return;
}
void MODIFY(int x,int v)
{
	for(int i=x;i;i=f[i])bit0[i].add(DIS(x,i)+1,v);
	for(int i=x;f[i];i=f[i])bit1[i].add(DIS(x,f[i])+1,v);
	return;
}
int main()
{
	int n,m;
	scanf("%d %d",&n,&m);
	for(int i=1;i<=n;i++)scanf("%d",&a[i]);
	for(int i=1;i<n;i++)
	{
		int u,v;
		scanf("%d %d",&u,&v);
		nodes[u].push_back(v);
		nodes[v].push_back(u);
	}
	dfs(1,0);
	init();
	rt=0,all=n;
	get_rt(1,0);
	build(rt);
	for(int i=1;i<=n;i++)MODIFY(i,a[i]);
	while(m--)
	{
		int op,x,y;
		scanf("%d %d %d",&op,&x,&y),x^=ans,y^=ans;
		if(op==0)QUERY(x,y);
		else MODIFY(x,y-a[x]),a[x]=y;
	}
	return 0;
}
2022/7/21 20:31
加载中...