m1une's library

This documentation is automatically generated by online-judge-tools/verification-helper

View on GitHub

:heavy_check_mark: Aho-Corasick
(string/aho_corasick.hpp)

Overview

AhoCorasick finds occurrences of many patterns in one text. It stores the patterns in a trie and adds failure links, allowing the text to be scanned in linear time plus the number of reported matches.

The alphabet must be a contiguous range of character codes. The default AhoCorasick<> accepts lowercase English letters. For decimal digits, use AhoCorasick<10, '0'>.

Construction

Insert every pattern, then call build():

m1une::string::AhoCorasick<> automaton;
int first_id = automaton.insert(std::string("he"));
int second_id = automaton.insert(std::string("she"));
automaton.build();

Pattern IDs are assigned in insertion order, starting from zero. Duplicate and empty patterns are supported and receive separate IDs.

No pattern may be inserted after build(). Call clear() to reuse the object with a new pattern set.

Methods

Let $K$ be the automaton node count, $P$ the pattern count, $\sigma$ the alphabet size, $N$ the text length, and $Z$ the number of reported occurrences.

Method Description Complexity    
insert(pattern) Inserts a pattern and returns its ID. $O( pattern )$
build() Builds failure links and all transitions. $O(K\sigma)$    
built() Returns whether build() has been called. $O(1)$    
pattern_count() Returns the number of inserted patterns. $O(1)$    
pattern_length(id) Returns a pattern’s length. $O(1)$    
pattern_node(id) Returns the terminal node of a pattern. $O(1)$    
node_count() Returns the number of trie nodes. $O(1)$    
root() Returns the root node ID. $O(1)$    
node(id) Returns a read-only node view. $O(1)$    
nodes() Returns a read-only view of the complete node array. $O(1)$    
bfs_order() Returns node IDs in failure-link BFS order. $O(1)$    
transition(state, symbol) Takes one automaton transition. $O(1)$    
for_each_output(state, callback) Reports patterns ending at a state. $O(1 + output)$    
match(text, callback) Reports every occurrence in the text. $O(N+Z)$    
count_occurrences(text) Returns an occurrence count for each pattern ID. $O(N+K+P)$    
reserve(node_capacity) Reserves trie-node storage before building. $O(K)$    
clear() Removes all patterns and returns to the insertion phase. $O(K)$    

match calls callback(end, pattern_id), where end is the exclusive end position. The occurrence starts at end - automaton.pattern_length(pattern_id).

An empty pattern occurs at every text boundary, including positions zero and text.size().

Node Data

Each Node exposes:

nodes() and bfs_order() make graph algorithms convenient without repeated accessor calls. For example, iterate bfs_order() in reverse to aggregate values from a node into its failure parent, or traverse failure_children to run a tree DP.

children, parent, and parent_symbol describe the sparse trie graph, while next describes the complete deterministic automaton graph. This distinction remains available after build() without storing two full transition tables per node.

Node IDs remain valid until clear(). References and iterators into nodes() may be invalidated by insert() or reserve() before construction finishes, so retain node IDs across insertions.

Example

#include "string/aho_corasick.hpp"

#include <iostream>
#include <string>

int main() {
    m1une::string::AhoCorasick<> automaton;
    automaton.insert(std::string("aba"));
    automaton.insert(std::string("ba"));
    automaton.build();

    automaton.match(std::string("ababa"), [](int end, int pattern_id) {
        std::cout << pattern_id << ' ' << end << "\n";
    });
}

Required by

Verified with

Code

#ifndef M1UNE_STRING_AHO_CORASICK_HPP
#define M1UNE_STRING_AHO_CORASICK_HPP 1

#include <array>
#include <cassert>
#include <cstddef>
#include <limits>
#include <queue>
#include <vector>

namespace m1une {
namespace string {

// Aho-Corasick automaton for a contiguous character alphabet.
template <int AlphabetSize = 26, int FirstCharacter = 'a'>
struct AhoCorasick {
    static_assert(0 < AlphabetSize);

    using node_id = int;
    static constexpr node_id null_node = -1;

    struct Node {
        // Completed automaton transitions. Valid after build().
        std::array<node_id, AlphabetSize> next;
        node_id failure;
        node_id output_link;
        node_id parent;
        int parent_symbol;
        int depth;
        std::vector<node_id> children;
        std::vector<node_id> failure_children;
        std::vector<int> pattern_ids;

        Node(
            node_id parent_value = null_node,
            int parent_symbol_value = -1,
            int depth_value = 0
        ) : failure(0),
            output_link(null_node),
            parent(parent_value),
            parent_symbol(parent_symbol_value),
            depth(depth_value) {
            next.fill(null_node);
        }
    };

