Skip to content

mergeable_frequency_tree.hpp

SECTIONData Structure INCLUDEnoya/mergeable_frequency_tree.hpp

维护值域频次线段树并支持破坏性合并;适合树上启发式合并、子树频率和顺序统计。

Complexity: Time: O(log U) update/query; O(min(nodes_a,nodes_b)) destructive merge. Space: O(A), where A is the total number of allocated nodes.

跳到代码 · GitHub ↗

Implementation

当前头文件,省略 include guard;依赖见 #include

/// @complexity Time: O(log U) update/query; O(min(nodes_a,nodes_b)) destructive merge.
/// Space: O(A), where A is the total number of allocated nodes.

#include <cassert>
#include <numeric>
#include <vector>

namespace noya {

/// @brief Dynamic frequency segment-tree pool supporting point updates, rank
/// queries, destructive merging, and interval extraction in O(log U).
template <class Count = int> struct mergeable_frequency_tree {
  struct node {
    int ls = 0;
    int rs = 0;
    Count sum{};
  };

  int low = 0;
  int big = 0;
  std::vector<node> tr{{}};

  mergeable_frequency_tree() = default;
  mergeable_frequency_tree(int lo, int hi) : low(lo), big(hi) {
    assert(low < big);
  }

  Count total(int rt) const {
    check_root(rt);
    return tr[rt].sum;
  }

  /// @brief Add delta at position, creating root and path nodes as needed.
  void add(int &rt, int pos, Count dif) {
    assert(low <= pos && pos < big);
    add_impl(rt, low, big, pos, dif);
  }

  /// @brief Return the total frequency in [ql, qr).
  Count count(int rt, int ql, int qr) const {
    check_root(rt);
    assert(low <= ql && ql <= qr && qr <= big);
    return count_impl(rt, low, big, ql, qr);
  }

  /// @brief Return the coordinate containing zero-indexed rank k.
  int kth(int rt, Count k) const {
    check_root(rt);
    assert(Count{} <= k && k < tr[rt].sum);
    int l = low;
    int r = big;
    while (r - l > 1) {
      Count sl = tr[tr[rt].ls].sum;
      int mid = std::midpoint(l, r);
      if (k < sl) {
        rt = tr[rt].ls;
        r = mid;
      } else {
        k -= sl;
        rt = tr[rt].rs;
        l = mid;
      }
    }
    return l;
  }

  /// @brief Move every frequency in [ql, qr) from root into
  /// a new root without duplicating leaves.
  int extract(int &rt, int ql, int qr) {
    check_root(rt);
    assert(low <= ql && ql <= qr && qr <= big);
    return extract_impl(rt, low, big, ql, qr);
  }

  /// @brief Destructively merge source into destination and clear source.
  void merge_into(int &dst, int &src) {
    check_root(dst);
    check_root(src);
    if (dst == src && dst != 0) {
      assert(false && "cannot merge a root into itself");
    }
    dst = merge_impl(dst, src);
    src = 0;
  }

private:
  void check_root(int rt) const { assert(0 <= rt && rt < int(tr.size())); }

  int make_node() {
    tr.push_back({});
    return int(tr.size()) - 1;
  }

  void pull(int rt) { tr[rt].sum = tr[tr[rt].ls].sum + tr[tr[rt].rs].sum; }

  void add_impl(int &rt, int l, int r, int pos, Count dif) {
    if (rt == 0) {
      rt = make_node();
    }
    if (r - l == 1) {
      tr[rt].sum += dif;
      assert(tr[rt].sum >= Count{});
      return;
    }
    int mid = std::midpoint(l, r);
    if (pos < mid) {
      int ch = tr[rt].ls;
      add_impl(ch, l, mid, pos, dif);
      tr[rt].ls = ch;
    } else {
      int ch = tr[rt].rs;
      add_impl(ch, mid, r, pos, dif);
      tr[rt].rs = ch;
    }
    pull(rt);
  }

  Count count_impl(int rt, int l, int r, int ql, int qr) const {
    if (rt == 0 || qr <= l || r <= ql) {
      return Count{};
    }
    if (ql <= l && r <= qr) {
      return tr[rt].sum;
    }
    int mid = std::midpoint(l, r);
    return count_impl(tr[rt].ls, l, mid, ql, qr) +
           count_impl(tr[rt].rs, mid, r, ql, qr);
  }

  int extract_impl(int &rt, int l, int r, int ql, int qr) {
    if (rt == 0 || qr <= l || r <= ql) {
      return 0;
    }
    if (ql <= l && r <= qr) {
      int res = rt;
      rt = 0;
      return res;
    }
    int res = make_node();
    int mid = std::midpoint(l, r);
    int lc = tr[rt].ls;
    int rc = tr[rt].rs;
    tr[res].ls = extract_impl(lc, l, mid, ql, qr);
    tr[res].rs = extract_impl(rc, mid, r, ql, qr);
    tr[rt].ls = lc;
    tr[rt].rs = rc;
    pull(rt);
    pull(res);
    if (tr[res].sum == Count{}) {
      return 0;
    }
    return res;
  }

  int merge_impl(int dst, int src) {
    if (dst == 0 || src == 0) {
      return dst | src;
    }
    int dl = tr[dst].ls;
    int dr = tr[dst].rs;
    dl = merge_impl(dl, tr[src].ls);
    dr = merge_impl(dr, tr[src].rs);
    tr[dst].ls = dl;
    tr[dst].rs = dr;
    if (dl == 0 && dr == 0) {
      tr[dst].sum += tr[src].sum;
    } else {
      pull(dst);
    }
    return dst;
  }
};

} // namespace noya
#ifndef NOYA_MERGEABLE_FREQUENCY_TREE_HPP
#define NOYA_MERGEABLE_FREQUENCY_TREE_HPP 1

