KMP自动机与数位DP:高效统计n位整数中‘2023‘子串出现次数
前几天刷算法题的时候撞见这么一道题给定n和m问所有n位十进制整数里恰好出现m个2023子串的数有多少个。我一开始还想得很天真直接枚举所有n位数每个数转成字符串然后数一下2023出现几次不就行了直到我意识到n稍微大一点比如n15要枚举10^14个数按每秒一亿次也得跑到世界末日。这道题的本质其实是一个“构造数字串统计指定子串出现次数”的计数问题典型的解法是KMP自动机状态机加数位DP。这篇文章就把我从暴力到正解、从朴素DP到矩阵加速的完整过程写下来给同样被“子串出现次数”类计数题卡住的人做个参考。1. 先试暴力为什么n到12就彻底跑不动了1.1 暴力版本能跑多大先写一版最无脑的代码枚举[10^(n-1), 10^n-1]里的所有整数把数字转成字符串用字符串查找统计2023出现的次数等于m就计入答案。这个逻辑本身完全没错问题在于规模。nn位数的个数暴力大概耗时49,000忽略不计890,000,000几秒到几十秒1090亿几分钟129000亿显然不可行159×10^14灾难所以暴力的意义只有一个用来验证小数据答案。n≤6的时候暴力跑得飞快可以作为后面DP程序的“裁判”。我实际写题的时候第一步永远是先把暴力版本写好后面优化版本跑出来的答案必须和暴力完全一致否则就是转移表写错了。1.2 手算几个小例子建立直观在写DP之前先用小手算算心里对答案有个数n4m1四位整数只有一个2023答案是1。n4m0四位整数总共9000个减去出现2023的那1个答案是8999。n5m1先看2023出现在5位数的哪里。它只能出现在第1到第4位或第2到第5位。 第1到第4位是2023xx取0~9共10个第2到第5位是x2023x取1~9共9个因为x0时02023是4位数而不是5位数。两种情况不重叠所以答案是19。n3m0三位数内不可能出现四位长的2023答案就是全部三位数900个。这些例子看起来简单但能帮我们抓住一个关键点“恰好出现m个”意味着必须处理“出现2次以上”的情况比如20232023这个8位数里出现了2次在统计m1时不能算进去。用容斥硬算很麻烦DP反而自然。1.3 这个题难在哪难点在于光知道“当前已经出现了几个2023”不够还得知道“当前构造到一半的数字串末尾到底匹配了模式串的哪一部分”。举个例子已经构造到1202末尾两个字符02并不匹配2023前缀但末尾的202刚好匹配了前三个字符如果再填一个3就凑成完整2023如果填2则重新积累到匹配长度为1。这种动态变化的信息正是KMP自动机擅长处理的。2. 把字符串匹配变成自动机状态2.1 为什么不能只记“已经出现了几次”如果DP只设计成两个维度当前填到第几位、已经出现多少次那么转移时无法知道填下一位数字后是否会形成新的一次2023。比如当前末尾是20这时填2后变成202还没完成匹配当前末尾是202填3后就完成了一次匹配。这两种情况下虽然都已出现的次数相同但“匹配到一半的状态”完全不同。所以必须把KMP的“当前最长匹配前缀长度”也作为一个状态维度。对于模式串2023可能的最长匹配长度只有0、1、2、3四种。这样状态空间很小转移起来非常舒服。2.2 手工推导2023的失配跳转这里先算模式串2023的next数组也叫pi数组。pi[i]表示“模式串的前i1个字符构成的子串中最长的相等真前后缀长度”pi[0]对应2没有真前后缀0。pi[1]对应20前缀2和后缀0不同0。pi[2]对应202前缀2和后缀2相同长度1。pi[3]对应2023前缀2不等于后缀320不等于23202不等于023所以是0。也就是pi [0, 0, 1, 0]。这个数组看起来不起眼但它在失配和匹配完成时会决定跳转目标。2.3 完整的4×10转移表从KMP自动机的角度构造一张表当前状态j0到3读入一个数字d后新状态是多少以及是否触发一次完整匹配。核心转移如下只列会触发状态变化的数字其他数字一律回到状态0当前状态当前匹配情况读入数字新状态是否完成一次20230无匹配21否1末尾匹配202否1末尾匹配221否2末尾匹配2023否3末尾匹配20230是3末尾匹配20221否注意两个容易懵的地方。一个是状态1时读入2新状态仍然是1比如构造到22末尾一个2又算是2023的第一个字符所以匹配长度不归零而是保留1。另一个是状态3时读入2新状态是1而不是0因为2022末尾的2可以作为新匹配的起点。2.4 2023没有border所以匹配完成后直接回到0匹配完成一次2023后当前新的匹配长度不应该是0吗在这个题里恰好是0因为2023的border为0也就是说它的真后缀中没有任何一个同时是它的前缀。但如果模式串换成0101pi[3]2完整匹配一次后应当回到状态2否则就会漏掉010101里重叠出现的第二个0101。所以这里我建议直接用KMP构建自动机而不是手工写死转移表。手动写死当前模版后面换题必踩坑。3. 三维DP状态定义、初始化、转移细节3.1 状态定义定义dp[i][j][k]表示已经构造了i位数字当前KMP匹配状态为jj0,1,2,3并且已经完整出现k次2023的数的个数。i的取值范围是1到n。k的取值范围是0到m。如果m大于n/4直接返回0因为每出现一次2023至少要占用4位数字n位数字最多出现floor(n/4)次。3.2 初始化第一位必须是1~9因为题目要求的是n位十进制整数第一位不能是0。第一位从1枚举到9第一位是2进入状态1dp[1][1][0]加1。第一位是1、3、4、5、6、7、8、9进入状态0dp[1][0][0]加8。这里有个细节第一位不可能直接完成一次2023因为模式串长度为4所以不需要处理k1的情况。但代码里还是建议写上对 nj L 的通用判断方便以后换模式串。3.3 转移过程与代码骨架转移时从当前状态j枚举下一位数字d通过转移表得到新状态nj。如果nj4说明刚完成一次完整匹配那么k加1并且新状态要变成pi[3]0否则新状态就是njk不变。Python实现如下def count_numbers(n, m, pattern2023): L len(pattern) if m n // L: return 0 # 构建KMP的pi数组 pi [0] * L for i in range(1, L): j pi[i - 1] while j 0 and pattern[i] ! pattern[j]: j pi[j - 1] if pattern[i] pattern[j]: j 1 pi[i] j # 构建自动机转移表状态j读入数字d后到哪 trans [[0] * 10 for _ in range(L)] for j in range(L): for d in range(10): ch chr(ord(0) d) k j while k 0 and ch ! pattern[k]: k pi[k - 1] if ch pattern[k]: k 1 trans[j][d] k # k可能等于L表示匹配完成 dp [[[0] * (m 1) for _ in range(L)] for __ in range(n 1)] # 首位不能为0 for first in range(1, 10): nj trans[0][first] if nj L: if m 1: dp[1][pi[L - 1]][1] 1 else: dp[1][nj][0] 1 # 逐位转移 for i in range(1, n): for j in range(L): for k in range(m 1): val dp[i][j][k] if val 0: continue for d in range(10): nj trans[j][d] if nj L: if k 1 m: dp[i 1][pi[L - 1]][k 1] val else: dp[i 1][nj][k] val return sum(dp[n][j][m] for j in range(L))复杂度是O(n × 4 × (m1) × 10)对于n5000、m100这个量级已经非常稳。空间上可以滚动数组压到O(4×(m1))但三维写法更直观先保证正确再优化。3.4 我踩过的一个坑把“匹配完成状态”当成常驻状态我第一版代码里简单粗暴地给自动机加了状态4表示“刚出现过一次2023”转移时如果nj4就让k加1但下一位继续从状态4去转移。结果20232023这个数只被统计为出现1次。原因很简单第一个2023出现后下一位是2它只是新一段匹配的开头不应该从状态4继续往前走。正确做法是匹配完成后立刻回到pi[L-1]本题是0。如果你以后遇到有border的模式串这个回跳目标还要小心设置。这个bug不实际跑20232023这种自连数据根本发现不了非常阴。4. 高精度与边界答案会大到什么程度4.1 方案数爆炸非常快n100、m25这种输入答案位数往往有三四十位。Python的int是无限精度直接用没问题。但如果是C选手又要求输出精确值那就得手写大数加法或者用boost::multiprecision::cpp_int。很多竞赛题会改成“对某个质数取模输出”这反而省事只需要在每次加法后取模。这里建议先看题目输出要求再决定代码形态。如果只验证思路用Python最舒服如果需要在限制较严的平台上跑那就要考虑用滚动数组降低空间再把加法运算抠细一点。4.2 m0的时候别想当然m0统计的是“完全不包含2023的n位数”。这个情况不能用总数量减去简单公式去凑最稳的还是让DP跑一遍只不过k的维度只有0。当n4时任何n位数都不可能出现2023因此输入输出n1, m09n2, m090n3, m0900n3, m10n4, m20这些边界写在代码最前面既省时间又避免数组越界。4.3 千万不要用浮点数估算替代精确输出组合计数类问题答案必须是确切整数。有人想当然用“总的n位数减去含一个2023的数”之类的容斥公式结果在重叠情况上一算就错还有人试图用浮点近似更是直接跑偏。老老实实走整数DP每一步都是精确的最后答案绝不会因为有浮点误差而多出几个数。5. n特别大时矩阵快速幂要不要上5.1 什么时候考虑矩阵加速普通DP是O(n)的n10^7已经很难跑n10^18更是直接不可能。但转移方程是和i无关的线性递推从(i, j, k)到(i1, nj, nk)的转移系数固定只由j、k和读入的数字d决定。这种情况下可以把整个递推写成行向量左乘转移矩阵的形式然后用矩阵快速幂把O(n)压成O(log n)。这里只建议在n超过千万、m比较小的时候用矩阵幂。如果n只是几千老老实实跑DP别为了秀操作把代码写复杂。5.2 状态压缩与矩阵构造把二维状态(j, k)压成一维idx j * (m1) k状态总数S4×(m1)。初始行向量v[1]首位为2时idx 1×(m1)0 处加1首位为其他非零数字时idx 0处加8。转移矩阵T大小为S×ST[old][new]表示从old状态经过一次数字转移到达new状态的字符数0到10中的一个整数。最终答案v[n] v[1] × T^(n-1)然后取所有j对应的第m列之和。矩阵快速幂的核心代码片段def mat_mul(A, B): S len(A) C [[0] * S for _ in range(S)] for i in range(S): Ai A[i] Ci C[i] for k in range(S): if Ai[k]: Bk B[k] for j in range(S): Ci[j] Ai[k] * Bk[j] return C def mat_pow(T, p): S len(T) R [[1 if i j else 0 for j in range(S)] for i in range(S)] while p: if p 1: R mat_mul(R, T) T mat_mul(T, T) p 1 return R注意S如果超过几百矩阵乘法O(S^3 log n)会非常吃力。所以这套方案只适用于m比较小的场景。如果m也大到几百那就不要硬上矩阵不如回到普通DP或者从生成函数、组合计数的角度另找路子。5.3 什么时候不建议上矩阵我的判断标准很简单如果n只有10^5级别普通DP只要一秒钟矩阵幂光建矩阵就要半天如果n到了10^9但m超过50矩阵规模突破200×200Python矩阵乘法的系数会让人怀疑人生。遇到这种输入先看一眼m的限制再决定路线。矩阵快速幂的价值主要是思路完整它告诉我们这类自动机DP本质上是一个有限状态线性递推和n的关系是幂次增长的。这样哪怕不写矩阵代码也能对问题的理论复杂度心中有数。6. 完整参考代码与测试验证6.1 可直接运行的Python完整代码下面这份代码按前面的思路组织包含了KMP自动机构建、DP计算和几个测试用例。def count_numbers(n, m, pattern2023): L len(pattern) if m n // L: return 0 pi [0] * L for i in range(1, L): j pi[i - 1] while j 0 and pattern[i] ! pattern[j]: j pi[j - 1] if pattern[i] pattern[j]: j 1 pi[i] j trans [[0] * 10 for _ in range(L)] for j in range(L): for d in range(10): ch chr(ord(0) d) k j while k 0 and ch ! pattern[k]: k pi[k - 1] if ch pattern[k]: k 1 trans[j][d] k dp [[[0] * (m 1) for _ in range(L)] for __ in range(n 1)] for first in range(1, 10): nj trans[0][first] if nj L: if m 1: dp[1][pi[L - 1]][1] 1 else: dp[1][nj][0] 1 for i in range(1, n): for j in range(L): for k in range(m 1): val dp[i][j][k] if val 0: continue for d in range(10): nj trans[j][d] if nj L: if k 1 m: dp[i 1][pi[L - 1]][k 1] val else: dp[i 1][nj][k] val return sum(dp[n][j][m] for j in range(L)) if __name__ __main__: print(count_numbers(4, 1)) # 1 print(count_numbers(4, 0)) # 8999 print(count_numbers(5, 1)) # 19 print(count_numbers(3, 0)) # 900 print(count_numbers(8, 2)) # 继续扩大规模测试这份代码里有两点建议第一用KMP构建自动机后换模式串只需要改pattern变量第二如果n很大可以把dp数组改成滚动版只保留两层的二维数组内存直接从O(n×4×(m1))变成O(4×(m1))。6.2 用暴力程序对拍写完DP后我习惯再补一个暴力对拍函数专门用来验证n≤6的情况。def brute(n, m, pattern2023): cnt 0 start 10 ** (n - 1) end 10 ** n for x in range(start, end): s str(x) if s.count(pattern) m: cnt 1 return cnt这里必须用s.count(pattern)而不是s.find因为count会统计重叠情况吗Python的str.count默认是不重叠统计。对于2023这种无border的模式串不重叠统计和真实重叠次数一致但如果以后换成0101这类有重叠的模式串暴力对拍时就要用循环find去数重叠出现次数。拿暴力程序跑出n5、m1的答案19再和DP跑出的答案对一下心就放下一半。然后再跑n8、m2这种稍微大一点的输入观察结果是否和直觉一致。6.3 把模式串换成别的怎么改把代码里的pattern从2023改成0101或者2024函数主体完全不用动KMP自动机会自动生成新的转移表和pi数组。这就是用自动机而不是硬编码的好处。唯一要注意的是pi数组在匹配完成后的回跳长度不再一定是0比如0101的pi[3]2此时DP里的回跳目标pi[L-1]会自动变成2不需要额外修改。所以这道题虽然名字叫“数2023”但掌握的是一类问题的解法任意模式串、任意出现次数、任意位数的计数都可以用同一个框架解决。我在本地做这道题时最大的体会是不要把KMP自动机想得太玄乎它就是一张“当前匹配到哪里遇到某个字符后去哪里”的查询表。把表建好DP就是在表上走路表错了后面全是白搭。因此建议所有读者拿到这类题第一件事先用暴力小数据把自动机转移表验证一遍再开始优化。