Skip to content

aho_corasick.hpp

SECTIONString INCLUDEnoya/aho_corasick.hpp

Aho-Corasick automaton for an integer alphabet [0, sigma).

Verified by aho_corasick.

把多个模式串建成自动机,一次扫描文本即可找出所有模式出现位置或累计匹配贡献。

Implementation

View on GitHub

#ifndef NOYA_AHO_CORASICK_HPP
#define NOYA_AHO_CORASICK_HPP 1

/// @complexity Time: O(S sigma) build and O(n + z) matching, for S states and z matches.
/// Space: O(S sigma + z).

#include <array>
#include <cassert>
#include <queue>
#include <string>
#include <utility>
#include <vector>

namespace noya {

/// @brief Aho-Corasick automaton for an integer alphabet [0, sigma).
template <int sigma = 26> struct aho_corasick {
  struct node {
    std::array<int, sigma> next;
    int parent = -1;
    int link = 0;
    int dictionary_link = -1;
    std::vector<int> output;

    node() { next.fill(-1); }
  };

  std::vector<node> nodes = {node{}};
  std::vector<int> pattern_node;
  std::vector<int> bfs_order;
  bool built = false;

  /// @brief Add a nonempty pattern and return its pattern id.
  int add(const std::vector<int> &pattern) {
    assert(!built);
    assert(!pattern.empty());
    int state = 0;
    for (int c : pattern) {
      assert(0 <= c && c < sigma);
      if (nodes[state].next[c] == -1) {
        nodes[state].next[c] = int(nodes.size());
        nodes.emplace_back();
        nodes.back().parent = state;
      }
      state = nodes[state].next[c];
    }
    int id = int(pattern_node.size());
    pattern_node.push_back(state);
    nodes[state].output.push_back(id);
    return id;
  }

  int add(const std::string &pattern, char first = 'a') {
    std::vector<int> converted;
    converted.reserve(pattern.size());
    for (unsigned char c : pattern) {
      converted.push_back(int(c) - int(static_cast<unsigned char>(first)));
    }
    return add(converted);
  }

  /// @brief Build failure links and total transitions in O(number_of_nodes *
  /// sigma).
  void build() {
    assert(!built);
    built = true;
    std::queue<int> queue;
    bfs_order = {0};
    for (int c = 0; c < sigma; c++) {
      int child = nodes[0].next[c];
      if (child == -1) {
        nodes[0].next[c] = 0;
      } else {
        nodes[child].link = 0;
        nodes[child].dictionary_link = -1;
        queue.push(child);
      }
    }
    while (!queue.empty()) {
      int u = queue.front();
      queue.pop();
      bfs_order.push_back(u);
      for (int c = 0; c < sigma; c++) {
        int child = nodes[u].next[c];
        if (child == -1) {
          nodes[u].next[c] = nodes[nodes[u].link].next[c];
          continue;
        }
        int link = nodes[nodes[u].link].next[c];
        nodes[child].link = link;
        nodes[child].dictionary_link =
            nodes[link].output.empty() ? nodes[link].dictionary_link : link;
        queue.push(child);
      }
    }
  }

  int transition(int state, int c) const {
    assert(built);
    assert(0 <= state && state < int(nodes.size()));
    assert(0 <= c && c < sigma);
    return nodes[state].next[c];
  }

  /// @brief Enumerate matches as (inclusive end position, pattern id).
  std::vector<std::pair<int, int>> match(const std::vector<int> &text) const {
    assert(built);
    std::vector<std::pair<int, int>> result;
    int state = 0;
    for (int i = 0; i < int(text.size()); i++) {
      state = transition(state, text[i]);
      for (int u = state; u != -1; u = nodes[u].dictionary_link) {
        for (int pattern_id : nodes[u].output) {
          result.emplace_back(i, pattern_id);
        }
      }
    }
    return result;
  }

  std::vector<std::pair<int, int>> match(const std::string &text,
                                         char first = 'a') const {
    std::vector<int> converted;
    converted.reserve(text.size());
    for (unsigned char c : text) {
      converted.push_back(int(c) - int(static_cast<unsigned char>(first)));
    }
    return match(converted);
  }

  /// @brief Count occurrences of every inserted pattern in O(text + nodes +
  /// patterns).
  std::vector<long long> count_matches(const std::vector<int> &text) const {
    assert(built);
    std::vector<long long> visits(nodes.size());
    int state = 0;
    for (int c : text) {
      state = transition(state, c);
      visits[state]++;
    }
    for (int i = int(bfs_order.size()) - 1; i > 0; i--) {
      int u = bfs_order[i];
      visits[nodes[u].link] += visits[u];
    }
    std::vector<long long> result(pattern_node.size());
    for (int id = 0; id < int(pattern_node.size()); id++) {
      result[id] = visits[pattern_node[id]];
    }
    return result;
  }

  std::vector<long long> count_matches(const std::string &text,
                                       char first = 'a') const {
    std::vector<int> converted;
    converted.reserve(text.size());
    for (unsigned char c : text) {
      converted.push_back(int(c) - int(static_cast<unsigned char>(first)));
    }
    return count_matches(converted);
  }
};

} // namespace noya

