题解:P16084 [ICPC 2024 NAC] Comparator

· · 题解

题意简述

一个比较函数依次执行若干规则。每条规则读取两个 k 位字各自指定的一位,计算布尔表达式;条件为真时立即返回指定结果,否则继续。全部条件均为假时返回默认值。

统计函数对全部字的一元组、有序二元组和有序三元组中,分别违反以下条件的数量:f(x,x)=0f(x,y)=1 时应有 f(y,x)=0f(x,y)=f(y,z)=1 时应有 f(x,z)=1

解题思路

最多有 2^{10}=1024 个不同的字,可以完整求出比较关系。瓶颈是对每个字对都执行至多 2\times 10^5 条规则,因此先压缩表达式和规则,再统计关系矩阵中的违例。

每个表达式实际仅接收两个比特 x,y,输入只有 (0,0),(0,1),(1,0),(1,1) 四种可能。按这个顺序,将四种输入的结果依次存入一个四位整数的第 0\sim 3 位,就能完整表示表达式。

于是,变量 x 对应掩码 12,变量 y 对应 10,常量 0 对应 0,常量 1 对应 15。注意常量真要让四种输入都为真,不能编码为整数 1

布尔运算可以一次作用于四种输入:与、或、异或直接使用对应的按位运算;非运算异或 15;相等运算先按位异或,再异或 15。运算后仍是四位真值表,不需要分别解析四次表达式。

使用数值栈和运算符栈求值。优先级必须按题面设置为 !=&|^ 依次降低,不能使用 C++ 自身的优先级。遇到二元运算符时,先计算栈顶优先级不低于它的运算符,从而实现二元运算的左结合;遇到前缀 ! 时直接入栈,使连续的非运算从内向外执行。左括号阻断栈顶运算,右括号将对应括号内的运算全部完成。

两个栈都显式存储,没有递归调用,能够处理长度接近 10^6 的连续非运算或深层括号。每个字符仅入栈、出栈常数次,表达式处理的总开销与输入总长度成正比。

再压缩规则。固定第一、第二个字读取的位置 a,b,以及这两个比特的赋值 t=2x+y,记录满足这一赋值的最早规则编号及其返回值,分别保存在 fst[a][b][t]val[a][b][t] 中。每读入一条规则,检查其真值表的四位,仅填写尚未出现过的情况。

对于同一组 (a,b,t),更晚的规则不可能成为函数实际执行时的首次触发:只要当前输入使它的条件成立,已经保留的更早规则也成立,并会提前返回。因此每组仅需保留一条,全部至多有 4k^2 组。

这里必须同时保留返回 0 的规则。条件为真且返回 0,仍会立即结束函数;它与条件不成立、继续执行后续规则完全不同。

枚举两个字 u,v,遍历全部 k^2 个位置对。对于每组位置,从两个字中提取比特,查找对应赋值的最早规则,最后选择其中编号最小的一条。原规则序列的所有可触发规则都属于某个位置对,所以各组最早规则中的最小者,就是原程序真正执行的第一条规则。若所有组均不存在匹配规则,使用默认返回值。

字从左到右按 1\sim k 编号,代码用整数 0\sim 2^k-1 表示字,因此位置 a 对应从低位开始的第 k-a 位。前导零也自然包含在这个表示中。

M=2^k,求出全部 f(u,v) 后,构造两个位集。G_u 的第 v 位表示 f(u,v)H_u 的第 v 位表示 f(v,u),即 H 保存关系矩阵的转置。

第一项违例是 f(u,u)=1,直接累加矩阵对角线。

第二项违例要求 f(u,v)=f(v,u)=1。固定 u,满足条件的 v 正好是 G_uH_u 的交集,所以累加 G[u]&H[u] 中的置位数即可。题目统计有序对,(u,v)(v,u) 分别计数,不能再除以 2;当 u=v 时也要照常统计。

