P16231 [蓝桥杯 2026 省 A] 基因研究 solution

题目分析

通过观察性质,我们会发现这是比较典型的自动机DP,具体如何发现的呢?

这是因为,对于一个串来说,在他的后面添加新的字符之后能产生的合法新串与不合法新串,与其所在的位置无关,换句话说:线性时不变转移

所以说对于一个串,其之后添加新碱基后,其变成的新串与易感串的最长匹配前缀只与该串有关,因此我们只需要考虑如何构造出这个匹配矩阵就可以了。

考虑动态规划存储该构造该匹配矩阵,定义数组 $ch[i][j] = $ 当前串已经匹配了易感串 $s$ 的前 $i$ 个字符时候,最后加入碱基 $j$ 时,产生的新串与易感串 $s$ 的最长匹配进度。

但是这里我们要思考,如何快速的计算在加入了新碱基之后的最长匹配进度?

这个过程可以使用 Z 函数(扩展 KMP)算法来解决。

Z 函数

其定义是 $Z[i] := LCP(s, s[i:])$,对于字符串 $s$ 来说,$z[i]$ 就是 $s$ 从 $i$ 位置开始的后缀与 $s$ 本身的最长相同前缀。

而这个算法可以 $O(|s|)$ 求出,具体算法可以自行搜索学习。

有了这个东西,怎么做到上述的匹配呢,可以这样思考:
对于一个已经匹配了 $j$ 长度的串,我们可以枚举其结尾添加的碱基,然后从大到小枚举现在串后缀易感串的最长匹配前缀长度 $l$,但是这里要注意 $l$ 是 $[0, l + 1]$ 这个取值区间。

但是我们不可能对于每一个添加了碱基的新串都跑一次匹配,那么怎么高效的做这个过程呢?
可以考虑把匹配本身分开,原本已经匹配了 $j$ 长度的串 $t$,对于一个添加了新碱基的串 $t'$,我们把其分为两部分分别匹配:一部分是 $t$,另一部分是新加入的碱基。
简单思考就会发现,如果想要 $t'$ 匹配,那么 $t$ 本身要匹配,然后新加入的碱基要匹配。

这样以来就可以写出代码了:

#include <bits/stdc++.h>
using namespace std;
using Int = long long;
#define rep(i, a, b) for(Int i = (a); i < (b); i ++)
using vi = vector<Int>;
#define sz(s) ((Int)size(s))
const Int mod = 998244353;

vi Z(string s) {
	Int n = sz(s);
	vi z(n, n);
	Int l = -1, r = -1;
	rep(i, 1, n) {
		Int &x = z[i] = i < r ? min(r - i, z[i - l]) : 0;
		while(i + x < n && s[i + x] == s[x]) x ++;
		if(i + x > r) l = i, r = i + x;
	}
	return z;
}

map<char, int> mp;

int main() {
	int n, m; cin >> n >> m;
	string s; cin >> s;
	mp['A'] = 0, mp['T'] = 1, mp['G'] = 2, mp['C'] = 3;

	auto z = Z(s);

	vector<vi> ch(m + 1, vi(4, 0));

	rep(j, 0, m) rep(k, 0, 4) for(int l = j + 1; l > 0; l --) {
		if(mp[s[l - 1]] != k) continue; 
        //分开匹配新加入的碱基对与原本碱基对的前缀.
		if(l == 1) { //没有前缀
			ch[j][k] = 1;
			break;
		}
		Int idx = j - l + 1;
		if(idx == 0) {
			ch[j][k] = l;
			break;
		} else if(idx > 0 && idx < m) {
			if(z[idx] >= l - 1) {
				ch[j][k] = l;
				break;
			}
		}
	}
	rep(k, 0, 4) ch[m][k] = m;

	vector<vi> f(n + 1, vi(m + 1, 0));
	f[0][0] = 1;
	vi p(4);
	rep(i, 1, n + 1) {
		rep(j, 0, 4) cin >> p[j];

		rep(j, 0, m + 1) {
			if(!f[i - 1][j]) continue;
			rep(k, 0, 4) {
				f[i][ch[j][k]] = (f[i][ch[j][k]] + f[i - 1][j] * p[k]) % mod;
			}
		}
	}
	cout << f[n][m] << '\n';
}