m1une's library

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

View on GitHub

:heavy_check_mark: Directed Minimum Spanning Tree
(graph/directed_mst.hpp)

Overview

directed_mst finds a minimum-cost spanning arborescence. It can either use a specified root or choose the root that minimizes the total cost. Every vertex must be reachable from the selected root using active directed edges. If no such arborescence exists, it returns std::nullopt.

The implementation uses the Chu-Liu/Edmonds algorithm with lazy meldable heaps, disjoint-set contraction, and a contraction forest for edge reconstruction.

Use Graph<T>::add_directed_edge to add edges. Parallel edges and self-loops are supported, and inactive edges are ignored.

Requirements

The cost type T must support T(0), addition, subtraction, and comparison with <. All input costs, reduced costs, and the final answer must fit in T. Negative edge costs are supported.

Interface

template <class T>
struct DirectedMinimumSpanningTree {
    T cost;
    std::vector<int> parent;
    std::vector<int> parent_edge;
    std::vector<Edge<T>> edges;
    int root;
};

template <class T>
std::optional<DirectedMinimumSpanningTree<T>>
directed_mst(const Graph<T>& graph, int root);

template <class T>
std::optional<DirectedMinimumSpanningTree<T>>
directed_mst(const Graph<T>& graph);

Result

Member Description
cost Sum of the selected edge costs.
parent[v] Parent of v; parent[root] == root.
parent_edge[v] ID of the selected edge entering v; -1 for the root.
edges The N - 1 selected original edges, ordered by destination vertex except for the root.
root The specified root vertex, or the root selected by the root-free overload.

Operations

Function Description Complexity
directed_mst(const Graph<T>& graph, int root) Returns a minimum rooted spanning arborescence, or std::nullopt if none exists. It does not mutate graph. Amortized $O((N + M)\log M)$
directed_mst(const Graph<T>& graph) Chooses the root giving the minimum-cost spanning arborescence. Returns std::nullopt for an empty graph or if no single root can span every vertex. It does not mutate graph. Amortized $O((N + M)\log (N + M))$

Complexity

For N vertices and M stored edges, the running time is O((N + M) log M) amortized and the memory usage is O(N + M). The implementation is iterative.

The root-free overload uses a lexicographic artificial-root cost. It minimizes the number of artificial edges before the original cost, so it does not require a numeric infinity or a large penalty value in T.

Example

m1une::graph::Graph<long long> graph(3);
graph.add_directed_edge(0, 1, 2);
graph.add_directed_edge(0, 2, 7);
graph.add_directed_edge(1, 2, 3);

auto answer = m1une::graph::directed_mst(graph, 0);
assert(answer.has_value());
assert(answer->cost == 5);
assert(answer->parent[1] == 0);
assert(answer->parent[2] == 1);

Depends on

Required by

Verified with

Code

#ifndef M1UNE_GRAPH_DIRECTED_MST_HPP
#define M1UNE_GRAPH_DIRECTED_MST_HPP 1

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

#include "graph.hpp"

namespace m1une {
namespace graph {

template <class T>
struct DirectedMinimumSpanningTree {
    T cost;
    std::vector<int> parent;
    std::vector<int> parent_edge;
    std::vector<Edge<T>> edges;
    int root;
};

namespace internal {

template <class T>
struct DirectedMstEdge {
    int from = -1;
    int to = -1;
    T cost = T(0);
    int id = -1;
};

template <class T>
struct DirectedMstHeapPool {
    using StoredEdge = DirectedMstEdge<T>;

    struct Node {
        StoredEdge edge;
        T offset = T(0);
        int child = -1;
        int sibling = -1;
    };

    struct Heap {
        int root = -1;
        int size = 0;
    };

    std::vector<Node> nodes;

    explicit DirectedMstHeapPool(int capacity = 0) {
        nodes.reserve(capacity);
    }

    T key(int node) const {
        return nodes[node].edge.cost + nodes[node].offset;
    }

    int meld_roots(int first, int second) {
        if (first == -1) return second;
        if (second == -1) return first;
        if (key(second) < key(first)) std::swap(first, second);
        nodes[second].offset -= nodes[first].offset;
        nodes[second].sibling = nodes[first].child;
        nodes[first].child = second;
        return first;
    }

    void push(Heap& heap, const StoredEdge& edge) {
        const int node = int(nodes.size());
        nodes.push_back(Node{edge, T(0), -1, -1});
        heap.root = meld_roots(heap.root, node);
        heap.size++;
    }

