m1une's library

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

View on GitHub

:heavy_check_mark: Matrix-Tree Theorem
(graph/matrix_tree_theorem.hpp)

Overview

matrix_tree_theorem.hpp counts spanning trees with a determinant of a Laplacian cofactor. It supports undirected spanning trees and both orientations of rooted directed spanning trees.

The functions are weighted. The weight of one tree is the product of its edge weights, and the answer is the sum of those products over every valid tree. Using the default graph edge weight 1 gives the ordinary number of trees.

Parallel edges are distinct choices and are supported. Self-loops never belong to a spanning tree and are ignored. Inactive edges are also ignored. A disconnected graph, or a directed graph with no valid rooted arborescence, returns 0.

Requirements

Field must provide exact field arithmetic: construction from an edge weight, addition, subtraction, multiplication, division by a nonzero value, and exact comparison with zero. A modular integer with a prime modulus, such as m1une::math::modint998244353, is the usual choice. Plain integer types are not suitable because Gaussian elimination divides by pivot values.

The graph must contain at least one vertex. count_spanning_trees expects edges created with Graph::add_edge. The directed functions expect edges created with Graph::add_directed_edge. Debug builds assert these storage conventions and the validity of root.

All intermediate field values and the result must be representable by Field.

Directed Orientation

count_out_arborescences(graph, root) counts directed spanning trees whose edges point away from root. Every vertex is reachable from root, and every non-root vertex has exactly one selected incoming edge.

count_in_arborescences(graph, root) counts directed spanning trees whose edges point toward root. The root is reachable from every vertex, and every non-root vertex has exactly one selected outgoing edge.

Reversing every edge swaps these two answers.

Interface

All functions are in namespace m1une::graph.

Function Exact signature Description Complexity
count_spanning_trees template <class Field, class Weight> Field count_spanning_trees(const Graph<Weight>& graph) Returns the total weight of all undirected spanning trees. $O(N^3 + M)$ time, $O(N^2 + M)$ auxiliary memory in debug builds and $O(N^2)$ in release builds.
count_out_arborescences template <class Field, class Weight> Field count_out_arborescences(const Graph<Weight>& graph, int root) Returns the total weight of all directed spanning trees rooted outward at root. $O(N^3 + M)$ time, $O(N^2 + M)$ auxiliary memory in debug builds and $O(N^2)$ in release builds.
count_in_arborescences template <class Field, class Weight> Field count_in_arborescences(const Graph<Weight>& graph, int root) Returns the total weight of all directed spanning trees rooted inward at root. $O(N^3 + M)$ time, $O(N^2 + M)$ auxiliary memory in debug builds and $O(N^2)$ in release builds.

The functions do not mutate the graph. For a one-vertex graph, each function returns 1: the empty edge set is the unique spanning tree.

Example

#include "graph/graph.hpp"
#include "graph/matrix_tree_theorem.hpp"
#include "math/modint.hpp"

#include <iostream>

int main() {
    using mint = m1une::math::modint998244353;

    m1une::graph::Graph<int> graph(3);
    graph.add_edge(0, 1);
    graph.add_edge(1, 2);
    graph.add_edge(2, 0);

    mint answer = m1une::graph::count_spanning_trees<mint>(graph);
    std::cout << answer << '\n';  // 3
}

Depends on

Required by

Verified with

Code

#ifndef M1UNE_GRAPH_MATRIX_TREE_THEOREM_HPP
#define M1UNE_GRAPH_MATRIX_TREE_THEOREM_HPP 1

#include <cassert>
#include <utility>
#include <vector>

#include "../math/matrix/linear_algebra.hpp"
#include "graph.hpp"

