Luzhiled's Library

This documentation is automatically generated by competitive-verifier/competitive-verifier

View the Project on GitHub ei1333/library

:heavy_check_mark: test/unittest/priority-sum-structure.test.cpp

Depends on

Code

// competitive-verifier: STANDALONE

#include "../../structure/others/priority-sum-structure.hpp"

#include <algorithm>
#include <cassert>
#include <random>
#include <vector>

int main() {
  MaximumSum<long long> maximum(0);
  MinimumSum<long long> minimum(0);
  std::vector<long long> values;
  std::mt19937 rng(123456789);

  for (int iteration = 0; iteration < 5000; iteration++) {
    if (values.empty() || rng() % 3 != 0) {
      long long value = (int)(rng() % 201) - 100;
      values.push_back(value);
      maximum.insert(value);
      minimum.insert(value);
    } else {
      int index = rng() % values.size();
      maximum.erase(values[index]);
      minimum.erase(values[index]);
      values.erase(values.begin() + index);
    }

    std::sort(values.begin(), values.end());
    std::size_t k = values.empty() ? 0 : rng() % (values.size() + 1);
    maximum.set_k(k);
    minimum.set_k(k);

    long long expected_minimum = 0;
    long long expected_maximum = 0;
    for (std::size_t i = 0; i < k; i++) {
      expected_minimum += values[i];
      expected_maximum += values[values.size() - 1 - i];
    }
    assert(minimum.query() == expected_minimum);
    assert(maximum.query() == expected_maximum);
    assert(minimum.size() == values.size());
    assert(maximum.size() == values.size());
    if (k > 0) {
      assert(minimum.kth_element() == values[k - 1]);
      assert(maximum.kth_element() == values[values.size() - k]);
    }
  }
}
#line 1 "test/unittest/priority-sum-structure.test.cpp"
// competitive-verifier: STANDALONE

#line 2 "structure/others/priority-sum-structure.hpp"

#include <cassert>
#include <cstddef>
#include <functional>
#include <queue>
#include <vector>

template <typename T, typename Compare = std::less<T>,
          typename RCompare = std::greater<T>>
struct PrioritySumStructure {
  std::size_t k;
  T sum;

  std::priority_queue<T, std::vector<T>, Compare> in, d_in;
  std::priority_queue<T, std::vector<T>, RCompare> out, d_out;

  PrioritySumStructure(int k) : k(k), sum(0) {}

  void modify() {
    while (in.size() - d_in.size() < k && !out.empty()) {
      auto p = out.top();
      out.pop();
      if (!d_out.empty() && p == d_out.top()) {
        d_out.pop();
      } else {
        sum += p;
        in.emplace(p);
      }
    }
    while (in.size() - d_in.size() > k) {
      auto p = in.top();
      in.pop();
      if (!d_in.empty() && p == d_in.top()) {
        d_in.pop();
      } else {
        sum -= p;
        out.emplace(p);
      }
    }
    while (!d_in.empty() && in.top() == d_in.top()) {
      in.pop();
      d_in.pop();
    }
  }

  T query() const { return sum; }

  T kth_element() {
    assert(0 < k && k <= size());
    modify();
    return in.top();
  }

  void insert(T x) {
    in.emplace(x);
    sum += x;
    modify();
  }

  void erase(T x) {
    assert(size());
    if (!in.empty() && in.top() == x) {
      sum -= x;
      in.pop();
    } else if (!in.empty() && RCompare()(in.top(), x)) {
      sum -= x;
      d_in.emplace(x);
    } else {
      d_out.emplace(x);
    }
    modify();
  }

  void set_k(std::size_t kk) {
    k = kk;
    modify();
  }

  std::size_t get_k() const { return k; }

  std::size_t size() const {
    return in.size() + out.size() - d_in.size() - d_out.size();
  }
};

template <typename T>
using MaximumSum = PrioritySumStructure<T, std::greater<T>, std::less<T>>;

template <typename T>
using MinimumSum = PrioritySumStructure<T, std::less<T>, std::greater<T>>;
#line 4 "test/unittest/priority-sum-structure.test.cpp"

#include <algorithm>
#line 7 "test/unittest/priority-sum-structure.test.cpp"
#include <random>
#line 9 "test/unittest/priority-sum-structure.test.cpp"

int main() {
  MaximumSum<long long> maximum(0);
  MinimumSum<long long> minimum(0);
  std::vector<long long> values;
  std::mt19937 rng(123456789);

  for (int iteration = 0; iteration < 5000; iteration++) {
    if (values.empty() || rng() % 3 != 0) {
      long long value = (int)(rng() % 201) - 100;
      values.push_back(value);
      maximum.insert(value);
      minimum.insert(value);
    } else {
      int index = rng() % values.size();
      maximum.erase(values[index]);
      minimum.erase(values[index]);
      values.erase(values.begin() + index);
    }

    std::sort(values.begin(), values.end());
    std::size_t k = values.empty() ? 0 : rng() % (values.size() + 1);
    maximum.set_k(k);
    minimum.set_k(k);

    long long expected_minimum = 0;
    long long expected_maximum = 0;
    for (std::size_t i = 0; i < k; i++) {
      expected_minimum += values[i];
      expected_maximum += values[values.size() - 1 - i];
    }
    assert(minimum.query() == expected_minimum);
    assert(maximum.query() == expected_maximum);
    assert(minimum.size() == values.size());
    assert(maximum.size() == values.size());
    if (k > 0) {
      assert(minimum.kth_element() == values[k - 1]);
      assert(maximum.kth_element() == values[values.size() - k]);
    }
  }
}
Back to top page