目录

信竞学习笔记:字符串

发表于
更新于
4 2.5~3.2 分钟 1108

注:这是本人的课堂笔记。

字符串是由字符连接而成的序列。

KMP 算法

例题:P3375

KMP 算法是一种字符串匹配算法,可以在 O(\lvert s_1 \rvert + \lvert s_2 \rvert) 的时间复杂度内完成字符串匹配。

KMP 算法的核心是利用已经匹配的信息来避免重复匹配。以下讲解字符串下标均从 1 开始。

具体地说,KMP 算法预处理出一个 next 数组。next_i 代表模式串 [1, i] 区间内子串的最长相等前后缀的长度(这里说的最长公共前后缀不包含字符串本身),即 [1. i] 区间内子串的 border。例如,字符串 abborder0;字符串 ababorder1,即字符 a

要求 next 数组:首先,next_1 必定为 0。在求后续 next_i 值的时候,可以借助 next_{i-1} 的信息。注意到,设模式串为 t,当 t_{next_{i-1}+1}t_i 相等时,next_i \leftarrow next_{i-1} + 1;当 t_{next_{i-1}+1}t_i 不相等时,可能有更短的公共前后缀,于是,我们可以对 j \leftarrow next_{i-1} 继续迭代求 next_j 的值,直到满足 t_{next_{j}+1}t_i 相等,推出 next_i \leftarrow next_j + 1。由此,有如下代码:

ne[1] = 0;
for (int i=2, j=0; i<=m; i++) { // s2 为模式串
    while (j && s2[i] != s2[j+1]) j = ne[j];
    if (s2[i] == s2[j+1]) j++;
    ne[i] = j;
}

求出 next 数组后,我们进行字符串匹配。设待匹配串为 s,模式串为 t。当 s_it_{j+1} 相等时,可以将 j 向右移动一位(j \leftarrow j+1);否则,由于已经求得 next 数组,容易注意到,t_{1 \ldots j} 有长度为 next_j 的公共前后缀,而 s_i 向前 j 个字符与 t_{1 \ldots j} 是相同的,所以其也具有公共前后缀。所以可以发现,s_i 向前 j 个字符的公共前后缀中的后缀,与 t_{1 \ldots j} 的公共前后缀的前缀是相同的,可以不用再重复匹配,于是,不断将 j \leftarrow next_j,迭代找到满足 s_i = t_{j+1} 的最长公共前后缀长度,再进行匹配,即可避免重复匹配。当 j = \lvert t \rvert 时,代表匹配完成。由此,有如下代码:

for (int i=1, j=0; i<=n; i++) {
    while (j && s1[i] != s2[j+1]) j = ne[j];
    if (s1[i] == s2[j+1]) j++;
    if (j == m) {
        cout << i - m + 1 << '\n';
        j = ne[j];
    }
}

Manacher 算法

例题:P3805

Manacher 算法可以做到在 O(n) 的时间复杂度内求出最长回文子串的长度。

定义:奇回文串为长度为奇数的回文串。对于奇回文串,其中央字符位置为其回文中心。奇回文串开头到回文中心的长度为回文半径

为了使所有回文串都具有回文中心,我们在回文串的每个字符中间增加一个特殊字符(如 #,首尾也加),这样每个回文串都成为了奇回文串,都拥有了回文中心。则对于任意回文串,其回文串长度就是增加字符后的回文半径减去 1

设有一个右端点 R 最靠右(不一定最长)的回文串回文中心为 C,左端点为 L。要计算以 i 为中心的回文串的最大回文半径 p_i:若 i < r,则在 C 的另一侧存在一个 i 的对应点 i' = 2C - i,而 i' 肯定在 i 之前计算过了,所以可以由 i' 的信息推出 i 的信息。当 i' 回文串在上述大回文串 s_{L \dots R} 的范围内时,可以全部代换,否则只能代换到这个回文串内部的部分。超出 R 的部分进行暴力扩展。最后更新 RC 和答案。代码如下:

#include <bits/stdc++.h>
using namespace std;

const int N = 2.5e7;
int n, p[N], ans;
char s[N], S[N];

int main() {
	
	ios::sync_with_stdio(false);
	cin.tie(nullptr); cout.tie(nullptr);
	
	cin >> s;
	
	n = strlen(s);
	S[0] = '$', S[1] = '#';
	int k = 2;
	for (int i=0; i<n; i++) S[k++] = s[i], S[k++] = '#';
	S[k] = '&';
	n = k;
	
	int R = 1, C = 1;
	for (int i=1; i<n; i++) {
		if (i <= R) p[i] = min(p[2*C-i], R-i+1);
		else p[i] = 1;
		while (S[i-p[i]] == S[i+p[i]]) p[i]++;
		if (i + p[i] > R) R = i + p[i] - 1, C = i;
		ans = max(ans, p[i]);
	}
	
	cout << ans - 1;
	
	return 0;
}