/// @complexity Time: O(log U) update/query; O(min(nodes_a,nodes_b)) destructive merge.
/// Space: O(A), where A is the total number of allocated nodes.

#include <cassert>
#include <numeric>
#include <vector>

namespace noya {

/// @brief Dynamic frequency segment-tree pool supporting point updates, rank
/// queries, destructive merging, and interval extraction in O(log U).
template <class Count = int> struct mergeable_frequency_tree {
  struct node {
    int ls = 0;
    int rs = 0;
    Count sum{};
  };

  int low = 0;
  int big = 0;
  std::vector<node> tr{{}};

  mergeable_frequency_tree() = default;
  mergeable_frequency_tree(int lo, int hi) : low(lo), big(hi) {
    assert(low < big);
  }

  Count total(int rt) const {
    check_root(rt);
    return tr[rt].sum;
  }

  /// @brief Add delta at position, creating root and path nodes as needed.
  void add(int &rt, int pos, Count dif) {
    assert(low <= pos && pos < big);
    add_impl(rt, low, big, pos, dif);
  }

  /// @brief Return the total frequency in [ql, qr).
  Count count(int rt, int ql, int qr) const {
    check_root(rt);
    assert(low <= ql && ql <= qr && qr <= big);
    return count_impl(rt, low, big, ql, qr);
  }

  /// @brief Return the coordinate containing zero-indexed rank k.
  int kth(int rt, Count k) const {
    check_root(rt);
    assert(Count{} <= k && k < tr[rt].sum);
    int l = low;
    int r = big;
    while (r - l > 1) {
      Count sl = tr[tr[rt].ls].sum;
      int mid = std::midpoint(l, r);
      if (k < sl) {
        rt = tr[rt].ls;
        r = mid;
      } else {
        k -= sl;
        rt = tr[rt].rs;
        l = mid;
      }
    }
    return l;
  }

  /// @brief Move every frequency in [ql, qr) from root into
  /// a new root without duplicating leaves.
  int extract(int &rt, int ql, int qr) {
    check_root(rt);
    assert(low <= ql && ql <= qr && qr <= big);
    return extract_impl(rt, low, big, ql, qr);
  }

  /// @brief Destructively merge source into destination and clear source.
  void merge_into(int &dst, int &src) {
    check_root(dst);
    check_root(src);
    if (dst == src && dst != 0) {
      assert(false && "cannot merge a root into itself");
    }
    dst = merge_impl(dst, src);
    src = 0;
  }

private:
  void check_root(int rt) const { assert(0 <= rt && rt < int(tr.size())); }

  int make_node() {
    tr.push_back({});
    return int(tr.size()) - 1;
  }

  void pull(int rt) { tr[rt].sum = tr[tr[rt].ls].sum + tr[tr[rt].rs].sum; }

  void add_impl(int &rt, int l, int r, int pos, Count dif) {
    if (rt == 0) {
      rt = make_node();
    }
    if (r - l == 1) {
      tr[rt].sum += dif;
      assert(tr[rt].sum >= Count{});
      return;
    }
    int mid = std::midpoint(l, r);
    if (pos < mid) {
      int ch = tr[rt].ls;
      add_impl(ch, l, mid, pos, dif);
      tr[rt].ls = ch;
    } else {
      int ch = tr[rt].rs;
      add_impl(ch, mid, r, pos, dif);
      tr[rt].rs = ch;
    }
    pull(rt);
  }

  Count count_impl(int rt, int l, int r, int ql, int qr) const {
    if (rt == 0 || qr <= l || r <= ql) {
      return Count{};
    }
    if (ql <= l && r <= qr) {
      return tr[rt].sum;
    }
    int mid = std::midpoint(l, r);
    return count_impl(tr[rt].ls, l, mid, ql, qr) +
           count_impl(tr[rt].rs, mid, r, ql, qr);
  }

  int extract_impl(int &rt, int l, int r, int ql, int qr) {
    if (rt == 0 || qr <= l || r <= ql) {
      return 0;
    }
    if (ql <= l && r <= qr) {
      int res = rt;
      rt = 0;
      return res;
    }
    int res = make_node();
    int mid = std::midpoint(l, r);
    int lc = tr[rt].ls;
    int rc = tr[rt].rs;
    tr[res].ls = extract_impl(lc, l, mid, ql, qr);
    tr[res].rs = extract_impl(rc, mid, r, ql, qr);
    tr[rt].ls = lc;
    tr[rt].rs = rc;
    pull(rt);
    pull(res);
    if (tr[res].sum == Count{}) {
      return 0;
    }
    return res;
  }

  int merge_impl(int dst, int src) {
    if (dst == 0 || src == 0) {
      return dst | src;
    }
    int dl = tr[dst].ls;
    int dr = tr[dst].rs;
    dl = merge_impl(dl, tr[src].ls);
    dr = merge_impl(dr, tr[src].rs);
    tr[dst].ls = dl;
    tr[dst].rs = dr;
    if (dl == 0 && dr == 0) {
      tr[dst].sum += tr[src].sum;
    } else {
      pull(dst);
    }
    return dst;
  }
};

} // namespace noya

