全局平衡二叉树

题解1
2
3
推荐理解
推荐代码+理解

本质上是静态的 \(LCT\) 。 认父不认子。 维护后继。 每个重链 build 。 轻儿子的fa继承上来。

保留跳/修改是对的。

!!!!!!!!!一定要记得, 访问子树用的是 rt[x]!!!!!!!!!! rt[v]

ANIG 的代码:


ANIG

#include <bits/stdc++.h>
using namespace std;
#define int long long
#define ll __int128
const int N=5e5+5;
ll inf=1e18;
int t,n,m,fa[N],w[N],sz[N],siz[N],zx[N],a[N],b[N],dep[N],son[N][2],dy[N],nxt[N];
vector<int>p[N];
int exgcd(int a,int b,int &x,int &y){
	if(b==0){x=1,y=0;return a;}
	int d=exgcd(b,a%b,x,y);
	int z=x;x=y,y=z-y*(a/b);
	return d;
}
void solve(ll a,ll b,int &x,int &m){
    int bx,by;
    int k=exgcd(m,b,bx,by);
    int sa=a,sb=b;
    b/=k;
    bx=(((ll)bx)*((a-x)/k)%b+b)%b;
    ll tx=x;
    tx+=((ll)(bx))*m;
    if((tx%m+m)%m==x&&(tx%sb+sb)%sb==sa){
        ll tms=m*b;
        tx=(tx%tms+tms)%tms;
        if(tx>inf)x=inf;
        else x=tx;
        m=min(tms,inf+1);
    }else{
        m=inf+1;x=inf;
    }
}
void dfs(int x){
    sz[x]=1;zx[x]=0;
    int v=p[x].size();
    for(int i=0;i<p[x].size();i++){
        int c=p[x][i];
        dep[c]=dep[x]+w[c];
        a[c]=a[x];b[c]=b[x];
        solve(((i-dep[x])%v+v)%v,v,b[c],a[c]);
        dfs(c);
        sz[x]+=sz[c];
        if(sz[c]>sz[zx[x]])zx[x]=c;
    }
}
int reset(vector<int>&jl,int l,int r){
    if(l>r)return 0;
    int mid=l+r>>1,he=0;
    for(int i=l;i<=r;i++)he+=siz[i];
    for(int i=l,j=0;i<r;i++){
        j+=siz[i];
        if(2*j>=he){
            mid=i;
            break;
        }
    }
    son[jl[mid]][0]=reset(jl,l,mid-1);
    son[jl[mid]][1]=reset(jl,mid+1,r);
    return jl[mid];
}
void build(int x){
    vector<int>jl;
    for(int y=x;y;y=zx[y]){
        jl.push_back(y);
        for(auto c:p[y]){
            if(c==zx[y])continue;
            build(c);
        }
    }
    for(int i=0;i+1<jl.size();i++)siz[i]=sz[jl[i]]-sz[jl[i+1]];
    siz[jl.size()-1]=sz[jl.back()];
    for(int i=0;i+1<jl.size();i++)nxt[jl[i]]=jl[i+1];
    dy[x]=reset(jl,0,jl.size()-1);
}
int solve(int x,int y){
    if(!x)return 0;
    if(y%a[x]==b[x]){
        if(!p[x].size())return x;
        int ans=0;
        if(y%a[nxt[x]]==b[nxt[x]])ans=solve(son[x][1],y);
        if(ans)return ans;
        return solve(dy[p[x][(y+dep[x])%p[x].size()]],y);
    }
    return solve(son[x][0],y);
}
signed main(){
    ios::sync_with_stdio(0);
    cin.tie(0);cout.tie(0);
    cin>>t;inf++;
    while(t--){
        cin>>n>>m;
        for(int i=1;i<=n;i++)p[i].clear(),son[i][0]=son[i][1]=0;
        for(int i=2;i<=n;i++)cin>>fa[i],p[fa[i]].push_back(i);
        for(int i=2;i<=n;i++)cin>>w[i];
        a[1]=1;b[1]=0;
        dfs(1);build(1);
        while(m--){
            int x;
            cin>>x;
            cout<<solve(dy[1],x)<<" ";
        }
        cout<<"\n";
    }
    cerr<<clock()*1.0/CLOCKS_PER_SEC;
}
mycode:

