cp-book

ecnerwala's competitive programming library

View the Project on GitHub ecnerwala/cp-book

:heavy_check_mark: verify/min_plus_convolution_convex_arbitrary.test.cpp

Depends on

Code

// competitive-verifier: PROBLEM https://judge.yosupo.jp/problem/min_plus_convolution_convex_arbitrary

#include <bits/stdc++.h>
#include <cassert>

#include "smawk.hpp"

int main() {
	std::ios_base::sync_with_stdio(false), std::cin.tie(nullptr);

	int N, M; std::cin >> N >> M;
	std::vector<int> A(N); for (auto& a : A) std::cin >> a;
	std::vector<int> B(M); for (auto& b : B) std::cin >> b;
	const int INF = *std::max_element(A.begin(), A.end()) + *std::max_element(B.begin(), B.end()) + 1;
	auto res = smawk::smawk(N+M-1, M, [&](int row, int col) -> int {
		int i = row - col;
		return i < 0 ? INF + M + (~i) : i >= N ? INF + (i-N) : (A[i] + B[col]);
	}, [&]([[maybe_unused]] int row, smawk::value_t<int> col1, smawk::value_t<int> col2) -> bool {
		return col1.v < col2.v ? 0 : 1;
	});
	for (int i = 0; i < N+M-1; i++) {
		std::cout << res[i].v << " \n"[i+1==N+M-1];
	}

	return 0;
}
#include <bits/stdc++.h>
#line 1 "verify/min_plus_convolution_convex_arbitrary.test.cpp"
// competitive-verifier: PROBLEM https://judge.yosupo.jp/problem/min_plus_convolution_convex_arbitrary

#line 5 "verify/min_plus_convolution_convex_arbitrary.test.cpp"

#line 2 "src/smawk.hpp"

#line 7 "src/smawk.hpp"

namespace smawk {

template <typename T> struct value_t {
	T v;
	int col;
};

// Get(int row, int col) -> T
// Select(int row, const value_t<T>& opt_0, const value_t<T>& opt_1) returns 0 or 1 for which is better
template <typename T, typename Get, typename Select> concept totally_monotone_matrix_oracle =
	std::default_initializable<T> && std::movable<T>
	&& std::invocable<Get, int, int> && std::convertible_to<std::invoke_result_t<Get, int, int>, T>
	&& std::predicate<Select, int, const value_t<T>&, const value_t<T>&>;


template <typename Get, typename Select, typename T = std::invoke_result_t<Get, int, int>>
requires totally_monotone_matrix_oracle<T, Get, Select>
class LARSCH {
public:
	int N;
	Get get;
	Select select;
	int L;
	int num_rows;

	std::vector<std::vector<value_t<T>>> stk;
	std::vector<std::pair<value_t<T>, int>> bests;
	LARSCH() {}
	LARSCH(int N_, Get&& get_, Select&& select_) : N(N_), get(std::forward<Get>(get_)), select(std::forward<Select>(select_)) {
		L = N ? 31 - __builtin_clz(N) : 0;
		stk.resize(L);
		bests.resize(L);
		// N >> L == 1, unless N == 0
		for (int i = 0; i < L; i++) {
			stk[i].reserve(N >> (i+1));
		}
		num_rows = 0;
	}

