m1une's library

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

View on GitHub

:heavy_check_mark: Persistent DSU
(ds/dsu/persistent_dsu.hpp)

Overview

PersistentDsu is a persistent Union-Find data structure. Merge operations return a new version and leave the old version available.

It uses union by size without path compression, because path compression mutates the search path. Parent and size values are stored in a persistent array, so each merge shares most nodes with older versions. Reference counting recycles internal array nodes after their final dependent version and parent are released.

merge returns a new version. merge_inplace mutates this handle with copy-on-write and returns whether two previously separate components were joined. Other live versions remain unchanged, and uniquely owned persistent array paths are reused.

Complexity Notation

Methods

Method Description Complexity
PersistentDsu() Creates an empty DSU. $O(1)$
explicit PersistentDsu(int n) Creates n singleton sets. $O(N)$
int size() const Returns the number of elements. $O(1)$
bool empty() const Returns whether the DSU has no elements. $O(1)$
void release() Releases this version immediately and makes this handle empty. $O(F)$
std::size_t node_count() const Returns live internal nodes in the shared version family. $O(1)$
PersistentDsu merge(int a, int b) const Returns a new version where the sets containing a and b are merged. $O(\log^2 N)$
bool merge_inplace(int a, int b) Merges in this version using copy-on-write and returns whether a merge occurred. $O(\log^2 N)$
bool same(int a, int b) const Returns whether a and b are in the same set. $O(\log^2 N)$
int leader(int a) const Returns the representative of the set containing a. $O(\log^2 N)$
int group_size(int a) const, int size(int a) const Returns the size of the set containing a. $O(\log^2 N)$
int get(int p) const Returns the internal parent-or-size value at index p. Roots store negative component sizes; non-roots store parent indices. $O(\log N)$
std::vector<std::vector<int>> groups() const Returns all sets as vectors of element indices. $O(N \log^2 N)$

Here $F$ is the number of internal nodes that become unreachable. Destruction and assignment release roots automatically.

Example

#include "ds/dsu/persistent_dsu.hpp"

#include <iostream>

using namespace m1une::ds;

int main() {
    PersistentDsu dsu(5);

    PersistentDsu a = dsu.merge(0, 1);
    PersistentDsu b = a.merge(1, 2);

    std::cout << dsu.same(0, 2) << "\n"; // 0
    std::cout << a.same(0, 2) << "\n";   // 0
    std::cout << b.same(0, 2) << "\n";   // 1
    std::cout << b.size(0) << "\n";       // 3
}

Depends on

Verified with

Code

#ifndef M1UNE_PERSISTENT_DSU_HPP
#define M1UNE_PERSISTENT_DSU_HPP 1

#include <algorithm>
#include <cassert>
#include <cstddef>
#include <memory>
#include <utility>
#include <vector>

#include "../detail/persistent_binary_node_pool.hpp"

namespace m1une {
namespace ds {

struct PersistentDsu {
   private:
    struct Node {
        int val;
        int l, r;

        Node() : val(0), l(0), r(0) {}
        explicit Node(int value) : val(value), l(0), r(0) {}
        Node(int value, int left, int right) : val(value), l(left), r(right) {}
    };

    int _n;
    int _root;
    using Pool = detail::PersistentBinaryNodePool<Node, 0>;

    std::shared_ptr<Pool> _pool;

    explicit PersistentDsu(int n, int root, std::shared_ptr<Pool> pool)
        : _n(n), _root(root), _pool(std::move(pool)) {
        _pool->retain(_root);
    }

    int new_node(const Node& node) const {
        return _pool->emplace(node);
    }

    int new_node(Node&& node) const {
        return _pool->emplace(std::move(node));
    }

    int build(int l, int r) const {
        if (l == r) return 0;
        if (r - l == 1) return new_node(Node(-1));
        int m = (l + r) >> 1;
        int left = build(l, m);
        int right = build(m, r);
        return new_node(Node(0, left, right));
    }