   private:
    std::vector<Node> _nodes;
    std::vector<int> _pattern_length;
    std::vector<node_id> _pattern_node;
    std::vector<node_id> _bfs_order;
    bool _built;

    template <class Symbol>
    static int symbol_index(const Symbol& symbol) {
        int index = int(symbol) - FirstCharacter;
        assert(0 <= index && index < AlphabetSize);
        return index;
    }

    node_id new_node(node_id parent, int parent_symbol) {
        assert(_nodes.size() < std::size_t(std::numeric_limits<int>::max()));
        assert(_nodes[parent].depth < std::numeric_limits<int>::max());
        _nodes.emplace_back(parent, parent_symbol, _nodes[parent].depth + 1);
        return int(_nodes.size()) - 1;
    }

   public:
    AhoCorasick() : _nodes(1), _built(false) {}

    node_id root() const {
        return 0;
    }

    bool built() const {
        return _built;
    }

    int pattern_count() const {
        return int(_pattern_length.size());
    }

    int pattern_length(int pattern_id) const {
        assert(0 <= pattern_id && pattern_id < pattern_count());
        return _pattern_length[pattern_id];
    }

    node_id pattern_node(int pattern_id) const {
        assert(0 <= pattern_id && pattern_id < pattern_count());
        return _pattern_node[pattern_id];
    }

    std::size_t node_count() const {
        return _nodes.size();
    }

    const std::vector<Node>& nodes() const {
        return _nodes;
    }

    const Node& node(node_id id) const {
        assert(0 <= id && std::size_t(id) < _nodes.size());
        return _nodes[id];
    }

    // Returns nodes in failure-link BFS order, beginning with the root.
    const std::vector<node_id>& bfs_order() const {
        assert(_built);
        return _bfs_order;
    }

    void reserve(std::size_t node_capacity) {
        assert(!_built);
        _nodes.reserve(node_capacity);
    }

    void clear() {
        _nodes.clear();
        _nodes.emplace_back();
        _pattern_length.clear();
        _pattern_node.clear();
        _bfs_order.clear();
        _built = false;
    }

    // Inserts a pattern and returns its insertion-order ID.
    template <class Sequence>
    int insert(const Sequence& pattern) {
        assert(!_built);
        int pattern_id = pattern_count();
        int length = 0;
        node_id state = root();
        for (const auto& symbol : pattern) {
            assert(length < std::numeric_limits<int>::max());
            int index = symbol_index(symbol);
            if (_nodes[state].next[index] == null_node) {
                node_id child = new_node(state, index);
                _nodes[state].next[index] = child;
                _nodes[state].children.push_back(child);
            }
            state = _nodes[state].next[index];
            length++;
        }
        _nodes[state].pattern_ids.push_back(pattern_id);
        _pattern_length.push_back(length);
        _pattern_node.push_back(state);
        return pattern_id;
    }

    // Builds failure links and completes every automaton transition.
    void build() {
        assert(!_built);
        std::queue<node_id> queue;
        _bfs_order.clear();
        _bfs_order.reserve(_nodes.size());
        _bfs_order.push_back(root());

        for (int symbol = 0; symbol < AlphabetSize; ++symbol) {
            node_id child = _nodes[root()].next[symbol];
            if (child == null_node) {
                _nodes[root()].next[symbol] = root();
            } else {
                _nodes[root()].next[symbol] = child;
                _nodes[child].failure = root();
                _nodes[child].output_link =
                    _nodes[root()].pattern_ids.empty() ? null_node : root();
                _nodes[root()].failure_children.push_back(child);
                queue.push(child);
            }
        }

        while (!queue.empty()) {
            node_id state = queue.front();
            queue.pop();
            _bfs_order.push_back(state);

            for (int symbol = 0; symbol < AlphabetSize; ++symbol) {
                node_id child = _nodes[state].next[symbol];
                if (child == null_node) {
                    _nodes[state].next[symbol] =
                        _nodes[_nodes[state].failure].next[symbol];
                    continue;
                }

                _nodes[state].next[symbol] = child;
                node_id failure =
                    _nodes[_nodes[state].failure].next[symbol];
                _nodes[child].failure = failure;
                _nodes[child].output_link =
                    _nodes[failure].pattern_ids.empty()
                        ? _nodes[failure].output_link
                        : failure;
                _nodes[failure].failure_children.push_back(child);
                queue.push(child);
            }
        }
        _built = true;
    }

    template <class Symbol>
    node_id transition(node_id state, const Symbol& symbol) const {
        assert(_built);
        assert(0 <= state && std::size_t(state) < _nodes.size());
        return _nodes[state].next[symbol_index(symbol)];
    }