	value_t<T> push_and_query_next() {
		assert(num_rows < N);
		int inp_row = num_rows++;

		int l = 0;
		value_t<T> nbest;
		while (true) {
			int r = inp_row >> l;
			int col = l == 0 ? inp_row : stk[l-1][r].col;
			if (l == L) {
				// special case: just return the unique element
				assert(r == 0);
				assert(inp_row == 0);
				int row = ((r+1) << l) - 1;
				if (l == 0) nbest = {get(row, col), col};
				else nbest = {std::move(stk[l-1][r].v), stk[l-1][r].col};
				l--;
				break;
			}
			assert(l < L);

			if (r & 1) {
				int row = ((r+1) << l) - 1;
				value_t<T> prv_col_top;
				if (l == 0) prv_col_top = {get(row, col), col};
				else prv_col_top = {std::move(stk[l-1][r].v), stk[l-1][r].col};

				assert(bests[l].second <= r);

				// just check this guy at this row, and then push it into the next layer, but don't query yet
				if (bests[l].second == r || select(row, bests[l].first, prv_col_top)) {
					// prv_col_top is better here
					bests[l].first = std::move(prv_col_top);
					bests[l].second = r;
				}
			}
			if (bests[l].second == r) {
				// We've committed to the new column being the best, so let's prune the stack and propagate to ensure consistency.
				assert(int(stk[l].size()) >= (r+1)/2);
				stk[l].resize((r+1)/2);
				// We can just set .second only, since the only time it's read is the r&1 case, which will override bests[l].first
				if (l+1 < L) bests[l+1].second = (r+1)/2;
			}
			std::optional<value_t<T>> to_push;
			while (int(stk[l].size()) > (r+1)/2) {
				int row = (int(stk[l].size()) << (l+1)) - 1;
				value_t<T> nv{get(row, col), col};
				if (select(row, stk[l].back(), nv)) {
					stk[l].pop_back();
					to_push = std::move(nv);
				} else {
					break;
				}
			}
			if (to_push) {
				stk[l].emplace_back(std::move(*to_push));
			} else {
				int row = (int(stk[l].size()+1) << (l+1)) - 1;
				if (row < N) stk[l].emplace_back(get(row, col), col);
			}
			if (r & 1) {
				// go return
				nbest = std::move(bests[l].first);
				l--;
				break;
			} else if (((r+2) << l) - 1 >= N) {
				// go return
				nbest.col = col;
				break;
			} else {
				l++;
				continue;
			}
			assert(false);
		}
		for (; l >= 0; l--) {
			int r = inp_row >> l;
			assert(!(r & 1));
			int row = ((r+1) << l) - 1;
			bests[l].first = std::move(nbest);
			bool did_set = false;
			while (true) {
				int idx = bests[l].second;
				assert(idx <= r);
				int col = (l == 0 ? idx : stk[l-1][idx].col);
				assert(col <= bests[l].first.col);
				value_t<T> cnd;
				if (l > 0 && idx == r) cnd = {std::move(stk[l-1][r].v), col};
				else cnd = {get(row, col), col};
				if (!did_set || select(row, nbest, cnd)) {
					did_set = true;
					nbest = std::move(cnd);
				}
				if (col == bests[l].first.col) break;
				bests[l].second++;
			}
		}
		assert(l == -1);
		return nbest;
	}
};

template <typename Get, typename Select, typename T = std::invoke_result_t<Get&&, int, int>>
requires totally_monotone_matrix_oracle<T, Get&&, Select&&>
std::vector<value_t<T>> smawk(int N, int M, Get&& get, Select&& select) {
	// TODO: If M >> N, then we should do an extra layer of column filter on the outside. The cutoff should be M > 2N or so.
	std::vector<value_t<T>> res(N);
	for (int i = 0; i < N; i++) res[i].col = -1;
	std::vector<int> stks(N);
	int L = N ? 31 - __builtin_clz(N) : 0;
	std::vector<int> stk_ends(L+1);
	stk_ends[0] = 0;
	for (int l = 0; l < L; l++) {
		int sz = 0;
		auto check_col = [&](int col, int min_sz) -> void {
			while (sz > min_sz) {
				int row = (sz << (l+1)) - 1;
				value_t<T> cnd(get(row, col), col);
				if (select(row, res[row], cnd)) {
					// we prefer cnd, save this
					res[row] = std::move(cnd);
					sz--;
				} else {
					break;
				}
			}

			if (sz < (N >> (l+1))) {
				int row = ((sz+1) << (l+1)) - 1;
				if (res[row].col == col) {
					stks[stk_ends[l] + sz] = col;
					sz++;
				} else {
					value_t<T> cnd(get(row, col), col);
					// This is a legal optimization, but I'm not sure it buys anything real, so just stub it out with true ||
					if (true || res[row].col == -1 || res[row].col < col || !select(row, cnd, res[row])) {
						res[row] = std::move(cnd);
						stks[stk_ends[l] + sz] = col;
						sz++;
					}
				}
			}
		};
		if (l == 0) {
			for (int col = 0; col < M; col++) {
				check_col(col, 0);
			}
		} else {
			for (int z = stk_ends[l-1]; z < stk_ends[l]; z++) {
				check_col(stks[z], (z - stk_ends[l-1]) / 2);
			}
		}
		assert(sz <= (N >> (l+1)));
		stk_ends[l+1] = stk_ends[l] + sz;
	}
	for (int l = L; l >= 0; l--) {
		int z = l == 0 ? 0 : stk_ends[l-1];
		for (int r = 0; r < (N >> l); r += 2) {
			int row = ((r+1) << l) - 1;
			// TODO: You could not reset this? Not sure if it buys anything real.
			res[row].col = -1;
			for (; z < (l == 0 ? M : stk_ends[l]); z++) {
				int col = l == 0 ? z : stks[z];
				value_t<T> cnd = {get(row, col), col};
				if (res[row].col == -1 || select(row, res[row], cnd)) {
					res[row] = std::move(cnd);
				}
				if ((r+1) < (N >> l) && col == res[((r+2) << l) - 1].col) break;
			}
			assert(res[row].col != -1);
		}
	}
	return res;
}

// namespace smawk
}
#line 7 "verify/min_plus_convolution_convex_arbitrary.test.cpp"

