Skip to content

persistent_affine_segtree.hpp

SECTIONData Structure INCLUDEnoya/persistent_affine_segtree.hpp

Persistent lazy segment tree for range affine transforms and sums. Every update clones only its two boundary paths. A range copy descends in a destination version and a source version together, replacing fully covered destination nodes by the corresponding immutable source nodes. Lazy tags are pushed only into freshly cloned nodes, so every older root remains unchanged.

Verified by persistent_range_affine_range_sum.

保留每次区间仿射修改后的线段树版本,并查询任一版本的区间和。

Implementation

View on GitHub

#ifndef NOYA_PERSISTENT_AFFINE_SEGTREE_HPP
#define NOYA_PERSISTENT_AFFINE_SEGTREE_HPP 1

/// @complexity Time: O(log n) per range affine update, range copy, or query.
/// Space: O(n + q log n) nodes for q derived versions.

#include <cassert>
#include <vector>

namespace noya {

/// @brief Persistent lazy segment tree for range affine transforms and sums.
/// Every update clones only its two boundary paths. A range copy descends in a
/// destination version and a source version together, replacing fully covered
/// destination nodes by the corresponding immutable source nodes. Lazy tags
/// are pushed only into freshly cloned nodes, so every older root remains
/// unchanged.
template <class T> class persistent_affine_segtree {
public:
  persistent_affine_segtree() = default;

  explicit persistent_affine_segtree(const std::vector<T> &values) {
    build(values);
  }

  void build(const std::vector<T> &values) {
    assert(!values.empty());
    size_ = int(values.size());
    nodes_.clear();
    nodes_.reserve(size_ * 2);
    initial_root_ = build_node(values, 0, size_);
  }

  void reserve(int node_count) { nodes_.reserve(node_count); }

  int initial_root() const { return initial_root_; }

  int apply(int root, int left, int right, T multiplier, T addition) {
    assert(valid_root(root));
    assert(0 <= left && left <= right && right <= size_);
    return apply_node(root, 0, size_, left, right, multiplier, addition);
  }

  int copy_range(int destination_root, int source_root, int left, int right) {
    assert(valid_root(destination_root) && valid_root(source_root));
    assert(0 <= left && left <= right && right <= size_);
    return copy_node(destination_root, source_root, 0, size_, left, right);
  }

  T prod(int root, int left, int right) const {
    assert(valid_root(root));
    assert(0 <= left && left <= right && right <= size_);
    return prod_node(root, 0, size_, left, right, T(1), T(0));
  }

private:
  struct node {
    T sum{};
    T multiplier = T(1);
    T addition{};
    int left = -1;
    int right = -1;
  };

  int size_ = 0;
  int initial_root_ = -1;
  std::vector<node> nodes_;

  bool valid_root(int root) const {
    return 0 <= root && root < int(nodes_.size());
  }

  int make_node(const node &value) {
    nodes_.push_back(value);
    return int(nodes_.size()) - 1;
  }

  int clone_node(int id) { return make_node(nodes_[id]); }

  int build_node(const std::vector<T> &values, int left, int right) {
    if (right - left == 1) {
      return make_node(node{.sum = values[left]});
    }
    int middle = (left + right) / 2;
    int first = build_node(values, left, middle);
    int second = build_node(values, middle, right);
    return make_node(node{.sum = nodes_[first].sum + nodes_[second].sum,
                          .left = first,
                          .right = second});
  }

  void apply_here(int id, int length, T multiplier, T addition) {
    nodes_[id].sum = multiplier * nodes_[id].sum + addition * T(length);
    nodes_[id].multiplier *= multiplier;
    nodes_[id].addition = multiplier * nodes_[id].addition + addition;
  }

  void push_cloned(int id, int left_length, int right_length) {
    if (nodes_[id].left < 0) {
      return;
    }
    T multiplier = nodes_[id].multiplier;
    T addition = nodes_[id].addition;
    if (multiplier == T(1) && addition == T(0)) {
      return;
    }
    int first = clone_node(nodes_[id].left);
    int second = clone_node(nodes_[id].right);
    nodes_[id].left = first;
    nodes_[id].right = second;
    apply_here(first, left_length, multiplier, addition);
    apply_here(second, right_length, multiplier, addition);
    nodes_[id].multiplier = T(1);
    nodes_[id].addition = T(0);
  }

  void pull(int id) {
    nodes_[id].sum =
        nodes_[nodes_[id].left].sum + nodes_[nodes_[id].right].sum;
  }

  int apply_node(int id, int left, int right, int query_left, int query_right,
                 T multiplier, T addition) {
    if (query_right <= left || right <= query_left) {
      return id;
    }
    int result = clone_node(id);
    if (query_left <= left && right <= query_right) {
      apply_here(result, right - left, multiplier, addition);
      return result;
    }
    int middle = (left + right) / 2;
    push_cloned(result, middle - left, right - middle);
    int first = apply_node(nodes_[result].left, left, middle, query_left,
                           query_right, multiplier, addition);
    int second = apply_node(nodes_[result].right, middle, right, query_left,
                            query_right, multiplier, addition);
    nodes_[result].left = first;
    nodes_[result].right = second;
    pull(result);
    return result;
  }

  int normalized_clone(int id, int left_length, int right_length) {
    int result = clone_node(id);
    push_cloned(result, left_length, right_length);
    return result;
  }

  int copy_node(int destination, int source, int left, int right,
                int query_left, int query_right) {
    if (query_right <= left || right <= query_left) {
      return destination;
    }
    if (query_left <= left && right <= query_right) {
      return source;
    }
    int middle = (left + right) / 2;
    int destination_clone =
        normalized_clone(destination, middle - left, right - middle);
    int source_clone =
        normalized_clone(source, middle - left, right - middle);
    int first = copy_node(nodes_[destination_clone].left,
                          nodes_[source_clone].left, left, middle, query_left,
                          query_right);
    int second = copy_node(nodes_[destination_clone].right,
                           nodes_[source_clone].right, middle, right,
                           query_left, query_right);
    nodes_[destination_clone].left = first;
    nodes_[destination_clone].right = second;
    pull(destination_clone);
    return destination_clone;
  }

  T prod_node(int id, int left, int right, int query_left, int query_right,
              T ancestor_multiplier, T ancestor_addition) const {
    if (query_right <= left || right <= query_left) {
      return T(0);
    }
    if (query_left <= left && right <= query_right) {
      return ancestor_multiplier * nodes_[id].sum +
             ancestor_addition * T(right - left);
    }
    T multiplier = ancestor_multiplier * nodes_[id].multiplier;
    T addition =
        ancestor_multiplier * nodes_[id].addition + ancestor_addition;
    int middle = (left + right) / 2;
    return prod_node(nodes_[id].left, left, middle, query_left, query_right,
                     multiplier, addition) +
           prod_node(nodes_[id].right, middle, right, query_left, query_right,
                     multiplier, addition);
  }
};

} // namespace noya

