树剖+线段树全WA求助
查看原帖
树剖+线段树全WA求助
684249
__kd楼主2023/9/29 14:05
#include<bits/stdc++.h>
using namespace std;
#define int long long
int n;
int top[200005],id[200005],max_son[200005],tim[200005];
int ok,siz[200005],deep[200005],fa[200005],w[200005];
vector<int> a[200005];
inline int read(){
	register int x=0,t=0;
	static char ch=getchar();
	while(!isdigit(ch)) t|=(ch=='-'),ch=getchar();
	while(isdigit(ch)){x=(x<<1)+(x<<3)+(ch^48);ch=getchar();}
	return t?-x:x;
}
struct Tree{
	int L,R,Mid,sum,lazy,left,right;
}tmp;
struct TTree{
	public:
		Tree p[800005];
		inline void pushup(int k){
			p[k].left=p[k<<1].left;
			p[k].right=p[k<<1|1].right;
			p[k].sum=p[k<<1].sum+p[k<<1|1].sum-(p[k<<1].right==p[k<<1|1].left);
		}
		inline void doit(int k,int x){
			p[k].lazy=x;
			p[k].sum=1;
			p[k].left=p[k].right=x;
		}
		inline void pushdown(int k){
			if(p[k].lazy){
				doit(k<<1,p[k].lazy);
				doit(k<<1|1,p[k].lazy);
				p[k].lazy=0;
			}
		}
		inline void build(int l,int r,int k){
			p[k].L=l;p[k].R=r;
			p[k].Mid=l+r>>1;
			p[k].lazy=0;
			if(l==r){
				p[k].left=p[k].right=w[tim[l]];
				p[k].sum=1;
				return;
			}
			build(l,p[k].Mid,k<<1);
			build(p[k].Mid+1,r,k<<1|1);
			pushup(k);
		}
		inline void change(int l,int r,int x,int k){
			if(l<=p[k].L&&p[k].R<=r){
				doit(k,x);
				return;
			}
			pushdown(k);
			if(l<=p[k].Mid) change(l,r,x,k<<1);
			if(r>p[k].Mid) change(l,r,x,k<<1|1);
			pushup(k);
		}
		inline Tree ask_sum(int l,int r,int k){
			if(l<=p[k].L&&p[k].R<=r){
				return p[k];
			}
			if(l>p[k].Mid) return ask_sum(l,r,k<<1|1);
			if(r<=p[k].Mid) return ask_sum(l,r,k<<1);
			Tree lt=ask_sum(l,r,k<<1),rt=ask_sum(l,r,k<<1|1);
			Tree ans;
			ans.left=lt.left;
			ans.right=rt.right;
			ans.sum=lt.sum+rt.sum-(lt.right==rt.left);
			return ans;
		}
}tr;
struct shupou{
	private:
		inline void dfs1(int x,int last){
			fa[x]=last;
			deep[x]=deep[last]+1;
			siz[x]=1;
			for(int y:a[x]){
				if(y==last) continue;
				dfs1(y,x);
				siz[x]+=siz[y];
				if(siz[max_son[x]]<siz[y]) max_son[x]=y;
			}
		}
		inline void dfs2(int x,int ff){
			top[x]=ff;
			id[x]=++ok;
			tim[ok]=x;
			if(max_son[x]) dfs2(max_son[x],ff);
			for(int y:a[x]){
				if(y==max_son[x]||y==fa[x]) continue;
				dfs2(y,y);
			}
		}
		inline void swap(int &x,int &y){x^=y^=x^=y;}
	public:
		inline void init(int x){
			dfs1(x,0);dfs2(x,x);
			tr.build(1,n,1);
		}
		inline void chang(int x,int y,int k){
			while(top[x]^top[y]){
				if(deep[top[x]]<deep[top[y]]) swap(x,y);
				tr.change(id[top[x]],id[x],k,1);
				x=fa[top[x]];
			}
			if(deep[x]>deep[y]) swap(x,y);
			tr.change(id[x],id[y],k,1);
		}
		inline int find_sum(int x,int y){
			int ans=0,last_x=0,last_y=0;
			while(top[x]^top[y]){
				if(deep[top[x]]<deep[top[y]]) swap(x,y),swap(last_x,last_y);
				tmp=tr.ask_sum(id[top[x]],id[x],1);
				ans+=tmp.sum-(tmp.right==last_x);
				last_x=tmp.left;
				x=fa[top[x]];
			}
			if(deep[x]>deep[y]) swap(x,y),swap(last_x,last_y);
			tmp=tr.ask_sum(id[x],id[y],1);
			return ans+tmp.sum-(tmp.left==last_x)-(tmp.right==last_y);
		}
}Tr;
signed main(){
	n=read();
	int q=read();
	for(register int i=1;i<=n;i++) w[i]=read();
	for(register int i=1;i<n;i++){
		int x=read(),y=read();
		a[x].push_back(y);
		a[y].push_back(x);
	}
	Tr.init(1);
	while(q--){
		char op;
		cin>>op;
		if(op=='C'){
			int x=read(),y=read(),c=read();
			Tr.chang(x,y,c);
		}
		else{
			int x=read(),y=read();
			printf("%lld\n",Tr.find_sum(x,y));
		}
	}
	return 0;
}
2023/9/29 14:05
加载中...