int main() {
	std::ios_base::sync_with_stdio(false), std::cin.tie(nullptr);

	int N, M; std::cin >> N >> M;
	std::vector<int> A(N); for (auto& a : A) std::cin >> a;
	std::vector<int> B(M); for (auto& b : B) std::cin >> b;
	const int INF = *std::max_element(A.begin(), A.end()) + *std::max_element(B.begin(), B.end()) + 1;
	auto res = smawk::smawk(N+M-1, M, [&](int row, int col) -> int {
		int i = row - col;
		return i < 0 ? INF + M + (~i) : i >= N ? INF + (i-N) : (A[i] + B[col]);
	}, [&]([[maybe_unused]] int row, smawk::value_t<int> col1, smawk::value_t<int> col2) -> bool {
		return col1.v < col2.v ? 0 : 1;
	});
	for (int i = 0; i < N+M-1; i++) {
		std::cout << res[i].v << " \n"[i+1==N+M-1];
	}

	return 0;
}
// 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>
// src/smawk.hpp
namespace smawk{
template<typename T>struct value_t{
T v;
int col;
};
template<typename T,typename Get,typename Select>concept totally_monotone_matrix_oracle=
std::default_initializable<T>&&std::movable<T>
&&std::invocable<Get,int,int>&&std::convertible_to<std::invoke_result_t<Get,int,int>,T>
&&std::predicate<Select,int,const value_t<T>&,const value_t<T>&>;
template<typename Get,typename Select,typename T=std::invoke_result_t<Get,int,int>>
requires totally_monotone_matrix_oracle<T,Get,Select>
class LARSCH{
public:
int N;
Get get;
Select select;
int L;
int num_rows;
std::vector<std::vector<value_t<T>>>stk;
std::vector<std::pair<value_t<T>,int>>bests;
LARSCH(){}
LARSCH(int N_,Get&&get_,Select&&select_):N(N_),get(std::forward<Get>(get_)),select(std::forward<Select>(select_)){
L=N?31-__builtin_clz(N):0;
stk.resize(L);
bests.resize(L);
for(int i=0;i<L;i++){
stk[i].reserve(N>>(i+1));
}
num_rows=0;
}
value_t<T>push_and_query_next(){
assert(num_rows<N);
int inp_row=num_rows++;
int l=0;
value_t<T>nbest;
while(true){
int r=inp_row>>l;
int col=l==0?inp_row:stk[l-1][r].col;
if(l==L){
assert(r==0);
assert(inp_row==0);
int row=((r+1)<<l)-1;
if(l==0)nbest={get(row,col),col};
else nbest={std::move(stk[l-1][r].v),stk[l-1][r].col};
l--;
break;
}
assert(l<L);
if(r&1){
int row=((r+1)<<l)-1;
value_t<T>prv_col_top;
if(l==0)prv_col_top={get(row,col),col};
else prv_col_top={std::move(stk[l-1][r].v),stk[l-1][r].col};
assert(bests[l].second<=r);
if(bests[l].second==r||select(row,bests[l].first,prv_col_top)){
bests[l].first=std::move(prv_col_top);
bests[l].second=r;
}
}
if(bests[l].second==r){
assert(int(stk[l].size())>=(r+1)/2);
stk[l].resize((r+1)/2);
if(l+1<L)bests[l+1].second=(r+1)/2;
}
std::optional<value_t<T>>to_push;
while(int(stk[l].size())>(r+1)/2){
int row=(int(stk[l].size())<<(l+1))-1;
value_t<T>nv{get(row,col),col};
if(select(row,stk[l].back(),nv)){
stk[l].pop_back();
to_push=std::move(nv);
}else{
break;
}
}
if(to_push){
stk[l].emplace_back(std::move(*to_push));
}else{
int row=(int(stk[l].size()+1)<<(l+1))-1;
if(row<N)stk[l].emplace_back(get(row,col),col);
}
if(r&1){
nbest=std::move(bests[l].first);
l--;
break;
}else if(((r+2)<<l)-1>=N){
nbest.col=col;
break;
}else{
l++;
continue;
}
assert(false);
}
for(;l>=0;l--){
int r=inp_row>>l;
assert(!(r&1));
int row=((r+1)<<l)-1;
bests[l].first=std::move(nbest);
bool did_set=false;
while(true){
int idx=bests[l].second;
assert(idx<=r);
int col=(l==0?idx:stk[l-1][idx].col);
assert(col<=bests[l].first.col);
value_t<T>cnd;
if(l>0&&idx==r)cnd={std::move(stk[l-1][r].v),col};
else cnd={get(row,col),col};
if(!did_set||select(row,nbest,cnd)){
did_set=true;
nbest=std::move(cnd);
}
if(col==bests[l].first.col)break;
bests[l].second++;
}
}
assert(l==-1);
return nbest;
}
};
template<typename Get,typename Select,typename T=std::invoke_result_t<Get&&,int,int>>
requires totally_monotone_matrix_oracle<T,Get&&,Select&&>
std::vector<value_t<T>>smawk(int N,int M,Get&&get,Select&&select){
std::vector<value_t<T>>res(N);
for(int i=0;i<N;i++)res[i].col=-1;
std::vector<int>stks(N);
int L=N?31-__builtin_clz(N):0;
std::vector<int>stk_ends(L+1);
stk_ends[0]=0;
for(int l=0;l<L;l++){
int sz=0;
auto check_col=[&](int col,int min_sz)->void{
while(sz>min_sz){
int row=(sz<<(l+1))-1;
value_t<T>cnd(get(row,col),col);
if(select(row,res[row],cnd)){
res[row]=std::move(cnd);
sz--;
}else{
break;
}
}
if(sz<(N>>(l+1))){
int row=((sz+1)<<(l+1))-1;
if(res[row].col==col){
stks[stk_ends[l]+sz]=col;
sz++;
}else{
value_t<T>cnd(get(row,col),col);
if(true||res[row].col==-1||res[row].col<col||!select(row,cnd,res[row])){
res[row]=std::move(cnd);
stks[stk_ends[l]+sz]=col;
sz++;
}
}
}
};
if(l==0){
for(int col=0;col<M;col++){
check_col(col,0);
}
}else{
for(int z=stk_ends[l-1];z<stk_ends[l];z++){
check_col(stks[z],(z-stk_ends[l-1])/2);
}
}
assert(sz<=(N>>(l+1)));
stk_ends[l+1]=stk_ends[l]+sz;
}
for(int l=L;l>=0;l--){
int z=l==0?0:stk_ends[l-1];
for(int r=0;r<(N>>l);r+=2){
int row=((r+1)<<l)-1;
res[row].col=-1;
for(;z<(l==0?M:stk_ends[l]);z++){
int col=l==0?z:stks[z];
value_t<T>cnd={get(row,col),col};
if(res[row].col==-1||select(row,res[row],cnd)){
res[row]=std::move(cnd);
}
if((r+1)<(N>>l)&&col==res[((r+2)<<l)-1].col)break;
}
assert(res[row].col!=-1);
}
}
return res;
}
}
// verify/min_plus_convolution_convex_arbitrary.test.cpp
int main(){
std::ios_base::sync_with_stdio(false),std::cin.tie(nullptr);
int N,M;std::cin>>N>>M;
std::vector<int>A(N);for(auto&a:A)std::cin>>a;
std::vector<int>B(M);for(auto&b:B)std::cin>>b;
const int INF=*std::max_element(A.begin(),A.end())+*std::max_element(B.begin(),B.end())+1;
auto res=smawk::smawk(N+M-1,M,[&](int row,int col)->int{
int i=row-col;
return i<0?INF+M+(~i):i>=N?INF+(i-N):(A[i]+B[col]);
},[&]([[maybe_unused]]int row,smawk::value_t<int>col1,smawk::value_t<int>col2)->bool{
return col1.v<col2.v?0:1;
});
for(int i=0;i<N+M-1;i++){
std::cout<<res[i].v<<" \n"[i+1==N+M-1];
}
return 0;
}
#pragma GCC diagnostic pop
// clang-format on
// @formatter:on

