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¶
#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