顺着第一篇题解的思路写的,但不知道哪里错了。
#include<bits/stdc++.h>
#define int long long
using namespace std;
inline int read()
{
int s=0,w=1;
char c=getchar();
while(c<'0'||c>'9')
{
if(c=='-')w=-1;
c=getchar();
}
while(c>='0'&&c<='9')s=(s<<3)+(s<<1)+(c^48),c=getchar();
return s*w;
}
inline void print(int x)
{
if(x<0)x=-x,putchar('-');
if(x>=10)print(x/10);
putchar(x%10+48);
}
struct node{
int lc,rc,sum,mul;//没有加操作,这不得偷个懒。
}C[10000010];
const int mod=998244353;
inline void exgcd(int a,int b,int &x,int &y)
{
if(!b)
{
x=1,y=0;
return;
}
exgcd(b,a%b,x,y);
int tmp=x;
x=y;
y=tmp-(a/b)*y;
}
inline int inv(int a)
{
int x=0,y=0;
exgcd(a,mod,x,y);
x%=mod;
x+=mod;
x%=mod;
return x;
}
int fa[1000010],rt[1000010],p[1000010],num[1000010],son[1000010][2],d[1000010],cnt,n,tot=0;
inline void pushup(int id)
{
C[id].sum=C[C[id].lc].sum+C[C[id].rc].sum;
C[id].sum%=mod;
}
inline void update(int &id,int l,int r,int x,int v)
{
if(!id)id=++tot,C[id].mul=1;
if(l==r)
{
C[id].sum=v;
return;
}
int mid=l+r>>1;
if(x<=mid)
update(C[id].lc,l,mid,x,v);
else update(C[id].rc,mid+1,r,x,v);
pushup(id);
}
inline void mu(int id,int b)
{
if(!id)return;
C[id].sum*=b;
C[id].mul*=b;
C[id].sum%=mod;
C[id].mul%=mod;
}
inline void pushdown(int id)
{
if(C[id].mul==1)return;
if(C[id].lc)
mu(C[id].lc,C[id].mul);
if(C[id].rc)
mu(C[id].rc,C[id].mul);
C[id].mul=1;
}
inline int merge(int a,int b,int l,int r,int xmul,int ymul,int v)
{
if(!a)
{
mu(b,ymul);
return b;
}
if(!b)
{
mu(a,xmul);
return a;
}
pushdown(a);
pushdown(b);
int mid=l+r>>1;
int lsx=C[C[a].lc].sum,lsy=C[C[b].lc].sum,rsx=C[C[a].rc].sum,rsy=C[C[b].rc].sum;
C[a].lc=merge(C[a].lc,C[b].lc,l,mid,(xmul+rsy*(1-v+mod)%mod)%mod,(ymul+rsx*(1-v+mod))%mod,v);
C[a].rc=merge(C[a].rc,C[b].rc,mid+1,r,(xmul+lsy*v%mod)%mod,(ymul+lsx*v%mod)%mod,v);
pushup(a);
return a;
}
inline int query(int id,int l,int r,int x)
{
if(!id)return 0;
if(l==r)return C[id].sum;
pushdown(id);
int mid=l+r>>1;
if(x<=mid)return query(C[id].lc,l,mid,x);
else return query(C[id].rc,mid+1,r,x);
}
inline void dfs(int u)
{
if(!num[u])
{
update(rt[u],1,cnt,p[u],1);
return;
}
if(num[u]==1)
{
dfs(son[u][0]);
rt[u]=rt[son[u][0]];
return;
}
if(num[u]==2)
{
dfs(son[u][0]);
dfs(son[u][1]);
rt[u]=merge(rt[son[u][0]],rt[son[u][1]],1,cnt,0,0,p[u]);
}
}
signed main()
{
n=read();
for(int i=1;i<=n;i++)
{
fa[i]=read();
if(fa[i])son[fa[i]][num[fa[i]]++]=i;
}
int c=inv(10000);
for(int i=1;i<=n;i++)
{
p[i]=read();
if(num[i])
{
p[i]*=c;
p[i]%=mod;
}
else
{
d[++cnt]=p[i];
}
}
sort(d+1,d+cnt+1);
for(int i=1;i<=n;i++)
{
if(!num[i])
p[i]=lower_bound(d+1,d+n+1,p[i])-d;
}
dfs(1);
int ans=0;
for(int i=1;i<=cnt;i++)
{
ans+=i*d[i]%mod*query(rt[1],1,cnt,i)%mod*query(rt[1],1,cnt,i)%mod;
ans%=mod;
}
print(ans);
}