KMP 算法
1 KMP 算法简介
Knuth-Morris-Pratt 算法(简称 KMP)是一种字符串匹配算法,可以在 O(n + m) 的时间复杂度内完成字符串匹配。
2 前缀函数
字符串前缀是指从串首开始到某一位置
字符串后缀是指从某一位置
前缀函数:给定一个长度为
- 如果子串
有一对相等的真前缀与真后缀,即 和 ,则前缀函数 。 是所有对相等的真前缀与真后缀长度的最大值。 - 如果不存在相等的一对真前缀与真后缀,则
。
注:
2.1 朴素算法求解前缀函数
cpp
void getnxt(char *p, int len) {
nxt[1] = 0;
for (int i = 2; i <= len; i++) {
for (int j = i - 1; j >= 0; j--) { // 从最大的真前缀长度开始尝试
bool flag = true;
for (int x = 1, y = i - j + 1; x <= j; x++, y++) {
if (p[x] != p[y]) {
flag = false;
}
}
if (flag) { // 找到一对则终止
nxt[i] = j;
break;
}
}
}
}2.2 第一个优化
相邻的前缀函数值至多增加 1。当前仅当
cpp
void getnxt(char *p, int len) {
nxt[1] = 0;
for (int i = 2; i <= len; i++) {
for (int j = nxt[i - 1] + 1; j >= 0; j--) { // 修改 j = i - 1 -> j = nxt[i - 1] + 1
bool flag = true;
for (int x = 1, y = i - j + 1; x <= j; x++, y++) {
if (p[x] != p[y]) {
flag = false;
}
}
if (flag) {
nxt[i] = j;
break;
}
}
}
}2.3 第二个优化
当
如果我们找到了这样的长度
由上图可知有
最后我们把代码写一下:
cpp
void getnxt(char *p, int len) {
nxt[1] = 0;
for (int i = 2, j = 0; i <= len; i++) {
while (j && p[i] != p[j + 1]) j = nxt[j];
if (p[i] == p[j + 1]) j++;
nxt[i] = j;
}
}3 KMP 算法
给定一个主串
cpp
#include <iostream>
using namespace std;
const int N = 1e5 + 5, M = 1e6 + 5;
char s[M], p[N]; // 主串,模式串
int nxt[N], n, m;
void getnxt(char *p, int len) {
nxt[1] = 0;
for (int i = 2, j = 0; i <= len; i++) {
while (j && p[i] != p[j + 1]) j = nxt[j];
if (p[i] == p[j + 1]) j++;
nxt[i] = j;
}
}
int main() {
cin >> n >> p + 1 >> m >> s + 1;
getnxt(p, n);
for (int i = 1, j = 0; i <= m; i++) {
while (j && s[i] != p[j + 1]) j = nxt[j];
if (s[i] == p[j + 1]) j++;
if (j == n) { // 匹配成功
cout << i - n << endl;
j = nxt[j];
}
}
return 0;
}4 练习题目
cpp
@author lllyouo
@date 2024-09-24
@problem Jouier 2196. Power Strings
#include <iostream>
#include <cstring>
using namespace std;
const int N = 1e6 + 5;
char s[N];
int nxt[N], n;
void getnxt(char *p, int len) {
nxt[1] = 0;
for (int i = 2, j = 0; i <= len; i++) {
while (j && p[i] != p[j + 1]) j = nxt[j];
if (p[i] == p[j + 1]) j++;
nxt[i] = j;
}
}
int main() {
while (cin >> s + 1) {
if (s[1] == '.' && strlen(s + 1) == 1) break;
n = strlen(s + 1);
getnxt(s, n);
if (n % (n - nxt[n]) == 0) {
cout << n / (n - nxt[n]) << endl;
} else {
cout << 1 << endl;
}
}
return 0;
}