ecnerwala's competitive programming library
#include <catch2/catch_test_macros.hpp>
#include <catch2/catch_get_random_seed.hpp>
#include "rmq.hpp"
#include <vector>
#include <random>
TEST_CASE("RangeMinQuery", "[rmq]") {
std::mt19937 mt(Catch::getSeed());
for (int N : {1, 2, 3, 5, 10, 20, 33, 48, 100, 163, 512}) {
std::vector<std::pair<int, int>> data(N);
for (int i = 0; i < N; i++) {
data[i] = {mt(), i};
}
RangeMinQuery<std::pair<int, int>> minQ(data);
RangeMaxQuery<std::pair<int, int>> maxQ(data);
for (int l = 0; l < N; l++) {
std::pair<int, int> cur_min = data[l];
std::pair<int, int> cur_max = data[l];
for (int r = l; r < N; r++) {
cur_min = min(cur_min, data[r]);
REQUIRE(minQ.query(l, r) == cur_min);
cur_max = max(cur_max, data[r]);
REQUIRE(maxQ.query(l, r) == cur_max);
}
}
}
}
#include <catch2/catch_test_macros.hpp>
#include <catch2/catch_get_random_seed.hpp>
#include <functional>
#include <vector>
#include <cassert>
#include <cstdint>
#include <random>
#line 3 "src/rmq.test.cpp"
#line 2 "src/rmq.hpp"
#line 7 "src/rmq.hpp"
template <typename T, class Compare = std::less<T>> class RangeMinQuery : private Compare {
static const int BUCKET_SIZE = 32;
static const int BUCKET_SIZE_LOG = 5;
static_assert(BUCKET_SIZE == (1 << BUCKET_SIZE_LOG), "BUCKET_SIZE should be a power of 2");
static const int CACHE_LINE_ALIGNMENT = 64;
int n = 0;
std::vector<T> data;
std::vector<T> pref_data;
std::vector<T> suff_data;
std::vector<T> sparse_table;
std::vector<uint32_t> range_mask;
private:
int num_buckets() const {
return n >> BUCKET_SIZE_LOG;
}
int num_levels() const {
return num_buckets() ? 32 - __builtin_clz(num_buckets()) : 0;
}
int sparse_table_size() const {
return num_buckets() * num_levels();
}
private:
const T& min(const T& a, const T& b) const {
return Compare::operator()(a, b) ? a : b;
}
void setmin(T& a, const T& b) const {
if (Compare::operator()(b, a)) a = b;
}
template <typename Vec> static int get_size(const Vec& v) { using std::size; return int(size(v)); }
public:
RangeMinQuery() {}
template <typename Vec> explicit RangeMinQuery(const Vec& data_, const Compare& comp_ = Compare())
: Compare(comp_)
, n(get_size(data_))
, data(n)
, pref_data(n)
, suff_data(n)
, sparse_table(sparse_table_size())
, range_mask(n)
{
for (int i = 0; i < n; i++) data[i] = data_[i];
for (int i = 0; i < n; i++) {
if (i & (BUCKET_SIZE-1)) {
uint32_t m = range_mask[i-1];
while (m && !Compare::operator()(data[(i | (BUCKET_SIZE-1)) - __builtin_clz(m)], data[i])) {
m -= uint32_t(1) << (BUCKET_SIZE - 1 - __builtin_clz(m));
}
m |= uint32_t(1) << (i & (BUCKET_SIZE - 1));
range_mask[i] = m;
} else {
range_mask[i] = 1;
}
}
for (int i = 0; i < n; i++) {
pref_data[i] = data[i];
if (i & (BUCKET_SIZE-1)) {
setmin(pref_data[i], pref_data[i-1]);
}
}
for (int i = n-1; i >= 0; i--) {
suff_data[i] = data[i];
if (i+1 < n && ((i+1) & (BUCKET_SIZE-1))) {
setmin(suff_data[i], suff_data[i+1]);
}
}
for (int i = 0; i < num_buckets(); i++) {
sparse_table[i] = data[i * BUCKET_SIZE];
for (int v = 1; v < BUCKET_SIZE; v++) {
setmin(sparse_table[i], data[i * BUCKET_SIZE + v]);
}
}
for (int l = 0; l+1 < num_levels(); l++) {
for (int i = 0; i + (1 << (l+1)) <= num_buckets(); i++) {
sparse_table[(l+1) * num_buckets() + i] = min(sparse_table[l * num_buckets() + i], sparse_table[l * num_buckets() + i + (1 << l)]);
}
}
}
T query(int l, int r) const {
assert(l <= r);
int bucket_l = (l >> BUCKET_SIZE_LOG);
int bucket_r = (r >> BUCKET_SIZE_LOG);
if (bucket_l == bucket_r) {
uint32_t msk = range_mask[r] & ~((uint32_t(1) << (l & (BUCKET_SIZE-1))) - 1);
int ind = (l & ~(BUCKET_SIZE-1)) + __builtin_ctz(msk);
return data[ind];
} else {
T ans = min(suff_data[l], pref_data[r]);
bucket_l++;
if (bucket_l < bucket_r) {
int level = (32 - __builtin_clz(bucket_r - bucket_l)) - 1;
setmin(ans, sparse_table[level * num_buckets() + bucket_l]);
setmin(ans, sparse_table[level * num_buckets() + bucket_r - (1 << level)]);
}
return ans;
}
}
};
template <typename T> using RangeMaxQuery = RangeMinQuery<T, std::greater<T>>;
#line 5 "src/rmq.test.cpp"
#line 8 "src/rmq.test.cpp"
TEST_CASE("RangeMinQuery", "[rmq]") {
std::mt19937 mt(Catch::getSeed());
for (int N : {1, 2, 3, 5, 10, 20, 33, 48, 100, 163, 512}) {
std::vector<std::pair<int, int>> data(N);
for (int i = 0; i < N; i++) {
data[i] = {mt(), i};
}
RangeMinQuery<std::pair<int, int>> minQ(data);
RangeMaxQuery<std::pair<int, int>> maxQ(data);
for (int l = 0; l < N; l++) {
std::pair<int, int> cur_min = data[l];
std::pair<int, int> cur_max = data[l];
for (int r = l; r < N; r++) {
cur_min = min(cur_min, data[r]);
REQUIRE(minQ.query(l, r) == cur_min);
cur_max = max(cur_max, data[r]);
REQUIRE(maxQ.query(l, r) == cur_max);
}
}
}
}
// clang-format off
// @formatter:off
#pragma GCC diagnostic push
#pragma GCC diagnostic ignored "-Wpragmas"
#pragma GCC diagnostic ignored "-Wunknown-warning-option"
#pragma GCC diagnostic ignored "-Wmisleading-indentation"
#pragma GCC diagnostic ignored "-Wmultistatement-macros"
#include <bits/stdc++.h>
#include <catch2/catch_test_macros.hpp>
#include <catch2/catch_get_random_seed.hpp>
#include <cassert>
// src/rmq.hpp
template<typename T,class Compare=std::less<T>>class RangeMinQuery:private Compare{
static const int BUCKET_SIZE=32;
static const int BUCKET_SIZE_LOG=5;
static_assert(BUCKET_SIZE==(1<<BUCKET_SIZE_LOG),"BUCKET_SIZE should be a power of 2");
static const int CACHE_LINE_ALIGNMENT=64;
int n=0;
std::vector<T>data;
std::vector<T>pref_data;
std::vector<T>suff_data;
std::vector<T>sparse_table;
std::vector<uint32_t>range_mask;
private:
int num_buckets()const{
return n>>BUCKET_SIZE_LOG;
}
int num_levels()const{
return num_buckets()?32-__builtin_clz(num_buckets()):0;
}
int sparse_table_size()const{
return num_buckets()*num_levels();
}
private:
const T&min(const T&a,const T&b)const{
return Compare::operator()(a,b)?a:b;
}
void setmin(T&a,const T&b)const{
if(Compare::operator()(b,a))a=b;
}
template<typename Vec>static int get_size(const Vec&v){using std::size;return int(size(v));}
public:
RangeMinQuery(){}
template<typename Vec>explicit RangeMinQuery(const Vec&data_,const Compare&comp_=Compare())
:Compare(comp_)
,n(get_size(data_))
,data(n)
,pref_data(n)
,suff_data(n)
,sparse_table(sparse_table_size())
,range_mask(n)
{
for(int i=0;i<n;i++)data[i]=data_[i];
for(int i=0;i<n;i++){
if(i&(BUCKET_SIZE-1)){
uint32_t m=range_mask[i-1];
while(m&&!Compare::operator()(data[(i|(BUCKET_SIZE-1))-__builtin_clz(m)],data[i])){
m-=uint32_t(1)<<(BUCKET_SIZE-1-__builtin_clz(m));
}
m|=uint32_t(1)<<(i&(BUCKET_SIZE-1));
range_mask[i]=m;
}else{
range_mask[i]=1;
}
}
for(int i=0;i<n;i++){
pref_data[i]=data[i];
if(i&(BUCKET_SIZE-1)){
setmin(pref_data[i],pref_data[i-1]);
}
}
for(int i=n-1;i>=0;i--){
suff_data[i]=data[i];
if(i+1<n&&((i+1)&(BUCKET_SIZE-1))){
setmin(suff_data[i],suff_data[i+1]);
}
}
for(int i=0;i<num_buckets();i++){
sparse_table[i]=data[i*BUCKET_SIZE];
for(int v=1;v<BUCKET_SIZE;v++){
setmin(sparse_table[i],data[i*BUCKET_SIZE+v]);
}
}
for(int l=0;l+1<num_levels();l++){
for(int i=0;i+(1<<(l+1))<=num_buckets();i++){
sparse_table[(l+1)*num_buckets()+i]=min(sparse_table[l*num_buckets()+i],sparse_table[l*num_buckets()+i+(1<<l)]);
}
}
}
T query(int l,int r)const{
assert(l<=r);
int bucket_l=(l>>BUCKET_SIZE_LOG);
int bucket_r=(r>>BUCKET_SIZE_LOG);
if(bucket_l==bucket_r){
uint32_t msk=range_mask[r]&~((uint32_t(1)<<(l&(BUCKET_SIZE-1)))-1);
int ind=(l&~(BUCKET_SIZE-1))+__builtin_ctz(msk);
return data[ind];
}else{
T ans=min(suff_data[l],pref_data[r]);
bucket_l++;
if(bucket_l<bucket_r){
int level=(32-__builtin_clz(bucket_r-bucket_l))-1;
setmin(ans,sparse_table[level*num_buckets()+bucket_l]);
setmin(ans,sparse_table[level*num_buckets()+bucket_r-(1<<level)]);
}
return ans;
}
}
};
template<typename T>using RangeMaxQuery=RangeMinQuery<T,std::greater<T>>;
// src/rmq.test.cpp
TEST_CASE("RangeMinQuery","[rmq]"){
std::mt19937 mt(Catch::getSeed());
for(int N:{1,2,3,5,10,20,33,48,100,163,512}){
std::vector<std::pair<int,int>>data(N);
for(int i=0;i<N;i++){
data[i]={mt(),i};
}
RangeMinQuery<std::pair<int,int>>minQ(data);
RangeMaxQuery<std::pair<int,int>>maxQ(data);
for(int l=0;l<N;l++){
std::pair<int,int>cur_min=data[l];
std::pair<int,int>cur_max=data[l];
for(int r=l;r<N;r++){
cur_min=min(cur_min,data[r]);
REQUIRE(minQ.query(l,r)==cur_min);
cur_max=max(cur_max,data[r]);
REQUIRE(maxQ.query(l,r)==cur_max);
}
}
}
}
#pragma GCC diagnostic pop
// clang-format on
// @formatter:on