第 23 届浙江省大学生程序设计竞赛 C. 回文串文回
题目描述
给定一个字符串 s 和一个字符 c。
每次操作:从所有尚未被选择过的位置中等概率随机选一个位置 i,将 s[i] 修改为 c(即使原本就是 c 也算一次操作)。
重复该操作,直到字符串 s 变成回文串时停止。
求停止时操作次数的期望值,答案对 998244353 取模。
如果初始字符串已经是回文串,则期望为 0。
输入格式
第一行一个字符串 s
第二行一个字符 c。
输出格式
输出一个整数,表示期望操作次数对 998244353 取模后的结果。
题解
考虑设字符串是 \(S[0 \sim n - 1]\),将回文串的位置两两匹配:
-
如果 \(S[i] = S[n - i - 1] = c\),则两个位置都是无关位置。
-
如果 \(S[i] = S[n - i - 1] \neq c\),则两个位置要么同时被修改,要么同时不被修改。
-
如果 \(S[i] \neq S[n - i - 1]\) 且有一个字符为 \(c\),则有一个无关位置和一个必须修改的位置。
-
如果 \(S[i] \neq S[n - i - 1]\) 且两个字符都不为 \(c\),则两个位置都必须被修改。
如上,我们可以对所有的位置个数分类,记无关位置个数为 \(notbind\),记必须被同时修改的位置对个数为 \(free\_pairs\),必须被修改的位置为 \(need\),显然我们的答案最难求的部分在于 \(free\_pairs\) 的贡献的期望。
不考虑无关位置的贡献,设 \(f_i\) 为恰好选择 \(i\) 个对的方案数,显然 \(f_0 = need!\),对于 \(f_i\),有 \((2 \times i + need)!\) 种方案数,减去恰好选择 \(j\) (\(0 \le j < i\)) 个对的方案数 \(\dbinom{i}{j}(2 \times (i - j))!f_j\),有:
记 \(F_i = \dfrac{f_i}{i!}, \, G_i = \dfrac{(2i)!}{i!}, \, H_i = \dfrac{(2i + need)!}{i!}\),则有 \(F_i = H_i - \sum\limits_{j}F_j \cdot G_{i-j}\),可通过分治 NTT 求得答案。
接下来考虑已知 \(f_i\) 如何求出期望,枚举选择的自由对个数 \(i\),在不考虑无关位置时的方案数为 \(f_i \times (2 \times (free\_pairs - i))!\),此时的期望次数为 \(2i + need\)。
考虑加入所有无关位置,记 \(M = 2 \times free\_pairs + need\),显然已经分配的 \(M\) 个位置划分出了 \(M + 1\) 个区间段,考虑每个无关位置出现在区间段内的概率,则只有前 \(2i + need\) 个位置是可以贡献答案的,所以每个无关位置对答案的期望有 \(\dfrac{2i + need}{M + 1}\) 的贡献,即总贡献为 \(\dfrac{2i + need}{M + 1} \times notbind\),再加上相关位置的贡献则有:
而出现这种情况的概率可以计算:
因此答案为 \(\sum\limits_{i = 0}^{free\_pairs}(2i + need)\dfrac{n + 1}{M + 1} \times P(X = i)\)。
时间复杂度 \(O(n\log^2{n})\)。
#include <bits/stdc++.h>
using namespace std;
#define ll long long
#define ull unsigned long long
#define i128 __int128
#define ld long double
#define pb push_back
#define PII pair<ll, ll>
#define PPI pair<ll, PII>
#define vi vector<ll>
#define vvi vector<vector<ll>>
#define clr(f, n) memset(f, 0, sizeof(int) * (n))
#define cpy(f, g, n) memcpy(f, g, sizeof(int) * (n))
#define rev(f, n) reverse(f, f + (n))
const int N = 6e5 + 10, mod = 998244353, INF = 1e9;
const int _G = 3;
ll qpow(ll a, ll k = mod - 2) {
ll res = 1;
while (k) {
if (k & 1) res = res * a % mod;
k >>= 1;
a = a * a % mod;
}
return res;
}
const int invG = qpow(_G), inv2 = qpow(2);
int rev[N << 1], rev_len;
void rev_init(int n) {
if (rev_len == n) return;
for (int i = 0; i < n; i ++ ) rev[i] = (rev[i >> 1] >> 1) | (i & 1 ? n >> 1 : 0);
rev_len = n;
}
void NTT(int *g, int op, int n) {
rev_init(n);
static ull f[N << 1], Gk[N << 1] = {1};
for (int i = 0; i < n; i ++ ) f[i] = g[rev[i]];
for (int k = 1; k < n; k <<= 1) {
int G1 = qpow(~op ? _G : invG, (mod - 1) / (k << 1));
for (int i = 1; i < k; i ++ ) Gk[i] = Gk[i - 1] * G1 % mod;
for (int i = 0; i < n; i += k << 1) {
for (int j = 0; j < k; j ++ ) {
int tmp = Gk[j] * f[i | j | k] % mod;
f[i | j | k] = f[i | j] + mod - tmp;
f[i | j] += tmp;
}
}
if (k == (1 << 10)) for (int i = 0; i < n; i ++ ) f[i] %= mod;
}
if (~op) for (int i = 0; i < n; i ++ ) g[i] = f[i] % mod;
else {
int invn = qpow(n);
for (int i = 0; i < n; i ++ ) g[i] = f[i] % mod * invn % mod;
}
}
void px(int *f, int *g, int n) {
for (int i = 0; i < n; i ++ ) f[i] = 1ll * f[i] * g[i] % mod;
}
void covolution(int *f, int *g, int len, int lim) {
static int sav[N << 1];
int n; for (n = 1; n < len << 1; n <<= 1);
clr(sav, n); cpy(sav, g, n);
NTT(sav, 1, n); NTT(f, 1, n);
px(f, sav, n); NTT(f, -1, n);
clr(f + lim, n - lim); clr(sav, n);
}
int n; char c;
string s;
int notbind, need, free_pairs;
ll fact[N], inv_fact[N];
int F[N], G[N];
ll C(int n, int m) {
if (n < m) return 0;
return 1ll * fact[n] * inv_fact[m] % mod * inv_fact[n - m] % mod;
}
void CDQ(int *f, int *g, int l, int r) {
static int b1[N << 1], b2[N << 1];
if (l == r) {
(f[l] += fact[2 * l + need] * inv_fact[l] % mod) %= mod;
return;
}
int mid = l + r >> 1;
CDQ(f, g, l, mid);
int n; for (n = 1; n <= r - l + 1; n <<= 1);
cpy(b1, f + l, mid - l + 1); clr(b1 + mid - l + 1, n - (mid - l));
cpy(b2, g, r - l + 1); clr(b2 + r - l + 1, n - (r - l));
NTT(b1, 1, n); NTT(b2, 1, n); px(b1, b2, n); NTT(b1, -1, n);
for (int i = mid + 1; i <= r; i ++ ) f[i] = (f[i] - b1[i - l] + mod) % mod;
clr(b1, n); clr(b2, n);
CDQ(f, g, mid + 1, r);
}
void solve() {
cin >> s; cin >> c; n = s.length();
string t = s; reverse(t.begin(), t.end());
if (s == t) return cout << "0\n", void();
fact[0] = inv_fact[0] = 1;
for (int i = 1; i <= n; i ++ ) fact[i] = fact[i - 1] * i % mod;
inv_fact[n] = qpow(fact[n]);
for (int i = n - 1; i; i -- ) inv_fact[i] = inv_fact[i + 1] * (i + 1) % mod;
if (n & 1) notbind ++ ;
for (int i = 0, j = n - 1; i < j; i ++ , j -- ) {
if (s[i] == s[j]) {
if (s[i] == c) notbind += 2;
else free_pairs ++ ;
} else {
if (s[i] == c || s[j] == c) notbind ++ , need ++ ;
else need += 2;
}
}
for (int i = 0; i <= free_pairs; i ++ ) G[i] = inv_fact[i] * fact[i << 1] % mod;
CDQ(F, G, 0, free_pairs);
for (int i = 0; i <= n; i ++ ) F[i] = F[i] * fact[i] % mod;
ll total = need + 2 * free_pairs;
ll res = 0, coff = inv_fact[total];
for (int i = 0; i <= free_pairs; i ++ ) {
ll p = F[i] * C(free_pairs, i) % mod * fact[2 * (free_pairs - i)] % mod * coff % mod;
ll cnt = need + 2 * i;
ll E = cnt * (n + 1) % mod * qpow(total + 1) % mod;
(res += p * E % mod) %= mod;
}
cout << res << "\n";
}
int main() {
ios::sync_with_stdio(false);
cin.tie(0), cout.tie(0);
int T = 1;
// cin >> T;
while (T -- ) solve();
return 0;
}

浙公网安备 33010602011771号