    int set_node(int t, int l, int r, int p, int value, bool copy_on_write = false) const {
        if (copy_on_write) t = _pool->clone_if_shared(t);
        if (r - l == 1) {
            if (copy_on_write) {
                (*_pool)[t].val = value;
                return t;
            }
            return new_node(Node(value));
        }
        int m = (l + r) >> 1;
        int left = (*_pool)[t].l;
        int right = (*_pool)[t].r;
        if (p < m) {
            left = set_node(left, l, m, p, value, copy_on_write);
        } else {
            right = set_node(right, m, r, p, value, copy_on_write);
        }
        if (copy_on_write) {
            _pool->replace((*_pool)[t].l, left);
            _pool->replace((*_pool)[t].r, right);
            return t;
        }
        return new_node(Node(0, left, right));
    }

    PersistentDsu make_version(int root) const {
        PersistentDsu result(_n, root, _pool);
        _pool->discard_unreferenced();
        return result;
    }

    int get_node(int t, int l, int r, int p) const {
        while (r - l > 1) {
            int m = (l + r) >> 1;
            if (p < m) {
                t = (*_pool)[t].l;
                r = m;
            } else {
                t = (*_pool)[t].r;
                l = m;
            }
        }
        return (*_pool)[t].val;
    }

   public:
    PersistentDsu() : PersistentDsu(0) {}

    explicit PersistentDsu(int n) : _n(n), _root(0), _pool(std::make_shared<Pool>()) {
        assert(0 <= n);
        _pool->reserve(n * 2 + 1);
        if (_n > 0) _root = build(0, _n);
        _pool->retain(_root);
        _pool->discard_unreferenced();
    }

    PersistentDsu(const PersistentDsu& other) : _n(other._n), _root(other._root), _pool(other._pool) {
        if (_pool) _pool->retain(_root);
    }

    PersistentDsu(PersistentDsu&& other) noexcept
        : _n(other._n), _root(other._root), _pool(std::move(other._pool)) {
        other._n = 0;
        other._root = 0;
    }

    PersistentDsu& operator=(const PersistentDsu& other) {
        if (this == &other) return *this;
        if (other._pool) other._pool->retain(other._root);
        if (_pool) _pool->release(_root);
        _n = other._n;
        _root = other._root;
        _pool = other._pool;
        return *this;
    }

    PersistentDsu& operator=(PersistentDsu&& other) noexcept {
        if (this == &other) return *this;
        if (_pool) _pool->release(_root);
        _n = other._n;
        _root = other._root;
        _pool = std::move(other._pool);
        other._n = 0;
        other._root = 0;
        return *this;
    }

    ~PersistentDsu() {
        if (_pool) _pool->release(_root);
    }

    int size() const {
        return _n;
    }

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

    void release() {
        if (_pool) _pool->release(_root);
        _n = 0;
        _root = 0;
        _pool = std::make_shared<Pool>();
    }

    std::size_t node_count() const { return _pool ? _pool->size() : 0; }

    int leader(int a) const {
        assert(0 <= a && a < _n);
        int x = a;
        int p = get(x);
        while (p >= 0) {
            x = p;
            p = get(x);
        }
        return x;
    }

    bool same(int a, int b) const {
        assert(0 <= a && a < _n);
        assert(0 <= b && b < _n);
        return leader(a) == leader(b);
    }

    int group_size(int a) const {
        assert(0 <= a && a < _n);
        return -get(leader(a));
    }

    int size(int a) const {
        return group_size(a);
    }

    int get(int p) const {
        assert(0 <= p && p < _n);
        return get_node(_root, 0, _n, p);
    }

    PersistentDsu merge(int a, int b) const {
        assert(0 <= a && a < _n);
        assert(0 <= b && b < _n);
        int x = leader(a), y = leader(b);
        if (x == y) return *this;
        int sx = -get(x), sy = -get(y);
        if (sx < sy) {
            std::swap(x, y);
            std::swap(sx, sy);
        }
        int root = set_node(_root, 0, _n, x, -(sx + sy));
        root = set_node(root, 0, _n, y, x);
        return make_version(root);
    }

