aho_corasick.hpp¶
Aho-Corasick automaton for an integer alphabet [0, sigma).
Verified by aho_corasick.
把多个模式串建成自动机,一次扫描文本即可找出所有模式出现位置或累计匹配贡献。
Implementation¶
#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