    void meld(Heap& destination, Heap& source) {
        destination.root = meld_roots(destination.root, source.root);
        destination.size += source.size;
        source.root = -1;
        source.size = 0;
    }

    const StoredEdge& top(const Heap& heap) const {
        assert(heap.root != -1);
        return nodes[heap.root].edge;
    }

    T top_key(const Heap& heap) const {
        assert(heap.root != -1);
        return key(heap.root);
    }

    void add_all(Heap& heap, const T& delta) {
        assert(heap.root != -1);
        nodes[heap.root].offset += delta;
    }

    void pop(Heap& heap) {
        assert(heap.root != -1 && heap.size > 0);
        const int old_root = heap.root;
        int child = nodes[old_root].child;
        std::vector<int> pairs;
        while (child != -1) {
            int first = child;
            child = nodes[first].sibling;
            nodes[first].sibling = -1;
            nodes[first].offset += nodes[old_root].offset;

            if (child != -1) {
                int second = child;
                child = nodes[second].sibling;
                nodes[second].sibling = -1;
                nodes[second].offset += nodes[old_root].offset;
                first = meld_roots(first, second);
            }
            pairs.push_back(first);
        }

        heap.root = -1;
        for (auto it = pairs.rbegin(); it != pairs.rend(); ++it) {
            heap.root = meld_roots(*it, heap.root);
        }
        heap.size--;
    }
};

struct DirectedMstDsu {
    std::vector<int> parent;

    explicit DirectedMstDsu(int n) : parent(n, -1) {}

    int leader(int vertex) {
        int root = vertex;
        while (parent[root] != -1) root = parent[root];
        while (vertex != root) {
            int next = parent[vertex];
            parent[vertex] = root;
            vertex = next;
        }
        return root;
    }
};

template <class T>
struct DirectedMstRootlessCost {
    int artificial_edges;
    T original_cost;

    DirectedMstRootlessCost() : artificial_edges(0), original_cost(T(0)) {}
    explicit DirectedMstRootlessCost(int zero)
        : artificial_edges(zero), original_cost(T(0)) {
        assert(zero == 0);
    }
    DirectedMstRootlessCost(int artificial_edges_, const T& original_cost_)
        : artificial_edges(artificial_edges_), original_cost(original_cost_) {}

    DirectedMstRootlessCost& operator+=(const DirectedMstRootlessCost& other) {
        artificial_edges += other.artificial_edges;
        original_cost += other.original_cost;
        return *this;
    }

    DirectedMstRootlessCost& operator-=(const DirectedMstRootlessCost& other) {
        artificial_edges -= other.artificial_edges;
        original_cost -= other.original_cost;
        return *this;
    }

    friend DirectedMstRootlessCost operator+(
        DirectedMstRootlessCost first,
        const DirectedMstRootlessCost& second
    ) {
        return first += second;
    }

    friend DirectedMstRootlessCost operator-(
        DirectedMstRootlessCost first,
        const DirectedMstRootlessCost& second
    ) {
        return first -= second;
    }