    // Calls callback(pattern_id) for every pattern ending at `state`.
    template <class Callback>
    void for_each_output(node_id state, Callback callback) const {
        assert(_built);
        assert(0 <= state && std::size_t(state) < _nodes.size());
        while (state != null_node) {
            for (int pattern_id : _nodes[state].pattern_ids) {
                callback(pattern_id);
            }
            state = _nodes[state].output_link;
        }
    }

    // Calls callback(end, pattern_id) for every occurrence. `end` is the
    // exclusive end position. Empty patterns occur at every text boundary.
    template <class Sequence, class Callback>
    void match(const Sequence& text, Callback callback) const {
        assert(_built);
        node_id state = root();
        for_each_output(state, [&callback](int pattern_id) {
            callback(0, pattern_id);
        });

        int end = 0;
        for (const auto& symbol : text) {
            state = transition(state, symbol);
            end++;
            for_each_output(state, [&callback, end](int pattern_id) {
                callback(end, pattern_id);
            });
        }
    }

    // Counts occurrences of every inserted pattern in linear time.
    template <class Sequence>
    std::vector<long long> count_occurrences(const Sequence& text) const {
        assert(_built);
        std::vector<long long> visits(_nodes.size(), 0);
        node_id state = root();
        visits[root()]++;
        for (const auto& symbol : text) {
            state = transition(state, symbol);
            visits[state]++;
        }

        for (std::size_t index = _bfs_order.size(); index-- > 1;) {
            node_id current = _bfs_order[index];
            visits[_nodes[current].failure] += visits[current];
        }

        std::vector<long long> result(pattern_count(), 0);
        for (node_id current : _bfs_order) {
            for (int pattern_id : _nodes[current].pattern_ids) {
                result[pattern_id] = visits[current];
            }
        }
        return result;
    }
};

}  // namespace string
}  // namespace m1une

#endif  // M1UNE_STRING_AHO_CORASICK_HPP
#line 1 "string/aho_corasick.hpp"



#include <array>
#include <cassert>
#include <cstddef>
#include <limits>
#include <queue>
#include <vector>

namespace m1une {
namespace string {

// Aho-Corasick automaton for a contiguous character alphabet.
template <int AlphabetSize = 26, int FirstCharacter = 'a'>
struct AhoCorasick {
    static_assert(0 < AlphabetSize);

    using node_id = int;
    static constexpr node_id null_node = -1;

    struct Node {
        // Completed automaton transitions. Valid after build().
        std::array<node_id, AlphabetSize> next;
        node_id failure;
        node_id output_link;
        node_id parent;
        int parent_symbol;
        int depth;
        std::vector<node_id> children;
        std::vector<node_id> failure_children;
        std::vector<int> pattern_ids;

        Node(
            node_id parent_value = null_node,
            int parent_symbol_value = -1,
            int depth_value = 0
        ) : failure(0),
            output_link(null_node),
            parent(parent_value),
            parent_symbol(parent_symbol_value),
            depth(depth_value) {
            next.fill(null_node);
        }
    };

   private:
    std::vector<Node> _nodes;
    std::vector<int> _pattern_length;
    std::vector<node_id> _pattern_node;
    std::vector<node_id> _bfs_order;
    bool _built;

    template <class Symbol>
    static int symbol_index(const Symbol& symbol) {
        int index = int(symbol) - FirstCharacter;
        assert(0 <= index && index < AlphabetSize);
        return index;
    }

    node_id new_node(node_id parent, int parent_symbol) {
        assert(_nodes.size() < std::size_t(std::numeric_limits<int>::max()));
        assert(_nodes[parent].depth < std::numeric_limits<int>::max());
        _nodes.emplace_back(parent, parent_symbol, _nodes[parent].depth + 1);
        return int(_nodes.size()) - 1;
    }

   public:
    AhoCorasick() : _nodes(1), _built(false) {}

    node_id root() const {
        return 0;
    }

    bool built() const {
        return _built;
    }

    int pattern_count() const {
        return int(_pattern_length.size());
    }

    int pattern_length(int pattern_id) const {
        assert(0 <= pattern_id && pattern_id < pattern_count());
        return _pattern_length[pattern_id];
    }

    node_id pattern_node(int pattern_id) const {
        assert(0 <= pattern_id && pattern_id < pattern_count());
        return _pattern_node[pattern_id];
    }

    std::size_t node_count() const {
        return _nodes.size();
    }

    const std::vector<Node>& nodes() const {
        return _nodes;
    }

    const Node& node(node_id id) const {
        assert(0 <= id && std::size_t(id) < _nodes.size());
        return _nodes[id];
    }

    // Returns nodes in failure-link BFS order, beginning with the root.
    const std::vector<node_id>& bfs_order() const {
        assert(_built);
        return _bfs_order;
    }