#endif // NOYA_MERGEABLE_FREQUENCY_TREE_HPP
#include <cassert>
#include <numeric>
#include <vector>

/// @complexity Time: O(log U) update/query; O(min(nodes_a,nodes_b)) destructive merge.
/// Space: O(A), where A is the total number of allocated nodes.

namespace noya {

/// @brief Dynamic frequency segment-tree pool supporting point updates, rank
/// queries, destructive merging, and interval extraction in O(log U).
template <class Count = int> struct mergeable_frequency_tree {
  struct node {
    int ls = 0;
    int rs = 0;
    Count sum{};
  };

  int low = 0;
  int big = 0;
  std::vector<node> tr{{}};

  mergeable_frequency_tree() = default;
  mergeable_frequency_tree(int lo, int hi) : low(lo), big(hi) {
    assert(low < big);
  }

  Count total(int rt) const {
    check_root(rt);
    return tr[rt].sum;
  }

  /// @brief Add delta at position, creating root and path nodes as needed.
  void add(int &rt, int pos, Count dif) {
    assert(low <= pos && pos < big);
    add_impl(rt, low, big, pos, dif);
  }

  /// @brief Return the total frequency in [ql, qr).
  Count count(int rt, int ql, int qr) const {
    check_root(rt);
    assert(low <= ql && ql <= qr && qr <= big);
    return count_impl(rt, low, big, ql, qr);
  }

  /// @brief Return the coordinate containing zero-indexed rank k.
  int kth(int rt, Count k) const {
    check_root(rt);
    assert(Count{} <= k && k < tr[rt].sum);
    int l = low;
    int r = big;
    while (r - l > 1) {
      Count sl = tr[tr[rt].ls].sum;
      int mid = std::midpoint(l, r);
      if (k < sl) {
        rt = tr[rt].ls;
        r = mid;
      } else {
        k -= sl;
        rt = tr[rt].rs;
        l = mid;
      }
    }
    return l;
  }

  /// @brief Move every frequency in [ql, qr) from root into
  /// a new root without duplicating leaves.
  int extract(int &rt, int ql, int qr) {
    check_root(rt);
    assert(low <= ql && ql <= qr && qr <= big);
    return extract_impl(rt, low, big, ql, qr);
  }

  /// @brief Destructively merge source into destination and clear source.
  void merge_into(int &dst, int &src) {
    check_root(dst);
    check_root(src);
    if (dst == src && dst != 0) {
      assert(false && "cannot merge a root into itself");
    }
    dst = merge_impl(dst, src);
    src = 0;
  }

private:
  void check_root(int rt) const { assert(0 <= rt && rt < int(tr.size())); }

  int make_node() {
    tr.push_back({});
    return int(tr.size()) - 1;
  }

  void pull(int rt) { tr[rt].sum = tr[tr[rt].ls].sum + tr[tr[rt].rs].sum; }

  void add_impl(int &rt, int l, int r, int pos, Count dif) {
    if (rt == 0) {
      rt = make_node();
    }
    if (r - l == 1) {
      tr[rt].sum += dif;
      assert(tr[rt].sum >= Count{});
      return;
    }
    int mid = std::midpoint(l, r);
    if (pos < mid) {
      int ch = tr[rt].ls;
      add_impl(ch, l, mid, pos, dif);
      tr[rt].ls = ch;
    } else {
      int ch = tr[rt].rs;
      add_impl(ch, mid, r, pos, dif);
      tr[rt].rs = ch;
    }
    pull(rt);
  }

  Count count_impl(int rt, int l, int r, int ql, int qr) const {
    if (rt == 0 || qr <= l || r <= ql) {
      return Count{};
    }
    if (ql <= l && r <= qr) {
      return tr[rt].sum;
    }
    int mid = std::midpoint(l, r);
    return count_impl(tr[rt].ls, l, mid, ql, qr) +
           count_impl(tr[rt].rs, mid, r, ql, qr);
  }

  int extract_impl(int &rt, int l, int r, int ql, int qr) {
    if (rt == 0 || qr <= l || r <= ql) {
      return 0;
    }
    if (ql <= l && r <= qr) {
      int res = rt;
      rt = 0;
      return res;
    }
    int res = make_node();
    int mid = std::midpoint(l, r);
    int lc = tr[rt].ls;
    int rc = tr[rt].rs;
    tr[res].ls = extract_impl(lc, l, mid, ql, qr);
    tr[res].rs = extract_impl(rc, mid, r, ql, qr);
    tr[rt].ls = lc;
    tr[rt].rs = rc;
    pull(rt);
    pull(res);
    if (tr[res].sum == Count{}) {
      return 0;
    }
    return res;
  }

  int merge_impl(int dst, int src) {
    if (dst == 0 || src == 0) {
      return dst | src;
    }
    int dl = tr[dst].ls;
    int dr = tr[dst].rs;
    dl = merge_impl(dl, tr[src].ls);
    dr = merge_impl(dr, tr[src].rs);
    tr[dst].ls = dl;
    tr[dst].rs = dr;
    if (dl == 0 && dr == 0) {
      tr[dst].sum += tr[src].sum;
    } else {
      pull(dst);
    }
    return dst;
  }
};

} // namespace noya