DSU on Tree
(graph/tree/dsu_on_tree.hpp)
- View this file on GitHub
- Last update: 2026-08-13 01:41:40+09:00
- Include:
#include "graph/tree/dsu_on_tree.hpp"
Overview
DsuOnTree<T> implements the small-to-large subtree technique also known as
sack. It answers one query for every rooted subtree while maintaining a
user-defined data structure.
For each vertex, light-child data is discarded after processing, while the largest child’s data is kept and reused. If inserting and removing one vertex cost $O(F)$, all callbacks together take $O(N\log N\cdot F)$ time.
The implementation uses an explicit action stack rather than recursion, so it is safe on a path-shaped tree with many vertices.
Complexity Notation
-
Nis the number of vertices. -
Fis the cost of one user callback.
Construction
DsuOnTree(const Graph<T>& graph, int root = 0);
void build(const Graph<T>& graph, int root = 0);
The graph must be a connected undirected tree built with add_edge. Inactive
edges are ignored. The chosen root determines every queried subtree.
Construction takes $O(N)$ time and memory.
Methods and Metadata
The object exposes:
| Member | Description |
|---|---|
n, root
|
Number of vertices and chosen root. |
parent, parent_edge, depth
|
Rooted-tree metadata. |
subtree_size |
Number of vertices in each subtree. |
heavy_child |
Largest child, or -1 for a leaf. |
children |
Children in the rooted tree. |
tin, tout, order
|
Preorder Euler intervals; subtree v is order[tin[v]..tout[v]). |
size(), empty(), and subtree_range(v) provide the corresponding basic
queries.
Running the Algorithm
dsu.run(add, remove, answer);
The callbacks receive a vertex index:
-
add(v)inserts vertexvinto the maintained state. -
remove(v)erases vertexv. -
answer(v)is called when the state contains exactly the vertices in the subtree ofv.
The structure may call add and remove for the same vertex several times.
They must therefore be mutually inverse operations. After run finishes, the
state contains the whole tree because the root’s sack is retained.
Example
This computes the number of distinct colors in every subtree:
#include "graph/graph.hpp"
#include "graph/tree/dsu_on_tree.hpp"
#include <vector>
int main() {
m1une::graph::Graph<int> graph(4);
graph.add_edge(0, 1);
graph.add_edge(0, 2);
graph.add_edge(1, 3);
std::vector<int> color = {0, 1, 0, 2};
std::vector<int> frequency(3);
std::vector<int> answer(4);
int distinct = 0;
m1une::tree::DsuOnTree dsu(graph, 0);
dsu.run(
[&](int vertex) {
if (frequency[color[vertex]]++ == 0) distinct++;
},
[&](int vertex) {
if (--frequency[color[vertex]] == 0) distinct--;
},
[&](int vertex) {
answer[vertex] = distinct;
}
);
}
Depends on
Required by
Verified with
verify/graph/cow_game.test.cpp
verify/graph/graph_algorithms.test.cpp
verify/graph/range_edge_graph.test.cpp
verify/graph/tree/dsu_on_tree.test.cpp
verify/graph/tree/tree_algorithms.test.cpp
Code
#ifndef M1UNE_TREE_DSU_ON_TREE_HPP
#define M1UNE_TREE_DSU_ON_TREE_HPP 1
#include <cassert>
#include <utility>
#include <vector>
#include "../graph.hpp"
namespace m1une {
namespace tree {
template <class T = int>
struct DsuOnTree {
int n;
int root;
std::vector<int> parent;
std::vector<int> parent_edge;
std::vector<int> depth;
std::vector<int> subtree_size;
std::vector<int> heavy_child;
std::vector<int> tin;
std::vector<int> tout;
std::vector<int> order;
std::vector<std::vector<int>> children;
DsuOnTree() : n(0), root(-1) {}
explicit DsuOnTree(
const m1une::graph::Graph<T>& graph,
int root_vertex = 0
) {
build(graph, root_vertex);
}
void build(
const m1une::graph::Graph<T>& graph,
int root_vertex = 0
) {
n = graph.size();
root = n == 0 ? -1 : root_vertex;
parent.assign(n, -2);
parent_edge.assign(n, -1);
depth.assign(n, 0);
subtree_size.assign(n, 1);
heavy_child.assign(n, -1);
tin.assign(n, -1);
tout.assign(n, -1);
order.clear();
order.reserve(n);
children.assign(n, {});
if (n == 0) return;
assert(0 <= root && root < n);
std::vector<int> stack;
stack.push_back(root);
parent[root] = -1;
while (!stack.empty()) {
int vertex = stack.back();
stack.pop_back();
tin[vertex] = int(order.size());
order.push_back(vertex);
for (const auto& edge : graph[vertex]) {
if (!edge.alive || parent[edge.to] != -2) continue;
parent[edge.to] = vertex;
parent_edge[edge.to] = edge.id;
depth[edge.to] = depth[vertex] + 1;
children[vertex].push_back(edge.to);
stack.push_back(edge.to);
}
}
assert(int(order.size()) == n);
for (int index = n - 1; index >= 0; --index) {
int vertex = order[index];
for (int child : children[vertex]) {
subtree_size[vertex] += subtree_size[child];
if (
heavy_child[vertex] == -1 ||
subtree_size[heavy_child[vertex]] < subtree_size[child]
) {
heavy_child[vertex] = child;
}
}
tout[vertex] = tin[vertex] + subtree_size[vertex];
}
}
int size() const {
return n;
}
bool empty() const {
return n == 0;
}
std::pair<int, int> subtree_range(int vertex) const {
assert(0 <= vertex && vertex < n);
return {tin[vertex], tout[vertex]};
}
// Runs DSU on tree. `add(v)` inserts one vertex into the maintained state,
// `remove(v)` erases it, and `answer(v)` observes the state for subtree(v).
template <class Add, class Remove, class Answer>
void run(Add add, Remove remove, Answer answer) const {
if (n == 0) return;
enum ActionType {
Process,
AddSubtree,
AddVertex,
AnswerVertex,
RemoveSubtree,
};
struct Action {
ActionType type;
int vertex;
bool keep;
};
std::vector<Action> actions;
actions.reserve(3 * std::size_t(n));
actions.push_back(Action{Process, root, true});
while (!actions.empty()) {
Action action = actions.back();
actions.pop_back();
int vertex = action.vertex;
if (action.type == AddSubtree) {
for (int index = tin[vertex]; index < tout[vertex]; ++index) {
add(order[index]);
}
} else if (action.type == AddVertex) {
add(vertex);
} else if (action.type == AnswerVertex) {
answer(vertex);
} else if (action.type == RemoveSubtree) {
for (int index = tin[vertex]; index < tout[vertex]; ++index) {
remove(order[index]);
}
} else {
if (!action.keep) {
actions.push_back(Action{
RemoveSubtree,
vertex,
false,
});
}
actions.push_back(Action{AnswerVertex, vertex, false});
actions.push_back(Action{AddVertex, vertex, false});
for (int child : children[vertex]) {
if (child != heavy_child[vertex]) {
actions.push_back(Action{
AddSubtree,
child,
false,
});
}
}
if (heavy_child[vertex] != -1) {
actions.push_back(Action{
Process,
heavy_child[vertex],
true,
});
}
for (int child : children[vertex]) {
if (child != heavy_child[vertex]) {
actions.push_back(Action{Process, child, false});
}
}
}
}
}
};
} // namespace tree
} // namespace m1une
#endif // M1UNE_TREE_DSU_ON_TREE_HPP#line 1 "graph/tree/dsu_on_tree.hpp"
#include <cassert>
#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 9 "graph/tree/dsu_on_tree.hpp"
namespace m1une {
namespace tree {
template <class T = int>
struct DsuOnTree {
int n;
int root;
std::vector<int> parent;
std::vector<int> parent_edge;
std::vector<int> depth;
std::vector<int> subtree_size;
std::vector<int> heavy_child;
std::vector<int> tin;
std::vector<int> tout;
std::vector<int> order;
std::vector<std::vector<int>> children;
DsuOnTree() : n(0), root(-1) {}
explicit DsuOnTree(
const m1une::graph::Graph<T>& graph,
int root_vertex = 0
) {
build(graph, root_vertex);
}
void build(
const m1une::graph::Graph<T>& graph,
int root_vertex = 0
) {
n = graph.size();
root = n == 0 ? -1 : root_vertex;
parent.assign(n, -2);
parent_edge.assign(n, -1);
depth.assign(n, 0);
subtree_size.assign(n, 1);
heavy_child.assign(n, -1);
tin.assign(n, -1);
tout.assign(n, -1);
order.clear();
order.reserve(n);
children.assign(n, {});
if (n == 0) return;
assert(0 <= root && root < n);
std::vector<int> stack;
stack.push_back(root);
parent[root] = -1;
while (!stack.empty()) {
int vertex = stack.back();
stack.pop_back();
tin[vertex] = int(order.size());
order.push_back(vertex);
for (const auto& edge : graph[vertex]) {
if (!edge.alive || parent[edge.to] != -2) continue;
parent[edge.to] = vertex;
parent_edge[edge.to] = edge.id;
depth[edge.to] = depth[vertex] + 1;
children[vertex].push_back(edge.to);
stack.push_back(edge.to);
}
}
assert(int(order.size()) == n);
for (int index = n - 1; index >= 0; --index) {
int vertex = order[index];
for (int child : children[vertex]) {
subtree_size[vertex] += subtree_size[child];
if (
heavy_child[vertex] == -1 ||
subtree_size[heavy_child[vertex]] < subtree_size[child]
) {
heavy_child[vertex] = child;
}
}
tout[vertex] = tin[vertex] + subtree_size[vertex];
}
}
int size() const {
return n;
}
bool empty() const {
return n == 0;
}
std::pair<int, int> subtree_range(int vertex) const {
assert(0 <= vertex && vertex < n);
return {tin[vertex], tout[vertex]};
}
// Runs DSU on tree. `add(v)` inserts one vertex into the maintained state,
// `remove(v)` erases it, and `answer(v)` observes the state for subtree(v).
template <class Add, class Remove, class Answer>
void run(Add add, Remove remove, Answer answer) const {
if (n == 0) return;
enum ActionType {
Process,
AddSubtree,
AddVertex,
AnswerVertex,
RemoveSubtree,
};
struct Action {
ActionType type;
int vertex;
bool keep;
};
std::vector<Action> actions;
actions.reserve(3 * std::size_t(n));
actions.push_back(Action{Process, root, true});
while (!actions.empty()) {
Action action = actions.back();
actions.pop_back();
int vertex = action.vertex;
if (action.type == AddSubtree) {
for (int index = tin[vertex]; index < tout[vertex]; ++index) {
add(order[index]);
}
} else if (action.type == AddVertex) {
add(vertex);
} else if (action.type == AnswerVertex) {
answer(vertex);
} else if (action.type == RemoveSubtree) {
for (int index = tin[vertex]; index < tout[vertex]; ++index) {
remove(order[index]);
}
} else {
if (!action.keep) {
actions.push_back(Action{
RemoveSubtree,
vertex,
false,
});
}
actions.push_back(Action{AnswerVertex, vertex, false});
actions.push_back(Action{AddVertex, vertex, false});
for (int child : children[vertex]) {
if (child != heavy_child[vertex]) {
actions.push_back(Action{
AddSubtree,
child,
false,
});
}
}
if (heavy_child[vertex] != -1) {
actions.push_back(Action{
Process,
heavy_child[vertex],
true,
});
}
for (int child : children[vertex]) {
if (child != heavy_child[vertex]) {
actions.push_back(Action{Process, child, false});
}
}
}
}
}
};
} // namespace tree
} // namespace m1une