Test cases

Env Name Status Elapsed Memory
g++-sanitizer example_00 :heavy_check_mark: AC 13 ms 8 MB
g++-sanitizer hack_00 :heavy_check_mark: AC 13 ms 8 MB
g++-sanitizer large_small_00 :heavy_check_mark: AC 105 ms 18 MB
g++-sanitizer large_small_01 :heavy_check_mark: AC 100 ms 20 MB
g++-sanitizer large_small_02 :heavy_check_mark: AC 192 ms 21 MB
g++-sanitizer large_small_03 :heavy_check_mark: AC 187 ms 21 MB
g++-sanitizer max_random_00 :heavy_check_mark: AC 276 ms 30 MB
g++-sanitizer max_random_01 :heavy_check_mark: AC 274 ms 30 MB
g++-sanitizer max_random_02 :heavy_check_mark: AC 273 ms 30 MB
g++-sanitizer med_random_00 :heavy_check_mark: AC 17 ms 9 MB
g++-sanitizer med_random_01 :heavy_check_mark: AC 22 ms 9 MB
g++-sanitizer med_random_02 :heavy_check_mark: AC 17 ms 9 MB
g++-sanitizer monotone_00 :heavy_check_mark: AC 316 ms 30 MB
g++-sanitizer monotone_01 :heavy_check_mark: AC 308 ms 30 MB
g++-sanitizer monotone_02 :heavy_check_mark: AC 271 ms 30 MB
g++-sanitizer monotone_03 :heavy_check_mark: AC 310 ms 30 MB
g++-sanitizer near_power_of_2_00 :heavy_check_mark: AC 149 ms 21 MB
g++-sanitizer near_power_of_2_01 :heavy_check_mark: AC 147 ms 21 MB
g++-sanitizer near_power_of_2_02 :heavy_check_mark: AC 148 ms 21 MB
g++-sanitizer near_power_of_2_03 :heavy_check_mark: AC 149 ms 21 MB
g++-sanitizer near_power_of_2_04 :heavy_check_mark: AC 154 ms 21 MB
g++-sanitizer near_power_of_2_05 :heavy_check_mark: AC 145 ms 21 MB
g++-sanitizer near_power_of_2_06 :heavy_check_mark: AC 148 ms 21 MB
g++-sanitizer near_power_of_2_07 :heavy_check_mark: AC 150 ms 21 MB
g++-sanitizer near_power_of_2_08 :heavy_check_mark: AC 154 ms 21 MB
g++-sanitizer only_first_small_00 :heavy_check_mark: AC 280 ms 30 MB
g++-sanitizer only_first_small_01 :heavy_check_mark: AC 275 ms 30 MB
g++-sanitizer random_00 :heavy_check_mark: AC 218 ms 26 MB
g++-sanitizer random_01 :heavy_check_mark: AC 226 ms 27 MB
g++-sanitizer random_02 :heavy_check_mark: AC 154 ms 19 MB
g++-sanitizer small_00 :heavy_check_mark: AC 16 ms 8 MB
g++-sanitizer small_01 :heavy_check_mark: AC 13 ms 8 MB
g++-sanitizer small_02 :heavy_check_mark: AC 13 ms 8 MB
g++-sanitizer small_03 :heavy_check_mark: AC 15 ms 8 MB
g++-sanitizer small_04 :heavy_check_mark: AC 16 ms 8 MB
g++-sanitizer small_05 :heavy_check_mark: AC 16 ms 8 MB
g++-sanitizer small_06 :heavy_check_mark: AC 17 ms 8 MB
g++-sanitizer small_07 :heavy_check_mark: AC 16 ms 8 MB
g++-sanitizer small_08 :heavy_check_mark: AC 17 ms 8 MB
g++-sanitizer small_slopes_00 :heavy_check_mark: AC 292 ms 29 MB
g++-sanitizer small_slopes_01 :heavy_check_mark: AC 296 ms 30 MB
g++ example_00 :heavy_check_mark: AC 3 ms 4 MB
g++ hack_00 :heavy_check_mark: AC 2 ms 4 MB
g++ large_small_00 :heavy_check_mark: AC 69 ms 12 MB
g++ large_small_01 :heavy_check_mark: AC 68 ms 12 MB
g++ large_small_02 :heavy_check_mark: AC 83 ms 12 MB
g++ large_small_03 :heavy_check_mark: AC 80 ms 12 MB
g++ max_random_00 :heavy_check_mark: AC 154 ms 20 MB
g++ max_random_01 :heavy_check_mark: AC 141 ms 20 MB
g++ max_random_02 :heavy_check_mark: AC 139 ms 20 MB
g++ med_random_00 :heavy_check_mark: AC 3 ms 4 MB
g++ med_random_01 :heavy_check_mark: AC 2 ms 4 MB
g++ med_random_02 :heavy_check_mark: AC 2 ms 4 MB
g++ monotone_00 :heavy_check_mark: AC 155 ms 20 MB
g++ monotone_01 :heavy_check_mark: AC 153 ms 20 MB
g++ monotone_02 :heavy_check_mark: AC 136 ms 20 MB
g++ monotone_03 :heavy_check_mark: AC 158 ms 20 MB
g++ near_power_of_2_00 :heavy_check_mark: AC 70 ms 12 MB
g++ near_power_of_2_01 :heavy_check_mark: AC 72 ms 12 MB
g++ near_power_of_2_02 :heavy_check_mark: AC 70 ms 12 MB
g++ near_power_of_2_03 :heavy_check_mark: AC 70 ms 12 MB
g++ near_power_of_2_04 :heavy_check_mark: AC 75 ms 12 MB
g++ near_power_of_2_05 :heavy_check_mark: AC 78 ms 12 MB
g++ near_power_of_2_06 :heavy_check_mark: AC 74 ms 12 MB
g++ near_power_of_2_07 :heavy_check_mark: AC 77 ms 12 MB
g++ near_power_of_2_08 :heavy_check_mark: AC 77 ms 12 MB
g++ only_first_small_00 :heavy_check_mark: AC 146 ms 20 MB
g++ only_first_small_01 :heavy_check_mark: AC 149 ms 20 MB
g++ random_00 :heavy_check_mark: AC 108 ms 16 MB
g++ random_01 :heavy_check_mark: AC 118 ms 17 MB
g++ random_02 :heavy_check_mark: AC 63 ms 10 MB
g++ small_00 :heavy_check_mark: AC 3 ms 4 MB
g++ small_01 :heavy_check_mark: AC 2 ms 4 MB
g++ small_02 :heavy_check_mark: AC 2 ms 4 MB
g++ small_03 :heavy_check_mark: AC 2 ms 4 MB
g++ small_04 :heavy_check_mark: AC 2 ms 4 MB
g++ small_05 :heavy_check_mark: AC 2 ms 4 MB
g++ small_06 :heavy_check_mark: AC 2 ms 4 MB
g++ small_07 :heavy_check_mark: AC 2 ms 4 MB
g++ small_08 :heavy_check_mark: AC 2 ms 4 MB
g++ small_slopes_00 :heavy_check_mark: AC 142 ms 20 MB
g++ small_slopes_01 :heavy_check_mark: AC 152 ms 20 MB
Back to top page