m1une's library

This documentation is automatically generated by online-judge-tools/verification-helper

View on GitHub

:heavy_check_mark: Static Rectangle Sum
(ds/range_query/static_rectangle_sum.hpp)

Overview

StaticRectangleSum<X, Y, Sum> stores an immutable set of weighted points and returns the total weight inside axis-aligned half-open rectangles:

\[[x_l,x_r)\times[y_l,y_r).\]

Points are sorted by their x-coordinate. Y-coordinates are compressed and queried with a weighted wavelet matrix. Multiple input points may have identical coordinates; every weight is included independently.

Requirements

X and Y must provide a strict weak ordering through <. Sum{} must be the additive identity, and Sum must support addition and subtraction. All weight sums must fit in Sum.

Interface

template <class X, class Y = X, class Sum = long long>
class StaticRectangleSum {
public:
    using x_type = X;
    using y_type = Y;
    using sum_type = Sum;
    using weighted_point_type = std::tuple<X, Y, Sum>;

    StaticRectangleSum();

    explicit StaticRectangleSum(
        const std::vector<weighted_point_type>& points
    );

    StaticRectangleSum(
        const std::vector<X>& x_coordinates,
        const std::vector<Y>& y_coordinates,
        const std::vector<Sum>& weights
    );

    void build(std::vector<weighted_point_type> points);

    void build(
        const std::vector<X>& x_coordinates,
        const std::vector<Y>& y_coordinates,
        const std::vector<Sum>& weights
    );

    int size() const;
    bool empty() const;

    Sum sum(
        const X& left,
        const X& right,
        const Y& lower,
        const Y& upper
    ) const;
};

Operations

Method Description Complexity
StaticRectangleSum() Constructs an empty structure. $O(1)$
explicit StaticRectangleSum(const std::vector<weighted_point_type>& points) Builds from (x, y, weight) tuples. $O(N\log N + NB)$
StaticRectangleSum(const std::vector<X>& xs, const std::vector<Y>& ys, const std::vector<Sum>& weights) Builds from parallel coordinate and weight vectors. $O(N\log N + NB)$
void build(std::vector<weighted_point_type> points) Replaces the stored points. $O(N\log N + NB)$
void build(const std::vector<X>& xs, const std::vector<Y>& ys, const std::vector<Sum>& weights) Replaces the stored points from parallel vectors. $O(N\log N + NB)$
int size() const Returns the number of input points, including duplicates. $O(1)$
bool empty() const Returns whether no points are stored. $O(1)$
Sum sum(const X& left, const X& right, const Y& lower, const Y& upper) const Returns the weight sum in [left,right) x [lower,upper). $O(\log N+B)$

Here B = 32, the bit width of the compressed y-rank type. Construction uses $O(NB)$ memory. Calls to sum do not mutate the structure. Empty rectangles and rectangles outside all stored coordinates return Sum{}.

Example

#include "ds/range_query/static_rectangle_sum.hpp"

#include <cassert>
#include <tuple>
#include <vector>

int main() {
    using Query = m1une::ds::StaticRectangleSum<int, int, long long>;
    std::vector<Query::weighted_point_type> points;
    points.emplace_back(1, 2, 5);
    points.emplace_back(3, 4, 7);
    points.emplace_back(1, 2, 11);

    Query query(points);
    assert(query.sum(1, 2, 2, 3) == 16);
    assert(query.sum(0, 4, 0, 5) == 23);
}

Depends on

Verified with

Code

#ifndef M1UNE_STATIC_RECTANGLE_SUM_HPP
#define M1UNE_STATIC_RECTANGLE_SUM_HPP 1

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

#include "../wavelet_matrix/wavelet_matrix_sum.hpp"

namespace m1une {
namespace ds {

template <class X, class Y = X, class Sum = long long>
class StaticRectangleSum {
   public:
    using x_type = X;
    using y_type = Y;
    using sum_type = Sum;
    using weighted_point_type = std::tuple<X, Y, Sum>;

   private:
    std::vector<X> _x_coordinates;
    std::vector<Y> _y_coordinates;
    WaveletMatrixSum<int, Sum> _matrix;

   public:
    StaticRectangleSum() = default;

    explicit StaticRectangleSum(
        const std::vector<weighted_point_type>& points
    ) {
        build(points);
    }

