题解:P9151 计数题
大家都太有观察力了,我来一个啥子做法。
题意:给出一个 01 序列 \(a\),你可以把相邻三个数删掉并在原位置塞入一个他们三个数的中位数,求可以得到的子序列。\(n\le 5\times 10^6\)。
做法:
先把操作转化为删掉相邻三个相同的两个或者删掉一对相邻的 \(01\)。
这种题其实相当套路,找能有的子序列个数,考虑一个子序列怎么判定是否合法,发现贪心地匹配下一位,也就是每次删掉能删的一段,去找下一位刚好和要匹配的一样的位置,这样做就是对的,因为假设最优为 \(i\),但是我匹配了后面的 \(j\),我可以删掉 \([i+1,j]\) 这一段。
那么我们考虑建出来一个子序列自动机,每个点有两条边,分别连向下一个可以匹配的 \(0/1\),那么答案就是在这个子序列自动机上做 dp,然后如果一个后缀可以被删除就贡献答案。
那么现在考虑怎么找下一个 \(0/1\),看其他题解都需要观察性质,其实有一个很简单的做法。我们考虑比如当前为 \(0\),找下一个 \(1\),那么我可以让中间变成 \(0001,0011\) 这样就可以了,我只需要找下一个 \(0\),后面的事情交给下一个 \(0\) 去找就行,找下一个 \(1\) 同理。那这里有人就会问了,欸你找 \(1\) 需要下一个 \(0\),找 \(0\) 需要下一个 \(1\),那不是循环转移了。这里其实很简单,其中有一个肯定是 \(i+1\) 这样最优了,转移另外一个就行。
那么求答案就在自动机上 dp 到当前位置结束有多少个串,要求后缀可删除,同样讨论这个位置后面的 \(01\) 情况删掉,看后面的后缀是否可以删了就行,复杂度线性。
代码:
#include <bits/stdc++.h>
using namespace std;
#define int long long
const int maxn = 5e6 + 5, mod = 998244353;
int n; string s;
int nxt[maxn][2], dp[maxn], use[maxn];
void solve() {
cin >> s; n = s.size(), s = ' ' + s;
nxt[n][0] = nxt[n][1] = n + 1;
for (int i = n - 1; i >= 0; i--) {
nxt[i][0] = nxt[i][1] = n + 2;
nxt[i][s[i + 1] - '0'] = i + 1;
if(s[i] == '0') {
if(s[i + 1] == '0') {
if(nxt[i][0] <= n && nxt[nxt[i][0]][0] <= n)
nxt[i][1] = min(nxt[nxt[nxt[i][0]][0]][1], nxt[i][1]);
}
}
else if(s[i] == '1') {
if(s[i + 1] == '1') {
if(nxt[i][1] <= n && nxt[nxt[i][1]][1] <= n)
nxt[i][0] = min(nxt[nxt[nxt[i][1]][1]][0], nxt[i][0]);
}
}
if(s[i + 1] == '0' && nxt[i][0] <= n && nxt[nxt[i][0]][1] <= n)
nxt[i][1] = min(nxt[nxt[nxt[i][0]][1]][1], nxt[i][1]);
if(s[i + 1] == '1' && nxt[i][1] <= n && nxt[nxt[i][1]][0] <= n)
nxt[i][0] = min(nxt[nxt[nxt[i][1]][0]][0], nxt[i][0]);
// cout << nxt[i][0] << " " << nxt[i][1] << endl;
}
for (int i = 1; i <= n; i++)
dp[i] = 0;
dp[0] = 1;
int ans = 0;
use[n] = 1;
for (int i = n - 1; i >= 0; i--) {
use[i] = 0;
if(nxt[i][0] <= n && nxt[nxt[i][0]][1] <= n)
use[i] |= use[nxt[nxt[i][0]][1]];
if(nxt[i][1] <= n && nxt[nxt[i][1]][0] <= n)
use[i] |= use[nxt[nxt[i][1]][0]];
if(i) {
int id = s[i] - '0';
if(nxt[i][id] <= n && nxt[nxt[i][id]][id] <= n)
use[i] |= use[nxt[nxt[i][id]][id]];
}
}
for (int i = 0; i <= n; i++) {
if(nxt[i][0] <= n)
dp[nxt[i][0]] = (dp[nxt[i][0]] + dp[i]) % mod;
if(nxt[i][1] <= n)
dp[nxt[i][1]] = (dp[nxt[i][1]] + dp[i]) % mod;
if(i && use[i])
ans = (ans + dp[i]) % mod;
}
cout << ans << endl;
}
signed main() {
int T; cin >> T;
while(T--)
solve();
return 0;
}

浙公网安备 33010602011771号