LOJ3833 「IOI2022」数字电路 题解
题目描述
给定一棵 \(n+m\) 个节点的树,其中有 \(n\) 个非叶节点和 \(m\) 个叶节点。
每个节点有 \(0\) 和 \(1\) 两种权值,初始所有叶节点的权值已给定。
每个非叶节点 \(u\) 有一个阈值 \(c_u\) ,取值范围为 \(1\sim\) 子节点总数。如果子节点中 \(1\) 的个数\(\ge c_u\) ,那么 \(u\) 的权值为 \(1\) ,否则为 \(0\) 。
接下来 \(q\) 次操作,给定 \(l,r\) ,翻转第 \([l,r]\) 个叶节点权值。
每次操作结束后,求有多少种给非叶节点确定阈值的方法,使得根节点权值为 \(1\) ,对\(10^9+2022\) 取模。
数据范围
- \(n,m,q\le 10^5\) 。
时间限制 \(\texttt{2s}\) ,空间限制 \(\texttt{2GB}\) 。
分析
对于一个非叶节点 \(u\) ,假设它有 \(x\) 个子节点权值为 \(1\) , \(y\) 个为 \(0\) 。
那么共有 \(x\) 种确定 \(c_u\) 的方式使得 \(u\) 的权值为 \(1\) ,这等价于 \(u\) 可以任选一个子节点并继承它的权值。
预处理对于每个叶节点,根节点继承它的权值的方案数,记为 \(f_i\) 。
那么答案为 \(\sum_{i=1}^nf_i\cdot val_i\) ,同时用线段树可以很方便地维护区间翻转。
\(f_i\) 等于所有不在它到根路径上的点,子节点个数的乘积。
由于模数没有逆元,所以在 \(dfs\) 的过程中预处理一下前后缀乘积即可。
时间复杂度 \(\mathcal O(n+m+q\log m)\) 。
#include<bits/stdc++.h>
#include"circuit.h"
#define ls p<<1
#define rs p<<1|1
using namespace std;
const int maxn=2e5+5,mod=1e9+2022;
int m,n,cnt;
int w[maxn],val[maxn];
bool op[maxn];
vector<int> g[maxn];
struct node
{
int l,r,all,sum,tag;
}f[4*maxn];
void dfs1(int u)
{
val[u]=u<=n?g[u].size():1;
for(auto v:g[u]) dfs1(v),val[u]=1ll*val[u]*val[v]%mod;
}
void dfs2(int u,int cur)
{
if(u>n) return w[u-n]=cur,void();
int cnt=g[u].size();
vector<int> pre(cnt+2),suf(cnt+2);
pre[0]=1,suf[cnt+1]=1;
for(int i=1;i<=cnt;i++) pre[i]=1ll*pre[i-1]*val[g[u][i-1]]%mod;
for(int i=cnt;i>=1;i--) suf[i]=1ll*suf[i+1]*val[g[u][i-1]]%mod;
for(int i=1;i<=cnt;i++) dfs2(g[u][i-1],1ll*cur*pre[i-1]%mod*suf[i+1]%mod);
}
void pushup(int p)
{
f[p].sum=(f[ls].sum+f[rs].sum)%mod;
}
void pushtag(int p)
{
f[p].tag^=1,f[p].sum=(f[p].all-f[p].sum+mod)%mod;
}
void pushdown(int p)
{
if(!f[p].tag) return ;
pushtag(ls),pushtag(rs),f[p].tag=0;
}
void build(int p,int l,int r)
{
f[p].l=l,f[p].r=r;
if(l==r) return f[p].all=w[l],f[p].sum=op[l]*w[l],void();
int mid=(l+r)/2;
build(ls,l,mid);
build(rs,mid+1,r);
f[p].all=(f[ls].all+f[rs].all)%mod,pushup(p);
}
void modify(int p,int l,int r)
{
if(l<=f[p].l&&f[p].r<=r) return pushtag(p);
if(l>f[p].r||r<f[p].l) return ;
pushdown(p);
modify(ls,l,r);
modify(rs,l,r);
pushup(p);
}
void init(int _n,int _m,vector<int> p,vector<int> a)
{
n=_n,m=_m;
for(int i=2;i<=n+m;i++) g[p[i-1]+1].push_back(i);
dfs1(1),dfs2(1,1);
for(int i=1;i<=m;i++) op[i]=a[i-1];
build(1,1,m);
}
int count_ways(int l,int r)
{
modify(1,l-n+1,r-n+1);
return f[1].sum;
}
本文来自博客园,作者:peiwenjun,转载请注明原文链接:https://www.cnblogs.com/peiwenjun/p/17059909.html
浙公网安备 33010602011771号