    StaticRectangleSum(
        const std::vector<X>& x_coordinates,
        const std::vector<Y>& y_coordinates,
        const std::vector<Sum>& weights
    ) {
        build(x_coordinates, y_coordinates, weights);
    }

    void build(std::vector<weighted_point_type> points) {
        std::sort(
            points.begin(),
            points.end(),
            [](const weighted_point_type& first,
               const weighted_point_type& second) {
                if (std::get<0>(first) < std::get<0>(second)) return true;
                if (std::get<0>(second) < std::get<0>(first)) return false;
                return std::get<1>(first) < std::get<1>(second);
            }
        );

        const int n = int(points.size());
        _x_coordinates.resize(n);
        _y_coordinates.clear();
        _y_coordinates.reserve(n);
        for (const auto& point : points) {
            _y_coordinates.push_back(std::get<1>(point));
        }
        std::sort(_y_coordinates.begin(), _y_coordinates.end());
        _y_coordinates.erase(
            std::unique(
                _y_coordinates.begin(),
                _y_coordinates.end(),
                [](const Y& first, const Y& second) {
                    return !(first < second) && !(second < first);
                }
            ),
            _y_coordinates.end()
        );

        std::vector<int> y_rank(n);
        std::vector<Sum> weights(n);
        for (int index = 0; index < n; index++) {
            _x_coordinates[index] = std::get<0>(points[index]);
            y_rank[index] = int(
                std::lower_bound(
                    _y_coordinates.begin(),
                    _y_coordinates.end(),
                    std::get<1>(points[index])
                ) - _y_coordinates.begin()
            );
            weights[index] = std::get<2>(points[index]);
        }
        _matrix = WaveletMatrixSum<int, Sum>(y_rank, weights);
    }

    void build(
        const std::vector<X>& x_coordinates,
        const std::vector<Y>& y_coordinates,
        const std::vector<Sum>& weights
    ) {
        assert(x_coordinates.size() == y_coordinates.size());
        assert(x_coordinates.size() == weights.size());
        std::vector<weighted_point_type> points;
        points.reserve(x_coordinates.size());
        for (int index = 0; index < int(x_coordinates.size()); index++) {
            points.emplace_back(
                x_coordinates[index],
                y_coordinates[index],
                weights[index]
            );
        }
        build(std::move(points));
    }

    int size() const {
        return int(_x_coordinates.size());
    }

    bool empty() const {
        return _x_coordinates.empty();
    }

    Sum sum(
        const X& left,
        const X& right,
        const Y& lower,
        const Y& upper
    ) const {
        assert(!(right < left));
        assert(!(upper < lower));
        int x_left = int(
            std::lower_bound(
                _x_coordinates.begin(),
                _x_coordinates.end(),
                left
            ) - _x_coordinates.begin()
        );
        int x_right = int(
            std::lower_bound(
                _x_coordinates.begin(),
                _x_coordinates.end(),
                right
            ) - _x_coordinates.begin()
        );
        int y_lower = int(
            std::lower_bound(
                _y_coordinates.begin(),
                _y_coordinates.end(),
                lower
            ) - _y_coordinates.begin()
        );
        int y_upper = int(
            std::lower_bound(
                _y_coordinates.begin(),
                _y_coordinates.end(),
                upper
            ) - _y_coordinates.begin()
        );
        return _matrix.range_sum(x_left, x_right, y_lower, y_upper);
    }
};

}  // namespace ds
}  // namespace m1une

#endif  // M1UNE_STATIC_RECTANGLE_SUM_HPP
#line 1 "ds/range_query/static_rectangle_sum.hpp"



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

#line 1 "ds/wavelet_matrix/wavelet_matrix_sum.hpp"



#line 5 "ds/wavelet_matrix/wavelet_matrix_sum.hpp"
#include <bit>
#line 7 "ds/wavelet_matrix/wavelet_matrix_sum.hpp"
#include <concepts>
#include <cstdint>
#include <limits>
#include <optional>
#include <type_traits>
#line 14 "ds/wavelet_matrix/wavelet_matrix_sum.hpp"

#if defined(__AVX2__) || defined(__BMI2__)
#include <immintrin.h>
#endif

