peiwenjun's blog 没有知识的荒原

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;
}

posted on 2023-01-18 15:21  peiwenjun  阅读(10)  评论(0)    收藏  举报

导航