KMP算法完全解析:从原理到Java实现

前言

字符串匹配是编程中最常见的问题之一。比如在文本编辑器中查找某个单词、在日志文件中搜索特定关键字、在DNA序列中寻找特定片段,这些都离不开字符串匹配。

最简单的办法是一个字符一个字符地比,但这种方法在数据量大时非常慢。KMP算法就是为了解决这个问题而生的,它能在线性时间内完成匹配,效率远超暴力方法。


一、从暴力匹配说起

1.1 什么是暴力匹配

暴力匹配的思路非常直接:

  1. 从主串的每个位置开始

  2. 和模式串逐个字符比较

  3. 如果全部匹配,就找到了

  4. 如果有字符不匹配,就移到下一个位置重新开始

用代码表示就是:

public static int bruteForce(String text, String pattern) {
    int n = text.length();
    int m = pattern.length();
    
    for (int i = 0; i <= n - m; i++) {
        int j = 0;
        while (j < m && text.charAt(i + j) == pattern.charAt(j)) {
            j++;
        }
        if (j == m) {
            return i;  // 匹配成功,返回起始位置
        }
    }
    return -1;  // 没找到
}

1.2 暴力匹配的效率问题

看一个具体例子:

主串:"AAAAABAAAAABAAAAAC"
模式串:`"AAAAAC""

匹配过程:

  • 从位置0开始:比较前5个"A"都匹配,第6个字符主串是"B",模式串是"C",不匹配

  • 从位置1开始:又比较前5个"A",第6个字符还是"B" vs "C",不匹配

  • 从位置2开始:同样的过程再来一遍

  • ...一直重复

问题出在哪?

当匹配到第6个字符失败时,我们已经知道了前面5个字符都是"A"。但暴力算法完全抛弃了这些信息,把模式串只往后移一位,导致大量重复比较。

1.3 暴力匹配的时间复杂度

  • 最好情况:O(n)(第一个字符就匹配不上)

  • 最坏情况:O(n × m)(比如上面的例子)

当n和m都很大时,O(n×m)是无法接受的。


二、KMP算法的核心思想

2.1 利用已知信息

KMP算法的核心是:当匹配失败时,利用已经匹配成功的部分,避免从头开始比较。

关键问题是:当匹配失败时,模式串应该往后移动多少位

2.2 一个直观的例子

主串:"ABABABC"
模式串:"ABABC"

匹配过程:

位置:   0 1 2 3 4 5 6
主串:   A B A B A B C
模式串: A B A B C
         ↑ ↑ ↑ ↑ ↑
         比到第4个字符发现不匹配

前4个字符 "ABAB" 已经匹配成功了。这时候我们能不能利用这4个字符的信息呢?

看 "ABAB" 这个字符串:

  • 前缀:"A""AB""ABA"

  • 后缀:"B""AB""BAB"

  • 相同的:"AB"

最长相同的前后缀是 "AB",长度是2。

这意味着:"ABAB" 的前缀"AB"等于后缀"AB"。所以我们可以把模式串往右移2位,让前缀"AB"对准主串中已经匹配的后缀"AB"位置:

位置:   0 1 2 3 4 5 6
主串:   A B A B A B C
模式串:     A B A B C
             ↑ ↑
             从这继续比

主串的指针不需要回退,模式串从位置2开始继续比较。

2.3 为什么要用"最长相等前后缀"

用最长相等前后缀可以保证:

  1. 安全:不会跳过可能匹配的位置

  2. 高效:跳过的距离最大


三、前缀、后缀、部分匹配表

3.1 前缀和后缀

前缀:从开头开始的子串,不包含最后一个字符
后缀:从结尾开始的子串,不包含第一个字符

以字符串 "ABCAB" 为例:

前缀:"A""AB""ABC""ABCA"
后缀:"B""AB""CAB""BCAB"
相同的:"AB"
长度:2

3.2 部分匹配表(PMT)

PMT的每个值记录了:以当前位置结尾的子串的最长相等前后缀长度

以模式串 "ABABC" 为例:

位置 字符 子串 前缀 后缀 相同部分 长度
0 A A 0
1 B AB {A} {B} 0
2 A ABA {A,AB} {A,BA} {A} 1
3 B ABAB {A,AB,ABA} {B,AB,BAB} {AB} 2
4 C ABABC {A,AB,ABA,ABAB} {C,BC,ABC,BABC} 0

PMT = [0, 0, 1, 2, 0]

3.3 Next数组

实际编程中,我们把PMT右移一位,并在开头补-1,得到next数组:

next[0] = -1
next[j] = PMT[j-1] (j ≥ 1)

所以 "ABABC" 的next数组是:[-1, 0, 0, 1, 2]

next[j]的含义:当模式串第j个字符匹配失败时,模式串指针应该跳到next[j]位置。

比如匹配到第4个字符(j=4,字符C)失败,next[4]=2,说明模式串跳到位置2继续匹配。


四、Java代码实现

4.1 计算Next数组

public class KMP {
    
    /**
     * 构建next数组
     * 
     * 核心思想:模式串自己和自己匹配
     * 
     * @param pattern 模式串
     * @return next数组
     */
    public static int[] buildNext(String pattern) {
        int m = pattern.length();
        int[] next = new int[m];
        
        // 第一个字符的next值是-1
        next[0] = -1;
        
        // i:当前要计算的位置
        // j:已经匹配的前缀长度
        int i = 0;
        int j = -1;
        
        while (i < m - 1) {
            if (j == -1 || pattern.charAt(i) == pattern.charAt(j)) {
                // 匹配成功,继续往后
                i++;
                j++;
                next[i] = j;
            } else {
                // 匹配失败,j回溯
                j = next[j];
            }
        }
        
        return next;
    }
    
    /**
     * 优化版:构建next数组
     * 
     * 当字符相同时,再跳一步,减少比较次数
     */
    public static int[] buildNextOptimized(String pattern) {
        int m = pattern.length();
        int[] next = new int[m];
        next[0] = -1;
        
        int i = 0;
        int j = -1;
        
        while (i < m - 1) {
            if (j == -1 || pattern.charAt(i) == pattern.charAt(j)) {
                i++;
                j++;
                // 优化:如果下一个字符还相同,再跳一步
                if (pattern.charAt(i) == pattern.charAt(j)) {
                    next[i] = next[j];
                } else {
                    next[i] = j;
                }
            } else {
                j = next[j];
            }
        }
        
        return next;
    }
}

4.2 KMP主匹配算法

public class KMP {
    
    /**
     * KMP字符串匹配
     * 
     * @param text 主串
     * @param pattern 模式串
     * @return 第一次匹配的位置,没找到返回-1
     */
    public static int search(String text, String pattern) {
        // 空串处理
        if (pattern == null || pattern.isEmpty()) {
            return 0;
        }
        if (text == null || text.isEmpty()) {
            return -1;
        }
        
        int n = text.length();
        int m = pattern.length();
        
        // 模式串比主串还长,不可能匹配
        if (m > n) {
            return -1;
        }
        
        // 构建next数组
        int[] next = buildNext(pattern);
        
        // i:主串指针
        // j:模式串指针
        int i = 0;
        int j = 0;
        
        while (i < n && j < m) {
            if (j == -1 || text.charAt(i) == pattern.charAt(j)) {
                // 匹配成功,两个指针都前进
                i++;
                j++;
            } else {
                // 匹配失败,模式串指针跳到next[j]
                j = next[j];
            }
        }
        
        // 如果j等于m,说明完全匹配
        return j == m ? i - j : -1;
    }
}

4.3 查找所有匹配位置

public class KMP {
    
    /**
     * 查找所有匹配位置
     * 
     * @param text 主串
     * @param pattern 模式串
     * @return 所有匹配位置的列表
     */
    public static List<Integer> searchAll(String text, String pattern) {
        List<Integer> positions = new ArrayList<>();
        
        if (pattern == null || pattern.isEmpty() || text == null || text.isEmpty()) {
            return positions;
        }
        
        int n = text.length();
        int m = pattern.length();
        
        if (m > n) {
            return positions;
        }
        
        int[] next = buildNext(pattern);
        int i = 0;
        int j = 0;
        
        while (i < n) {
            if (j == -1 || text.charAt(i) == pattern.charAt(j)) {
                i++;
                j++;
                
                if (j == m) {
                    // 找到一个匹配
                    positions.add(i - j);
                    // 继续查找下一个
                    j = next[j - 1] + 1;
                }
            } else {
                j = next[j];
            }
        }
        
        return positions;
    }
}

五、完整代码

import java.util.ArrayList;
import java.util.List;

public class KMP {
    
    // ==================== 构建next数组 ====================
    
    public static int[] buildNext(String pattern) {
        int m = pattern.length();
        int[] next = new int[m];
        next[0] = -1;
        
        int i = 0;
        int j = -1;
        
        while (i < m - 1) {
            if (j == -1 || pattern.charAt(i) == pattern.charAt(j)) {
                i++;
                j++;
                next[i] = j;
            } else {
                j = next[j];
            }
        }
        
        return next;
    }
    
    public static int[] buildNextOptimized(String pattern) {
        int m = pattern.length();
        int[] next = new int[m];
        next[0] = -1;
        
        int i = 0;
        int j = -1;
        
        while (i < m - 1) {
            if (j == -1 || pattern.charAt(i) == pattern.charAt(j)) {
                i++;
                j++;
                if (pattern.charAt(i) == pattern.charAt(j)) {
                    next[i] = next[j];
                } else {
                    next[i] = j;
                }
            } else {
                j = next[j];
            }
        }
        
        return next;
    }
    
    // ==================== 匹配方法 ====================
    
    public static int search(String text, String pattern) {
        if (pattern == null || pattern.isEmpty()) {
            return 0;
        }
        if (text == null || text.isEmpty()) {
            return -1;
        }
        
        int n = text.length();
        int m = pattern.length();
        
        if (m > n) {
            return -1;
        }
        
        int[] next = buildNext(pattern);
        int i = 0;
        int j = 0;
        
        while (i < n && j < m) {
            if (j == -1 || text.charAt(i) == pattern.charAt(j)) {
                i++;
                j++;
            } else {
                j = next[j];
            }
        }
        
        return j == m ? i - j : -1;
    }
    
    public static List<Integer> searchAll(String text, String pattern) {
        List<Integer> positions = new ArrayList<>();
        
        if (pattern == null || pattern.isEmpty() || text == null || text.isEmpty()) {
            return positions;
        }
        
        int n = text.length();
        int m = pattern.length();
        
        if (m > n) {
            return positions;
        }
        
        int[] next = buildNext(pattern);
        int i = 0;
        int j = 0;
        
        while (i < n) {
            if (j == -1 || text.charAt(i) == pattern.charAt(j)) {
                i++;
                j++;
                
                if (j == m) {
                    positions.add(i - j);
                    j = next[j - 1] + 1;
                }
            } else {
                j = next[j];
            }
        }
        
        return positions;
    }
    
    // ==================== 辅助方法 ====================
    
    public static void printNext(String pattern) {
        int[] next = buildNext(pattern);
        System.out.print("模式串: " + pattern + "\nnext: ");
        for (int v : next) {
            System.out.print(v + " ");
        }
        System.out.println();
    }
    
    // ==================== 测试 ====================
    
    public static void main(String[] args) {
        
        // 测试1:基本匹配
        System.out.println("=== 测试1:基本匹配 ===");
        String text1 = "BBC ABCDAB ABCDABCDABDE";
        String pattern1 = "ABCDABD";
        int pos1 = search(text1, pattern1);
        System.out.println("主串: " + text1);
        System.out.println("模式串: " + pattern1);
        System.out.println("匹配位置: " + pos1);
        System.out.println();
        
        // 测试2:打印next数组
        System.out.println("=== 测试2:打印next数组 ===");
        printNext("ABCDABD");
        printNext("AAAAAB");
        printNext("ABABC");
        System.out.println();
        
        // 测试3:查找所有匹配
        System.out.println("=== 测试3:查找所有匹配 ===");
        String text2 = "ABABABAB";
        String pattern2 = "ABAB";
        List<Integer> positions = searchAll(text2, pattern2);
        System.out.println("主串: " + text2);
        System.out.println("模式串: " + pattern2);
        System.out.println("所有匹配位置: " + positions);
        System.out.println();
        
        // 测试4:找不到的情况
        System.out.println("=== 测试4:找不到 ===");
        String text3 = "hello world";
        String pattern3 = "java";
        int pos3 = search(text3, pattern3);
        System.out.println("主串: " + text3);
        System.out.println("模式串: " + pattern3);
        System.out.println("匹配位置: " + pos3);
        System.out.println();
        
        // 测试5:空串
        System.out.println("=== 测试5:空串 ===");
        System.out.println("模式串为空: " + search("abc", ""));
        System.out.println("主串为空: " + search("", "abc"));
        System.out.println();
        
        // 测试6:优化版对比
        System.out.println("=== 测试6:优化版next数组 ===");
        String p = "AAAAAB";
        int[] next1 = buildNext(p);
        int[] next2 = buildNextOptimized(p);
        System.out.print("标准版next: ");
        for (int v : next1) System.out.print(v + " ");
        System.out.println();
        System.out.print("优化版next: ");
        for (int v : next2) System.out.print(v + " ");
        System.out.println();
    }
}

六、运行结果

=== 测试1:基本匹配 ===
主串: BBC ABCDAB ABCDABCDABDE
模式串: ABCDABD
匹配位置: 15

=== 测试2:打印next数组 ===
模式串: ABCDABD
next: -1 0 0 0 0 1 2 
模式串: AAAAAB
next: -1 0 1 2 3 4 
模式串: ABABC
next: -1 0 0 1 2 

=== 测试3:查找所有匹配 ===
主串: ABABABAB
模式串: ABAB
所有匹配位置: [0, 2, 4]

=== 测试4:找不到 ===
主串: hello world
模式串: java
匹配位置: -1

=== 测试5:空串 ===
模式串为空: 0
主串为空: -1

=== 测试6:优化版next数组 ===
标准版next: -1 0 1 2 3 4 
优化版next: -1 0 0 0 0 4 

七、时间复杂度分析

操作 时间复杂度
构建next数组 O(m)
KMP匹配 O(n)
总体 O(n + m)

其中n是主串长度,m是模式串长度。

相比暴力匹配的O(n×m),KMP的效率提升是巨大的。


八、总结

记住这几点就够了:

  1. KMP解决的问题:避免字符串匹配中的重复比较

  2. 核心思想:匹配失败时,利用已匹配部分的信息,让主串指针不回溯

  3. 关键概念

    • 前缀、后缀

    • 最长相等前后缀

    • next数组

  4. next数组含义:next[j]表示当第j个字符匹配失败时,模式串指针应该跳到哪个位置

  5. 时间复杂度:O(n + m)

  6. 适用场景:所有需要高效字符串匹配的地方

KMP算法的本质是:用预处理时间(O(m))换取匹配时间(O(n)),总时间复杂度从O(n×m)降为O(n+m)。

Logo

一站式 AI 云服务平台

更多推荐