namespace m1une {
namespace ds {

// A static wavelet matrix with additive weights.
// By default, each value is also used as its weight.
template <std::integral T, typename Sum = T>
requires(!std::same_as<std::remove_cv_t<T>, bool>)
struct WaveletMatrixSum {
    using value_type = T;
    using sum_type = Sum;
    using unsigned_type = std::make_unsigned_t<T>;

   private:
    static constexpr int value_bit_width =
        std::numeric_limits<unsigned_type>::digits;
    static constexpr unsigned_type sign_mask = [] {
        if constexpr (std::signed_integral<T>) {
            return unsigned_type(1) << (value_bit_width - 1);
        } else {
            return unsigned_type(0);
        }
    }();

    struct BitVector {
        std::vector<std::uint64_t> bits;
        std::vector<int> prefix;

        BitVector() = default;

        explicit BitVector(int n)
            : bits(((std::size_t(n) + 63) >> 6) + 1, 0),
              prefix(bits.size(), 0) {}

        void build() {
            for (std::size_t i = 0; i + 1 < bits.size(); i++) {
                prefix[i + 1] = prefix[i] + std::popcount(bits[i]);
            }
        }

        bool get(int p) const {
            return (bits[std::size_t(p) >> 6] >> (p & 63)) & 1;
        }

        int rank1(int r) const {
            std::size_t word = std::size_t(r) >> 6;
            int offset = r & 63;
            int result = prefix[word];
#if defined(__BMI2__)
            result += std::popcount(
                _bzhi_u64(bits[word], static_cast<unsigned int>(offset))
            );
#else
            if (offset != 0) {
                result += std::popcount(
                    bits[word] & ((std::uint64_t(1) << offset) - 1)
                );
            }
#endif
            return result;
        }
    };

    int _n;
    int _log;
    unsigned_type _key_prefix;
    unsigned_type _min_key;
    unsigned_type _max_key;
    std::vector<BitVector> _matrix;
    std::vector<int> _zero_count;
    std::vector<std::vector<Sum>> _zero_prefix;
    std::vector<Sum> _original_prefix;
    std::vector<Sum> _final_prefix;

    static unsigned_type encode(T value) {
        unsigned_type bits;
        if constexpr (std::signed_integral<T>) {
            bits = std::bit_cast<unsigned_type>(value);
        } else {
            bits = value;
        }
        return bits ^ sign_mask;
    }

    static T decode(unsigned_type key) {
        unsigned_type bits = key ^ sign_mask;
        if constexpr (std::signed_integral<T>) {
            return std::bit_cast<T>(bits);
        } else {
            return bits;
        }
    }

    bool bit(unsigned_type value, int level) const {
        return (value >> (_log - 1 - level)) & unsigned_type(1);
    }

    static std::uint64_t extract_bits(
        const unsigned_type* values,
        int count,
        int shift
    ) {
        std::uint64_t result = 0;
        int i = 0;
#if defined(__AVX2__)
        if constexpr (sizeof(unsigned_type) == 8) {
            __m128i left = _mm_cvtsi32_si128(63 - shift);
            for (; i + 4 <= count; i += 4) {
                __m256i data = _mm256_loadu_si256(
                    reinterpret_cast<const __m256i*>(values + i)
                );
                data = _mm256_sll_epi64(data, left);
                int mask = _mm256_movemask_pd(_mm256_castsi256_pd(data));
                result |= std::uint64_t(mask) << i;
            }
        } else if constexpr (sizeof(unsigned_type) == 4) {
            __m128i left = _mm_cvtsi32_si128(31 - shift);
            for (; i + 8 <= count; i += 8) {
                __m256i data = _mm256_loadu_si256(
                    reinterpret_cast<const __m256i*>(values + i)
                );
                data = _mm256_sll_epi32(data, left);
                int mask = _mm256_movemask_ps(_mm256_castsi256_ps(data));
                result |= std::uint64_t(mask) << i;
            }
        }
#endif
        for (; i < count; i++) {
            result |= std::uint64_t((values[i] >> shift) & unsigned_type(1))
                      << i;
        }
        return result;
    }

    Sum zero_sum(int level, int l, int r) const {
        return _zero_prefix[level][r] - _zero_prefix[level][l];
    }

