未过样例,树链剖分求调qwq
查看原帖
未过样例,树链剖分求调qwq
760859
Let_Fly楼主2023/8/2 09:17
#include<bits/stdc++.h>
using namespace std;
#define int long long
const int N=5e5+5,mod=998244353;

int n,q,k,cnt;
vector<int> to[N];
int mi[N];
struct ques {
    int u, z, id;
    bool operator < (const ques x) const {return u < x.u;}
}que[N];
struct Tree{
    int tag,sum,cs;
}tr[N<<2];

int dep[N];
int fa[N];
int son[N];
int sz[N];
int top[N];
int id[N];
int fid[N];
// int a[N];
// int b[N];

int ans[N];

void dfs1(int u,int f){
    fa[u]=f;
    dep[u]=dep[f]+1;
    sz[u]=1;
    for(auto v:to[u]){
        if(v==f)continue;
        dfs1(v,u);
        sz[u]+=sz[v];
        if(sz[son[u]]<sz[v])son[u]=v;
    }
}

void dfs2(int u,int t){
    id[u]=++cnt;
    fid[cnt]=u;
    // a[cnt]=b[u];
	top[u]=t;
	if(son[u])dfs2(son[u],t);
	for(auto v:to[u]){
		if(v==fa[u]||v==son[u])continue;
		dfs2(v,v);
	}
}

int qpow(int a,int b){
    int r=1;
    while(b){
        if(b&1)r=r*a%mod;
        b>>=1,a=a*a%mod;
    }
    return r;
}

void pushup(int u){
    tr[u].sum=(tr[u<<1].sum+tr[u<<1|1].sum)%mod;
}

void pushdown(int u){
    if(!tr[u].tag)return;
    tr[u<<1].sum=(tr[u<<1].sum+(tr[u].tag*tr[u<<1].cs)%mod)%mod;
    tr[u<<1|1].sum=(tr[u<<1|1].sum+(tr[u].tag*tr[u<<1|1].cs)%mod)%mod;
    tr[u<<1].tag+=tr[u].tag%mod;
    tr[u<<1|1].tag+=tr[u].tag%mod;
    tr[u].tag=0;
}

void build(int u,int l,int r){
    if(l==r){
        tr[u].cs=(mi[dep[fid[l]]]-mi[dep[fid[l]]-1]+mod)%mod;
        // cout<<tr[u].cs;
        return;
    }
    int mid=l+r>>1;
    build(u<<1,l,mid);
    build(u<<1|1,mid+1,r);
    tr[u].cs=(tr[u<<1].cs+tr[u<<1|1].cs+mod)%mod;
}

void update(int u,int l,int r,int L,int R){
    if(l<=L&&R<=r){
        tr[u].sum=(tr[u].sum+tr[u].cs%mod)%mod;
        tr[u].tag++;
        return;
    }
    int mid=L+R>>1;
    pushdown(u);
    if(l<=mid)update(u<<1,l,r,L,mid);
    if(r>mid)update(u<<1|1,l,r,mid+1,R);
    pushdown(u<<1);
    pushdown(u<<1|1);
    pushup(u);
}

int query(int u,int l,int r,int L,int R){
    pushdown(u);
    if(l<=L&&R<=r){
        return tr[u].sum;
    }
    int res=0;
    int mid=L+R>>1;
    pushdown(u);
    if(l<=mid)res+=query(u<<1,l,r,L,mid);
    if(r>mid)res+=query(u<<1|1,l,r,mid+1,r);
    pushup(u);
    return res;
}

void ud(int u){
    while(top[u]) update(1,id[top[u]],id[u],1,n),u=fa[top[u]];
}

int qr(int u){
    int ans = 0;
    while(top[u]) ans=(ans+query(1,id[top[u]],id[u],1,n))%mod,u=fa[top[u]];
    return ans;
}

signed main(){
    cin>>n>>q>>k;
    k=k%(mod-1);
    for(int i=1;i<=n;i++)mi[i]=qpow(i,k);
    for(int i=2;i<=n;i++){
        int f;
        cin>>f;
        to[f].push_back(i);
    }
    for(int i=1;i<=q;i++){
        int x,y;
        cin>>x>>y;
        que[i]={x,y,i};
    }
    dfs1(1,0);
    dfs2(1,1);
    build(1,1,n);
    sort(que+1,que+q+1);
    int nw=1;
    for(int i=1;i<=n;i++){
        ud(i);
        // for(int i=1;i<=n;i++){
        //     cout<<"now sum "<<tr[i].sum<<'\n';
        // }
        // cout<<endl;
        // cout<<tr[i].sum<<'\n';
        while(que[nw].u<=i&&nw<=q){
            ans[que[nw].id]+=qr(que[nw].z);
            nw++;
        }
    }
    // ans[1]=query(1,2,4,1,n);
    // for(int i=1;i<=n;i++){
    //     cout<<"id is "<<tr[i].sum<<'\n';
    // }
    for(int i=1;i<=q;i++){
        cout<<ans[i]<<'\n';
    }
    return 0;
}
2023/8/2 09:17
加载中...