[R22F]子序列计数


题目

给定字符串 sstt,考虑 ss 的所有互不相同的重排列(多重集排列),对每个重排列统计 tt 作为子序列出现的次数,求所有这些统计结果之和,对 998244353998244353 取模。

对于 100%100\% 的数据,1s5001\le |s|\le 5001t31\le |t|\le 3,字符串 sstt 仅由小写字母组成。

思路

n=sn=|s|k=tk=|t|cac_a 为字母 aass 中出现的次数,rar_a 为字母 aatt 中出现的次数。

把「ss 的一个重排列」和「tt 在其中作为子序列出现的一组位置 1i1<i2<<ikn1\le i_1<i_2<\cdots<i_k\le n」作为一个整体计数,答案就是这个整体计数。换一种视角:先固定位置的选择,再数有多少个重排列在这些位置上恰好依次填入 t1,t2,,tkt_1,t_2,\dots,t_k

固定位置后,这 kk 个位置的字符已确定,剩余 nkn-k 个位置要填入剩余字母的多重集 {cara}\{c_a-r_a\},可填出的不同重排列数为

(nk)!a(cara)!,\frac{(n-k)!}{\prod_a (c_a-r_a)!},

当某个字母满足 ca<rac_a<r_a(即 tt 的字母多重集不是 ss 的子多重集)时该值为 00。这个数量只取决于 tt 的字母多重集,与具体位置无关;而位置的选择共有 (nk)\binom{n}{k} 种,因此答案为

(nk)(nk)!a(cara)!=n!k!a(cara)!.\binom{n}{k}\cdot\frac{(n-k)!}{\prod_a (c_a-r_a)!} =\frac{n!}{k!\prod_a (c_a-r_a)!}.

验证s=s= aabb 时重排列为 aabbabababbabaabbababbaat=t= ab 的出现次数分别为 4,3,2,2,1,04,3,2,2,1,0,总和 1212,而 4!/(2!1!1!)=124!/(2!\cdot1!\cdot1!)=12,两者一致。

由于 n500<998244353n\le 500<998244353 且模数是质数,分母中每个阶乘都非零、可逆。分子分母各自对模数取模后,用快速幂求 xp2x^{p-2} 作为 xx 的逆元(费马小定理),乘回分子即可。

复杂度

  • 时间复杂度:O(Σn)O(|\Sigma|\cdot n),其中 Σ=26|\Sigma|=26,另有 O(logp)O(\log p) 的快速幂求逆元。
  • 空间复杂度:O(Σ)O(|\Sigma|)

仓颉实现

import std.env.*
import std.convert.*

const MOD: Int64 = 998244353

func modpow(base0: Int64, exp0: Int64): Int64 {
    var base = base0 % MOD
    var exp = exp0
    var res: Int64 = 1
    while (exp > 0) {
        if ((exp & 1) == 1) {
            res = (res * base) % MOD
        }
        base = (base * base) % MOD
        exp = exp / 2
    }
    return res
}

main(): Int64 {
    let reader = getStdIn()
    let s = reader.readln().getOrThrow()
    let t = reader.readln().getOrThrow()

    let cs = Array<Int64>(26, { _ => 0 })
    for (ch in s) {
        cs[Int64(ch) - 97] += 1
    }
    let ct = Array<Int64>(26, { _ => 0 })
    for (ch in t) {
        ct[Int64(ch) - 97] += 1
    }

    let n = s.size
    let k = t.size

    var fact: Int64 = 1
    var i: Int64 = 2
    while (i <= n) {
        fact = (fact * i) % MOD
        i += 1
    }

    var denom: Int64 = 1
    i = 2
    while (i <= k) {
        denom = (denom * i) % MOD
        i += 1
    }
    var a: Int64 = 0
    while (a < 26) {
        let left = cs[a] - ct[a]
        if (left < 0) {
            println("0")
            return 0
        }
        i = 2
        while (i <= left) {
            denom = (denom * i) % MOD
            i += 1
        }
        a += 1
    }

    println(((fact * modpow(denom, MOD - 2)) % MOD).toString())
    return 0
}