    Sum sum_less_encoded(int l, int r, unsigned_type upper) const {
        if (_n == 0 || upper <= _min_key) return Sum{};
        if (upper > _max_key) {
            return _original_prefix[r] - _original_prefix[l];
        }

        Sum result{};
        for (int level = 0; level < _log; level++) {
            int l1 = _matrix[level].rank1(l);
            int r1 = _matrix[level].rank1(r);
            if (bit(upper, level)) {
                result = result + zero_sum(level, l, r);
                l = _zero_count[level] + l1;
                r = _zero_count[level] + r1;
            } else {
                l -= l1;
                r -= r1;
            }
        }
        return result;
    }

    int count_less_encoded(int l, int r, unsigned_type upper) const {
        if (_n == 0 || upper <= _min_key) return 0;
        if (upper > _max_key) return r - l;

        int result = 0;
        for (int level = 0; level < _log; level++) {
            int l1 = _matrix[level].rank1(l);
            int r1 = _matrix[level].rank1(r);
            if (bit(upper, level)) {
                result += (r - l) - (r1 - l1);
                l = _zero_count[level] + l1;
                r = _zero_count[level] + r1;
            } else {
                l -= l1;
                r -= r1;
            }
        }
        return result;
    }

    void build(const std::vector<T>& values, const std::vector<Sum>& weights) {
        assert(values.size() == weights.size());

        std::vector<unsigned_type> current_keys(_n);
        std::vector<unsigned_type> next_keys(_n);
        std::vector<Sum> current_weights(weights);
        std::vector<Sum> next_weights(_n);
        for (int i = 0; i < _n; i++) current_keys[i] = encode(values[i]);

        _original_prefix.assign(std::size_t(_n) + 1, Sum{});
        for (int i = 0; i < _n; i++) {
            _original_prefix[i + 1] = _original_prefix[i] + weights[i];
        }
        if (_n == 0) {
            _final_prefix.assign(1, Sum{});
            return;
        }

        _min_key = current_keys[0];
        _max_key = current_keys[0];
        for (unsigned_type key : current_keys) {
            if (key < _min_key) _min_key = key;
            if (_max_key < key) _max_key = key;
        }
        _log = int(std::bit_width(unsigned_type(_min_key ^ _max_key)));
        if (_log != value_bit_width) {
            _key_prefix = unsigned_type((_min_key >> _log) << _log);
        }
        _zero_count.assign(_log, 0);

        _matrix.reserve(_log);
        _zero_prefix.reserve(_log);
        for (int level = 0; level < _log; level++) {
            _matrix.emplace_back(_n);
            _zero_prefix.emplace_back(std::size_t(_n) + 1, Sum{});
            BitVector& bit_vector = _matrix.back();
            int shift = _log - 1 - level;
            int zeros = 0;
            for (int base = 0; base < _n; base += 64) {
                int count = std::min(64, _n - base);
                std::uint64_t word = extract_bits(
                    current_keys.data() + base,
                    count,
                    shift
                );
                bit_vector.bits[std::size_t(base) >> 6] = word;
                zeros += count - std::popcount(word);
            }
            bit_vector.build();

            std::vector<Sum>& prefix = _zero_prefix.back();
            for (int base = 0; base < _n; base += 64) {
                int count = std::min(64, _n - base);
                std::uint64_t ones = bit_vector.bits[std::size_t(base) >> 6];
                for (int offset = 0; offset < count; offset++) {
                    int i = base + offset;
                    prefix[i + 1] = prefix[i];
                    if (((ones >> offset) & 1) == 0) {
                        prefix[i + 1] = prefix[i + 1] + current_weights[i];
                    }
                }
            }

            _zero_count[level] = zeros;
            int zero_pos = 0;
            int one_pos = zeros;
            for (int base = 0; base < _n; base += 64) {
                int count = std::min(64, _n - base);
                std::uint64_t ones = bit_vector.bits[std::size_t(base) >> 6];
                std::uint64_t valid = count == 64
                                          ? ~std::uint64_t(0)
                                          : (std::uint64_t(1) << count) - 1;
                std::uint64_t zeroes = (~ones) & valid;
                while (zeroes != 0) {
                    int offset = std::countr_zero(zeroes);
                    next_keys[zero_pos] = current_keys[base + offset];
                    next_weights[zero_pos] = current_weights[base + offset];
                    zero_pos++;
                    zeroes &= zeroes - 1;
                }
                while (ones != 0) {
                    int offset = std::countr_zero(ones);
                    next_keys[one_pos] = current_keys[base + offset];
                    next_weights[one_pos] = current_weights[base + offset];
                    one_pos++;
                    ones &= ones - 1;
                }
            }
            current_keys.swap(next_keys);
            current_weights.swap(next_weights);
        }

        _final_prefix.assign(std::size_t(_n) + 1, Sum{});
        for (int i = 0; i < _n; i++) {
            _final_prefix[i + 1] = _final_prefix[i] + current_weights[i];
        }
    }

