T了第二和第九个点。
ACwing 和 Loj 上的树套树板子题能过。
不知道是不是我的 Treap 板子太慢了,平衡树的板子题跑 400 多毫秒。
有点长,码风也丑···
// Problem: P3380 【模板】二逼平衡树(树套树)
// Contest: Luogu
// URL: https://www.luogu.com.cn/problem/P3380
// Memory Limit: 512 MB
// Time Limit: 2000 ms
//
// Powered by CP Editor (https://cpeditor.org)
// created on Lucian Xu's Laptop
#include <cstdio>
#include <iostream>
#include <cstring>
#include <cmath>
#include <algorithm>
#include <vector>
#include <stack>
#include <map>
#include <set>
#include <queue>
#include <utility>
#include <unordered_map>
#include <ctime>
#include <random>
// using namespace std;
typedef unsigned int UI;
typedef unsigned long long ULL;
typedef long long LL;
typedef unsigned long long ULL;
typedef std::pair<int, int> PII;
typedef std::pair<int, LL> PIL;
typedef std::pair<LL, int> PLI;
typedef std::pair<LL, LL> PLL;
#define rep(i, l, r) for(auto i = (l); i <= (r); i++)
#define per(i, r, l) for(auto i = (r); i >= (l); i--)
#define ff first
#define ss second
#define makepair make_pair
#define pushback push_back
#define endl '\n'
#define all(v) v.begin(), v.end()
#define rall(v) v.rbegin(), v.rend()
const int N = 2e5+10;
const int mod = 998244353;
const int inf = 2147483647;
const LL INF = 1e18;
const double pi = acos(-1.0);
const double eps = 1e-6;
ULL add(ULL a, ULL b) {return a + b < mod ? a + b : a + b - mod;}
ULL minus(ULL a, ULL b) {return a < b ? a + mod - b : a - b;}
ULL mul(ULL a, ULL b) {return a * b % mod;}
int n, m, op, l, r, pos, key, root_tot;
int a[N];
struct node2{
node2 *ch[2];
int key, val;
int cnt, size;
node2(int _key) : key(_key), cnt(1), size(1) {
ch[0] = ch[1] = nullptr;
val = rand();
}
// node2(node2 *_node2) {
// key = _node2->key, val = _node2->val, cnt = _node2->cnt, size = _node2->size;
// }
inline void push_up() {
size = cnt;
if(ch[0] != nullptr) size += ch[0] -> size;
if(ch[1] != nullptr) size += ch[1] -> size;
}
};
struct treap{
#define _2 second.first
#define _3 second.second
node2 *root;
std::pair<node2 *, node2 *> split(node2 *p, int key){
if(p == nullptr) return {nullptr, nullptr};
if(p -> key <= key) {
auto temp = split(p -> ch[1], key);
p -> ch[1] = temp.first;
p -> push_up();
return {p, temp.second};
}
else{
auto temp = split(p -> ch[0], key);
p -> ch[0] = temp.second;
p -> push_up();
return {temp.first, p};
}
}
std::pair<node2 *, std::pair<node2 *, node2 *> > split_by_rank(node2 *p, int rank) {
if (p == nullptr) return {nullptr, {nullptr, nullptr}};
int ls_size = p -> ch[0] == nullptr ? 0 : p -> ch[0] -> size;
if (rank <= ls_size) {
auto temp = split_by_rank(p -> ch[0], rank);
p -> ch[0] = temp._3;
p -> push_up();
return {temp.first, {temp._2, p}};
}
else if (rank <= ls_size + p -> cnt) {
node2 *lt = p -> ch[0];
node2 *rt = p -> ch[1];
p -> ch[0] = p -> ch[1] = nullptr;
return {lt, {p, rt}};
}
else {
auto temp = split_by_rank(p -> ch[1], rank - ls_size - p -> cnt);
p -> ch[1] = temp.first;
p -> push_up();
return {p, {temp._2, temp._3}};
}
}
node2 *merge(node2 *u, node2 *v) {
if(u == nullptr && v == nullptr) return nullptr;
if(u != nullptr && v == nullptr) return u;
if(v != nullptr && u == nullptr) return v;
if(u -> val < v -> val) {
u -> ch[1] = merge(u -> ch[1], v);
u -> push_up();
return u;
}
else{
v -> ch[0] = merge(u, v -> ch[0]);
v -> push_up();
return v;
}
}
void insert(int key) {
auto temp = split(root, key);
auto l_tr = split(temp.first, key - 1);
node2 *new_node2;
if (l_tr.second == nullptr) new_node2 = new node2(key);
else {
l_tr.second -> cnt++;
l_tr.second -> push_up();
}
node2 *l_tr_combined =
merge(l_tr.first, l_tr.second == nullptr ? new_node2 : l_tr.second);
root = merge(l_tr_combined, temp.second);
}
void remove(int key) {
auto temp = split(root, key);
auto l_tr = split(temp.first, key - 1);
if(l_tr.second -> cnt > 1) {
l_tr.second -> cnt--;
l_tr.second -> push_up();
l_tr.first = merge(l_tr.first, l_tr.second);
}
else{
if(temp.first == l_tr.second) temp.first = nullptr;
delete l_tr.second;
l_tr.second = nullptr;
}
root = merge(l_tr.first, temp.second);
}
int get_rank_by_key(node2 *p, int key) {
auto temp = split(p, key - 1);
int ret = (temp.first == nullptr ? 0 : temp.first -> size) + 1;
root = merge(temp.first, temp.second);
return ret;
}
int get_key_by_rank(node2 *p, int rank) {
auto temp = split_by_rank(p, rank);
int ret = temp._2 -> key;
root = merge(temp.first, merge(temp._2, temp._3));
return ret;
}
int get_prev(int key) {
auto temp = split(root, key - 1);
int ret = get_key_by_rank(temp.first, temp.first -> size);
root = merge(temp.first, temp.second);
return ret;
}
int get_nex(int key) {
auto temp = split(root, key);
int ret = get_key_by_rank(temp.second, 1);
root = merge(temp.first, temp.second);
return ret;
}
};
treap tr2[N << 4];
struct node1{
int l, r, root;
}tr1[N << 4];
void build(int u, int l, int r){
tr1[u] = {l, r, u};
root_tot = std::max(root_tot, u);
if(l == r) return;
int mid = (l + r) >> 1;
build(u << 1, l, mid), build(u << 1 | 1, mid + 1, r);
}
void modify(int u, int pos, int key){
tr2[u].insert(key);
if(tr1[u].l == tr1[u].r) return;
int mid = (tr1[u].l + tr1[u].r) >> 1;
if(pos <= mid) modify(u << 1, pos, key);
else modify(u << 1 | 1, pos, key);
}
int get_rank_by_key_in_interval(int u, int l, int r, int key){
if(l <= tr1[u].l && tr1[u].r <= r)
return tr2[u].get_rank_by_key(tr2[u].root, key) - 2;
int mid = (tr1[u].l + tr1[u].r) >> 1, ans = 0;
if(l <= mid) ans += get_rank_by_key_in_interval(u << 1, l, r, key);
if(mid < r) ans += get_rank_by_key_in_interval(u << 1 | 1, l, r, key);
return ans;
}
int get_key_by_rank_in_interval(int u, int l, int r, int rank){
int L = 0, R = 1e8;
while(L < R){
int mid = (L + R + 1) / 2;
if(get_rank_by_key_in_interval(1, l, r, mid) < rank) L = mid;
else R = mid - 1;
}
return L;
}
void change(int u, int pos, int pre_key, int key){
tr2[u].remove(pre_key);
tr2[u].insert(key);
if(tr1[u].l == tr1[u].r) return;
int mid = (tr1[u].l + tr1[u].r) >> 1;
if(pos <= mid) change(u << 1, pos, pre_key, key);
else change(u << 1 | 1, pos, pre_key, key);
}
int get_prev_in_interval(int u, int l, int r, int key){
if(l <= tr1[u].l && tr1[u].r <= r)
return tr2[u].get_prev(key);
int mid = (tr1[u].l + tr1[u].r) >> 1, ans = -inf;
if(l <= mid) ans = std::max(ans, get_prev_in_interval(u << 1, l, r, key));
if(mid < r) ans = std::max(ans, get_prev_in_interval(u << 1 | 1, l, r, key));
return ans;
}
int get_nex_in_interval(int u, int l, int r, int key){
if(l <= tr1[u].l && tr1[u].r <= r)
return tr2[u].get_nex(key);
int mid = (tr1[u].l + tr1[u].r) >> 1, ans = inf;
if(l <= mid) ans = std::min(ans, get_nex_in_interval(u << 1, l, r, key));
if(mid < r) ans = std::min(ans, get_nex_in_interval(u << 1 | 1, l, r, key));
return ans;
}
int main(){
std::ios::sync_with_stdio(false);
std::cin.tie(0);
std::cout.tie(0);
srand(time(0));
std::cin >> n >> m;
build(1, 1, n);
rep(i, 1, n){
std::cin >> a[i];
modify(1, i, a[i]);
}
rep(i, 1, root_tot){
tr2[i].insert(inf), tr2[i].insert(-inf);
}
rep(i, 1, m){
std::cin >> op;
if(op == 1){
std::cin >> l >> r >> key;
std::cout << get_rank_by_key_in_interval(1, l, r, key) + 1 << endl;
}
if(op == 2){
std::cin >> l >> r >> key;
std::cout << get_key_by_rank_in_interval(1, l, r, key) << endl;
}
if(op == 3){
std::cin >> pos >> key;
change(1, pos, a[pos], key);
a[pos] = key;
}
if(op == 4){
std::cin >> l >> r >> key;
std::cout << get_prev_in_interval(1, l, r, key) << endl;
}
if(op == 5){
std::cin >> l >> r >> key;
std::cout << get_nex_in_interval(1, l, r, key) << endl;
}
}
return 0;
}