*题解:P5298 [PKUWC2018] Minimax
解析
首先注意读题时不要漏条件,题目说的是每个结点最多有两个子结点。
设 \(f_{i,j}\) 表示点 \(i\) 的权值为 \(j\) 的概率,那么如果只考虑左儿子:
\[f_{i,j}=p_i \cdot f_{ls,j} \cdot \sum_{k < j}f_{rs,k} +(1 - p_i) \cdot f_{ls_,j} \cdot \sum_{k > j} f_{rs,k}
\]
合并同类项得:
\[f_{i,j}=f_{ls,j}(p_i\cdot \sum_{k < j}f_{rs,k} +(1 - p_i) \cdot \sum_{k > j} f_{rs,k})
\]
分别表示 \(j\) 作为最大值和最小值的情况。同理可推得右儿子的转移式。
考虑如何优化转移。首先肯定要将权值离散化,将求和变为求前后缀和。然而这样还不够,我们发现题目的时空限制不支持我们枚举每一个权值,所以考虑线段树合并好吧我知道这很牵强,刷题量太少总结不出经验导致的。
合并过程中主要要求出 \(p_i\cdot \sum_{k < j}f_{rs,k} +(1 - p_i) \cdot \sum_{k > j} f_{rs,k}\),即一段前后缀的权值。考虑线段树往下分治的过程,设当前所在区间为 \([l,r]\),且 \([1,l)\) 和 \((r,n]\) 的权值已经求好,接下来就只需要根据 \(j\) 的位置来决定新加入的区间是 \([l,mid]\) 还是 \([mid + 1,r]\)。在左半边就加入后者,在右半边就加入前者。两棵树合并时如果走到一棵树独有的点,对其打乘法标记即可,其值就为求出的前后缀权值。
时间复杂度 \(O(n \log n)\)。
代码
#include <bits/stdc++.h>
#define mid ((l + r) >> 1)
using namespace std;
const int N = 3e5 + 5,M = 20,mod = 998244353;
typedef long long ll;
typedef pair<int,int> pii;
vector<int> son[N];
int p[N],rt[N];
int m;
int sum[N * M],mul[N * M],ls[N * M],rs[N * M],cnt;
vector<int> val;
int qmi(int a,int b){
int res = 1;
while(b){
if(b & 1) res = 1ll * res * a % mod;
b >>= 1;
a = 1ll * a * a % mod;
}
return res;
}
int inv(int a){
return qmi(a,mod - 2);
}
void push_up(int p){
sum[p] = (sum[ls[p]] + sum[rs[p]]) % mod;
}
void add_tag(int p,int k){
if(!k) return;
sum[p] = 1ll * sum[p] * k % mod;
mul[p] = 1ll * mul[p] * k % mod;
}
void push_down(int p){
if(mul[p] == 1) return;
add_tag(ls[p],mul[p]);
add_tag(rs[p],mul[p]);
mul[p] = 1;
}
int ask(int p,int l,int r,int L,int R){
if(!p || l > R || r < L){
return 0;
}
if(l >= L && r <= R){
return sum[p];
}
push_down(p);
return (ask(ls[p],l,mid,L,R) + ask(rs[p],mid + 1,r,L,R)) % mod;
}
int merge(int x,int y,int l,int r,int vx,int vy,int u){
if(!x && !y) return 0;
if(!x){
add_tag(y,vy);
return y;
}
if(!y){
add_tag(x,vx);
return x;
}
push_down(x),push_down(y);
int lsm[2] = {sum[ls[x]],sum[ls[y]]},rsm[2] = {sum[rs[x]],sum[rs[y]]};
ls[x] = merge(ls[x],ls[y],l,mid,(vx + 1ll * (1 + mod - p[u]) * rsm[1] % mod) % mod,(vy + 1ll * (1 + mod - p[u]) * rsm[0] % mod) % mod,u);
rs[x] = merge(rs[x],rs[y],mid + 1,r,(vx + 1ll * p[u] * lsm[1] % mod) % mod,(vy + 1ll * p[u] * lsm[0] % mod) % mod,u);
push_up(x);
return x;
}
void add(int &p,int l,int r,int k){
if(l > k || r < k) return;
if(!p) p = ++cnt,mul[p] = 1;
if(l == r){
sum[p] = 1;
return;
}
push_down(p);
add(ls[p],l,mid,k);
add(rs[p],mid + 1,r,k);
push_up(p);
}
void dfs(int x){
if(son[x].empty())
add(rt[x],1,m,p[x]);
for(int i=0;i<son[x].size();i++){
dfs(son[x][i]);
rt[x] = merge(rt[x],rt[son[x][i]],1,m,0,0,x);
}
}
int main(){
ios::sync_with_stdio(false);
cin.tie(0);
// freopen("in.txt","r",stdin);
// freopen("out1.txt","w",stdout);
int n;
cin>>n;
for(int i=1;i<=n;i++){
int x;
cin>>x;
son[x].push_back(i);
}
int iv = inv(10000);
for(int i=1;i<=n;i++){
int x;
cin>>x;
if(son[i].empty()){
p[i] = x;
val.push_back(x);
}else{
p[i] = 1ll * x * iv % mod;
}
}
sort(val.begin(),val.end());
m = val.size();
for(int i=1;i<=n;i++){
if(son[i].empty()){
p[i] = lower_bound(val.begin(),val.end(),p[i]) - val.begin() + 1;
}
}
dfs(1);
int res = 0;
for(int i=1;i<=val.size();i++){
int x = ask(rt[1],1,val.size(),i,i);
// cout<<i<<" "<<val[i - 1]<<" "<<x<<'\n';
res = (1ll * i * val[i - 1] % mod * qmi(x,2) % mod + res) % mod;
}
cout<<res;
return 0;
}

浙公网安备 33010602011771号