每日算法 — 使用java实现Rabin-Karp:滚动哈希与多模式字符串匹配

Rabin-Karp算法是字符串匹配领域的经典算法之一,由Michael O. Rabin和Richard M. Karp于1987年提出。与KMP算法基于前缀函数的思路不同,Rabin-Karp巧妙地利用哈希函数将字符串比较转化为数值比较,通过滚动哈希技术在平均情况下实现线性时间复杂度。本文将从哈希设计原理出发,逐步推导出完整的Java实现,并扩展至多模式匹配场景。

算法核心思想

传统暴力匹配中,每移动一个位置就要进行最多 $m$ 次字符比较($m$ 为模式串长度)。Rabin-Karp的核心洞察在于:如果两个字符串相等,那么它们的哈希值必然相等。虽然哈希值相等不能反向推出字符串相等(存在哈希碰撞),但可以先通过 $O(1)$ 的哈希比较快速排除大量不匹配位置,仅在哈希冲突时才进行完整的字符比对。

这种”先筛选、后验证”的策略,使得算法在平均情况下只需 $O(n)$ 时间即可完成文本扫描,其中 $n$ 为文本串长度。

滚动哈希的数学原理

多项式滚动哈希

将长度为 $m$ 的模式串视为一个以字符ASCII码为系数、以基数 $b$ 为底的 $m-1$ 次多项式:

$$H(s) = s_0 \cdot b^{m-1} + s_1 \cdot b^{m-2} + \cdots + s_{m-1} \cdot b^0$$

为避免整数溢出,计算结果对一个足够大的质数 $p$ 取模。基数 $b$ 通常选择256(覆盖扩展ASCII)或一个大于字符集大小的质数如101。

滚动更新的关键推导

当窗口从位置 $i$ 滑动到 $i+1$ 时,新哈希值可由旧哈希值在 $O(1)$ 时间内推导:

$$H_{new} = \left( (H_{old} – s_i \cdot b^{m-1}) \cdot b + s_{i+m} \right) \mod p$$

其中 $s_i \cdot b^{m-1}$ 是移出窗口的最高位贡献,乘以 $b$ 实现整体左移,加上 $s_{i+m}$ 补入新的最低位。为防止中间结果为负,取模前加上 $p$ 的倍数:

$$H_{new} = \left( (H_{old} – s_i \cdot b^{m-1} \mod p + p) \cdot b + s_{i+m} \right) \mod p$$

完整Java实现

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

/**
 * Rabin-Karp 滚动哈希字符串匹配算法实现
 * 支持单模式匹配与多模式匹配场景
 */
public class RabinKarpMatcher {

    /** 哈希基数,选择大于字符集大小的质数可减少碰撞 */
    private static final int BASE = 256;

    /** 大质数模数,用于防止哈希值溢出并降低碰撞概率 */
    private static final long MOD = 1_000_000_007L;

    /** 预计算的 BASE^(m-1) % MOD,用于滚动时移除最高位 */
    private long highestPower;

    /** 模式串长度 */
    private final int patternLength;

    /** 模式串的哈希值 */
    private final long patternHash;

    /** 模式串本身,用于哈希冲突时的精确比对 */
    private final String pattern;

    /**
     * 构造函数:初始化模式串并预计算相关参数
     * @param pattern 待匹配的模式串
     */
    public RabinKarpMatcher(String pattern) {
        if (pattern == null || pattern.isEmpty()) {
            throw new IllegalArgumentException("模式串不能为空");
        }
        this.pattern = pattern;
        this.patternLength = pattern.length();
        // 预计算 BASE^(m-1) % MOD,这是滚动哈希的关键系数
        this.highestPower = computeHighestPower(patternLength);
        // 计算模式串的初始哈希值
        this.patternHash = computeHash(pattern, 0, patternLength);
    }

    /**
     * 计算 BASE^(length-1) % MOD
     * 使用快速幂思想,但此处是线性累积,因为长度通常不大
     */
    private long computeHighestPower(int length) {
        long result = 1;
        for (int i = 1; i < length; i++) {
            result = (result * BASE) % MOD;
        }
        return result;
    }

    /**
     * 计算子串 text[start, end) 的哈希值
     * 多项式哈希:H = sum(text[i] * BASE^(end-1-i)) % MOD
     */
    private long computeHash(String text, int start, int end) {
        long hash = 0;
        for (int i = start; i < end; i++) {
            hash = (hash * BASE + text.charAt(i)) % MOD;
        }
        return hash;
    }

