wa on 7代码求调
查看原帖
wa on 7代码求调
378346
expnoi楼主2023/6/12 13:56

顺着第一篇题解的思路写的,但不知道哪里错了。

#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);
}
2023/6/12 13:56
加载中...