Directed Minimum Spanning Tree
(graph/directed_mst.hpp)
- View this file on GitHub
- Last update: 2026-08-13 01:41:40+09:00
- Include:
#include "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
verify/graph/cow_game.test.cpp
verify/graph/directed_mst.test.cpp
verify/graph/graph_algorithms.test.cpp
verify/graph/range_edge_graph.test.cpp
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