Skip to content

tree_contour.hpp

SECTIONGraph INCLUDEnoya/tree_contour.hpp

在静态树上按到某点的距离区间做更新或查询;适合查询距离满足 \(l\le x<r\) 的所有点。

Complexity: Time: O(n log^2 n) construction and O(log^2 n) per update or query. Space: O(n log n).

AC 记录:vertex_add_range_contour_sum_on_tree, vertex_get_range_contour_add_on_tree

跳到代码 · GitHub ↗

Implementation

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

/// @complexity Time: O(n log^2 n) construction and O(log^2 n) per update or
/// query. Space: O(n log n).

#include <algorithm>
#include <cassert>
#include <tuple>
#include <utility>
#include <vector>

namespace noya {

/// @brief Maintain values indexed by tree distance. Every vertex is recorded
/// in all centroid-ancestor buckets; querying adds all buckets and subtracts
/// the bucket of the branch containing the query vertex. Fenwick trees over
/// distances support both point-add/range-sum and the dual range-add/point-get.
class tree_contour {
  struct fenwick {
    std::vector<long long> dat;

    fenwick() = default;
    explicit fenwick(int siz) : dat(siz + 1) {}

    int size() const { return int(dat.size()) - 1; }

    void add(int i, long long val) {
      assert(0 <= i && i < size());
      for (i++; i < int(dat.size()); i += i & -i) {
        dat[i] += val;
      }
    }

    long long prefix_sum(int cnt) const {
      cnt = std::clamp(cnt, 0, size());
      long long res = 0;
      for (; cnt > 0; cnt -= cnt & -cnt) {
        res += dat[cnt];
      }
      return res;
    }
  };

  struct bucket {
    int nd = 0;
    fenwick bit;
    fenwick tag;

    void reset(int cnt) {
      nd = cnt;
      bit = fenwick(cnt);
      tag = fenwick(cnt + 1);
    }

    void point_add(int dis, long long val) { bit.add(dis, val); }

    long long point_prefix(int lim) const { return bit.prefix_sum(lim); }

    void range_add(int l, int r, long long val) {
      l = std::clamp(l, 0, nd);
      r = std::clamp(r, 0, nd);
      if (l >= r) {
        return;
      }
      tag.add(l, val);
      tag.add(r, -val);
    }

    long long point_tag(int dis) const {
      assert(0 <= dis && dis < nd);
      return tag.prefix_sum(dis + 1);
    }
  };

  struct ancestor_entry {
    int cen = -1;
    int dis = 0;
    int br = -1;
  };

  int n_ = 0;
  std::vector<std::vector<int>> g_;
  std::vector<char> rm_;
  std::vector<int> cp;
  std::vector<int> sz_;
  std::vector<std::vector<ancestor_entry>> up_;
  std::vector<bucket> a_;
  std::vector<std::vector<bucket>> br_;
  std::vector<long long> iv;

  int find_centroid(int s) {
    std::vector<int> ord;
    ord.push_back(s);
    cp[s] = -1;
    for (int i = 0; i < int(ord.size()); i++) {
      int u = ord[i];
      for (int nxt : g_[u]) {
        if (rm_[nxt] || nxt == cp[u]) {
          continue;
        }
        cp[nxt] = u;
        ord.push_back(nxt);
      }
    }

    for (int i = int(ord.size()) - 1; i >= 0; i--) {
      int u = ord[i];
      sz_[u] = 1;
      for (int nxt : g_[u]) {
        if (!rm_[nxt] && cp[nxt] == u) {
          sz_[u] += sz_[nxt];
        }
      }
    }

    int sz = int(ord.size());
    int cen = s;
    int bst = sz;
    for (int u : ord) {
      int mp = sz - sz_[u];
      for (int nxt : g_[u]) {
        if (!rm_[nxt] && cp[nxt] == u) {
          mp = std::max(mp, sz_[nxt]);
        }
      }
      if (mp < bst) {
        bst = mp;
        cen = u;
      }
    }
    return cen;
  }

