题解:P17125 [ICPC 2025 Shanghai R] Flower' s land 3

· · 题解

题意简述

给定 n 个长度为 m 的二进制串。对每个 i\ge 2,需要选择一个 p_i<i,并满足 s_is_{p_i} 的 Hamming 距离不超过 k

求所有合法序列 (p_2,p_3,\dots,p_n) 的数量。这里 n\le 5000m\le 15000k\le 3

解题思路

对每个 i,设前面与 s_i 距离不超过 k 的字符串数量为 c_i。每次选择彼此独立,所以答案为:

\prod_{i=2}^n c_i

困难在于快速筛出可能满足距离限制的字符串,并快速验证距离。

把每个字符串切成若干个 64 位整数。设实际块数为 W,再补零到不小于 W 的最小二次幂 B

在这些块上建立完全二叉树。叶子表示一个 64 位块,父节点表示左右两段的拼接。

对所有叶子值排序离散化,相同值赋相同编号。随后逐层把有序对 (l,r) 排序离散化。于是同层两个节点编号相同,当且仅当它们表示的整段完全相同。

这个过程不使用概率哈希,不存在碰撞风险。

利用规范化编号,可以递归计算截断 Hamming 距离。若两个节点编号相同,整段贡献为 0。到达叶子时,用 popcount 计算距离。

计算左子树后,若距离已经超过限制便立即停止。否则只把剩余限制传给右子树。对于限制 L,递归结果恰好为实际距离与 L+1 的较小值。

接下来筛选候选父亲。

B\ge 4 时,把整棵树均分成四段。对每一段,按照规范化编号保存此前具有相同段的字符串下标。

若两串的 Hamming 距离不超过 k\le 3,四段中至少一段完全相同。

因此,每个合法父亲一定出现在四个编号桶的并集中。

同一个下标可能匹配多个分块。用时间戳数组去重后,再调用截断距离递归验证即可。

B<4 时,字符串至多包含两个机器字,直接枚举全部此前字符串。

补零不会改变 Hamming 距离。补出的整段可能带来额外候选,但所有候选都会再次验证,所以不会影响正确性。重复字符串也仍按不同下标分别计数。

四分块编号已经是连续整数,所以直接用编号索引四组桶,不需要散列表。设实际产生的候选对数量为 C。建树时间为 O(nB\log(nB)),距离验证时间为 O(C(k+1)\log B),空间复杂度为 O(nB)

正确性证明

先证明规范化编号准确。叶子编号相同当且仅当对应整数相同。若结论对某层成立,则上一层编号相同当且仅当左右儿子编号分别相同。因此,归纳可知任意节点编号相同当且仅当对应二进制段相同。

距离递归遇到相同编号时跳过的整段距离为 0。不同节点最终会递归到所有发生差异的叶子。超过限制后返回 L+1,否则累加全部差异。因此,验证结果不超过 k 当且仅当真实 Hamming 距离不超过 k

对于任意合法父亲,全部差异至多为 3。把差异分配到四段,至少一段没有差异。由编号准确性,这一段的编号必然相同,所以合法父亲不会被候选筛选遗漏。

候选经过距离验证后,留下的下标恰好是所有合法父亲。每个 p_i 可以独立选择,故所有合法父亲数量的乘积就是答案。

参考代码

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

using ll=long long;
using ull=unsigned long long;
using pii=pair<int,int>;
const int mod=998244353;
int value(char c)
{
    if(c<='9')return c-'0';
    return c-'A'+10;
}
vector<int> compress_leaf(const vector<ull> &a,vector<ull> &val)
{
    size_t size=a.size();
    int n=(int)size;
    vector<pair<ull,int>> ord(n);
    for(int i=0;i<n;i++)ord[i]={a[i],i};
    sort(ord.begin(),ord.end());
    vector<int> id(n);
    val.clear();
    int now=-1;
    for(int i=0;i<n;i++)
    {
        if(i==0||ord[i].first!=ord[i-1].first)
        {
            val.push_back(ord[i].first);
            now++;
        }
        id[ord[i].second]=now;
    }
    return id;
}
vector<int> compress_node(const vector<int> &a,vector<pii> &son)
{
    size_t size=a.size();
    int n=(int)size/2;
    vector<pair<ull,int>> ord(n);
    for(int i=0;i<n;i++)
    {
        ull key=(ull)(unsigned int)a[i*2]<<32|(unsigned int)a[i*2+1];
        ord[i]={key,i};
    }
    sort(ord.begin(),ord.end());
    vector<int> id(n);
    son.clear();
    int now=-1;
    for(int i=0;i<n;i++)
    {
        if(i==0||ord[i].first!=ord[i-1].first)
        {
            int p=ord[i].second;
            son.push_back({a[p*2],a[p*2+1]});
            now++;
        }
        id[ord[i].second]=now;
    }
    return id;
}
int distance(int level,int x,int y,int limit,const vector<ull> &val,const vector<vector<pii>> &son)
{
    if(x==y)return 0;
    if(level==0)return min(limit+1,__builtin_popcountll(val[x]^val[y]));
    int ans=distance(level-1,son[level][x].first,son[level][y].first,limit,val,son);
    if(ans>limit)return ans;
    return ans+distance(level-1,son[level][x].second,son[level][y].second,limit-ans,val,son);
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    int n,m,limit;
    cin>>n>>m>>limit;
    int words=(m+63)/64;
    int base=1;
    int levels=0;
    while(base<words)
    {
        base<<=1;
        levels++;
    }
    vector<ull> raw((ll)n*base);
    for(int i=0;i<n;i++)
    {
        string s;
        cin>>s;
        for(int j=0;j<m/4;j++)
        {
            int bit=j*4;
            raw[(ll)i*base+bit/64]|=(ull)value(s[j])<<(bit%64);
        }
    }
    vector<ull> leaf;
    vector<int> cur=compress_leaf(raw,leaf);
    vector<vector<pii>> son(levels+1);
    vector<array<int,4>> part(n);
    int width=base;
    for(int i=1;i<=levels;i++)
    {
        if(width==4)
        {
            for(int j=0;j<n;j++)
                for(int k=0;k<4;k++)part[j][k]=cur[j*4+k];
        }
        cur=compress_node(cur,son[i]);
        width>>=1;
    }
    vector<int> root(n);
    for(int i=0;i<n;i++)root[i]=cur[i];
    vector<int> seen(n,-1),candidate;
    int types=0;
    for(auto v:part)
        for(int x:v)types=max(types,x+1);
    array<vector<vector<int>>,4> bucket;
    for(int j=0;j<4;j++)bucket[j].resize(types);
    ll ans=1;
    for(int i=0;i<n;i++)
    {
        candidate.clear();
        if(base<4)
        {
            for(int j=0;j<i;j++)candidate.push_back(j);
        }
        else
        {
            for(int j=0;j<4;j++)
            {
                for(int k:bucket[j][part[i][j]])
                {
                    if(seen[k]==i)continue;
                    seen[k]=i;
                    candidate.push_back(k);
                }
            }
        }
        if(i)
        {
            int cnt=0;
            for(int j:candidate)
            {
                if(distance(levels,root[i],root[j],limit,leaf,son)<=limit)cnt++;
            }
            if(cnt==0)
            {
                cout<<0<<'\n';
                return 0;
            }
            ans=ans*cnt%mod;
        }
        if(base>=4)
        {
            for(int j=0;j<4;j++)bucket[j][part[i][j]].push_back(i);
        }
    }
    cout<<ans<<'\n';
    return 0;
}