    bool merge_inplace(int a, int b) {
        assert(0 <= a && a < _n);
        assert(0 <= b && b < _n);
        int x = leader(a), y = leader(b);
        if (x == y) return false;
        int sx = -get(x), sy = -get(y);
        if (sx < sy) {
            std::swap(x, y);
            std::swap(sx, sy);
        }
        int root = set_node(_root, 0, _n, x, -(sx + sy), true);
        _pool->replace(_root, root);
        root = set_node(_root, 0, _n, y, x, true);
        _pool->replace(_root, root);
        _pool->discard_unreferenced();
        return true;
    }

    std::vector<std::vector<int>> groups() const {
        std::vector<int> leader_buf(_n), group_size(_n);
        for (int i = 0; i < _n; i++) {
            leader_buf[i] = leader(i);
            group_size[leader_buf[i]]++;
        }
        std::vector<std::vector<int>> result(_n);
        for (int i = 0; i < _n; i++) {
            result[i].reserve(group_size[i]);
        }
        for (int i = 0; i < _n; i++) {
            result[leader_buf[i]].push_back(i);
        }
        result.erase(std::remove_if(result.begin(), result.end(), [&](const std::vector<int>& v) { return v.empty(); }),
                     result.end());
        return result;
    }
};

}  // namespace ds
}  // namespace m1une

#endif  // M1UNE_PERSISTENT_DSU_HPP
#line 1 "ds/dsu/persistent_dsu.hpp"



#include <algorithm>
#include <cassert>
#include <cstddef>
#include <memory>
#include <utility>
#include <vector>

#line 1 "ds/detail/persistent_binary_node_pool.hpp"



#line 6 "ds/detail/persistent_binary_node_pool.hpp"
#include <deque>
#include <limits>
#include <optional>
#line 11 "ds/detail/persistent_binary_node_pool.hpp"

namespace m1une {
namespace ds {
namespace detail {

// Node must have integer `l` and `r` members. New nodes initially have no
// owner; discard_unreferenced() removes temporary path-copy nodes after the
// result roots have been retained.
template <class Node, int null_node = -1>
struct PersistentBinaryNodePool {
   private:
    std::deque<std::optional<Node>> _nodes;
    std::vector<int> _references;
    std::vector<int> _next_free;
    std::vector<int> _unowned;
    int _first_free = -1;
    std::size_t _live_nodes = 0;

    void release_zero(int node) {
        assert(node != null_node && _nodes[node].has_value());
        int left = (*_nodes[node]).l;
        int right = (*_nodes[node]).r;
        _nodes[node].reset();
        _next_free[node] = _first_free;
        _first_free = node;
        --_live_nodes;
        if (left != null_node && --_references[left] == 0) release_zero(left);
        if (right != null_node && --_references[right] == 0) release_zero(right);
    }

   public:
    PersistentBinaryNodePool() {
        if constexpr (null_node == 0) {
            _nodes.emplace_back();
            _references.push_back(0);
            _next_free.push_back(-1);
        }
    }

    Node& operator[](int node) {
        assert(node != null_node && _nodes[node].has_value());
        return *_nodes[node];
    }

    const Node& operator[](int node) const {
        assert(node != null_node && _nodes[node].has_value());
        return *_nodes[node];
    }

    template <class... Args>
    int emplace(Args&&... args) {
        int result;
        if (_first_free == -1) {
            assert(_nodes.size() < std::size_t(std::numeric_limits<int>::max()));
            result = int(_nodes.size());
            _nodes.emplace_back(std::in_place, std::forward<Args>(args)...);
            _references.push_back(0);
            _next_free.push_back(-1);
        } else {
            result = _first_free;
            _first_free = _next_free[result];
            _nodes[result].emplace(std::forward<Args>(args)...);
            _references[result] = 0;
        }
        retain((*_nodes[result]).l);
        retain((*_nodes[result]).r);
        _unowned.push_back(result);
        ++_live_nodes;
        return result;
    }