  void decompose(int s) {
    int cen = find_centroid(s);
    rm_[cen] = true;
    up_[cen].push_back({cen, 0, -1});

    int mxd = 0;
    for (int v : g_[cen]) {
      if (rm_[v]) {
        continue;
      }
      int br = int(br_[cen].size());
      int bmd = 0;
      std::vector<std::tuple<int, int, int>> stk = {{v, cen, 1}};
      while (!stk.empty()) {
        auto [u, fa, dis] = stk.back();
        stk.pop_back();
        up_[u].push_back({cen, dis, br});
        bmd = std::max(bmd, dis);
        for (int nxt : g_[u]) {
          if (!rm_[nxt] && nxt != fa) {
            stk.push_back({nxt, u, dis + 1});
          }
        }
      }
      br_[cen].push_back({});
      br_[cen].back().reset(bmd + 1);
      mxd = std::max(mxd, bmd);
    }
    a_[cen].reset(mxd + 1);

    for (int v : g_[cen]) {
      if (!rm_[v]) {
        decompose(v);
      }
    }
  }

  long long prefix_contour_sum(int u, int lim) const {
    long long res = 0;
    for (const ancestor_entry &ent : up_[u]) {
      int rem = lim - ent.dis;
      res += a_[ent.cen].point_prefix(rem);
      if (ent.br != -1) {
        res -= br_[ent.cen][ent.br].point_prefix(rem);
      }
    }
    return res;
  }

public:
  explicit tree_contour(const std::vector<std::vector<int>> &g,
                        const std::vector<long long> &iv1 = {})
      : n_(int(g.size())), g_(g), rm_(n_), cp(n_), sz_(n_), up_(n_), a_(n_),
        br_(n_), iv(iv1.empty() ? std::vector<long long>(n_) : iv1) {
    assert(int(iv.size()) == n_);
    if (n_ == 0) {
      return;
    }
    decompose(0);
    for (int u = 0; u < n_; u++) {
      point_add(u, iv[u]);
    }
  }

  /// @brief Add value to one vertex for later contour-sum queries.
  void point_add(int u, long long val) {
    assert(0 <= u && u < n_);
    for (const ancestor_entry &ent : up_[u]) {
      a_[ent.cen].point_add(ent.dis, val);
      if (ent.br != -1) {
        br_[ent.cen][ent.br].point_add(ent.dis, val);
      }
    }
  }

  /// @brief Sum values at vertices whose distance from vertex is in [l,r).
  long long range_sum(int u, int l, int r) const {
    assert(0 <= u && u < n_ && l <= r);
    return prefix_contour_sum(u, r) - prefix_contour_sum(u, l);
  }

  /// @brief Add value to vertices whose distance from vertex is in [l,r).
  void range_add(int u, int l, int r, long long val) {
    assert(0 <= u && u < n_ && l <= r);
    for (const ancestor_entry &ent : up_[u]) {
      int ql = l - ent.dis;
      int qr = r - ent.dis;
      a_[ent.cen].range_add(ql, qr, val);
      if (ent.br != -1) {
        br_[ent.cen][ent.br].range_add(ql, qr, val);
      }
    }
  }

  /// @brief Return the initial value plus every contour-range addition that
  /// contains this vertex.
  long long point_get(int u) const {
    assert(0 <= u && u < n_);
    long long res = iv[u];
    for (const ancestor_entry &ent : up_[u]) {
      res += a_[ent.cen].point_tag(ent.dis);
      if (ent.br != -1) {
        res -= br_[ent.cen][ent.br].point_tag(ent.dis);
      }
    }
    return res;
  }
};

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

/// @complexity Time: O(n log^2 n) construction and O(log^2 n) per update or
/// query. Space: O(n log n).

#include <algorithm>
#include <cassert>
#include <tuple>
#include <utility>
#include <vector>

namespace noya {

/// @brief Maintain values indexed by tree distance. Every vertex is recorded
/// in all centroid-ancestor buckets; querying adds all buckets and subtracts
/// the bucket of the branch containing the query vertex. Fenwick trees over
/// distances support both point-add/range-sum and the dual range-add/point-get.
class tree_contour {
  struct fenwick {
    std::vector<long long> dat;

    fenwick() = default;
    explicit fenwick(int siz) : dat(siz + 1) {}

    int size() const { return int(dat.size()) - 1; }

    void add(int i, long long val) {
      assert(0 <= i && i < size());
      for (i++; i < int(dat.size()); i += i & -i) {
        dat[i] += val;
      }
    }

    long long prefix_sum(int cnt) const {
      cnt = std::clamp(cnt, 0, size());
      long long res = 0;
      for (; cnt > 0; cnt -= cnt & -cnt) {
        res += dat[cnt];
      }
      return res;
    }
  };