namespace m1une {
namespace graph {

namespace matrix_tree_detail {

inline int minor_index(int vertex, int removed) {
    assert(vertex != removed);
    return vertex < removed ? vertex : vertex - 1;
}

template <class Weight>
void assert_edge_incidence(const Graph<Weight>& graph, int expected) {
#ifndef NDEBUG
    std::vector<int> incidence(graph.edge_count(), 0);
    for (int vertex = 0; vertex < graph.size(); vertex++) {
        for (const Edge<Weight>& edge : graph[vertex]) {
            if (!edge.alive) continue;
            assert(0 <= edge.id && edge.id < graph.edge_count());
            incidence[edge.id]++;
        }
    }
    for (int count : incidence) {
        if (count != 0) assert(count == expected);
    }
#else
    (void)graph;
    (void)expected;
#endif
}

template <class Field, class Weight>
Field count_arborescences(
    const Graph<Weight>& graph,
    int root,
    bool outward
) {
    const int n = graph.size();
    assert(0 <= root && root < n);
    assert_edge_incidence(graph, 1);

    matrix::Matrix<Field> minor(n - 1, n - 1);
    for (int vertex = 0; vertex < n; vertex++) {
        for (const Edge<Weight>& edge : graph[vertex]) {
            if (!edge.alive || edge.from == edge.to) continue;
            const int row = outward ? edge.to : edge.from;
            const int col = outward ? edge.from : edge.to;
            if (row == root) continue;

            const Field weight(edge.cost);
            const int reduced_row = minor_index(row, root);
            minor[reduced_row][reduced_row] += weight;
            if (col != root) {
                minor[reduced_row][minor_index(col, root)] -= weight;
            }
        }
    }
    return matrix::determinant(std::move(minor));
}

}  // namespace matrix_tree_detail

// Returns the total weight of all undirected spanning trees. The weight of a
// tree is the product of its edge costs.
template <class Field, class Weight>
Field count_spanning_trees(const Graph<Weight>& graph) {
    const int n = graph.size();
    assert(n > 0);
    matrix_tree_detail::assert_edge_incidence(graph, 2);

    const int removed = n - 1;
    matrix::Matrix<Field> minor(n - 1, n - 1);
    for (int vertex = 0; vertex < n; vertex++) {
        for (const Edge<Weight>& edge : graph[vertex]) {
            if (!edge.alive || edge.from >= edge.to) continue;
            const int from = edge.from;
            const int to = edge.to;
            const Field weight(edge.cost);

            if (from != removed) {
                const int reduced_from = matrix_tree_detail::minor_index(from, removed);
                minor[reduced_from][reduced_from] += weight;
            }
            if (to != removed) {
                const int reduced_to = matrix_tree_detail::minor_index(to, removed);
                minor[reduced_to][reduced_to] += weight;
            }
            if (from != removed && to != removed) {
                const int reduced_from = matrix_tree_detail::minor_index(from, removed);
                const int reduced_to = matrix_tree_detail::minor_index(to, removed);
                minor[reduced_from][reduced_to] -= weight;
                minor[reduced_to][reduced_from] -= weight;
            }
        }
    }
    return matrix::determinant(std::move(minor));
}

// Counts directed spanning trees whose edges point away from root, so every
// vertex is reachable from root.
template <class Field, class Weight>
Field count_out_arborescences(const Graph<Weight>& graph, int root) {
    return matrix_tree_detail::count_arborescences<Field>(graph, root, true);
}

// Counts directed spanning trees whose edges point toward root, so root is
// reachable from every vertex.
template <class Field, class Weight>
Field count_in_arborescences(const Graph<Weight>& graph, int root) {
    return matrix_tree_detail::count_arborescences<Field>(graph, root, false);
}

}  // namespace graph
}  // namespace m1une

#endif  // M1UNE_GRAPH_MATRIX_TREE_THEOREM_HPP
#line 1 "graph/matrix_tree_theorem.hpp"



#include <cassert>
#include <utility>
#include <vector>

#line 1 "math/matrix/linear_algebra.hpp"



#include <optional>
#include <type_traits>
#line 7 "math/matrix/linear_algebra.hpp"

#line 1 "math/matrix/matrix.hpp"



#line 5 "math/matrix/matrix.hpp"
#include <cstddef>
#include <cstdint>
#line 9 "math/matrix/matrix.hpp"

namespace m1une {
namespace matrix {

template <class T>
class Matrix {
   private:
    int _rows;
    int _cols;
    std::vector<T> _data;

    static std::size_t storage_size(int rows, int cols) {
        assert(rows >= 0);
        assert(cols >= 0);
        return std::size_t(rows) * std::size_t(cols);
    }

   public:
    using value_type = T;

    Matrix() : _rows(0), _cols(0) {}

    Matrix(int rows, int cols, const T& value = T())
        : _rows(rows), _cols(cols), _data(storage_size(rows, cols), value) {}

    Matrix(int rows, int cols, std::vector<T> values)
        : _rows(rows), _cols(cols), _data(std::move(values)) {
        assert(rows >= 0);
        assert(cols >= 0);
        assert(_data.size() == std::size_t(rows) * std::size_t(cols));
    }