    /**
     * 滚动更新哈希值:移除旧最高位,整体左移,补入新最低位
     * @param oldHash 当前窗口的哈希值
     * @param outgoing 移出窗口的字符(最高位)
     * @param incoming 进入窗口的字符(最低位)
     * @return 更新后的哈希值
     */
    private long rollHash(long oldHash, char outgoing, char incoming) {
        // 先移除最高位的贡献,加上 MOD 防止负数
        long removed = (oldHash - (outgoing * highestPower) % MOD + MOD) % MOD;
        // 整体左移一位(乘以BASE),然后加上新字符
        long shifted = (removed * BASE) % MOD;
        return (shifted + incoming) % MOD;
    }

    /**
     * 在文本中搜索模式串的所有匹配位置
     * @param text 待搜索的文本
     * @return 所有匹配起始位置的列表
     */
    public List<Integer> search(String text) {
        List<Integer> matches = new ArrayList<>();
        if (text == null || text.length() < patternLength) {
            return matches;
        }

        int textLength = text.length();
        // 计算文本第一个窗口的哈希值
        long textHash = computeHash(text, 0, patternLength);

        // 检查第一个窗口是否匹配
        if (textHash == patternHash && verifyMatch(text, 0)) {
            matches.add(0);
        }

        // 滚动遍历后续窗口
        for (int i = 1; i <= textLength - patternLength; i++) {
            char outgoing = text.charAt(i - 1);      // 移出窗口的字符
            char incoming = text.charAt(i + patternLength - 1); // 进入窗口的字符

            textHash = rollHash(textHash, outgoing, incoming);

            // 哈希值相等时进行精确比对(处理哈希碰撞)
            if (textHash == patternHash && verifyMatch(text, i)) {
                matches.add(i);
            }
        }

        return matches;
    }

    /**
     * 精确验证:当哈希值相等时,逐字符比对确认是否真正匹配
     * 这是处理哈希碰撞的必要步骤
     */
    private boolean verifyMatch(String text, int start) {
        for (int i = 0; i < patternLength; i++) {
            if (text.charAt(start + i) != pattern.charAt(i)) {
                return false;
            }
        }
        return true;
    }

    /**
     * 单模式匹配便捷方法:返回首次匹配位置,未找到返回-1
     */
    public int firstMatch(String text) {
        List<Integer> matches = search(text);
        return matches.isEmpty() ? -1 : matches.get(0);
    }

    // ==================== 多模式匹配扩展 ====================

    /**
     * 多模式Rabin-Karp匹配器
     * 使用多个哈希函数进一步降低碰撞概率
     */
    public static class MultiPatternMatcher {

        /** 双哈希的第二个模数,与第一个互质 */
        private static final long MOD2 = 1_000_000_009L;
        private static final int BASE2 = 257;

        private final String[] patterns;
        private final long[] patternHashes1;
        private final long[] patternHashes2;
        private final int patternLength;
        private long highestPower1;
        private long highestPower2;

        public MultiPatternMatcher(String[] patterns) {
            if (patterns == null || patterns.length == 0) {
                throw new IllegalArgumentException("模式串数组不能为空");
            }
            this.patterns = patterns;
            this.patternLength = patterns[0].length();
            // 要求所有模式串长度相同(Rabin-Karp多模式匹配的常见约束)
            for (String p : patterns) {
                if (p.length() != patternLength) {
                    throw new IllegalArgumentException("所有模式串长度必须相同");
                }
            }
            this.highestPower1 = computeHighestPower(patternLength, BASE, MOD);
            this.highestPower2 = computeHighestPower(patternLength, BASE2, MOD2);
            this.patternHashes1 = new long[patterns.length];
            this.patternHashes2 = new long[patterns.length];
            for (int i = 0; i < patterns.length; i++) {
                patternHashes1[i] = computeHash(patterns[i], 0, patternLength, BASE, MOD);
                patternHashes2[i] = computeHash(patterns[i], 0, patternLength, BASE2, MOD2);
            }
        }

        private long computeHighestPower(int length, int base, long mod) {
            long result = 1;
            for (int i = 1; i < length; i++) {
                result = (result * base) % mod;
            }
            return result;
        }

