Skip to content

top_k_sum.hpp

SECTIONData Structure INCLUDEnoya/top_k_sum.hpp

Dynamic multiset maintaining the sum of the k largest values in O(log n) per insert, erase, or k change.

动态维护集合中最大的或最小的 k 个元素之和;适合排名阈值变化、选取固定数量最优元素的题目。

Implementation

View on GitHub

#ifndef NOYA_TOP_K_SUM_HPP
#define NOYA_TOP_K_SUM_HPP 1

/// @complexity Time: O(log n) insert/erase/set-k and O(1) sum query.
/// Space: O(n).

#include <algorithm>
#include <cassert>
#include <set>
#include <utility>

namespace noya {

/// @brief Dynamic multiset maintaining the sum of the k largest values in
/// O(log n) per insert, erase, or k change.
template <class T> struct top_k_sum {
  using entry = std::pair<T, int>;
  int k = 0;
  int next_id = 0;
  T sum{};
  std::set<entry> selected;
  std::set<entry> remaining;

  top_k_sum() = default;
  explicit top_k_sum(int k_) : k(k_) { assert(k >= 0); }

  /// @brief Insert value and return a stable handle for exact erasure.
  int insert(const T &value) {
    int id = next_id++;
    remaining.emplace(value, id);
    rebalance();
    return id;
  }

  /// @brief Erase the entry identified by (value, handle).
  void erase(const T &value, int id) {
    entry target{value, id};
    if (auto iterator = selected.find(target); iterator != selected.end()) {
      sum -= iterator->first;
      selected.erase(iterator);
    } else {
      auto remaining_iterator = remaining.find(target);
      assert(remaining_iterator != remaining.end());
      remaining.erase(remaining_iterator);
    }
    rebalance();
  }

  void set_k(int k_) {
    assert(k_ >= 0);
    k = k_;
    rebalance();
  }

  const T &query() const { return sum; }
  int size() const { return int(selected.size() + remaining.size()); }

private:
  void rebalance() {
    int target = std::min(k, size());
    while (int(selected.size()) < target) {
      auto iterator = std::prev(remaining.end());
      sum += iterator->first;
      selected.insert(*iterator);
      remaining.erase(iterator);
    }
    while (int(selected.size()) > target) {
      auto iterator = selected.begin();
      sum -= iterator->first;
      remaining.insert(*iterator);
      selected.erase(iterator);
    }
    while (!selected.empty() && !remaining.empty() &&
           *selected.begin() < *remaining.rbegin()) {
      entry small = *selected.begin();
      entry large = *remaining.rbegin();
      selected.erase(selected.begin());
      remaining.erase(std::prev(remaining.end()));
      selected.insert(large);
      remaining.insert(small);
      sum += large.first - small.first;
    }
  }
};

} // namespace noya

#endif // NOYA_TOP_K_SUM_HPP
#include <algorithm>
#include <cassert>
#include <set>
#include <utility>

/// @complexity Time: O(log n) insert/erase/set-k and O(1) sum query.
/// Space: O(n).

namespace noya {

/// @brief Dynamic multiset maintaining the sum of the k largest values in
/// O(log n) per insert, erase, or k change.
template <class T> struct top_k_sum {
  using entry = std::pair<T, int>;
  int k = 0;
  int next_id = 0;
  T sum{};
  std::set<entry> selected;
  std::set<entry> remaining;

  top_k_sum() = default;
  explicit top_k_sum(int k_) : k(k_) { assert(k >= 0); }

  /// @brief Insert value and return a stable handle for exact erasure.
  int insert(const T &value) {
    int id = next_id++;
    remaining.emplace(value, id);
    rebalance();
    return id;
  }

  /// @brief Erase the entry identified by (value, handle).
  void erase(const T &value, int id) {
    entry target{value, id};
    if (auto iterator = selected.find(target); iterator != selected.end()) {
      sum -= iterator->first;
      selected.erase(iterator);
    } else {
      auto remaining_iterator = remaining.find(target);
      assert(remaining_iterator != remaining.end());
      remaining.erase(remaining_iterator);
    }
    rebalance();
  }

  void set_k(int k_) {
    assert(k_ >= 0);
    k = k_;
    rebalance();
  }

  const T &query() const { return sum; }
  int size() const { return int(selected.size() + remaining.size()); }

private:
  void rebalance() {
    int target = std::min(k, size());
    while (int(selected.size()) < target) {
      auto iterator = std::prev(remaining.end());
      sum += iterator->first;
      selected.insert(*iterator);
      remaining.erase(iterator);
    }
    while (int(selected.size()) > target) {
      auto iterator = selected.begin();
      sum -= iterator->first;
      remaining.insert(*iterator);
      selected.erase(iterator);
    }
    while (!selected.empty() && !remaining.empty() &&
           *selected.begin() < *remaining.rbegin()) {
      entry small = *selected.begin();
      entry large = *remaining.rbegin();
      selected.erase(selected.begin());
      remaining.erase(std::prev(remaining.end()));
      selected.insert(large);
      remaining.insert(small);
      sum += large.first - small.first;
    }
  }
};

} // namespace noya