    explicit Matrix(const std::vector<std::vector<T>>& values)
        : _rows(int(values.size())), _cols(values.empty() ? 0 : int(values[0].size())),
          _data(storage_size(_rows, _cols)) {
        for (int row = 0; row < _rows; row++) {
            assert(int(values[std::size_t(row)].size()) == _cols);
            for (int col = 0; col < _cols; col++) {
                (*this)[row][col] = values[std::size_t(row)][std::size_t(col)];
            }
        }
    }

    int rows() const {
        return _rows;
    }

    int cols() const {
        return _cols;
    }

    bool empty() const {
        return _rows == 0 || _cols == 0;
    }

    std::vector<T>& data() {
        return _data;
    }

    const std::vector<T>& data() const {
        return _data;
    }

    T* operator[](int row) {
        assert(0 <= row && row < _rows);
        return _data.data() + std::size_t(row) * std::size_t(_cols);
    }

    const T* operator[](int row) const {
        assert(0 <= row && row < _rows);
        return _data.data() + std::size_t(row) * std::size_t(_cols);
    }

    T& operator()(int row, int col) {
        assert(0 <= col && col < _cols);
        return (*this)[row][col];
    }

    const T& operator()(int row, int col) const {
        assert(0 <= col && col < _cols);
        return (*this)[row][col];
    }

    static Matrix identity(int size) {
        assert(size >= 0);
        Matrix result(size, size);
        for (int i = 0; i < size; i++) result[i][i] = T(1);
        return result;
    }

    Matrix transposed() const {
        Matrix result(_cols, _rows);
        for (int row = 0; row < _rows; row++) {
            for (int col = 0; col < _cols; col++) {
                result[col][row] = (*this)[row][col];
            }
        }
        return result;
    }

    void swap_rows(int first, int second) {
        assert(0 <= first && first < _rows);
        assert(0 <= second && second < _rows);
        if (first == second) return;
        for (int col = 0; col < _cols; col++) {
            std::swap((*this)[first][col], (*this)[second][col]);
        }
    }

    Matrix& operator+=(const Matrix& rhs) {
        assert(_rows == rhs._rows && _cols == rhs._cols);
        for (std::size_t i = 0; i < _data.size(); i++) _data[i] += rhs._data[i];
        return *this;
    }

    Matrix& operator-=(const Matrix& rhs) {
        assert(_rows == rhs._rows && _cols == rhs._cols);
        for (std::size_t i = 0; i < _data.size(); i++) _data[i] -= rhs._data[i];
        return *this;
    }

    Matrix& operator*=(const T& scalar) {
        for (T& value : _data) value *= scalar;
        return *this;
    }

    Matrix& operator/=(const T& scalar) {
        for (T& value : _data) value /= scalar;
        return *this;
    }

    Matrix& operator*=(const Matrix& rhs) {
        return *this = *this * rhs;
    }

    Matrix operator+() const {
        return *this;
    }

    Matrix operator-() const {
        Matrix result = *this;
        for (T& value : result._data) value = T() - value;
        return result;
    }

    friend Matrix operator+(Matrix lhs, const Matrix& rhs) {
        return lhs += rhs;
    }

    friend Matrix operator-(Matrix lhs, const Matrix& rhs) {
        return lhs -= rhs;
    }

    friend Matrix operator*(Matrix lhs, const T& rhs) {
        return lhs *= rhs;
    }

    friend Matrix operator*(const T& lhs, Matrix rhs) {
        return rhs *= lhs;
    }

    friend Matrix operator/(Matrix lhs, const T& rhs) {
        return lhs /= rhs;
    }

    friend Matrix operator*(const Matrix& lhs, const Matrix& rhs) {
        assert(lhs._cols == rhs._rows);
        Matrix result(lhs._rows, rhs._cols);
        for (int row = 0; row < lhs._rows; row++) {
            T* output = result[row];
            for (int middle = 0; middle < lhs._cols; middle++) {
                const T coefficient = lhs[row][middle];
                if (coefficient == T()) continue;
                const T* input = rhs[middle];
                for (int col = 0; col < rhs._cols; col++) {
                    output[col] += coefficient * input[col];
                }
            }
        }
        return result;
    }

