萌新刚学长链剖分1s,求调(有注释
查看原帖
萌新刚学长链剖分1s,求调(有注释
388414
comcopy楼主2022/10/5 18:54
#include<bits/stdc++.h>
#define int long long
using namespace std;
const int N=500005;
int dep[N],lon[N],hei[N],fa[N][21],hf[N],htop[N],highbit[N];
// dep表示深度,lon表示链长,hei表示长链深度,fa用于倍增跳祖先,hf用于上一级祖先,htop表示链顶,highbit就是高度位
vector<int> g[N],up[N],dn[N];
//g存边,up存祖先,dn存dfn序
void dfs1(int u,int p){
//	cout<<u<<endl;
    dep[u]=dep[p]+1,lon[u]=0,fa[u][0]=p,hf[u]=p;//初始化
    for(int i=1;fa[u][i-1] && i<20;++i,fa[u][i]=fa[fa[u][i-1]][i-1]){fa[u][i]=fa[fa[u][i-1]][i-1];}//处理祖先节点
        for(vector<int>::iterator it=g[u].begin();it!=g[u].end();++it){
        	int v=*it;
            if(v!=p) dfs1(v,u);
            if(!lon[u] || hei[lon[u]]<hei[v]){
                lon[u]=v;
            }
            else
            lon[u]=0;
        }
    
    hei[u]=lon[u]?hei[lon[u]]+1:1;
}

void dfs3(int u,int p,int htp){//与重链剖分一样,优先遍历长儿子
    htop[u]=htp;
    if(u==htp){//若当前点是链顶,预处理出长儿子和祖先
        for(int v=u;v;v=lon[v])
            dn[u].push_back(v);//记录dfn序保证dfn序不变
        for(int v=u;v && up[u].size()<dn[u].size();v=hf[v])
            up[u].push_back(v);
    }
    if(lon[u]) dfs3(lon[u],u,htp);//优先遍历长儿子
    for(vector<int>::iterator it=g[u].begin();it!=g[u].end();++it){//再遍历其余点
    	int v=*it;
        if(v!=p && v!=lon[u])
            dfs3(v,u,v);
    }
}

int kthans(int u,int k){
    if(dep[u]<=k) return (0-0); //如果当前的深度小于等于 k ,说明祖先一定没有 k 级,下去也无意义了,直接返回 0 
    if(k==0) return u;//如果刚好是当前就返回当前 
    u=fa[u][highbit[k]],k-=1<<highbit[k];//找最接近的位数 
    int d=dep[u]-k-dep[htop[u]];//跳祖先,dep[u]-dep[h[top[u]]] 就是祖先与现在的距离,减去个 k 就是祖先与目标定的距离 
    return d>=0?dn[htop[u]][d]:dn[htop[u]][-d];//保证跳的时候 d>=0就行了 
}
int rt;
int ans;

#define ui unsigned int
ui s;

inline ui get(ui x) {
	x ^= x << 13;
	x ^= x >> 17;
	x ^= x << 5;
	return s = x; 
}
int n,m;

signed main(){
	cin>>n>>m;
	cin>>s;
	for(int i=2;i<=n;++i){
		highbit[i]=highbit[i>>1]+1;
	}
	rt=1;
	for(int i=1;i<=n;++i){
		int x;
		cin>>x;
		if(!x) rt=i;
		else g[x].push_back(i);
	}
	dep[0]=0;
	dfs1(rt,0);
	dfs3(rt,rt,rt);
	int lsans(0);
	for(int i=1;i<=m;++i){
		int x=(get(s)^lsans)%n+1,k=(get(s)^lsans)%dep[x];
		lsans=kthans(x,k);
		cout<<x<<' '<<k<<' '<<lsans<<endl;
		ans^=(i*lsans);
	}
	cout<<ans<<endl;
	return (0-0);
}
2022/10/5 18:54
加载中...