    friend bool operator<(
        const DirectedMstRootlessCost& first,
        const DirectedMstRootlessCost& second
    ) {
        if (first.artificial_edges != second.artificial_edges) {
            return first.artificial_edges < second.artificial_edges;
        }
        return first.original_cost < second.original_cost;
    }
};

}  // namespace internal

// Returns a minimum-cost spanning arborescence rooted at root, or nullopt when
// some vertex is unreachable from the root using active directed edges.
template <class T>
std::optional<DirectedMinimumSpanningTree<T>> directed_mst(
    const Graph<T>& graph,
    int root
) {
    const int n = graph.size();
    assert(0 <= root && root < n);
    const int maximum_node_count = 2 * n;

    int active_edge_count = 0;
#ifndef NDEBUG
    std::vector<int> incidence(graph.edge_count(), 0);
#endif
    for (int vertex = 0; vertex < n; vertex++) {
        for (const Edge<T>& edge : graph[vertex]) {
            if (!edge.alive) continue;
            assert(0 <= edge.id && edge.id < graph.edge_count());
#ifndef NDEBUG
            incidence[edge.id]++;
#endif
            active_edge_count++;
        }
    }
#ifndef NDEBUG
    for (int count : incidence) {
        if (count != 0) assert(count == 1);
    }
#endif

    using StoredEdge = internal::DirectedMstEdge<T>;
    using HeapPool = internal::DirectedMstHeapPool<T>;
    HeapPool pool(active_edge_count);
    std::vector<typename HeapPool::Heap> heaps(maximum_node_count);
    for (int vertex = 0; vertex < n; vertex++) {
        for (const Edge<T>& edge : graph[vertex]) {
            if (!edge.alive) continue;
            pool.push(heaps[edge.to], StoredEdge{edge.from, edge.to, edge.cost, edge.id});
        }
    }

    internal::DirectedMstDsu dsu(maximum_node_count);
    std::vector<int> contraction_parent(maximum_node_count, -1);
    std::vector<int> visited(maximum_node_count, 0);
    std::vector<StoredEdge> selected(maximum_node_count);
    int node_count = n;
    int visit_token = 1;
    visited[root] = 1;

    for (int start = 0; start < n; start++) {
        if (visited[start] != 0) continue;
        visit_token++;
        int component = start;
        while (visited[component] == 0 || visited[component] == visit_token) {
            if (visited[component] == visit_token) {
                if (node_count == maximum_node_count) return std::nullopt;
                const int contracted = node_count++;
                int current = component;
                do {
                    const T reduction = T(0) - pool.top_key(heaps[current]);
                    pool.add_all(heaps[current], reduction);
                    pool.meld(heaps[contracted], heaps[current]);
                    contraction_parent[current] = contracted;
                    dsu.parent[current] = contracted;
                    current = dsu.leader(selected[current].from);
                } while (current != contracted);
                component = contracted;
            }

            assert(visited[component] == 0);
            visited[component] = visit_token;
            while (heaps[component].size > 0 &&
                   dsu.leader(pool.top(heaps[component]).from) == component) {
                pool.pop(heaps[component]);
            }
            if (heaps[component].size == 0) return std::nullopt;
            selected[component] = pool.top(heaps[component]);
            component = dsu.leader(selected[component].from);
        }
    }

    DirectedMinimumSpanningTree<T> result;
    result.cost = T(0);
    result.parent.assign(n, -1);
    result.parent_edge.assign(n, -1);
    result.root = root;
    result.parent[root] = root;

    std::vector<char> expanded(node_count, false);
    std::vector<StoredEdge> chosen(n);
    for (int component = node_count - 1; component >= 0; component--) {
        if (component == root || expanded[component]) continue;
        const StoredEdge& edge = selected[component];
        if (edge.id == -1) return std::nullopt;
        int vertex = edge.to;
        while (vertex != -1 && !expanded[vertex]) {
            expanded[vertex] = true;
            vertex = contraction_parent[vertex];
        }
        result.cost += edge.cost;
        result.parent[edge.to] = edge.from;
        result.parent_edge[edge.to] = edge.id;
        chosen[edge.to] = edge;
    }

    result.edges.reserve(n - 1);
    for (int vertex = 0; vertex < n; vertex++) {
        if (vertex == root) continue;
        if (result.parent[vertex] == -1) return std::nullopt;
        const StoredEdge& edge = chosen[vertex];
        result.edges.emplace_back(edge.from, edge.to, edge.cost, edge.id, true);
    }
    return result;
}

// Chooses the root that gives a minimum-cost spanning arborescence.
template <class T>
std::optional<DirectedMinimumSpanningTree<T>> directed_mst(
    const Graph<T>& graph
) {
    const int n = graph.size();
    if (n == 0) return std::nullopt;

    using Cost = internal::DirectedMstRootlessCost<T>;
    Graph<Cost> augmented(n + 1);
    std::vector<int> original_edge_id;
    original_edge_id.reserve(graph.edge_count() + n);

#ifndef NDEBUG
    std::vector<int> incidence(graph.edge_count(), 0);
#endif
    for (int vertex = 0; vertex < n; vertex++) {
        for (const Edge<T>& edge : graph[vertex]) {
            if (!edge.alive) continue;
#ifndef NDEBUG
            assert(0 <= edge.id && edge.id < graph.edge_count());
            incidence[edge.id]++;
#endif
            augmented.add_directed_edge(
                edge.from,
                edge.to,
                Cost(0, edge.cost)
            );
            original_edge_id.push_back(edge.id);
        }
    }
#ifndef NDEBUG
    for (int count : incidence) {
        if (count != 0) assert(count == 1);
    }
#endif

    const int artificial_root = n;
    for (int vertex = 0; vertex < n; vertex++) {
        augmented.add_directed_edge(
            artificial_root,
            vertex,
            Cost(1, T(0))
        );
        original_edge_id.push_back(-1);
    }

    auto augmented_result = directed_mst(augmented, artificial_root);
    if (!augmented_result || augmented_result->cost.artificial_edges != 1) {
        return std::nullopt;
    }

    DirectedMinimumSpanningTree<T> result;
    result.cost = augmented_result->cost.original_cost;
    result.parent.assign(n, -1);
    result.parent_edge.assign(n, -1);
    result.root = -1;
    result.edges.reserve(n - 1);

    for (int vertex = 0; vertex < n; vertex++) {
        int augmented_edge_id = augmented_result->parent_edge[vertex];
        assert(0 <= augmented_edge_id &&
               augmented_edge_id < int(original_edge_id.size()));
        int edge_id = original_edge_id[augmented_edge_id];
        if (edge_id == -1) {
            assert(result.root == -1);
            result.root = vertex;
            result.parent[vertex] = vertex;
            continue;
        }

        result.parent[vertex] = augmented_result->parent[vertex];
        result.parent_edge[vertex] = edge_id;
        result.edges.emplace_back(
            result.parent[vertex],
            vertex,
            augmented_result->edges[vertex].cost.original_cost,
            edge_id,
            true
        );
    }
    assert(result.root != -1);
    return result;
}

}  // namespace graph
}  // namespace m1une

