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\) 的概率,转移方程:
对每个 \(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;
}
本文来自博客园,作者:peiwenjun,转载请注明原文链接:https://www.cnblogs.com/peiwenjun/p/16527766.html
浙公网安备 33010602011771号