  struct bucket {
    int nd = 0;
    fenwick bit;
    fenwick tag;

    void reset(int cnt) {
      nd = cnt;
      bit = fenwick(cnt);
      tag = fenwick(cnt + 1);
    }

    void point_add(int dis, long long val) { bit.add(dis, val); }

    long long point_prefix(int lim) const { return bit.prefix_sum(lim); }

    void range_add(int l, int r, long long val) {
      l = std::clamp(l, 0, nd);
      r = std::clamp(r, 0, nd);
      if (l >= r) {
        return;
      }
      tag.add(l, val);
      tag.add(r, -val);
    }

    long long point_tag(int dis) const {
      assert(0 <= dis && dis < nd);
      return tag.prefix_sum(dis + 1);
    }
  };

  struct ancestor_entry {
    int cen = -1;
    int dis = 0;
    int br = -1;
  };

  int n_ = 0;
  std::vector<std::vector<int>> g_;
  std::vector<char> rm_;
  std::vector<int> cp;
  std::vector<int> sz_;
  std::vector<std::vector<ancestor_entry>> up_;
  std::vector<bucket> a_;
  std::vector<std::vector<bucket>> br_;
  std::vector<long long> iv;

  int find_centroid(int s) {
    std::vector<int> ord;
    ord.push_back(s);
    cp[s] = -1;
    for (int i = 0; i < int(ord.size()); i++) {
      int u = ord[i];
      for (int nxt : g_[u]) {
        if (rm_[nxt] || nxt == cp[u]) {
          continue;
        }
        cp[nxt] = u;
        ord.push_back(nxt);
      }
    }

    for (int i = int(ord.size()) - 1; i >= 0; i--) {
      int u = ord[i];
      sz_[u] = 1;
      for (int nxt : g_[u]) {
        if (!rm_[nxt] && cp[nxt] == u) {
          sz_[u] += sz_[nxt];
        }
      }
    }

    int sz = int(ord.size());
    int cen = s;
    int bst = sz;
    for (int u : ord) {
      int mp = sz - sz_[u];
      for (int nxt : g_[u]) {
        if (!rm_[nxt] && cp[nxt] == u) {
          mp = std::max(mp, sz_[nxt]);
        }
      }
      if (mp < bst) {
        bst = mp;
        cen = u;
      }
    }
    return cen;
  }

  void decompose(int s) {
    int cen = find_centroid(s);
    rm_[cen] = true;
    up_[cen].push_back({cen, 0, -1});

    int mxd = 0;
    for (int v : g_[cen]) {
      if (rm_[v]) {
        continue;
      }
      int br = int(br_[cen].size());
      int bmd = 0;
      std::vector<std::tuple<int, int, int>> stk = {{v, cen, 1}};
      while (!stk.empty()) {
        auto [u, fa, dis] = stk.back();
        stk.pop_back();
        up_[u].push_back({cen, dis, br});
        bmd = std::max(bmd, dis);
        for (int nxt : g_[u]) {
          if (!rm_[nxt] && nxt != fa) {
            stk.push_back({nxt, u, dis + 1});
          }
        }
      }
      br_[cen].push_back({});
      br_[cen].back().reset(bmd + 1);
      mxd = std::max(mxd, bmd);
    }
    a_[cen].reset(mxd + 1);

    for (int v : g_[cen]) {
      if (!rm_[v]) {
        decompose(v);
      }
    }
  }

  long long prefix_contour_sum(int u, int lim) const {
    long long res = 0;
    for (const ancestor_entry &ent : up_[u]) {
      int rem = lim - ent.dis;
      res += a_[ent.cen].point_prefix(rem);
      if (ent.br != -1) {
        res -= br_[ent.cen][ent.br].point_prefix(rem);
      }
    }
    return res;
  }

public:
  explicit tree_contour(const std::vector<std::vector<int>> &g,
                        const std::vector<long long> &iv1 = {})
      : n_(int(g.size())), g_(g), rm_(n_), cp(n_), sz_(n_), up_(n_), a_(n_),
        br_(n_), iv(iv1.empty() ? std::vector<long long>(n_) : iv1) {
    assert(int(iv.size()) == n_);
    if (n_ == 0) {
      return;
    }
    decompose(0);
    for (int u = 0; u < n_; u++) {
      point_add(u, iv[u]);
    }
  }