#endif  // M1UNE_GRAPH_DIRECTED_MST_HPP
#line 1 "graph/directed_mst.hpp"



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

#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/directed_mst.hpp"

namespace m1une {
namespace graph {

template <class T>
struct DirectedMinimumSpanningTree {
    T cost;
    std::vector<int> parent;
    std::vector<int> parent_edge;
    std::vector<Edge<T>> edges;
    int root;
};

namespace internal {

template <class T>
struct DirectedMstEdge {
    int from = -1;
    int to = -1;
    T cost = T(0);
    int id = -1;
};

template <class T>
struct DirectedMstHeapPool {
    using StoredEdge = DirectedMstEdge<T>;

    struct Node {
        StoredEdge edge;
        T offset = T(0);
        int child = -1;
        int sibling = -1;
    };

    struct Heap {
        int root = -1;
        int size = 0;
    };

    std::vector<Node> nodes;

    explicit DirectedMstHeapPool(int capacity = 0) {
        nodes.reserve(capacity);
    }

    T key(int node) const {
        return nodes[node].edge.cost + nodes[node].offset;
    }

    int meld_roots(int first, int second) {
        if (first == -1) return second;
        if (second == -1) return first;
        if (key(second) < key(first)) std::swap(first, second);
        nodes[second].offset -= nodes[first].offset;
        nodes[second].sibling = nodes[first].child;
        nodes[first].child = second;
        return first;
    }

    void push(Heap& heap, const StoredEdge& edge) {
        const int node = int(nodes.size());
        nodes.push_back(Node{edge, T(0), -1, -1});
        heap.root = meld_roots(heap.root, node);
        heap.size++;
    }

    void meld(Heap& destination, Heap& source) {
        destination.root = meld_roots(destination.root, source.root);
        destination.size += source.size;
        source.root = -1;
        source.size = 0;
    }

    const StoredEdge& top(const Heap& heap) const {
        assert(heap.root != -1);
        return nodes[heap.root].edge;
    }

    T top_key(const Heap& heap) const {
        assert(heap.root != -1);
        return key(heap.root);
    }

    void add_all(Heap& heap, const T& delta) {
        assert(heap.root != -1);
        nodes[heap.root].offset += delta;
    }

    void pop(Heap& heap) {
        assert(heap.root != -1 && heap.size > 0);
        const int old_root = heap.root;
        int child = nodes[old_root].child;
        std::vector<int> pairs;
        while (child != -1) {
            int first = child;
            child = nodes[first].sibling;
            nodes[first].sibling = -1;
            nodes[first].offset += nodes[old_root].offset;

            if (child != -1) {
                int second = child;
                child = nodes[second].sibling;
                nodes[second].sibling = -1;
                nodes[second].offset += nodes[old_root].offset;
                first = meld_roots(first, second);
            }
            pairs.push_back(first);
        }

        heap.root = -1;
        for (auto it = pairs.rbegin(); it != pairs.rend(); ++it) {
            heap.root = meld_roots(*it, heap.root);
        }
        heap.size--;
    }
};

struct DirectedMstDsu {
    std::vector<int> parent;