#endif // NOYA_AHO_CORASICK_HPP
#include <array>
#include <cassert>
#include <queue>
#include <string>
#include <utility>
#include <vector>

/// @complexity Time: O(S sigma) build and O(n + z) matching, for S states and z matches.
/// Space: O(S sigma + z).

namespace noya {

/// @brief Aho-Corasick automaton for an integer alphabet [0, sigma).
template <int sigma = 26> struct aho_corasick {
  struct node {
    std::array<int, sigma> next;
    int parent = -1;
    int link = 0;
    int dictionary_link = -1;
    std::vector<int> output;

    node() { next.fill(-1); }
  };

  std::vector<node> nodes = {node{}};
  std::vector<int> pattern_node;
  std::vector<int> bfs_order;
  bool built = false;

  /// @brief Add a nonempty pattern and return its pattern id.
  int add(const std::vector<int> &pattern) {
    assert(!built);
    assert(!pattern.empty());
    int state = 0;
    for (int c : pattern) {
      assert(0 <= c && c < sigma);
      if (nodes[state].next[c] == -1) {
        nodes[state].next[c] = int(nodes.size());
        nodes.emplace_back();
        nodes.back().parent = state;
      }
      state = nodes[state].next[c];
    }
    int id = int(pattern_node.size());
    pattern_node.push_back(state);
    nodes[state].output.push_back(id);
    return id;
  }

  int add(const std::string &pattern, char first = 'a') {
    std::vector<int> converted;
    converted.reserve(pattern.size());
    for (unsigned char c : pattern) {
      converted.push_back(int(c) - int(static_cast<unsigned char>(first)));
    }
    return add(converted);
  }

  /// @brief Build failure links and total transitions in O(number_of_nodes *
  /// sigma).
  void build() {
    assert(!built);
    built = true;
    std::queue<int> queue;
    bfs_order = {0};
    for (int c = 0; c < sigma; c++) {
      int child = nodes[0].next[c];
      if (child == -1) {
        nodes[0].next[c] = 0;
      } else {
        nodes[child].link = 0;
        nodes[child].dictionary_link = -1;
        queue.push(child);
      }
    }
    while (!queue.empty()) {
      int u = queue.front();
      queue.pop();
      bfs_order.push_back(u);
      for (int c = 0; c < sigma; c++) {
        int child = nodes[u].next[c];
        if (child == -1) {
          nodes[u].next[c] = nodes[nodes[u].link].next[c];
          continue;
        }
        int link = nodes[nodes[u].link].next[c];
        nodes[child].link = link;
        nodes[child].dictionary_link =
            nodes[link].output.empty() ? nodes[link].dictionary_link : link;
        queue.push(child);
      }
    }
  }

  int transition(int state, int c) const {
    assert(built);
    assert(0 <= state && state < int(nodes.size()));
    assert(0 <= c && c < sigma);
    return nodes[state].next[c];
  }

  /// @brief Enumerate matches as (inclusive end position, pattern id).
  std::vector<std::pair<int, int>> match(const std::vector<int> &text) const {
    assert(built);
    std::vector<std::pair<int, int>> result;
    int state = 0;
    for (int i = 0; i < int(text.size()); i++) {
      state = transition(state, text[i]);
      for (int u = state; u != -1; u = nodes[u].dictionary_link) {
        for (int pattern_id : nodes[u].output) {
          result.emplace_back(i, pattern_id);
        }
      }
    }
    return result;
  }

  std::vector<std::pair<int, int>> match(const std::string &text,
                                         char first = 'a') const {
    std::vector<int> converted;
    converted.reserve(text.size());
    for (unsigned char c : text) {
      converted.push_back(int(c) - int(static_cast<unsigned char>(first)));
    }
    return match(converted);
  }

  /// @brief Count occurrences of every inserted pattern in O(text + nodes +
  /// patterns).
  std::vector<long long> count_matches(const std::vector<int> &text) const {
    assert(built);
    std::vector<long long> visits(nodes.size());
    int state = 0;
    for (int c : text) {
      state = transition(state, c);
      visits[state]++;
    }
    for (int i = int(bfs_order.size()) - 1; i > 0; i--) {
      int u = bfs_order[i];
      visits[nodes[u].link] += visits[u];
    }
    std::vector<long long> result(pattern_node.size());
    for (int id = 0; id < int(pattern_node.size()); id++) {
      result[id] = visits[pattern_node[id]];
    }
    return result;
  }

  std::vector<long long> count_matches(const std::string &text,
                                       char first = 'a') const {
    std::vector<int> converted;
    converted.reserve(text.size());
    for (unsigned char c : text) {
      converted.push_back(int(c) - int(static_cast<unsigned char>(first)));
    }
    return count_matches(converted);
  }
};

} // namespace noya