   public:
    WaveletMatrixSum()
        : _n(0),
          _log(0),
          _key_prefix(0),
          _min_key(0),
          _max_key(0),
          _original_prefix(1, Sum{}),
          _final_prefix(1, Sum{}) {}

    explicit WaveletMatrixSum(const std::vector<T>& values)
        requires std::convertible_to<T, Sum>
        : _n(int(values.size())),
          _log(0),
          _key_prefix(0),
          _min_key(0),
          _max_key(0) {
        std::vector<Sum> weights;
        weights.reserve(values.size());
        for (T value : values) weights.push_back(static_cast<Sum>(value));
        build(values, weights);
    }

    WaveletMatrixSum(
        const std::vector<T>& values,
        const std::vector<Sum>& weights
    ) : _n(int(values.size())),
        _log(0),
        _key_prefix(0),
        _min_key(0),
        _max_key(0) {
        build(values, weights);
    }

    int size() const {
        return _n;
    }

    bool empty() const {
        return _n == 0;
    }

    T access(int p) const {
        assert(0 <= p && p < _n);
        unsigned_type key = _key_prefix;
        for (int level = 0; level < _log; level++) {
            int ones_before = _matrix[level].rank1(p);
            bool one = _matrix[level].get(p);
            if (one) {
                key |= unsigned_type(1) << (_log - 1 - level);
                p = _zero_count[level] + ones_before;
            } else {
                p -= ones_before;
            }
        }
        return decode(key);
    }

    T operator[](int p) const {
        return access(p);
    }

    int rank(T value, int r) const {
        assert(0 <= r && r <= _n);
        return rank(value, 0, r);
    }

    int rank(T value, int l, int r) const {
        assert(0 <= l && l <= r && r <= _n);
        unsigned_type key = encode(value);
        if (_n == 0 || key < _min_key || _max_key < key) return 0;
        for (int level = 0; level < _log; level++) {
            int l1 = _matrix[level].rank1(l);
            int r1 = _matrix[level].rank1(r);
            if (bit(key, level)) {
                l = _zero_count[level] + l1;
                r = _zero_count[level] + r1;
            } else {
                l -= l1;
                r -= r1;
            }
        }
        return r - l;
    }

    T kth_smallest(int l, int r, int k) const {
        assert(0 <= l && l <= r && r <= _n);
        assert(0 <= k && k < r - l);
        unsigned_type key = _key_prefix;
        for (int level = 0; level < _log; level++) {
            int l1 = _matrix[level].rank1(l);
            int r1 = _matrix[level].rank1(r);
            int l0 = l - l1;
            int r0 = r - r1;
            int zeros = r0 - l0;
            if (k < zeros) {
                l = l0;
                r = r0;
            } else {
                k -= zeros;
                key |= unsigned_type(1) << (_log - 1 - level);
                l = _zero_count[level] + l1;
                r = _zero_count[level] + r1;
            }
        }
        return decode(key);
    }

    T kth_largest(int l, int r, int k) const {
        assert(0 <= l && l <= r && r <= _n);
        assert(0 <= k && k < r - l);
        return kth_smallest(l, r, r - l - 1 - k);
    }

    int range_freq(int l, int r, T upper) const {
        assert(0 <= l && l <= r && r <= _n);
        return count_less_encoded(l, r, encode(upper));
    }

    int range_freq(int l, int r, T lower, T upper) const {
        assert(0 <= l && l <= r && r <= _n);
        if (upper <= lower) return 0;
        return range_freq(l, r, upper) - range_freq(l, r, lower);
    }

