15pts求助
查看原帖
15pts求助
758680
hzy_____楼主2023/9/13 20:05
#include<bits/stdc++.h>
#define int long long
#define endl '\n'
using namespace std;
namespace lct{
    #define lc(n) t[n].s[0]
    #define rc(n) t[n].s[1]
    #define fa(n) t[n].fa
    #define ident(x,f) (rc(f)==x)
    #define connect(x,f,k) t[fa(x)=f].s[k]=x
    #define update(n) t[n].res=(t[n].val+t[lc(n)].res+t[rc(n)].res)%mod,t[n].size=t[lc(n)].size+t[rc(n)].size+1
    #define ntroot(n) (lc(fa(n))==n||rc(fa(n))==n)
    #define reverse(n) swap(lc(n),rc(n)),t[n].tag^=1
    const int N=3e5+10,mod=51061;
    int st[N];
    struct p{int fa,s[2],size,val,res,tag,add,mul;}t[N];
    void pushadd(int n,int x){
    	t[n].res=(t[n].res+x*t[n].size)%mod;
    	t[n].val=(t[n].val+x)%mod;
    	t[n].add=(t[n].add+x)%mod;
	}
	void pushmul(int n,int x){
		t[n].val=t[n].val*x%mod;
		t[n].res=t[n].res*x%mod;
		t[n].add=t[n].add*x%mod;
		t[n].mul=t[n].mul*x%mod;
	}
    void push(int n){
    	if(t[n].mul>1){
    		if(lc(n))pushmul(lc(n),t[n].mul);
    		if(rc(n))pushmul(rc(n),t[n].mul);
		}
		if(t[n].add){
			if(lc(n))pushadd(lc(n),t[n].add);
			if(rc(n))pushadd(rc(n),t[n].add);
		}
        if(t[n].tag){
            if(lc(n))reverse(lc(n));
            if(rc(n))reverse(rc(n));
        }
        t[n].tag=0;
        t[n].mul=1;
        t[n].add=0;
    }
    void rorate(int x){
        int f=fa(x),ff=fa(f),k=ident(x,f);
        connect(t[x].s[k^1],f,k);
        fa(x)=ff;
        if(ntroot(f))t[ff].s[ident(f,ff)]=x;
        connect(f,x,k^1);
        update(f),update(x);
    }
    void splay(int x){
    	int l=0,y=x;
    	st[++l]=y;
    	while(ntroot(y))st[++l]=y=fa(y);
    	while(l)push(st[l--]);
        while(ntroot(x)){
            int f=fa(x),ff=fa(f);
            if(ntroot(f))ident(f,ff)^ident(x,f)?rorate(x):rorate(f);
            rorate(x);
        }
        update(x);
    }
    void access(int x){
        for(int y=0;x;x=fa(y=x)){
            splay(x);
            rc(x)=y;
            update(x);
        }
    }
    void mkrt(int x){
        access(x);
        splay(x);
        reverse(x);
    }
    void link(int x,int y){
        mkrt(x);
        fa(x)=y;
    }
    void split(int x,int y){
        mkrt(x);
        access(y);
        splay(y);
    }
    void cut(int x,int y){
		split(x,y);
        fa(x)=lc(y)=0;
    }
    void init(int n){
    	for(int i=1;i<=n;i++)t[i].size=t[i].val=t[i].res=t[i].mul=1;
	}
}
using namespace lct;
int x,y,z,n,m;
char op;
void print(){
	for(int i=1;i<=n;i++)split(i,i),cout<<t[i].res<<" ";
	cout<<endl;
}
signed main(){
    ios::sync_with_stdio(0),cin.tie(0),cout.tie(0);
    cin>>n>>m;
    init(n);
    for(int i=1;i<n;i++){
    	cin>>x>>y;
    	link(x,y);
	}
//	print();
	while(m--){
		cin>>op;
		if(op=='-'){
			cin>>x>>y;
			cut(x,y);
			cin>>x>>y;
			link(x,y);
		}else if(op=='/'){
			cin>>x>>y;
			split(x,y);
			cout<<t[y].res<<endl;
		}else if(op=='+'){
			cin>>x>>y>>z;
			split(x,y);
			pushadd(y,z);
		}else{
			cin>>x>>y>>z;
			split(x,y);
			pushmul(y,z);
		}
//		print();
	}
}
2023/9/13 20:05
加载中...