    friend std::vector<T> operator*(const Matrix& lhs, const std::vector<T>& rhs) {
        assert(lhs._cols == int(rhs.size()));
        std::vector<T> result(std::size_t(lhs._rows));
        for (int row = 0; row < lhs._rows; row++) {
            T value = T();
            for (int col = 0; col < lhs._cols; col++) {
                value += lhs[row][col] * rhs[std::size_t(col)];
            }
            result[std::size_t(row)] = value;
        }
        return result;
    }

    friend std::vector<T> operator*(const std::vector<T>& lhs, const Matrix& rhs) {
        assert(int(lhs.size()) == rhs._rows);
        std::vector<T> result(std::size_t(rhs._cols));
        for (int row = 0; row < rhs._rows; row++) {
            if (lhs[std::size_t(row)] == T()) continue;
            for (int col = 0; col < rhs._cols; col++) {
                result[std::size_t(col)] += lhs[std::size_t(row)] * rhs[row][col];
            }
        }
        return result;
    }

    bool operator==(const Matrix& rhs) const {
        return _rows == rhs._rows && _cols == rhs._cols && _data == rhs._data;
    }

    bool operator!=(const Matrix& rhs) const {
        return !(*this == rhs);
    }

    Matrix pow(std::uint64_t exponent) const {
        assert(_rows == _cols);
        Matrix result = identity(_rows);
        Matrix base = *this;
        while (exponent > 0) {
            if (exponent & 1) result *= base;
            exponent >>= 1;
            if (exponent > 0) base *= base;
        }
        return result;
    }
};

}  // namespace matrix
}  // namespace m1une


#line 9 "math/matrix/linear_algebra.hpp"

namespace m1une {
namespace matrix {

template <class T>
constexpr T default_epsilon() {
    if constexpr (std::is_floating_point_v<T>) {
        return T(1e-10);
    } else {
        return T();
    }
}

namespace detail {

template <class T>
T matrix_abs(T value) {
    return value < T() ? T() - value : value;
}

template <class T>
bool is_zero(const T& value, const T& eps) {
    if constexpr (std::is_floating_point_v<T>) {
        return matrix_abs(value) <= eps;
    } else {
        (void)eps;
        return value == T();
    }
}

template <class T>
int choose_pivot(const Matrix<T>& matrix, int first_row, int col, const T& eps) {
    int pivot = -1;
    if constexpr (std::is_floating_point_v<T>) {
        for (int row = first_row; row < matrix.rows(); row++) {
            if (is_zero(matrix[row][col], eps)) continue;
            if (pivot == -1 || matrix_abs(matrix[pivot][col]) < matrix_abs(matrix[row][col])) {
                pivot = row;
            }
        }
    } else {
        for (int row = first_row; row < matrix.rows(); row++) {
            if (!is_zero(matrix[row][col], eps)) {
                pivot = row;
                break;
            }
        }
    }
    return pivot;
}

template <class T>
std::vector<int> row_reduce(Matrix<T>& matrix, int pivot_col_limit, const T& eps,
                            bool reduced) {
    std::vector<int> pivot_columns;
    int pivot_row = 0;
    for (int col = 0; col < pivot_col_limit && pivot_row < matrix.rows(); col++) {
        int pivot = choose_pivot(matrix, pivot_row, col, eps);
        if (pivot == -1) continue;
        matrix.swap_rows(pivot_row, pivot);

        const T pivot_value = matrix[pivot_row][col];
        if (reduced) {
            for (int j = col; j < matrix.cols(); j++) matrix[pivot_row][j] /= pivot_value;
        }

        const int first_row = reduced ? 0 : pivot_row + 1;
        for (int row = first_row; row < matrix.rows(); row++) {
            if (row == pivot_row || is_zero(matrix[row][col], eps)) continue;
            T factor = matrix[row][col];
            if (!reduced) factor /= pivot_value;
            matrix[row][col] = T();
            for (int j = col + 1; j < matrix.cols(); j++) {
                matrix[row][j] -= factor * matrix[pivot_row][j];
            }
        }

        pivot_columns.push_back(col);
        pivot_row++;
    }

    if constexpr (std::is_floating_point_v<T>) {
        for (T& value : matrix.data()) {
            if (is_zero(value, eps)) value = T();
        }
    }
    return pivot_columns;
}

}  // namespace detail

template <class T>
struct RowReduction {
    Matrix<T> matrix;
    std::vector<int> pivot_columns;