  /// @brief Add value to one vertex for later contour-sum queries.
  void point_add(int u, long long val) {
    assert(0 <= u && u < n_);
    for (const ancestor_entry &ent : up_[u]) {
      a_[ent.cen].point_add(ent.dis, val);
      if (ent.br != -1) {
        br_[ent.cen][ent.br].point_add(ent.dis, val);
      }
    }
  }

  /// @brief Sum values at vertices whose distance from vertex is in [l,r).
  long long range_sum(int u, int l, int r) const {
    assert(0 <= u && u < n_ && l <= r);
    return prefix_contour_sum(u, r) - prefix_contour_sum(u, l);
  }

  /// @brief Add value to vertices whose distance from vertex is in [l,r).
  void range_add(int u, int l, int r, long long val) {
    assert(0 <= u && u < n_ && l <= r);
    for (const ancestor_entry &ent : up_[u]) {
      int ql = l - ent.dis;
      int qr = r - ent.dis;
      a_[ent.cen].range_add(ql, qr, val);
      if (ent.br != -1) {
        br_[ent.cen][ent.br].range_add(ql, qr, val);
      }
    }
  }

  /// @brief Return the initial value plus every contour-range addition that
  /// contains this vertex.
  long long point_get(int u) const {
    assert(0 <= u && u < n_);
    long long res = iv[u];
    for (const ancestor_entry &ent : up_[u]) {
      res += a_[ent.cen].point_tag(ent.dis);
      if (ent.br != -1) {
        res -= br_[ent.cen][ent.br].point_tag(ent.dis);
      }
    }
    return res;
  }
};

} // namespace noya

#endif // NOYA_TREE_CONTOUR_HPP
#include <algorithm>
#include <cassert>
#include <tuple>
#include <utility>
#include <vector>

/// @complexity Time: O(n log^2 n) construction and O(log^2 n) per update or
/// query. Space: O(n log n).

namespace noya {

/// @brief Maintain values indexed by tree distance. Every vertex is recorded
/// in all centroid-ancestor buckets; querying adds all buckets and subtracts
/// the bucket of the branch containing the query vertex. Fenwick trees over
/// distances support both point-add/range-sum and the dual range-add/point-get.
class tree_contour {
  struct fenwick {
    std::vector<long long> dat;

    fenwick() = default;
    explicit fenwick(int siz) : dat(siz + 1) {}

    int size() const { return int(dat.size()) - 1; }

    void add(int i, long long val) {
      assert(0 <= i && i < size());
      for (i++; i < int(dat.size()); i += i & -i) {
        dat[i] += val;
      }
    }

    long long prefix_sum(int cnt) const {
      cnt = std::clamp(cnt, 0, size());
      long long res = 0;
      for (; cnt > 0; cnt -= cnt & -cnt) {
        res += dat[cnt];
      }
      return res;
    }
  };

  struct bucket {
    int nd = 0;
    fenwick bit;
    fenwick tag;

    void reset(int cnt) {
      nd = cnt;
      bit = fenwick(cnt);
      tag = fenwick(cnt + 1);
    }

    void point_add(int dis, long long val) { bit.add(dis, val); }

    long long point_prefix(int lim) const { return bit.prefix_sum(lim); }

    void range_add(int l, int r, long long val) {
      l = std::clamp(l, 0, nd);
      r = std::clamp(r, 0, nd);
      if (l >= r) {
        return;
      }
      tag.add(l, val);
      tag.add(r, -val);
    }

    long long point_tag(int dis) const {
      assert(0 <= dis && dis < nd);
      return tag.prefix_sum(dis + 1);
    }
  };

  struct ancestor_entry {
    int cen = -1;
    int dis = 0;
    int br = -1;
  };

  int n_ = 0;
  std::vector<std::vector<int>> g_;
  std::vector<char> rm_;
  std::vector<int> cp;
  std::vector<int> sz_;
  std::vector<std::vector<ancestor_entry>> up_;
  std::vector<bucket> a_;
  std::vector<std::vector<bucket>> br_;
  std::vector<long long> iv;

