Clustering
C++20 header-only: DBSCAN, HDBSCAN, k-means.
Loading...
Searching...
No Matches
dsu.h
Go to the documentation of this file.
1#pragma once
2
3#include <atomic>
4#include <cassert>
5#include <cstddef>
6#include <cstdint>
7#include <type_traits>
8#include <utility>
9#include <vector>
10
11namespace clustering {
12
24template <class Idx = std::uint32_t> class UnionFind {
25public:
26 static_assert(std::is_unsigned_v<Idx>, "UnionFind Idx must be an unsigned integer type");
27
33 explicit UnionFind(std::size_t n) : m_parent(n), m_rank(n, 0), m_size(n, 1), m_components(n) {
34 for (std::size_t i = 0; i < n; ++i) {
35 m_parent[i] = static_cast<Idx>(i);
36 }
37 }
38
49 Idx find(Idx x) noexcept {
50 assert(static_cast<std::size_t>(x) < m_parent.size() && "UnionFind::find index out of range");
51 Idx root = x;
52 while (m_parent[root] != root) {
53 root = m_parent[root];
54 }
55 Idx node = x;
56 while (m_parent[node] != root) {
57 const Idx next = m_parent[node];
58 m_parent[node] = root;
59 node = next;
60 }
61 return root;
62 }
63
75 bool unite(Idx a, Idx b) noexcept {
76 Idx ra = find(a);
77 Idx rb = find(b);
78 if (ra == rb) {
79 return false;
80 }
81 if (m_rank[ra] < m_rank[rb]) {
82 m_parent[ra] = rb;
83 m_size[rb] += m_size[ra];
84 } else if (m_rank[ra] > m_rank[rb]) {
85 m_parent[rb] = ra;
86 m_size[ra] += m_size[rb];
87 } else {
88 m_parent[rb] = ra;
89 m_size[ra] += m_size[rb];
90 ++m_rank[ra];
91 }
92 --m_components;
93 return true;
94 }
95
105 bool sameComponent(Idx a, Idx b) noexcept { return find(a) == find(b); }
106
114 [[nodiscard]] std::size_t countComponents() const noexcept { return m_components; }
115
121 [[nodiscard]] std::size_t size() const noexcept { return m_parent.size(); }
122
134 [[nodiscard]] std::size_t componentSize(Idx root) const noexcept {
135 assert(static_cast<std::size_t>(root) < m_parent.size() &&
136 "UnionFind::componentSize index out of range");
137 return m_size[root];
138 }
139
140private:
141 std::vector<Idx> m_parent;
142 std::vector<std::uint8_t> m_rank;
143 std::vector<std::size_t> m_size;
144 std::size_t m_components;
145};
146
162template <class Idx = std::uint32_t> class AtomicUnionFind {
163public:
164 static_assert(std::is_unsigned_v<Idx>, "AtomicUnionFind Idx must be an unsigned integer type");
165
167 explicit AtomicUnionFind(std::size_t n) : m_parent(n) {
168 for (std::size_t i = 0; i < n; ++i) {
169 m_parent[i].store(static_cast<Idx>(i), std::memory_order_relaxed);
170 }
171 }
172
183 Idx find(Idx x) noexcept {
184 assert(static_cast<std::size_t>(x) < m_parent.size() &&
185 "AtomicUnionFind::find index out of range");
186 while (true) {
187 Idx parent = m_parent[x].load(std::memory_order_acquire);
188 if (parent == x) {
189 return x;
190 }
191 const Idx grandparent = m_parent[parent].load(std::memory_order_acquire);
192 if (grandparent == parent) {
193 return parent;
194 }
195 // Halve the path: retarget x at its grandparent. A lost race means another thread
196 // already advanced the slot; the walk continues from the grandparent either way.
197 m_parent[x].compare_exchange_weak(parent, grandparent, std::memory_order_release,
198 std::memory_order_relaxed);
199 x = grandparent;
200 }
201 }
202
209 bool unite(Idx a, Idx b) noexcept {
210 while (true) {
211 Idx ra = find(a);
212 Idx rb = find(b);
213 if (ra == rb) {
214 return false;
215 }
216 if (rb < ra) {
217 std::swap(ra, rb);
218 }
219 Idx expected = rb;
220 if (m_parent[rb].compare_exchange_strong(expected, ra, std::memory_order_acq_rel,
221 std::memory_order_acquire)) {
222 return true;
223 }
224 // rb stopped being a root mid-flight; restart from the advanced positions.
225 a = ra;
226 b = expected;
227 }
228 }
229
231 [[nodiscard]] std::size_t size() const noexcept { return m_parent.size(); }
232
233private:
234 std::vector<std::atomic<Idx>> m_parent;
235};
236
237} // namespace clustering
AtomicUnionFind(std::size_t n)
Construct n singleton components numbered [0, n).
Definition dsu.h:167
bool unite(Idx a, Idx b) noexcept
Merge the components containing a and b; larger root links under smaller.
Definition dsu.h:209
Idx find(Idx x) noexcept
Root of the component containing x at some point during the call.
Definition dsu.h:183
std::size_t size() const noexcept
Total number of elements under management (fixed at construction).
Definition dsu.h:231
Idx find(Idx x) noexcept
Root of the component containing x, with path compression applied.
Definition dsu.h:49
std::size_t size() const noexcept
Total number of elements under management (fixed at construction).
Definition dsu.h:121
bool unite(Idx a, Idx b) noexcept
Merge the components containing a and b.
Definition dsu.h:75
std::size_t componentSize(Idx root) const noexcept
Population of the component whose root is root.
Definition dsu.h:134
UnionFind(std::size_t n)
Construct n singleton components numbered [0, n).
Definition dsu.h:33
bool sameComponent(Idx a, Idx b) noexcept
Whether a and b share a component.
Definition dsu.h:105
std::size_t countComponents() const noexcept
Current number of distinct components.
Definition dsu.h:114