    explicit DirectedMstDsu(int n) : parent(n, -1) {}

    int leader(int vertex) {
        int root = vertex;
        while (parent[root] != -1) root = parent[root];
        while (vertex != root) {
            int next = parent[vertex];
            parent[vertex] = root;
            vertex = next;
        }
        return root;
    }
};

template <class T>
struct DirectedMstRootlessCost {
    int artificial_edges;
    T original_cost;

    DirectedMstRootlessCost() : artificial_edges(0), original_cost(T(0)) {}
    explicit DirectedMstRootlessCost(int zero)
        : artificial_edges(zero), original_cost(T(0)) {
        assert(zero == 0);
    }
    DirectedMstRootlessCost(int artificial_edges_, const T& original_cost_)
        : artificial_edges(artificial_edges_), original_cost(original_cost_) {}

    DirectedMstRootlessCost& operator+=(const DirectedMstRootlessCost& other) {
        artificial_edges += other.artificial_edges;
        original_cost += other.original_cost;
        return *this;
    }

    DirectedMstRootlessCost& operator-=(const DirectedMstRootlessCost& other) {
        artificial_edges -= other.artificial_edges;
        original_cost -= other.original_cost;
        return *this;
    }

    friend DirectedMstRootlessCost operator+(
        DirectedMstRootlessCost first,
        const DirectedMstRootlessCost& second
    ) {
        return first += second;
    }

    friend DirectedMstRootlessCost operator-(
        DirectedMstRootlessCost first,
        const DirectedMstRootlessCost& second
    ) {
        return first -= second;
    }

