peiwenjun's blog 没有知识的荒原

P5298 [PKUWC2018]Minimax

题目描述

一棵 \(n\) 个点的有根树, \(1\) 号点为根,每个叶节点有权值,保证所有叶节点权值不同。

每个非叶节点 \(u\) 至多有 \(2\) 个子节点。 \(u\) 点权值有 \(p_u\) 概率是子节点权值的最大值,有 \(1-p_u\) 概率是子节点权值的最小值。

对每个 \(i\) ,求最终 \(1\) 号点权值为第 \(i\) 小的叶节点权值的概率。

数据范围

  • \(1\le n\le 3\cdot 10^5\) ,叶子个数 \(\le 10^4\) 。

时间限制 \(\texttt{1s}\) ,空间限制 \(\texttt{500MB}\) 。

分析

线段树合并优化 \(\texttt{dp}\) 。

先离散化, \(dp_{u,i}\) 表示 \(u\) 点权值为 \(i\) 的概率,转移方程:

\[dp_{u,i}=dp_{lc,i}\cdot(p_u\sum_{j\lt i}dp_{rc,j}+(1-p_u)\sum_{j\gt i}dp_{rc,j}) +dp_{rc,i}\cdot(p_u\sum_{j\lt i}dp_{lc,j}+(1-p_u)\sum_{j\gt i}dp_{lc,j})\\ \]

对每个 \(u\) 开一棵线段树维护 \(dp_{u,i}\) 。

如果 \(u\) 只有一个子节点 \(v\) ,直接继承 \(v\) 的线段树即可。

否则需要线段树合并,假设当前递归的两棵子树为 \(x,y\) ,值域 \([l,r]\) 。

我们希望维护 \(p_u+\sum_{j\lt l}dp_{rc,j}+(1-p_u)\sum_{j\gt r}dp_{rc,j}\) 的值, \(lc\) 同理。递归到叶子 \([i,i]\) 时用 \(dp_{lc/rc,i}\) 更新 \(dp_{u,i}\) 。

往左走就加上右边的贡献,往右走就加上左边的贡献。

int merge(int x,int y,int l,int r,int p1,int p2,const int &p)
{
    if(!x&&!y) return 0;///注意是&&
    if(!x) return pushmul(y,p2),y;
    if(!y) return pushmul(x,p1),x;
    int mid=(l+r)/2;
    pushdown(x),pushdown(y);
    int xl=f[f[x].ls].sum,xr=f[f[x].rs].sum,yl=f[f[y].ls].sum,yr=f[f[y].rs].sum;
    ///开long long!!!
    f[x].ls=merge(f[x].ls,f[y].ls,l,mid,(p1+(mod+1ll-p)*yr)%mod,(p2+(mod+1ll-p)*xr)%mod,p);
    f[x].rs=merge(f[x].rs,f[y].rs,mid+1,r,(p1+1ll*p*yl)%mod,(p2+1ll*p*xl)%mod,p);
    return pushup(x),x;
}

时间复杂度 \(\mathcal O(n\log n)\) 。

#include<bits/stdc++.h>
using namespace std;
const int maxn=3e5+5,mod=998244353;
int n,cnt,res,tot;
int c[maxn],p[maxn],rt[maxn];
vector<int> g[maxn];
struct node
{
    int ls,rs,mul,sum;
}f[20*maxn];
int read()
{
    int q=0;char ch=getchar();
    while(!isdigit(ch)) ch=getchar();
    while(isdigit(ch)) q=10*q+ch-'0',ch=getchar();
    return q;
}
int qpow(int a,int k)
{
    int ans=1;
    while(k)
    {
        if(k&1) ans=1ll*ans*a%mod;
        a=1ll*a*a%mod,k/=2;
    }
    return ans;
}
int newnode()
{
    f[++tot].mul=1;
    return tot;
}
void pushmul(int p,int v)
{
    if(!p) return ;
    f[p].mul=1ll*f[p].mul*v%mod,f[p].sum=1ll*f[p].sum*v%mod;
}
void pushdown(int p)
{
    if(f[p].mul==1) return ;
    pushmul(f[p].ls,f[p].mul),pushmul(f[p].rs,f[p].mul),f[p].mul=1;
}
void pushup(int p)
{
    f[p].sum=(f[f[p].ls].sum+f[f[p].rs].sum)%mod;
}
void insert(int &p,int l,int r,int pos,int val)
{
    if(!p) p=newnode();
    if(l==r) return f[p].sum=val,void();
    int mid=(l+r)/2;
    pushdown(p);
    if(pos<=mid) insert(f[p].ls,l,mid,pos,val);
    else insert(f[p].rs,mid+1,r,pos,val);
    pushup(p);
}
int merge(int x,int y,int l,int r,int p1,int p2,const int &p)
{
    if(!x&&!y) return 0;
    if(!x) return pushmul(y,p2),y;
    if(!y) return pushmul(x,p1),x;
    int mid=(l+r)/2;
    pushdown(x),pushdown(y);
    int xl=f[f[x].ls].sum,xr=f[f[x].rs].sum,yl=f[f[y].ls].sum,yr=f[f[y].rs].sum;
    f[x].ls=merge(f[x].ls,f[y].ls,l,mid,(p1+(mod+1ll-p)*yr)%mod,(p2+(mod+1ll-p)*xr)%mod,p);
    f[x].rs=merge(f[x].rs,f[y].rs,mid+1,r,(p1+1ll*p*yl)%mod,(p2+1ll*p*xl)%mod,p);
    return pushup(x),x;
}
void dfs1(int u)
{
    if(g[u].empty()) return ;
    if(g[u].size()==1) dfs1(g[u][0]),rt[u]=rt[g[u][0]];
    if(g[u].size()==2)
    {
        dfs1(g[u][0]),dfs1(g[u][1]);
        rt[u]=merge(rt[g[u][0]],rt[g[u][1]],1,cnt,0,0,p[u]);
    }
}
void dfs2(int p,int l,int r)
{
    if(l==r) return res=(res+1ll*l*c[l]%mod*f[p].sum%mod*f[p].sum)%mod,void();
    int mid=(l+r)/2;
    pushdown(p);
    dfs2(f[p].ls,l,mid),dfs2(f[p].rs,mid+1,r);
}
int main()
{
    n=read();
    for(int i=1;i<=n;i++) g[read()].push_back(i);
    for(int i=1;i<=n;i++)
    {
        p[i]=read();
        if(g[i].empty()) c[++cnt]=p[i];
        else p[i]=1ll*p[i]*qpow(10000,mod-2)%mod;
    }
    sort(c+1,c+cnt+1);
    for(int i=1;i<=n;i++)
        if(g[i].empty())
        {
            p[i]=lower_bound(c+1,c+cnt+1,p[i])-c;
            insert(rt[i],1,cnt,p[i],1);
        }
    dfs1(1),dfs2(rt[1],1,cnt);
    printf("%d\n",res);
    return 0;
}

posted on 2022-07-28 10:42  peiwenjun  阅读(20)  评论(0)    收藏  举报

导航