    void retain(int node) {
        if (node != null_node) {
            assert(_nodes[node].has_value());
            ++_references[node];
        }
    }

    void release(int node) {
        if (node == null_node) return;
        assert(_nodes[node].has_value() && _references[node] > 0);
        if (--_references[node] == 0) release_zero(node);
    }

    bool unique(int node) const {
        return node == null_node || _references[node] == 1;
    }

    int clone(int node) {
        assert(node != null_node && _nodes[node].has_value());
        return emplace(*_nodes[node]);
    }

    // Returns node itself when it has one owner, otherwise an unowned clone.
    // A returned clone becomes owned when a root or parent edge retains it.
    int clone_if_shared(int node) {
        if (unique(node)) return node;
        return clone(node);
    }

    void replace(int& edge, int node) {
        if (edge == node) return;
        retain(node);
        int old = edge;
        edge = node;
        release(old);
    }

    void discard_unreferenced() {
        while (!_unowned.empty()) {
            int node = _unowned.back();
            _unowned.pop_back();
            if (_nodes[node].has_value() && _references[node] == 0) release_zero(node);
        }
    }

    void reserve(std::size_t) {}

    int next_index() const { return _first_free == -1 ? int(_nodes.size()) : _first_free; }

    std::size_t size() const { return _live_nodes; }
};

}  // namespace detail
}  // namespace ds
}  // namespace m1une


#line 12 "ds/dsu/persistent_dsu.hpp"

namespace m1une {
namespace ds {

struct PersistentDsu {
   private:
    struct Node {
        int val;
        int l, r;

        Node() : val(0), l(0), r(0) {}
        explicit Node(int value) : val(value), l(0), r(0) {}
        Node(int value, int left, int right) : val(value), l(left), r(right) {}
    };

    int _n;
    int _root;
    using Pool = detail::PersistentBinaryNodePool<Node, 0>;

    std::shared_ptr<Pool> _pool;

    explicit PersistentDsu(int n, int root, std::shared_ptr<Pool> pool)
        : _n(n), _root(root), _pool(std::move(pool)) {
        _pool->retain(_root);
    }

    int new_node(const Node& node) const {
        return _pool->emplace(node);
    }

    int new_node(Node&& node) const {
        return _pool->emplace(std::move(node));
    }

    int build(int l, int r) const {
        if (l == r) return 0;
        if (r - l == 1) return new_node(Node(-1));
        int m = (l + r) >> 1;
        int left = build(l, m);
        int right = build(m, r);
        return new_node(Node(0, left, right));
    }

    int set_node(int t, int l, int r, int p, int value, bool copy_on_write = false) const {
        if (copy_on_write) t = _pool->clone_if_shared(t);
        if (r - l == 1) {
            if (copy_on_write) {
                (*_pool)[t].val = value;
                return t;
            }
            return new_node(Node(value));
        }
        int m = (l + r) >> 1;
        int left = (*_pool)[t].l;
        int right = (*_pool)[t].r;
        if (p < m) {
            left = set_node(left, l, m, p, value, copy_on_write);
        } else {
            right = set_node(right, m, r, p, value, copy_on_write);
        }
        if (copy_on_write) {
            _pool->replace((*_pool)[t].l, left);
            _pool->replace((*_pool)[t].r, right);
            return t;
        }
        return new_node(Node(0, left, right));
    }

    PersistentDsu make_version(int root) const {
        PersistentDsu result(_n, root, _pool);
        _pool->discard_unreferenced();
        return result;
    }

    int get_node(int t, int l, int r, int p) const {
        while (r - l > 1) {
            int m = (l + r) >> 1;
            if (p < m) {
                t = (*_pool)[t].l;
                r = m;
            } else {
                t = (*_pool)[t].r;
                l = m;
            }
        }
        return (*_pool)[t].val;
    }

   public:
    PersistentDsu() : PersistentDsu(0) {}