    std::optional<T> prev_value(int l, int r, T upper) const {
        assert(0 <= l && l <= r && r <= _n);
        int count = range_freq(l, r, upper);
        if (count == 0) return std::nullopt;
        return kth_smallest(l, r, count - 1);
    }

    std::optional<T> next_value(int l, int r, T lower) const {
        assert(0 <= l && l <= r && r <= _n);
        int count = range_freq(l, r, lower);
        if (count == r - l) return std::nullopt;
        return kth_smallest(l, r, count);
    }

    Sum range_sum(int l, int r) const {
        assert(0 <= l && l <= r && r <= _n);
        return _original_prefix[r] - _original_prefix[l];
    }

    Sum range_sum(int l, int r, T upper) const {
        assert(0 <= l && l <= r && r <= _n);
        return sum_less_encoded(l, r, encode(upper));
    }

    Sum range_sum(int l, int r, T lower, T upper) const {
        assert(0 <= l && l <= r && r <= _n);
        if (upper <= lower) return Sum{};
        return range_sum(l, r, upper) - range_sum(l, r, lower);
    }

    Sum sum_k_smallest(int l, int r, int k) const {
        assert(0 <= l && l <= r && r <= _n);
        assert(0 <= k && k <= r - l);
        Sum result{};
        for (int level = 0; level < _log; level++) {
            int l1 = _matrix[level].rank1(l);
            int r1 = _matrix[level].rank1(r);
            int l0 = l - l1;
            int r0 = r - r1;
            int zeros = r0 - l0;
            if (k < zeros) {
                l = l0;
                r = r0;
            } else {
                result = result + zero_sum(level, l, r);
                k -= zeros;
                l = _zero_count[level] + l1;
                r = _zero_count[level] + r1;
            }
        }
        return result + (_final_prefix[l + k] - _final_prefix[l]);
    }

    Sum sum_k_largest(int l, int r, int k) const {
        assert(0 <= l && l <= r && r <= _n);
        assert(0 <= k && k <= r - l);
        return range_sum(l, r) - sum_k_smallest(l, r, r - l - k);
    }

    template <class Predicate>
    int max_count_smallest(int l, int r, Predicate predicate) const {
        assert(0 <= l && l <= r && r <= _n);
        assert(predicate(Sum{}));
        Sum result{};
        int count = 0;
        for (int level = 0; level < _log; level++) {
            int l1 = _matrix[level].rank1(l);
            int r1 = _matrix[level].rank1(r);
            int l0 = l - l1;
            int r0 = r - r1;
            int zeros = r0 - l0;
            Sum zero_result = result + zero_sum(level, l, r);
            if (predicate(zero_result)) {
                result = zero_result;
                count += zeros;
                l = _zero_count[level] + l1;
                r = _zero_count[level] + r1;
            } else {
                l = l0;
                r = r0;
            }
        }

        int low = 0;
        int high = r - l;
        while (low < high) {
            int middle = low + (high - low + 1) / 2;
            Sum candidate =
                result + (_final_prefix[l + middle] - _final_prefix[l]);
            if (predicate(candidate)) {
                low = middle;
            } else {
                high = middle - 1;
            }
        }
        return count + low;
    }

    template <class Predicate>
    int max_count_largest(int l, int r, Predicate predicate) const {
        assert(0 <= l && l <= r && r <= _n);
        assert(predicate(Sum{}));
        Sum result{};
        Sum current_sum = range_sum(l, r);
        int count = 0;
        for (int level = 0; level < _log; level++) {
            int l1 = _matrix[level].rank1(l);
            int r1 = _matrix[level].rank1(r);
            int l0 = l - l1;
            int r0 = r - r1;
            int ones = r1 - l1;
            Sum zero_result = zero_sum(level, l, r);
            Sum one_result = current_sum - zero_result;
            Sum candidate = result + one_result;
            if (predicate(candidate)) {
                result = candidate;
                count += ones;
                current_sum = zero_result;
                l = l0;
                r = r0;
            } else {
                current_sum = one_result;
                l = _zero_count[level] + l1;
                r = _zero_count[level] + r1;
            }
        }

        int low = 0;
        int high = r - l;
        while (low < high) {
            int middle = low + (high - low + 1) / 2;
            Sum candidate =
                result + (_final_prefix[r] - _final_prefix[r - middle]);
            if (predicate(candidate)) {
                low = middle;
            } else {
                high = middle - 1;
            }
        }
        return count + low;
    }
};

}  // namespace ds
}  // namespace m1une