    friend bool operator<(
        const DirectedMstRootlessCost& first,
        const DirectedMstRootlessCost& second
    ) {
        if (first.artificial_edges != second.artificial_edges) {
            return first.artificial_edges < second.artificial_edges;
        }
        return first.original_cost < second.original_cost;
    }
};

}  // namespace internal

// Returns a minimum-cost spanning arborescence rooted at root, or nullopt when
// some vertex is unreachable from the root using active directed edges.
template <class T>
std::optional<DirectedMinimumSpanningTree<T>> directed_mst(
    const Graph<T>& graph,
    int root
) {
    const int n = graph.size();
    assert(0 <= root && root < n);
    const int maximum_node_count = 2 * n;

    int active_edge_count = 0;
#ifndef NDEBUG
    std::vector<int> incidence(graph.edge_count(), 0);
#endif
    for (int vertex = 0; vertex < n; vertex++) {
        for (const Edge<T>& edge : graph[vertex]) {
            if (!edge.alive) continue;
            assert(0 <= edge.id && edge.id < graph.edge_count());
#ifndef NDEBUG
            incidence[edge.id]++;
#endif
            active_edge_count++;
        }
    }
#ifndef NDEBUG
    for (int count : incidence) {
        if (count != 0) assert(count == 1);
    }
#endif

    using StoredEdge = internal::DirectedMstEdge<T>;
    using HeapPool = internal::DirectedMstHeapPool<T>;
    HeapPool pool(active_edge_count);
    std::vector<typename HeapPool::Heap> heaps(maximum_node_count);
    for (int vertex = 0; vertex < n; vertex++) {
        for (const Edge<T>& edge : graph[vertex]) {
            if (!edge.alive) continue;
            pool.push(heaps[edge.to], StoredEdge{edge.from, edge.to, edge.cost, edge.id});
        }
    }

    internal::DirectedMstDsu dsu(maximum_node_count);
    std::vector<int> contraction_parent(maximum_node_count, -1);
    std::vector<int> visited(maximum_node_count, 0);
    std::vector<StoredEdge> selected(maximum_node_count);
    int node_count = n;
    int visit_token = 1;
    visited[root] = 1;

    for (int start = 0; start < n; start++) {
        if (visited[start] != 0) continue;
        visit_token++;
        int component = start;
        while (visited[component] == 0 || visited[component] == visit_token) {
            if (visited[component] == visit_token) {
                if (node_count == maximum_node_count) return std::nullopt;
                const int contracted = node_count++;
                int current = component;
                do {
                    const T reduction = T(0) - pool.top_key(heaps[current]);
                    pool.add_all(heaps[current], reduction);
                    pool.meld(heaps[contracted], heaps[current]);
                    contraction_parent[current] = contracted;
                    dsu.parent[current] = contracted;
                    current = dsu.leader(selected[current].from);
                } while (current != contracted);
                component = contracted;
            }

            assert(visited[component] == 0);
            visited[component] = visit_token;
            while (heaps[component].size > 0 &&
                   dsu.leader(pool.top(heaps[component]).from) == component) {
                pool.pop(heaps[component]);
            }
            if (heaps[component].size == 0) return std::nullopt;
            selected[component] = pool.top(heaps[component]);
            component = dsu.leader(selected[component].from);
        }
    }

    DirectedMinimumSpanningTree<T> result;
    result.cost = T(0);
    result.parent.assign(n, -1);
    result.parent_edge.assign(n, -1);
    result.root = root;
    result.parent[root] = root;

    std::vector<char> expanded(node_count, false);
    std::vector<StoredEdge> chosen(n);
    for (int component = node_count - 1; component >= 0; component--) {
        if (component == root || expanded[component]) continue;
        const StoredEdge& edge = selected[component];
        if (edge.id == -1) return std::nullopt;
        int vertex = edge.to;
        while (vertex != -1 && !expanded[vertex]) {
            expanded[vertex] = true;
            vertex = contraction_parent[vertex];
        }
        result.cost += edge.cost;
        result.parent[edge.to] = edge.from;
        result.parent_edge[edge.to] = edge.id;
        chosen[edge.to] = edge;
    }

    result.edges.reserve(n - 1);
    for (int vertex = 0; vertex < n; vertex++) {
        if (vertex == root) continue;
        if (result.parent[vertex] == -1) return std::nullopt;
        const StoredEdge& edge = chosen[vertex];
        result.edges.emplace_back(edge.from, edge.to, edge.cost, edge.id, true);
    }
    return result;
}

// Chooses the root that gives a minimum-cost spanning arborescence.
template <class T>
std::optional<DirectedMinimumSpanningTree<T>> directed_mst(
    const Graph<T>& graph
) {
    const int n = graph.size();
    if (n == 0) return std::nullopt;

    using Cost = internal::DirectedMstRootlessCost<T>;
    Graph<Cost> augmented(n + 1);
    std::vector<int> original_edge_id;
    original_edge_id.reserve(graph.edge_count() + n);

#ifndef NDEBUG
    std::vector<int> incidence(graph.edge_count(), 0);
#endif
    for (int vertex = 0; vertex < n; vertex++) {
        for (const Edge<T>& edge : graph[vertex]) {
            if (!edge.alive) continue;
#ifndef NDEBUG
            assert(0 <= edge.id && edge.id < graph.edge_count());
            incidence[edge.id]++;
#endif
            augmented.add_directed_edge(
                edge.from,
                edge.to,
                Cost(0, edge.cost)
            );
            original_edge_id.push_back(edge.id);
        }
    }
#ifndef NDEBUG
    for (int count : incidence) {
        if (count != 0) assert(count == 1);
    }
#endif

    const int artificial_root = n;
    for (int vertex = 0; vertex < n; vertex++) {
        augmented.add_directed_edge(
            artificial_root,
            vertex,
            Cost(1, T(0))
        );
        original_edge_id.push_back(-1);
    }

    auto augmented_result = directed_mst(augmented, artificial_root);
    if (!augmented_result || augmented_result->cost.artificial_edges != 1) {
        return std::nullopt;
    }

    DirectedMinimumSpanningTree<T> result;
    result.cost = augmented_result->cost.original_cost;
    result.parent.assign(n, -1);
    result.parent_edge.assign(n, -1);
    result.root = -1;
    result.edges.reserve(n - 1);

    for (int vertex = 0; vertex < n; vertex++) {
        int augmented_edge_id = augmented_result->parent_edge[vertex];
        assert(0 <= augmented_edge_id &&
               augmented_edge_id < int(original_edge_id.size()));
        int edge_id = original_edge_id[augmented_edge_id];
        if (edge_id == -1) {
            assert(result.root == -1);
            result.root = vertex;
            result.parent[vertex] = vertex;
            continue;
        }

        result.parent[vertex] = augmented_result->parent[vertex];
        result.parent_edge[vertex] = edge_id;
        result.edges.emplace_back(
            result.parent[vertex],
            vertex,
            augmented_result->edges[vertex].cost.original_cost,
            edge_id,
            true
        );
    }
    assert(result.root != -1);
    return result;
}

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