        private long computeHash(String text, int start, int end, int base, long mod) {
            long hash = 0;
            for (int i = start; i < end; i++) {
                hash = (hash * base + text.charAt(i)) % mod;
            }
            return hash;
        }

        private long rollHash(long oldHash, char outgoing, char incoming,
                              long highestPow, int base, long mod) {
            long removed = (oldHash - (outgoing * highestPow) % mod + mod) % mod;
            long shifted = (removed * base) % mod;
            return (shifted + incoming) % mod;
        }

        /**
         * 在文本中搜索多个模式串,返回每个匹配位置及对应模式串索引
         */
        public List<MatchResult> searchMultiple(String text) {
            List<MatchResult> results = new ArrayList<>();
            if (text == null || text.length() < patternLength) {
                return results;
            }

            long textHash1 = computeHash(text, 0, patternLength, BASE, MOD);
            long textHash2 = computeHash(text, 0, patternLength, BASE2, MOD2);
            checkMatches(text, 0, textHash1, textHash2, results);

            for (int i = 1; i <= text.length() - patternLength; i++) {
                textHash1 = rollHash(textHash1, text.charAt(i - 1),
                        text.charAt(i + patternLength - 1), highestPower1, BASE, MOD);
                textHash2 = rollHash(textHash2, text.charAt(i - 1),
                        text.charAt(i + patternLength - 1), highestPower2, BASE2, MOD2);
                checkMatches(text, i, textHash1, textHash2, results);
            }
            return results;
        }

        private void checkMatches(String text, int pos, long hash1, long hash2,
                                   List<MatchResult> results) {
            for (int i = 0; i < patterns.length; i++) {
                if (hash1 == patternHashes1[i] && hash2 == patternHashes2[i]) {
                    // 双重哈希同时碰撞的概率极低,但仍做精确验证
                    if (verifyMatch(text, pos, patterns[i])) {
                        results.add(new MatchResult(pos, patterns[i]));
                    }
                }
            }
        }

        private boolean verifyMatch(String text, int start, String pattern) {
            for (int i = 0; i < pattern.length(); i++) {
                if (text.charAt(start + i) != pattern.charAt(i)) return false;
            }
            return true;
        }
    }

    /**
     * 多模式匹配结果封装
     */
    public static class MatchResult {
        public final int position;
        public final String pattern;

        public MatchResult(int position, String pattern) {
            this.position = position;
            this.pattern = pattern;
        }

        @Override
        public String toString() {
            return String.format("位置 %d: \"%s\"", position, pattern);
        }
    }

    // ==================== 测试演示 ====================

    public static void main(String[] args) {
        System.out.println("===== Rabin-Karp 单模式匹配测试 =====\n");

        // 测试1:基本匹配
        String text1 = "ABABDABACDABABCABAB";
        String pattern1 = "ABABCABAB";
        RabinKarpMatcher matcher1 = new RabinKarpMatcher(pattern1);
        List<Integer> matches1 = matcher1.search(text1);
        System.out.println("文本: " + text1);
        System.out.println("模式: " + pattern1);
        System.out.println("匹配位置: " + matches1); // 期望 [10]

        // 测试2:多位置匹配(重叠匹配)
        String text2 = "AAAAA";
        String pattern2 = "AA";
        RabinKarpMatcher matcher2 = new RabinKarpMatcher(pattern2);
        List<Integer> matches2 = matcher2.search(text2);
        System.out.println("\n文本: " + text2);
        System.out.println("模式: " + pattern2);
        System.out.println("匹配位置: " + matches2); // 期望 [0, 1, 2, 3]

        // 测试3:无匹配
        String text3 = "HELLO WORLD";
        String pattern3 = "TEST";
        RabinKarpMatcher matcher3 = new RabinKarpMatcher(pattern3);
        List<Integer> matches3 = matcher3.search(text3);
        System.out.println("\n文本: " + text3);
        System.out.println("模式: " + pattern3);
        System.out.println("匹配位置: " + matches3); // 期望 []

        // 测试4:长文本性能测试
        System.out.println("\n===== 性能测试 =====");
        StringBuilder longText = new StringBuilder();
        for (int i = 0; i < 1_000_000; i++) {
            longText.append((char) ('A' + (i % 26)));
        }
        String longPattern = "XYZABC";
        // 在末尾插入一个匹配
        longText.setCharAt(999_990, 'X');
        longText.setCharAt(999_991, 'Y');
        longText.setCharAt(999_992, 'Z');
        longText.setCharAt(999_993, 'A');
        longText.setCharAt(999_994, 'B');
        longText.setCharAt(999_995, 'C');

        RabinKarpMatcher perfMatcher = new RabinKarpMatcher(longPattern);
        long startTime = System.currentTimeMillis();
        List<Integer> perfMatches = perfMatcher.search(longText.toString());
        long endTime = System.currentTimeMillis();
        System.out.println("百万字符文本搜索耗时: " + (endTime - startTime) + "ms");
        System.out.println("匹配位置: " + perfMatches);

        // 测试5:多模式匹配
        System.out.println("\n===== 多模式匹配测试 =====");
        String[] multiPatterns = {"ABC", "BCA", "CAB"};
        String multiText = "ABCABCAABCAB";
        MultiPatternMatcher multiMatcher = new MultiPatternMatcher(multiPatterns);
        List<MatchResult> multiResults = multiMatcher.searchMultiple(multiText);
        System.out.println("文本: " + multiText);
        System.out.println("模式集: " + String.join(", ", multiPatterns));
        for (MatchResult r : multiResults) {
            System.out.println(r);
        }
    }
}

