题解: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;
}
posted @ 2026-06-26 17:57  LUlululu1616  阅读(14)  评论(0)    收藏  举报