#include<bits/stdc++.h>
#define int long long
#define siz(s) ((int)s.size())
#define int long long
#define pb push_back
#define P pair<int,int>
#define fi first
#define se second
#define el '\n'
#define a(x) array<int,x>
#define de(x) cerr<<#x<<'='<<x<<' '
#define i128 __int128
using namespace std;
bool stm;
template<class A> inline int cn(A&x,A y){
    if(y<x){
        x=y;return 1;
    }
    return 0;
}
template<class A> inline int cm(A&x,A y){
    if(y>x){
        x=y;return 1;
    }
    return 0;
}
const int N=1e6+10,inf=4e18;
int n,fa[N],w[N],q,s[N],son[N],sz[N],top[N],d[N],id[N];
int ch[N][2],rt[N],nxt[N],S[N];
struct msg{ 
    int m,x;
    inline int ck(int t){
        if(!m)return 0;
        if(t%m!=x)return 0;
        return 1;
    }
}a[N];  
vector<int> p[N],e[N];
i128 exgcd(i128 a,i128 b,i128 &x,i128 &y){
    if(!b){
        x=1;y=0;
        return a;   
    }
    i128 g=exgcd(b,a%b,y,x);
    y-=(a/b)*x; // ax+by = bx'+(a-br)y' // ax+by=ay' + b(x'-ry')
    return g;
}
msg mrg(msg a,msg b){
    if(a.m==0||b.m==0)return a;
    i128 m1=a.m,m2=b.m,x1=a.x,x2=b.x;
    //x1+m1x = x2+m2y  m1x+m2y=x2-x1  
    i128 x=0,y=0;
    i128 g=exgcd(m1,m2,x,y);
    if((x2-x1)%g)return {0,0};
    i128 mod=m1/g*m2;
    x*=(x2-x1)/g;x%=mod;
    i128 rs=((x1+m1*x)%mod+mod)%mod;
    if(rs>inf)return {0,0};
    // cn(mod,inf);
    if(mod>inf)mod=inf;
    return {mod,rs};
}
void dfs2(int u,int t){
    top[u]=t;p[t].pb(u);
    if(u==t){
        s[u]=0;
        a[u]={1,0};
    }else{
        s[u]=s[fa[u]]+w[u];
        a[u]=mrg(a[fa[u]],{d[fa[u]],(id[u]-s[fa[u]])%d[fa[u]]});//x+s[u]=id
    }
    if(!son[u])return;
    dfs2(son[u],t);
    for(auto v:e[u]){
        if(v==son[u])continue;
        dfs2(v,v);
    }
}   
void dfs1(int u){
    sz[u]=1;son[u]=0;
    int t=0;
    for(auto v:e[u]){
        id[v]=t++;
        dfs1(v);
        if(!son[u]||sz[son[u]]<sz[v])son[u]=v;
        sz[u]+=sz[v];
    }
}
int reset(const vector<int>&p,int l,int r){
    if(l>r)return 0;
    int sum=0,md=l,rs=0;
    for(int i=l;i<=r;i++)sum+=S[i];
    for(md=l;md<=r;md++){
        rs+=S[md];
        if(rs*2>=sum)break;
    }
    int u=p[md];
    ch[u][0]=reset(p,l,md-1);
    ch[u][1]=reset(p,md+1,r);
    // if(!ch[u][1]&&nxt[u]){
    //     de(nxt[u]);de(md);de(r);
    // }
    return u;
}
void bd(int x){
    vector<int> jl;
    for(int i=x;i;i=son[i]){
        jl.pb(i);
        for(auto v:e[i]){
            if(v==fa[i]||v==son[i])continue;
            bd(v);
        }
    }
    for(int i=0;i<siz(jl);i++)S[i]=sz[jl[i]]-sz[son[jl[i]]];
    for(int i=0;i+1<siz(jl);i++)nxt[jl[i]]=jl[i+1];
    // if(x==1){
    //     for(auto nw:jl)cout<<nw<<' ';
    //     cout<<el;
    // }
    nxt[jl.back()]=0;
    rt[x]=reset(jl,0,siz(jl)-1);
}
// int solve(int x,int m){
//     if(!x)return 0;
//     if(a[x].ck(m)){
//         if(!e[x].size())return x;
//         int ans=0;
//         if(a[nxt[x]].ck(m))ans=solve(ch[x][1],m);
//         if(ans)return ans;
//         return solve([p[x][(y+dep[x])%p[x].size()]],y);
//     }
//     return solve(son[x][0],y);
// }
int solve(int x,int m){
    // de(x);de(ch[x][0]);de(ch[x][1])<<el;
    // assert(x);
    if(!x)return 0;
    // if(!ch[x][0])assert(a[x].ck(m));
    if(!a[x].ck(m))return solve(ch[x][0],m);
    else{
        if(!siz(e[x]))return x;
        assert(nxt[x]);
        // assert(ch[x][1]);
        if(a[nxt[x]].ck(m)){
            int ans=solve(ch[x][1],m);
            if(ans)return ans;
        }
        int v=e[x][(m+s[x])%d[x]];
        return solve(rt[v],m+s[x]+w[v]); //!!!!!!!!!! rt[v]!!!!!!!!!!!
    }
}
void INIT(int lm){
    // int ch[N][2],rt[N],nxt[N],S[N];
    for(int i=0;i<=lm+2;i++){
        s[i]=son[i]=sz[i]=top[i]=d[i]=id[i]=ch[i][0]=ch[i][1]=rt[i]=nxt[i]=S[i]=0;
        a[i]={1,0};
        e[i].clear();p[i].clear();
    }
    // ...
}
void solve(){
    cin>>n>>q;
    INIT(n);
    for(int i=2;i<=n;i++){
        cin>>fa[i];e[fa[i]].pb(i);
    }
    for(int i=1;i<=n;i++)d[i]=siz(e[i]);
    for(int i=2;i<=n;i++)cin>>w[i];
    dfs1(1);dfs2(1,1);
    bd(1);
    while(q--){
        int m;
        cin>>m;
        int ans=solve(rt[1],m); //!!!!!!! 链接上来的是rt[i]!!!!!!!
        assert(ans);
        cout<<ans<<' ';
    }
    cout<<el;
}
bool edm;
signed main(){
    ios::sync_with_stdio(0);
    cin.tie(0);cout.tie(0);
    cerr<<&edm-&stm<<el;
    if(abs(&edm-&stm)>200000000)cerr<<"MLE!!"<<el;
    int stc=clock();
    int T;cin>>T;
    while(T--)solve();
    cerr<<el<<clock()-stc<<el;
    return 0;
}
/*
1
10 1
1 2 2 2 1 1 3 4 5
1 2 3 4 5 6 7 8 9
3

*/

posted @ 2026-07-06 23:19  Str_ywr  阅读(3)  评论(0)    收藏  举报