#endif // NOYA_PERSISTENT_AFFINE_SEGTREE_HPP
#include <cassert>
#include <vector>

/// @complexity Time: O(log n) per range affine update, range copy, or query.
/// Space: O(n + q log n) nodes for q derived versions.

namespace noya {

/// @brief Persistent lazy segment tree for range affine transforms and sums.
/// Every update clones only its two boundary paths. A range copy descends in a
/// destination version and a source version together, replacing fully covered
/// destination nodes by the corresponding immutable source nodes. Lazy tags
/// are pushed only into freshly cloned nodes, so every older root remains
/// unchanged.
template <class T> class persistent_affine_segtree {
public:
  persistent_affine_segtree() = default;

  explicit persistent_affine_segtree(const std::vector<T> &values) {
    build(values);
  }

  void build(const std::vector<T> &values) {
    assert(!values.empty());
    size_ = int(values.size());
    nodes_.clear();
    nodes_.reserve(size_ * 2);
    initial_root_ = build_node(values, 0, size_);
  }

  void reserve(int node_count) { nodes_.reserve(node_count); }

  int initial_root() const { return initial_root_; }

  int apply(int root, int left, int right, T multiplier, T addition) {
    assert(valid_root(root));
    assert(0 <= left && left <= right && right <= size_);
    return apply_node(root, 0, size_, left, right, multiplier, addition);
  }

  int copy_range(int destination_root, int source_root, int left, int right) {
    assert(valid_root(destination_root) && valid_root(source_root));
    assert(0 <= left && left <= right && right <= size_);
    return copy_node(destination_root, source_root, 0, size_, left, right);
  }

  T prod(int root, int left, int right) const {
    assert(valid_root(root));
    assert(0 <= left && left <= right && right <= size_);
    return prod_node(root, 0, size_, left, right, T(1), T(0));
  }

private:
  struct node {
    T sum{};
    T multiplier = T(1);
    T addition{};
    int left = -1;
    int right = -1;
  };

  int size_ = 0;
  int initial_root_ = -1;
  std::vector<node> nodes_;

  bool valid_root(int root) const {
    return 0 <= root && root < int(nodes_.size());
  }

  int make_node(const node &value) {
    nodes_.push_back(value);
    return int(nodes_.size()) - 1;
  }

  int clone_node(int id) { return make_node(nodes_[id]); }

  int build_node(const std::vector<T> &values, int left, int right) {
    if (right - left == 1) {
      return make_node(node{.sum = values[left]});
    }
    int middle = (left + right) / 2;
    int first = build_node(values, left, middle);
    int second = build_node(values, middle, right);
    return make_node(node{.sum = nodes_[first].sum + nodes_[second].sum,
                          .left = first,
                          .right = second});
  }

  void apply_here(int id, int length, T multiplier, T addition) {
    nodes_[id].sum = multiplier * nodes_[id].sum + addition * T(length);
    nodes_[id].multiplier *= multiplier;
    nodes_[id].addition = multiplier * nodes_[id].addition + addition;
  }

  void push_cloned(int id, int left_length, int right_length) {
    if (nodes_[id].left < 0) {
      return;
    }
    T multiplier = nodes_[id].multiplier;
    T addition = nodes_[id].addition;
    if (multiplier == T(1) && addition == T(0)) {
      return;
    }
    int first = clone_node(nodes_[id].left);
    int second = clone_node(nodes_[id].right);
    nodes_[id].left = first;
    nodes_[id].right = second;
    apply_here(first, left_length, multiplier, addition);
    apply_here(second, right_length, multiplier, addition);
    nodes_[id].multiplier = T(1);
    nodes_[id].addition = T(0);
  }

  void pull(int id) {
    nodes_[id].sum =
        nodes_[nodes_[id].left].sum + nodes_[nodes_[id].right].sum;
  }

  int apply_node(int id, int left, int right, int query_left, int query_right,
                 T multiplier, T addition) {
    if (query_right <= left || right <= query_left) {
      return id;
    }
    int result = clone_node(id);
    if (query_left <= left && right <= query_right) {
      apply_here(result, right - left, multiplier, addition);
      return result;
    }
    int middle = (left + right) / 2;
    push_cloned(result, middle - left, right - middle);
    int first = apply_node(nodes_[result].left, left, middle, query_left,
                           query_right, multiplier, addition);
    int second = apply_node(nodes_[result].right, middle, right, query_left,
                            query_right, multiplier, addition);
    nodes_[result].left = first;
    nodes_[result].right = second;
    pull(result);
    return result;
  }

  int normalized_clone(int id, int left_length, int right_length) {
    int result = clone_node(id);
    push_cloned(result, left_length, right_length);
    return result;
  }

  int copy_node(int destination, int source, int left, int right,
                int query_left, int query_right) {
    if (query_right <= left || right <= query_left) {
      return destination;
    }
    if (query_left <= left && right <= query_right) {
      return source;
    }
    int middle = (left + right) / 2;
    int destination_clone =
        normalized_clone(destination, middle - left, right - middle);
    int source_clone =
        normalized_clone(source, middle - left, right - middle);
    int first = copy_node(nodes_[destination_clone].left,
                          nodes_[source_clone].left, left, middle, query_left,
                          query_right);
    int second = copy_node(nodes_[destination_clone].right,
                           nodes_[source_clone].right, middle, right,
                           query_left, query_right);
    nodes_[destination_clone].left = first;
    nodes_[destination_clone].right = second;
    pull(destination_clone);
    return destination_clone;
  }

  T prod_node(int id, int left, int right, int query_left, int query_right,
              T ancestor_multiplier, T ancestor_addition) const {
    if (query_right <= left || right <= query_left) {
      return T(0);
    }
    if (query_left <= left && right <= query_right) {
      return ancestor_multiplier * nodes_[id].sum +
             ancestor_addition * T(right - left);
    }
    T multiplier = ancestor_multiplier * nodes_[id].multiplier;
    T addition =
        ancestor_multiplier * nodes_[id].addition + ancestor_addition;
    int middle = (left + right) / 2;
    return prod_node(nodes_[id].left, left, middle, query_left, query_right,
                     multiplier, addition) +
           prod_node(nodes_[id].right, middle, right, query_left, query_right,
                     multiplier, addition);
  }
};

} // namespace noya