    explicit PersistentDsu(int n) : _n(n), _root(0), _pool(std::make_shared<Pool>()) {
        assert(0 <= n);
        _pool->reserve(n * 2 + 1);
        if (_n > 0) _root = build(0, _n);
        _pool->retain(_root);
        _pool->discard_unreferenced();
    }

    PersistentDsu(const PersistentDsu& other) : _n(other._n), _root(other._root), _pool(other._pool) {
        if (_pool) _pool->retain(_root);
    }

    PersistentDsu(PersistentDsu&& other) noexcept
        : _n(other._n), _root(other._root), _pool(std::move(other._pool)) {
        other._n = 0;
        other._root = 0;
    }

    PersistentDsu& operator=(const PersistentDsu& other) {
        if (this == &other) return *this;
        if (other._pool) other._pool->retain(other._root);
        if (_pool) _pool->release(_root);
        _n = other._n;
        _root = other._root;
        _pool = other._pool;
        return *this;
    }

    PersistentDsu& operator=(PersistentDsu&& other) noexcept {
        if (this == &other) return *this;
        if (_pool) _pool->release(_root);
        _n = other._n;
        _root = other._root;
        _pool = std::move(other._pool);
        other._n = 0;
        other._root = 0;
        return *this;
    }

    ~PersistentDsu() {
        if (_pool) _pool->release(_root);
    }

    int size() const {
        return _n;
    }

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

    void release() {
        if (_pool) _pool->release(_root);
        _n = 0;
        _root = 0;
        _pool = std::make_shared<Pool>();
    }

    std::size_t node_count() const { return _pool ? _pool->size() : 0; }

    int leader(int a) const {
        assert(0 <= a && a < _n);
        int x = a;
        int p = get(x);
        while (p >= 0) {
            x = p;
            p = get(x);
        }
        return x;
    }

    bool same(int a, int b) const {
        assert(0 <= a && a < _n);
        assert(0 <= b && b < _n);
        return leader(a) == leader(b);
    }

    int group_size(int a) const {
        assert(0 <= a && a < _n);
        return -get(leader(a));
    }

    int size(int a) const {
        return group_size(a);
    }

    int get(int p) const {
        assert(0 <= p && p < _n);
        return get_node(_root, 0, _n, p);
    }

    PersistentDsu merge(int a, int b) const {
        assert(0 <= a && a < _n);
        assert(0 <= b && b < _n);
        int x = leader(a), y = leader(b);
        if (x == y) return *this;
        int sx = -get(x), sy = -get(y);
        if (sx < sy) {
            std::swap(x, y);
            std::swap(sx, sy);
        }
        int root = set_node(_root, 0, _n, x, -(sx + sy));
        root = set_node(root, 0, _n, y, x);
        return make_version(root);
    }

    bool merge_inplace(int a, int b) {
        assert(0 <= a && a < _n);
        assert(0 <= b && b < _n);
        int x = leader(a), y = leader(b);
        if (x == y) return false;
        int sx = -get(x), sy = -get(y);
        if (sx < sy) {
            std::swap(x, y);
            std::swap(sx, sy);
        }
        int root = set_node(_root, 0, _n, x, -(sx + sy), true);
        _pool->replace(_root, root);
        root = set_node(_root, 0, _n, y, x, true);
        _pool->replace(_root, root);
        _pool->discard_unreferenced();
        return true;
    }

    std::vector<std::vector<int>> groups() const {
        std::vector<int> leader_buf(_n), group_size(_n);
        for (int i = 0; i < _n; i++) {
            leader_buf[i] = leader(i);
            group_size[leader_buf[i]]++;
        }
        std::vector<std::vector<int>> result(_n);
        for (int i = 0; i < _n; i++) {
            result[i].reserve(group_size[i]);
        }
        for (int i = 0; i < _n; i++) {
            result[leader_buf[i]].push_back(i);
        }
        result.erase(std::remove_if(result.begin(), result.end(), [&](const std::vector<int>& v) { return v.empty(); }),
                     result.end());
        return result;
    }
};

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