第三项违例要求 f(u,v)=f(v,z)=1,但 f(u,z)=0。先枚举满足 f(u,v)=1 的有序对,合法的 z 正好属于集合差 G_v\setminus G_u,因此贡献为 G[v]&~G[u] 的置位数。每个有序三元组会在它对应的 (u,v) 处计算一次,没有重计或漏计,也不需要对三个字是否相同作额外限制。

位集按最大规模分配,多余位置始终没有出现在任何 G_v 中。虽然取反会将这些位置变成 1,随后与 G_v 求交仍会将它们清零,因此不会统计到超出当前字长的对象。

设所有表达式的总长度为 L,真值表处理需要 O(L),构建比较关系需要 O(k^2M^2)。若一个位集占 B 个机器字,违例统计需要 O(M^2B);本实现每个位集最多扫描 1764 位机器字。三个输出分别不超过 M,M^2,M^3,而 M^3\le 2^{30},均可以使用 int

参考代码

#include <bits/stdc++.h>
using namespace std;

const int N=1029;
const int M=1000005;
const int K=15;
const int inf=0x3f3f3f3f;
int top,cnt;
int stk[M],pr[128],fst[K][K][4];
bool val[K][K][4];
char op[M];
bitset<N> G[N],H[N];
void calc()
{
    cnt--;
    char c=op[cnt];
    if(c=='!'){stk[top-1]^=15;return;}
    top--;
    int y=stk[top];
    if(c=='=')stk[top-1]^=y^15;
    else if(c=='&')stk[top-1]&=y;
    else if(c=='|')stk[top-1]|=y;
    else if(c=='^')stk[top-1]^=y;
}
int parse(const string &s)
{
    top=cnt=0;
    for(auto c:s)
    {
        if(c=='x')stk[top++]=12;
        else if(c=='y')stk[top++]=10;
        else if(c=='0')stk[top++]=0;
        else if(c=='1')stk[top++]=15;
        else if(c=='('||c=='!')op[cnt++]=c;
        else if(c==')')
        {
            while(op[cnt-1]!='(')calc();
            cnt--;
        }
        else
        {
            while(cnt&&op[cnt-1]!='('&&pr[op[cnt-1]]>=pr[c])calc();
            op[cnt++]=c;
        }
    }
    while(cnt)calc();
    return stk[0];
}
bool cmp(int x,int y,int k,bool res)
{
    int pos=inf;
    for(int i=1;i<=k;i++)
    {
        for(int j=1;j<=k;j++)
        {
            int t=((x>>(k-i))&1)*2+((y>>(k-j))&1);
            if(fst[i][j][t]>=pos)continue;
            pos=fst[i][j][t];
            res=val[i][j][t];
        }
    }
    return res;
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    pr['!']=5;
    pr['=']=4;
    pr['&']=3;
    pr['|']=2;
    pr['^']=1;
    memset(fst,0x3f,sizeof fst);
    int n,k;
    cin>>n>>k;
    for(int i=1;i<=n;i++)
    {
        int a,b;
        string s;
        bool r;
        cin>>a>>b>>s>>r;
        int mask=parse(s);
        for(int j=0;j<4;j++)
        {
            if(!(mask>>j&1)||fst[a][b][j]!=inf)continue;
            fst[a][b][j]=i;
            val[a][b][j]=r;
        }
    }
    bool r;
    cin>>r;
    int m=1<<k;
    for(int i=0;i<m;i++)
    {
        for(int j=0;j<m;j++)if(cmp(i,j,k,r))G[i][j]=H[j][i]=1;
    }
    int ans[3]={};
    for(int i=0;i<m;i++)
    {
        ans[0]+=G[i][i];
        ans[1]+=(G[i]&H[i]).count();
        for(int j=0;j<m;j++)if(G[i][j])ans[2]+=(G[j]&~G[i]).count();
    }
    cout<<ans[0]<<' '<<ans[1]<<' '<<ans[2]<<'\n';
    return 0;
}