不开O2AC,开了RE
查看原帖
不开O2AC,开了RE
555287
_ANIG_楼主2022/11/13 14:44

RT。

O2

AC

#include <bits/stdc++.h>
using namespace std;
int mods;
#define int long long
namespace tr{
    struct node{int l,r,sum,add,mid,lth;}p[400005];
	int ys[100005];
	void upset(int x){p[x].sum=p[x<<1].sum+p[x<<1|1].sum;p[x].sum%=mods;}
	void gx(int x,int y){p[x].sum+=p[x].lth*y,p[x].add+=y;p[x].sum%=mods;p[x].add%=mods;}
	void dnset(int x){gx(x<<1,p[x].add);gx(x<<1|1,p[x].add);p[x].add=0;}
	void reset(int x,int l,int r){
		p[x].l=l,p[x].r=r,p[x].mid=l+r>>1,p[x].lth=r-l+1;
		if(l==r){p[x].sum=ys[l];return;}
		reset(x<<1,l,p[x].mid);
		reset(x<<1|1,p[x].mid+1,r);
		upset(x);
	}
	void add(int x,int l,int r,int sum){
		if(l<=p[x].l&&r>=p[x].r){gx(x,sum);return;}
		dnset(x);
		if(l<=p[x].mid)add(x<<1,l,r,sum);
		if(r>p[x].mid)add(x<<1|1,l,r,sum);
		upset(x);
	}
	int gets(int x,int l,int r){
		if(l<=p[x].l&&r>=p[x].r)return p[x].sum;
		dnset(x);
		int res=0;
		if(l<=p[x].mid)res+=gets(x<<1,l,r);
		if(r>p[x].mid)res+=gets(x<<1|1,l,r);
		upset(x);
		return res%mods;
	}
};
vector<int>p[100005];
int n,m,root,fa[100005],dfn[100005],siz[100005],zs[100005],dep[100005],mk[100005],d[100005],idx,dy[100005],eds[100005],w[100005];
void dfs1(int x){
	mk[x]=1;
	for(int i=0;i<p[x].size();i++){
		int c=p[x][i];
		if(mk[c])continue;
		dep[c]=dep[x]+1;
		dfs1(c);
		siz[x]+=siz[c],fa[c]=x;
		if(siz[c]>siz[zs[x]])zs[x]=c;
	}
	siz[x]++;
}
void dfs2(int x,int y){
	dfn[x]=++idx,dy[idx]=x;
	d[x]=y,mk[x]=1;
	if(zs[x])dfs2(zs[x],y);
	for(int i=0;i<p[x].size();i++){
		int c=p[x][i];
		if(mk[c])continue;
		dfs2(c,c);
	}
	eds[x]=idx;
}
void add(int a,int b,int sum){
	while(d[a]!=d[b]){
		if(dep[d[a]]>dep[d[b]])swap(a,b);
		tr::add(1,dfn[d[b]],dfn[b],sum);
		b=d[b],b=fa[b];
	}
	if(dep[a]>dep[b])swap(a,b);
	tr::add(1,dfn[a],dfn[b],sum);
}
int gets(int a,int b){
	int res=0;
	while(d[a]!=d[b]){
		if(dep[d[a]]>dep[d[b]])swap(a,b);
		res+=tr::gets(1,dfn[d[b]],dfn[b]);
		b=d[b],b=fa[b];
		res%=mods;
	}
	if(dep[a]>dep[b])swap(a,b);
	return (res+tr::gets(1,dfn[a],dfn[b]))%mods;
}
int gets(int x){return tr::gets(1,dfn[x],eds[x]);}
int add(int x,int sum){tr::add(1,dfn[x],eds[x],sum);}
signed main(){
	cin>>n>>m>>root>>mods;
	for(int i=1;i<=n;i++)scanf("%lld",&w[i]);
	for(int i=1;i<n;i++){
		int x,y;
		scanf("%lld%lld",&x,&y);
		p[x].push_back(y);
		p[y].push_back(x);
	}
	dfs1(root);memset(mk,0,sizeof(mk));dfs2(root,root);
	for(int i=1;i<=n;i++)tr::ys[i]=w[dy[i]];
	tr::reset(1,1,idx);
	while(m--){
		int op,x,y,z;
		scanf("%lld%lld",&op,&x);if(op!=4)scanf("%lld",&y);
		if(op==1){scanf("%lld",&z);add(x,y,z);}
		if(op==2)printf("%lld\n",gets(x,y));
		if(op==3)add(x,y);
		if(op==4)printf("%lld\n",gets(x));
	}
}
2022/11/13 14:44
加载中...