T两个点 求原因
查看原帖
T两个点 求原因
390898
app1eDog楼主2022/12/30 17:51

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;
}
2022/12/30 17:51
加载中...