    int rank() const {
        return int(pivot_columns.size());
    }
};

template <class T>
RowReduction<T> reduced_row_echelon_form(Matrix<T> matrix,
                                         T eps = default_epsilon<T>()) {
    RowReduction<T> result;
    result.pivot_columns = detail::row_reduce(matrix, matrix.cols(), eps, true);
    result.matrix = std::move(matrix);
    return result;
}

template <class T>
int matrix_rank(Matrix<T> matrix, T eps = default_epsilon<T>()) {
    return int(detail::row_reduce(matrix, matrix.cols(), eps, false).size());
}

template <class T>
T determinant(Matrix<T> matrix, T eps = default_epsilon<T>()) {
    assert(matrix.rows() == matrix.cols());
    const int size = matrix.rows();
    T result = T(1);
    bool negate = false;

    for (int col = 0; col < size; col++) {
        int pivot = detail::choose_pivot(matrix, col, col, eps);
        if (pivot == -1) return T();
        if (pivot != col) {
            matrix.swap_rows(pivot, col);
            negate = !negate;
        }

        const T pivot_value = matrix[col][col];
        result *= pivot_value;
        for (int row = col + 1; row < size; row++) {
            if (detail::is_zero(matrix[row][col], eps)) continue;
            const T factor = matrix[row][col] / pivot_value;
            matrix[row][col] = T();
            for (int j = col + 1; j < size; j++) {
                matrix[row][j] -= factor * matrix[col][j];
            }
        }
    }
    return negate ? T() - result : result;
}

template <class T>
std::optional<Matrix<T>> inverse(const Matrix<T>& matrix,
                                 T eps = default_epsilon<T>()) {
    assert(matrix.rows() == matrix.cols());
    const int size = matrix.rows();
    Matrix<T> augmented(size, size * 2);
    for (int row = 0; row < size; row++) {
        for (int col = 0; col < size; col++) {
            augmented[row][col] = matrix[row][col];
        }
        augmented[row][size + row] = T(1);
    }

    const std::vector<int> pivots = detail::row_reduce(augmented, size, eps, true);
    if (int(pivots.size()) != size) return std::nullopt;

    Matrix<T> result(size, size);
    for (int row = 0; row < size; row++) {
        for (int col = 0; col < size; col++) {
            result[row][col] = augmented[row][size + col];
        }
    }
    return result;
}

template <class T>
struct LinearSystemResult {
    bool consistent = false;
    std::vector<T> particular_solution;
    std::vector<std::vector<T>> nullspace_basis;
    std::vector<int> pivot_columns;

    int rank() const {
        return int(pivot_columns.size());
    }

    int nullity() const {
        return consistent ? int(nullspace_basis.size()) : 0;
    }

    bool has_unique_solution() const {
        return consistent && nullspace_basis.empty();
    }
};

template <class T>
LinearSystemResult<T> solve_linear_system(const Matrix<T>& coefficients,
                                          const std::vector<T>& constants,
                                          T eps = default_epsilon<T>()) {
    assert(coefficients.rows() == int(constants.size()));
    const int equation_count = coefficients.rows();
    const int variable_count = coefficients.cols();
    Matrix<T> augmented(equation_count, variable_count + 1);
    for (int row = 0; row < equation_count; row++) {
        for (int col = 0; col < variable_count; col++) {
            augmented[row][col] = coefficients[row][col];
        }
        augmented[row][variable_count] = constants[std::size_t(row)];
    }

    LinearSystemResult<T> result;
    result.pivot_columns =
        detail::row_reduce(augmented, variable_count, eps, true);

    for (int row = result.rank(); row < equation_count; row++) {
        bool zero_left = true;
        for (int col = 0; col < variable_count; col++) {
            if (!detail::is_zero(augmented[row][col], eps)) {
                zero_left = false;
                break;
            }
        }
        if (zero_left && !detail::is_zero(augmented[row][variable_count], eps)) {
            return result;
        }
    }

    result.consistent = true;
    result.particular_solution.assign(std::size_t(variable_count), T());
    std::vector<bool> is_pivot(std::size_t(variable_count), false);
    for (int row = 0; row < result.rank(); row++) {
        const int col = result.pivot_columns[std::size_t(row)];
        is_pivot[std::size_t(col)] = true;
        result.particular_solution[std::size_t(col)] = augmented[row][variable_count];
    }

    for (int free_col = 0; free_col < variable_count; free_col++) {
        if (is_pivot[std::size_t(free_col)]) continue;
        std::vector<T> direction(static_cast<std::size_t>(variable_count));
        direction[std::size_t(free_col)] = T(1);
        for (int row = 0; row < result.rank(); row++) {
            const int pivot_col = result.pivot_columns[std::size_t(row)];
            direction[std::size_t(pivot_col)] = T() - augmented[row][free_col];
        }
        result.nullspace_basis.push_back(std::move(direction));
    }
    return result;
}

}  // namespace matrix
}  // namespace m1une