双重哈希优化

单哈希函数的碰撞概率约为 $1/M$($M$ 为模数)。当处理大量数据或对抗性输入时,碰撞可能导致过多的精确比对回退。采用两个独立的哈希函数(不同基数和模数)可将碰撞概率降至约 $1/M^2$,这在实际应用中几乎可以忽略。

上述 MultiPatternMatcher 实现中,同时使用 $(BASE, MOD)$ 和 $(BASE2, MOD2)$ 两组参数,仅当两组哈希均相等时才触发精确验证。两个模数 $1,000,000,007$ 和 $1,000,000,009$ 均为大质数且互质,确保了哈希空间的独立性。

时间复杂度分析

场景 时间复杂度 说明
预处理 $O(m)$ 计算模式串哈希及 $BASE^{m-1}$
最佳情况 $O(n)$ 无哈希碰撞,每次比较均为 $O(1)$
平均情况 $O(n)$ 均匀分布文本,碰撞概率极低
最坏情况 $O(nm)$ 所有位置哈希均碰撞(如全相同字符)
空间复杂度 $O(1)$ 仅需常数额外空间

其中 $n$ 为文本长度,$m$ 为模式串长度。与KMP的严格 $O(n+m)$ 相比,Rabin-Karp的平均性能更优,且多模式扩展更加自然——只需维护一组模式串哈希值即可同时匹配多个模式。

实际应用场景

Rabin-Karp算法在工程实践中有着广泛的应用:

** plagiarism检测**:将文档分割为固定长度的指纹片段,通过哈希比对快速定位相似段落,这是Turnitin等系统的底层原理之一。

** 生物信息学**:DNA序列由A、C、G、T四种碱基组成,基数可选择4,模数选择大质数,在基因组比对中实现高效的短读段定位。

** 入侵检测系统**:Snort等网络入侵检测工具使用多模式Rabin-Karp变体,在数据包载荷中快速匹配大量已知攻击特征串。

** 代码查重**:将源代码按token或行进行哈希 fingerprint,通过滚动比对发现抄袭或重复代码块。

与KMP算法的对比

特性 Rabin-Karp KMP
核心思想 哈希筛选+精确验证 前缀函数跳过已匹配部分
最坏复杂度 $O(nm)$ 严格 $O(n+m)$
平均复杂度 $O(n)$ $O(n+m)$
多模式支持 天然支持 需构建AC自动机
空间占用 $O(1)$ $O(m)$
适用场景 大数据多模式过滤 确定性单模式精确匹配

选择策略上,当需要同时匹配大量模式串或文本数据量极大时,Rabin-Karp配合双哈希是更优的工程方案;而对于需要确定性复杂度保证的场景,KMP或Boyer-Moore更为稳妥。

总结

Rabin-Karp算法展示了哈希思想在字符串处理中的强大威力。通过将字符串映射为数值空间中的指纹,算法将线性字符比较转化为常数哈希比较,配合滚动更新技术实现了高效的滑动窗口扫描。双重哈希的引入进一步将碰撞概率压至极低水平,使其能够胜任生产环境中的大规模文本处理任务。理解这一算法的核心——多项式哈希与模运算的巧妙结合——对于掌握字符串算法乃至密码学中的散列设计都有着重要的启发意义。