首先考虑相邻字符均不同时的方案数。
设 $S_1$ 和 $S_2$ 为两个字符串,可以发现 $S_1$ 向 $S_2$ 当且仅当 $S_2$ 是 $S_1$ 去掉第一个或最后一个字符。
因此考虑走到的串是原串开头去掉 $p$ 个字符末尾去掉 $q$ 个字符的方案数。
组合意义。转化为在 $q$ 个二类操作的空隙中任意插入 $p$ 个一类操作。
这并不好做,但是可以发现:在 $p+1$ 个不同的盒子中放 $q$ 个相同的球,盒子可以为空的方案数,不用第二类而是用第一类插板法转化,即转化为该问题。
因此可以用第二类插板法算方案数。方案数即为 $\displaystyle{\binom{p+q}{p}}$。
枚举 $p,q$,总方案数为 $\displaystyle{\sum_{p=0}^{n-1}\sum_{q=0}^{n-p-1}\binom{p+q}{p}}$。
令 $t=p+q$,化为 $\displaystyle{\sum_{t=0}^{n-1}\sum_{p=0}^{t}\binom{t}{p}}$,也就是 $\displaystyle{\sum_{t=0}^{n-1}2^t}=2^n-1$。
那什么时候答案小于这个呢?考虑一个所有字符相同的串,假设长度为 $L$,那么这段实际上有多少种方案数呢,容易发现是 $L$,那么贡献即为 $L-2^L+1$。下令 $f_i=i-2^i+1$。
那么我们可以扫连续段。假设这一段长度为 $L$,左边有 $p$ 个字符,右边有 $q$ 个字符,贡献即为 $\displaystyle{\binom{p+q}{p}f_i}$。
那么,做完了吗?没有。考虑 baaaa,它可以从 baaa 到 aaa,而这个 aaa 也要算。
那这部分贡献怎么算呢,对于连续相同段枚举每一个前缀和每一个后缀。
不失一般性地,假设一个前缀长度为 $L$,如果整体的连续段左边有 $p$ 个字符,右边有 $q$ 个字符,我们不能从右边进去,不然肯定先走到了整体的连续段,这部分贡献已经算过了。因此我们统计走到这段前缀的方案数要包含连续段的前一个字符,也就是 $\displaystyle{\binom{p+q-1}{q}}$。然后拿它乘以 $f_L$。
如果这样太抽象,看一个例子:
对于串 baaa,考虑 aa 的贡献。
首先最后两个
a组成的aa无贡献,因为只能从aaa走过去,而aaa出发的贡献已经被去掉了。前面两个
a组成的aa有贡献,走到这个aa的方案数等于走到baa的方案数,因为它只能由baa去掉前面的b得到,根据乘法原理可知。
然后这样就做完了。值得注意的是, 对于整个的连续段,$p+q=n-L$。
#include<bits/stdc++.h>
#define int long long
using namespace std;
string s;
int n, pw[300005], f[300005], fac[300005], ifac[300005];
const int mod = 998244353;
int C(int n, int m){
if (n < 0 || m < 0)
return 0;
return fac[n] * ifac[m] % mod * ifac[n - m] % mod;
}
int qpow(int x, int k){
int ans = 1;
while (k){
if (k & 1)
ans = ans * x % mod;
x = x * x % mod;
k >>= 1;
}
return ans;
}
void ini(){
fac[0] = 1;
for (int i = 1; i <= n; i++)
fac[i] = fac[i - 1] * i % mod;
ifac[n] = qpow(fac[n], mod - 2);
for (int i = n - 1; ~i; i--)
ifac[i] = ifac[i + 1] * (i + 1) % mod;
pw[0] = 1;
for (int i = 1; i <= n; i++)
pw[i] = (pw[i - 1] << 1) % mod;
for (int i = 2; i <= n; i++)
f[i] = (mod + i - pw[i] + 1) % mod;
}
signed main(){
ios::sync_with_stdio(false);
cin.tie(0);cout.tie(0);
cin >> s;
n = s.size();
s = " " + s;
ini();
int ans = pw[n] + mod - 1;
ans %= mod;
for (int i = 1; i <= n; i++){
int pos = i;
while (pos < n && s[pos + 1] == s[i]) pos++;
if (pos > i){
int len = pos - i + 1, l = i - 1, r = n - pos;
ans = (ans + f[len] * C(n - len, l)) % mod;
for (int j = 1; j < len; j++)
ans = (ans + ((i == 1 ? 0 : C(n - j - 1, l - 1)) + (pos == n ? 0 : C(n - j - 1, r - 1))) % mod * f[j]) % mod;
i = pos;
}
}
cout << ans;
return 0;
}