#line 1 "graph/graph.hpp"



#include <array>
#line 8 "graph/graph.hpp"

namespace m1une {
namespace graph {

template <class T = int>
struct Edge {
    using cost_type = T;

    int from;
    int to;
    T cost;
    int id;
    bool alive;

    Edge() : from(-1), to(-1), cost(T()), id(-1), alive(true) {}
    Edge(int from_, int to_, T cost_ = T(1), int id_ = -1, bool alive_ = true)
        : from(from_), to(to_), cost(cost_), id(id_), alive(alive_) {}

    int other(int v) const {
        assert(v == from || v == to);
        return from ^ to ^ v;
    }
};

template <class T = int>
struct Graph {
    using edge_type = Edge<T>;
    using cost_type = T;

   private:
    struct EdgePositions {
        std::array<std::pair<int, int>, 2> value{};
        int size = 0;

        void push_back(std::pair<int, int> position) {
            assert(size < 2);
            value[size++] = position;
        }
    };

    int _n;
    int _edge_count;
    std::vector<std::vector<edge_type>> _g;
    std::vector<EdgePositions> _edge_positions;

   public:
    Graph() : _n(0), _edge_count(0) {}
    explicit Graph(int n) : _n(n), _edge_count(0), _g(n) {
        assert(0 <= n);
    }

    int size() const {
        return _n;
    }

    bool empty() const {
        return _n == 0;
    }

    int edge_count() const {
        return _edge_count;
    }

    int add_vertex() {
        _g.emplace_back();
        return _n++;
    }

    int add_directed_edge(int from, int to, T cost = T(1)) {
        assert(0 <= from && from < _n);
        assert(0 <= to && to < _n);
        int id = _edge_count++;
        int idx = int(_g[from].size());
        _g[from].push_back(edge_type(from, to, cost, id));
        _edge_positions.emplace_back();
        _edge_positions.back().push_back({from, idx});
        return id;
    }

    int add_edge(int u, int v, T cost = T(1)) {
        assert(0 <= u && u < _n);
        assert(0 <= v && v < _n);
        int id = _edge_count++;
        int u_idx = int(_g[u].size());
        _g[u].push_back(edge_type(u, v, cost, id));
        int v_idx = int(_g[v].size());
        _g[v].push_back(edge_type(v, u, cost, id));
        _edge_positions.emplace_back();
        _edge_positions.back().push_back({u, u_idx});
        _edge_positions.back().push_back({v, v_idx});
        return id;
    }

    void set_edge_alive(int id, bool alive) {
        assert(0 <= id && id < _edge_count);
        for (int i = 0; i < _edge_positions[id].size; ++i) {
            auto [v, idx] = _edge_positions[id].value[i];
            _g[v][idx].alive = alive;
        }
    }

    void erase_edge(int id) {
        set_edge_alive(id, false);
    }

    void revive_edge(int id) {
        set_edge_alive(id, true);
    }

    bool is_edge_alive(int id) const {
        assert(0 <= id && id < _edge_count);
        assert(_edge_positions[id].size != 0);
        auto [v, idx] = _edge_positions[id].value[0];
        return _g[v][idx].alive;
    }

    const std::vector<edge_type>& operator[](int v) const {
        assert(0 <= v && v < _n);
        return _g[v];
    }

    std::vector<edge_type>& operator[](int v) {
        assert(0 <= v && v < _n);
        return _g[v];
    }

    const std::vector<std::vector<edge_type>>& adjacency() const {
        return _g;
    }

    std::vector<std::vector<edge_type>>& adjacency() {
        return _g;
    }

    std::vector<edge_type> edges(bool include_inactive = false) const {
        std::vector<edge_type> result;
        result.reserve(_edge_count);
        std::vector<char> used(_edge_count, false);
        for (int v = 0; v < _n; v++) {
            for (const auto& e : _g[v]) {
                if (!include_inactive && !e.alive) continue;
                if (0 <= e.id && e.id < _edge_count) {
                    if (used[e.id]) continue;
                    used[e.id] = true;
                }
                result.push_back(e);
            }
        }
        return result;
    }

