1 条题解

  • 0
    @ 2026-6-9 0:51:15

    提供一种更为简洁的树剖做法。

    考虑对每一个节点,维护这样一条信息:

    • 该点与它父亲的颜色是否相同。

    若能成功求出这个信息的变化,则可以用线段树简单地维护 2、3 操作的答案。

    这个信息的变化次数是多少呢?其实对于单次 1 操作,只会有 O(logn)O(\log n) 个节点的信息发生变化。下面是一个简短的证明。

    将所有点分为两类:

    1. 该点与他颜色相同的儿子是它的重儿子。
    2. 该点与他颜色相同的儿子是某个轻儿子,或并没有与它颜色相同的儿子。

    这样,每次 1 操作会使得某些 1 类点变为 2 类点,或使 2 类点同色儿子切换,又或是将 2 类点变为 1 类点。前两个情况的总和是 O(logn)O(\log n) 的,因为只有跳轻边的时候可能会有前两个情况发生。从而,第三个情况均摊下来也是 O(logn)O(\log n ) 的。因此,我们说明了上信息的变化次数是 O(logn)O(\log n) 的。从而,只需要每次求出信息的变化,就可以通过这道题了。

    考虑如何维护,可以对于每一个重链开两个 set:

    1. 储存了若干组对 (x,y)(x,y),表示深度为 xx 的点 kk 轻儿子 yy 与他颜色相同。
    2. 储存了若干个 xx,表示深度为 xx 的点 kk 是它父亲的重儿子,且 kk 与其颜色不同。

    具体的实现建议自己思考,也可以直接参考代码。

    #include<bits/stdc++.h>
    const int N=200050;
    using namespace std;
    
    int n,m;
    
    #define _to e[i].to
    #define fore(k) for(int i=hd[k];i;i=e[i].nx)
    struct edge{
        int to,nx;
    }e[N];int hd[N],S;
    void add(int fr,int to){
        e[++S]=(edge){to,hd[fr]},hd[fr]=S;
    }
    
    int fa[N],dp[N],sg[N],bc[N],sz[N],sn[N],tp[N],cnt;
    void dfs(int k,int f){
        fa[k]=f,dp[k]=dp[f]+1,sz[k]=1;
        fore(k)if(_to!=f)
            dfs(_to,k),sz[k]+=sz[_to],sn[k]=sz[sn[k]]<sz[_to]?_to:sn[k];
    }
    basic_string<int> a[N];
    multiset<int> s1[N];
    multiset<pair<int,int>> s[N];
    void df5(int k,int t){
        tp[k]=t,sg[k]=++cnt,bc[cnt]=k;a[t]+=k;
        if(sn[k])df5(sn[k],t),s1[t].insert(dp[k]-dp[t]+1);
        fore(k)if(_to!=fa[k]&&_to!=sn[k])df5(_to,_to);
    }
    
    struct sgt{
    	#define ls k<<1
    	#define rs k<<1|1
    	#define mid ((l+r)>>1)
    	int ad[N<<2],mx[N<<2];
    	void pshad(int k,int p){
    		ad[k]+=p,mx[k]+=p;
    	}
    	void pshdn(int k){
    		pshad(ls,ad[k]),pshad(rs,ad[k]);
    		ad[k]=0;
    	}
    	void pshup(int k){mx[k]=max(mx[ls],mx[rs]);}
    	void build(int k,int l,int r){
    		if(l^r)build(ls,l,mid),build(rs,mid+1,r),pshup(k);
    		else pshad(k,dp[bc[l]]-1);
    	}
    	void add(int k,int l,int r,int x,int y,int p){
    		if(r<x||l>y)return ;
    		if(l>=x&&r<=y)return pshad(k,p);
    		pshdn(k);
    		add(ls,l,mid,x,y,p),add(rs,mid+1,r,x,y,p);
    		pshup(k);
    	}
    	int qmx(int k,int l,int r,int x,int y){
    		if(r<x||l>y)return 0;
    		if(l>=x&&r<=y)return mx[k];
    		pshdn(k);
    		return max(qmx(ls,l,mid,x,y),qmx(rs,mid+1,r,x,y));
    	} 
    	int qsm(int k,int l,int r,int x){
    		if(l^r){
    			pshdn(k);
    			if(x<=mid)return qsm(ls,l,mid,x);
    			else return qsm(rs,mid+1,r,x);
    		}else return ad[k];
    	}
    }T;
    
    int lca(int x,int y){
    	while(tp[x]!=tp[y]){
    		if(dp[tp[x]]<dp[tp[y]])swap(x,y);
    		x=fa[tp[x]];
    	}
    	return dp[x]<dp[y]?x:y;
    }
    
    int getv(int k){
    	return T.qsm(1,1,n,sg[k]);
    }
    
    //for most of node
    // - his heavy son equal to him
    // - other sons not equal to him
    //s a unheavy son equal to him
    //s1 a heavy son not equal to him
    
    void clear(int x){
    	int pr=0;
    	while(x){
    		int t=tp[x],dl=dp[x]-dp[t],nw;
    		while(s[t].size()&&s[t].begin()->first<=dl){
    			nw=s[t].begin()->second;
    			T.add(1,1,n,sg[nw],sg[nw]+sz[nw]-1,1);
    			s[t].erase(s[t].begin()); 
    		}
    		while(s1[t].size()&&*s1[t].begin()<=dl){
    			nw=a[t][*s1[t].begin()];
    			T.add(1,1,n,sg[nw],sg[nw]+sz[nw]-1,-1);
    			s1[t].erase(s1[t].begin());
    		}
    		if(pr){
    			T.add(1,1,n,sg[pr],sg[pr]+sz[pr]-1,-1);
    			s[t].insert({dl,pr});
    			// 轻儿子 pr 与 x 颜色相同 
    		}
    		if(sn[x]&&s1[t].find(dl+1)==s1[t].end()){
    			T.add(1,1,n,sg[sn[x]],sg[sn[x]]+sz[sn[x]]-1,1);
    			s1[t].insert(dl+1);
    			//重儿子 sn[x] 与 x 颜色不同 
    		}
    		pr=t;
    		x=fa[t];
    	}
    } 
    
    int getmx(int x){
    	return T.qmx(1,1,n,sg[x],sg[x]+sz[x]-1)+1;
    }
    
    void solve(){
    	cin>>n>>m;
    	for(int i=1;i<n;i++){
    		int x,y;cin>>x>>y;
    		add(x,y),add(y,x);
    	}
    	dfs(1,0);
    	df5(1,1);
    	T.build(1,1,n);
    	for(int i=1;i<=m;i++){
    		int x,y,z;cin>>x>>y;
    		if(x==1)clear(y);
    		if(x==2){
    			cin>>z;
    			cout<<getv(y)+getv(z)-getv(lca(y,z))*2+1<<'\n';
    		}
    		if(x==3)cout<<getmx(y)<<'\n';
    	}
    }
    
    main(){
    	ios::sync_with_stdio(false);
    	solve();
    }
    
    
    • 1

    信息

    ID
    6486
    时间
    1000ms
    内存
    256MiB
    难度
    10
    标签
    递交数
    6
    已通过
    4
    上传者