#line 11 "ds/range_query/static_rectangle_sum.hpp"

namespace m1une {
namespace ds {

template <class X, class Y = X, class Sum = long long>
class StaticRectangleSum {
   public:
    using x_type = X;
    using y_type = Y;
    using sum_type = Sum;
    using weighted_point_type = std::tuple<X, Y, Sum>;

   private:
    std::vector<X> _x_coordinates;
    std::vector<Y> _y_coordinates;
    WaveletMatrixSum<int, Sum> _matrix;

   public:
    StaticRectangleSum() = default;

    explicit StaticRectangleSum(
        const std::vector<weighted_point_type>& points
    ) {
        build(points);
    }

    StaticRectangleSum(
        const std::vector<X>& x_coordinates,
        const std::vector<Y>& y_coordinates,
        const std::vector<Sum>& weights
    ) {
        build(x_coordinates, y_coordinates, weights);
    }

    void build(std::vector<weighted_point_type> points) {
        std::sort(
            points.begin(),
            points.end(),
            [](const weighted_point_type& first,
               const weighted_point_type& second) {
                if (std::get<0>(first) < std::get<0>(second)) return true;
                if (std::get<0>(second) < std::get<0>(first)) return false;
                return std::get<1>(first) < std::get<1>(second);
            }
        );

        const int n = int(points.size());
        _x_coordinates.resize(n);
        _y_coordinates.clear();
        _y_coordinates.reserve(n);
        for (const auto& point : points) {
            _y_coordinates.push_back(std::get<1>(point));
        }
        std::sort(_y_coordinates.begin(), _y_coordinates.end());
        _y_coordinates.erase(
            std::unique(
                _y_coordinates.begin(),
                _y_coordinates.end(),
                [](const Y& first, const Y& second) {
                    return !(first < second) && !(second < first);
                }
            ),
            _y_coordinates.end()
        );

        std::vector<int> y_rank(n);
        std::vector<Sum> weights(n);
        for (int index = 0; index < n; index++) {
            _x_coordinates[index] = std::get<0>(points[index]);
            y_rank[index] = int(
                std::lower_bound(
                    _y_coordinates.begin(),
                    _y_coordinates.end(),
                    std::get<1>(points[index])
                ) - _y_coordinates.begin()
            );
            weights[index] = std::get<2>(points[index]);
        }
        _matrix = WaveletMatrixSum<int, Sum>(y_rank, weights);
    }

    void build(
        const std::vector<X>& x_coordinates,
        const std::vector<Y>& y_coordinates,
        const std::vector<Sum>& weights
    ) {
        assert(x_coordinates.size() == y_coordinates.size());
        assert(x_coordinates.size() == weights.size());
        std::vector<weighted_point_type> points;
        points.reserve(x_coordinates.size());
        for (int index = 0; index < int(x_coordinates.size()); index++) {
            points.emplace_back(
                x_coordinates[index],
                y_coordinates[index],
                weights[index]
            );
        }
        build(std::move(points));
    }

    int size() const {
        return int(_x_coordinates.size());
    }

    bool empty() const {
        return _x_coordinates.empty();
    }

    Sum sum(
        const X& left,
        const X& right,
        const Y& lower,
        const Y& upper
    ) const {
        assert(!(right < left));
        assert(!(upper < lower));
        int x_left = int(
            std::lower_bound(
                _x_coordinates.begin(),
                _x_coordinates.end(),
                left
            ) - _x_coordinates.begin()
        );
        int x_right = int(
            std::lower_bound(
                _x_coordinates.begin(),
                _x_coordinates.end(),
                right
            ) - _x_coordinates.begin()
        );
        int y_lower = int(
            std::lower_bound(
                _y_coordinates.begin(),
                _y_coordinates.end(),
                lower
            ) - _y_coordinates.begin()
        );
        int y_upper = int(
            std::lower_bound(
                _y_coordinates.begin(),
                _y_coordinates.end(),
                upper
            ) - _y_coordinates.begin()
        );
        return _matrix.range_sum(x_left, x_right, y_lower, y_upper);
    }
};

}  // namespace ds
}  // namespace m1une
Back to top page