    Graph reversed() const {
        Graph result(_n);
        result._edge_count = _edge_count;
        result._edge_positions.assign(_edge_count, {});
        for (int v = 0; v < _n; v++) {
            for (const auto& e : _g[v]) {
                int idx = int(result._g[e.to].size());
                result._g[e.to].push_back(edge_type(e.to, e.from, e.cost, e.id, e.alive));
                if (0 <= e.id && e.id < _edge_count) result._edge_positions[e.id].push_back({e.to, idx});
            }
        }
        return result;
    }
};

}  // namespace graph
}  // namespace m1une


#line 10 "graph/matrix_tree_theorem.hpp"

namespace m1une {
namespace graph {

namespace matrix_tree_detail {

inline int minor_index(int vertex, int removed) {
    assert(vertex != removed);
    return vertex < removed ? vertex : vertex - 1;
}

template <class Weight>
void assert_edge_incidence(const Graph<Weight>& graph, int expected) {
#ifndef NDEBUG
    std::vector<int> incidence(graph.edge_count(), 0);
    for (int vertex = 0; vertex < graph.size(); vertex++) {
        for (const Edge<Weight>& edge : graph[vertex]) {
            if (!edge.alive) continue;
            assert(0 <= edge.id && edge.id < graph.edge_count());
            incidence[edge.id]++;
        }
    }
    for (int count : incidence) {
        if (count != 0) assert(count == expected);
    }
#else
    (void)graph;
    (void)expected;
#endif
}

template <class Field, class Weight>
Field count_arborescences(
    const Graph<Weight>& graph,
    int root,
    bool outward
) {
    const int n = graph.size();
    assert(0 <= root && root < n);
    assert_edge_incidence(graph, 1);

    matrix::Matrix<Field> minor(n - 1, n - 1);
    for (int vertex = 0; vertex < n; vertex++) {
        for (const Edge<Weight>& edge : graph[vertex]) {
            if (!edge.alive || edge.from == edge.to) continue;
            const int row = outward ? edge.to : edge.from;
            const int col = outward ? edge.from : edge.to;
            if (row == root) continue;

            const Field weight(edge.cost);
            const int reduced_row = minor_index(row, root);
            minor[reduced_row][reduced_row] += weight;
            if (col != root) {
                minor[reduced_row][minor_index(col, root)] -= weight;
            }
        }
    }
    return matrix::determinant(std::move(minor));
}

}  // namespace matrix_tree_detail

// Returns the total weight of all undirected spanning trees. The weight of a
// tree is the product of its edge costs.
template <class Field, class Weight>
Field count_spanning_trees(const Graph<Weight>& graph) {
    const int n = graph.size();
    assert(n > 0);
    matrix_tree_detail::assert_edge_incidence(graph, 2);

    const int removed = n - 1;
    matrix::Matrix<Field> minor(n - 1, n - 1);
    for (int vertex = 0; vertex < n; vertex++) {
        for (const Edge<Weight>& edge : graph[vertex]) {
            if (!edge.alive || edge.from >= edge.to) continue;
            const int from = edge.from;
            const int to = edge.to;
            const Field weight(edge.cost);

            if (from != removed) {
                const int reduced_from = matrix_tree_detail::minor_index(from, removed);
                minor[reduced_from][reduced_from] += weight;
            }
            if (to != removed) {
                const int reduced_to = matrix_tree_detail::minor_index(to, removed);
                minor[reduced_to][reduced_to] += weight;
            }
            if (from != removed && to != removed) {
                const int reduced_from = matrix_tree_detail::minor_index(from, removed);
                const int reduced_to = matrix_tree_detail::minor_index(to, removed);
                minor[reduced_from][reduced_to] -= weight;
                minor[reduced_to][reduced_from] -= weight;
            }
        }
    }
    return matrix::determinant(std::move(minor));
}

// Counts directed spanning trees whose edges point away from root, so every
// vertex is reachable from root.
template <class Field, class Weight>
Field count_out_arborescences(const Graph<Weight>& graph, int root) {
    return matrix_tree_detail::count_arborescences<Field>(graph, root, true);
}

// Counts directed spanning trees whose edges point toward root, so root is
// reachable from every vertex.
template <class Field, class Weight>
Field count_in_arborescences(const Graph<Weight>& graph, int root) {
    return matrix_tree_detail::count_arborescences<Field>(graph, root, false);
}

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