    void reserve(std::size_t node_capacity) {
        assert(!_built);
        _nodes.reserve(node_capacity);
    }

    void clear() {
        _nodes.clear();
        _nodes.emplace_back();
        _pattern_length.clear();
        _pattern_node.clear();
        _bfs_order.clear();
        _built = false;
    }

    // Inserts a pattern and returns its insertion-order ID.
    template <class Sequence>
    int insert(const Sequence& pattern) {
        assert(!_built);
        int pattern_id = pattern_count();
        int length = 0;
        node_id state = root();
        for (const auto& symbol : pattern) {
            assert(length < std::numeric_limits<int>::max());
            int index = symbol_index(symbol);
            if (_nodes[state].next[index] == null_node) {
                node_id child = new_node(state, index);
                _nodes[state].next[index] = child;
                _nodes[state].children.push_back(child);
            }
            state = _nodes[state].next[index];
            length++;
        }
        _nodes[state].pattern_ids.push_back(pattern_id);
        _pattern_length.push_back(length);
        _pattern_node.push_back(state);
        return pattern_id;
    }

    // Builds failure links and completes every automaton transition.
    void build() {
        assert(!_built);
        std::queue<node_id> queue;
        _bfs_order.clear();
        _bfs_order.reserve(_nodes.size());
        _bfs_order.push_back(root());

        for (int symbol = 0; symbol < AlphabetSize; ++symbol) {
            node_id child = _nodes[root()].next[symbol];
            if (child == null_node) {
                _nodes[root()].next[symbol] = root();
            } else {
                _nodes[root()].next[symbol] = child;
                _nodes[child].failure = root();
                _nodes[child].output_link =
                    _nodes[root()].pattern_ids.empty() ? null_node : root();
                _nodes[root()].failure_children.push_back(child);
                queue.push(child);
            }
        }

        while (!queue.empty()) {
            node_id state = queue.front();
            queue.pop();
            _bfs_order.push_back(state);

            for (int symbol = 0; symbol < AlphabetSize; ++symbol) {
                node_id child = _nodes[state].next[symbol];
                if (child == null_node) {
                    _nodes[state].next[symbol] =
                        _nodes[_nodes[state].failure].next[symbol];
                    continue;
                }

                _nodes[state].next[symbol] = child;
                node_id failure =
                    _nodes[_nodes[state].failure].next[symbol];
                _nodes[child].failure = failure;
                _nodes[child].output_link =
                    _nodes[failure].pattern_ids.empty()
                        ? _nodes[failure].output_link
                        : failure;
                _nodes[failure].failure_children.push_back(child);
                queue.push(child);
            }
        }
        _built = true;
    }

    template <class Symbol>
    node_id transition(node_id state, const Symbol& symbol) const {
        assert(_built);
        assert(0 <= state && std::size_t(state) < _nodes.size());
        return _nodes[state].next[symbol_index(symbol)];
    }

    // Calls callback(pattern_id) for every pattern ending at `state`.
    template <class Callback>
    void for_each_output(node_id state, Callback callback) const {
        assert(_built);
        assert(0 <= state && std::size_t(state) < _nodes.size());
        while (state != null_node) {
            for (int pattern_id : _nodes[state].pattern_ids) {
                callback(pattern_id);
            }
            state = _nodes[state].output_link;
        }
    }

    // Calls callback(end, pattern_id) for every occurrence. `end` is the
    // exclusive end position. Empty patterns occur at every text boundary.
    template <class Sequence, class Callback>
    void match(const Sequence& text, Callback callback) const {
        assert(_built);
        node_id state = root();
        for_each_output(state, [&callback](int pattern_id) {
            callback(0, pattern_id);
        });

        int end = 0;
        for (const auto& symbol : text) {
            state = transition(state, symbol);
            end++;
            for_each_output(state, [&callback, end](int pattern_id) {
                callback(end, pattern_id);
            });
        }
    }

    // Counts occurrences of every inserted pattern in linear time.
    template <class Sequence>
    std::vector<long long> count_occurrences(const Sequence& text) const {
        assert(_built);
        std::vector<long long> visits(_nodes.size(), 0);
        node_id state = root();
        visits[root()]++;
        for (const auto& symbol : text) {
            state = transition(state, symbol);
            visits[state]++;
        }

        for (std::size_t index = _bfs_order.size(); index-- > 1;) {
            node_id current = _bfs_order[index];
            visits[_nodes[current].failure] += visits[current];
        }

        std::vector<long long> result(pattern_count(), 0);
        for (node_id current : _bfs_order) {
            for (int pattern_id : _nodes[current].pattern_ids) {
                result[pattern_id] = visits[current];
            }
        }
        return result;
    }
};

}  // namespace string
}  // namespace m1une
Back to top page