  int find_centroid(int s) {
    std::vector<int> ord;
    ord.push_back(s);
    cp[s] = -1;
    for (int i = 0; i < int(ord.size()); i++) {
      int u = ord[i];
      for (int nxt : g_[u]) {
        if (rm_[nxt] || nxt == cp[u]) {
          continue;
        }
        cp[nxt] = u;
        ord.push_back(nxt);
      }
    }

    for (int i = int(ord.size()) - 1; i >= 0; i--) {
      int u = ord[i];
      sz_[u] = 1;
      for (int nxt : g_[u]) {
        if (!rm_[nxt] && cp[nxt] == u) {
          sz_[u] += sz_[nxt];
        }
      }
    }

    int sz = int(ord.size());
    int cen = s;
    int bst = sz;
    for (int u : ord) {
      int mp = sz - sz_[u];
      for (int nxt : g_[u]) {
        if (!rm_[nxt] && cp[nxt] == u) {
          mp = std::max(mp, sz_[nxt]);
        }
      }
      if (mp < bst) {
        bst = mp;
        cen = u;
      }
    }
    return cen;
  }

  void decompose(int s) {
    int cen = find_centroid(s);
    rm_[cen] = true;
    up_[cen].push_back({cen, 0, -1});

    int mxd = 0;
    for (int v : g_[cen]) {
      if (rm_[v]) {
        continue;
      }
      int br = int(br_[cen].size());
      int bmd = 0;
      std::vector<std::tuple<int, int, int>> stk = {{v, cen, 1}};
      while (!stk.empty()) {
        auto [u, fa, dis] = stk.back();
        stk.pop_back();
        up_[u].push_back({cen, dis, br});
        bmd = std::max(bmd, dis);
        for (int nxt : g_[u]) {
          if (!rm_[nxt] && nxt != fa) {
            stk.push_back({nxt, u, dis + 1});
          }
        }
      }
      br_[cen].push_back({});
      br_[cen].back().reset(bmd + 1);
      mxd = std::max(mxd, bmd);
    }
    a_[cen].reset(mxd + 1);

    for (int v : g_[cen]) {
      if (!rm_[v]) {
        decompose(v);
      }
    }
  }

  long long prefix_contour_sum(int u, int lim) const {
    long long res = 0;
    for (const ancestor_entry &ent : up_[u]) {
      int rem = lim - ent.dis;
      res += a_[ent.cen].point_prefix(rem);
      if (ent.br != -1) {
        res -= br_[ent.cen][ent.br].point_prefix(rem);
      }
    }
    return res;
  }

public:
  explicit tree_contour(const std::vector<std::vector<int>> &g,
                        const std::vector<long long> &iv1 = {})
      : n_(int(g.size())), g_(g), rm_(n_), cp(n_), sz_(n_), up_(n_), a_(n_),
        br_(n_), iv(iv1.empty() ? std::vector<long long>(n_) : iv1) {
    assert(int(iv.size()) == n_);
    if (n_ == 0) {
      return;
    }
    decompose(0);
    for (int u = 0; u < n_; u++) {
      point_add(u, iv[u]);
    }
  }

  /// @brief Add value to one vertex for later contour-sum queries.
  void point_add(int u, long long val) {
    assert(0 <= u && u < n_);
    for (const ancestor_entry &ent : up_[u]) {
      a_[ent.cen].point_add(ent.dis, val);
      if (ent.br != -1) {
        br_[ent.cen][ent.br].point_add(ent.dis, val);
      }
    }
  }

  /// @brief Sum values at vertices whose distance from vertex is in [l,r).
  long long range_sum(int u, int l, int r) const {
    assert(0 <= u && u < n_ && l <= r);
    return prefix_contour_sum(u, r) - prefix_contour_sum(u, l);
  }

  /// @brief Add value to vertices whose distance from vertex is in [l,r).
  void range_add(int u, int l, int r, long long val) {
    assert(0 <= u && u < n_ && l <= r);
    for (const ancestor_entry &ent : up_[u]) {
      int ql = l - ent.dis;
      int qr = r - ent.dis;
      a_[ent.cen].range_add(ql, qr, val);
      if (ent.br != -1) {
        br_[ent.cen][ent.br].range_add(ql, qr, val);
      }
    }
  }

  /// @brief Return the initial value plus every contour-range addition that
  /// contains this vertex.
  long long point_get(int u) const {
    assert(0 <= u && u < n_);
    long long res = iv[u];
    for (const ancestor_entry &ent : up_[u]) {
      res += a_[ent.cen].point_tag(ent.dis);
      if (ent.br != -1) {
        res -= br_[ent.cen][ent.br].point_tag(ent.dis);
      }
    }
    return res;
  }
};

} // namespace noya