求问卡常
查看原帖
求问卡常
573341
MiniLong楼主2023/5/17 17:02

线段树套splay

有无大佬能帮忙卡卡splay,或者给点卡常建议

#include <bits/stdc++.h>
#define _rep(i, x, y) for(int i = x; i <= y; ++i)
#define _req(i, x, y) for(int i = x; i >= y; --i)
#define _rev(i, u) for(int i = head[u]; i; i = e[i].nxt)
#define pb(x) push_back(x)
#define mst(f, i) memset(f, i, sizeof f)
using namespace std;
#ifdef ONLINE_JUDGE
#define debug(...) 0
#else
#define debug(...) fprintf(stderr, __VA_ARGS__), fflush(stderr)
#endif
namespace fastio{
    char ibuf[50007],*p1 = ibuf, *p2 = ibuf;
    #ifdef ONLINE_JUDGE
    #define get() p1 == p2 && (p2 = (p1 = ibuf) + fread(ibuf, 1, 50007, stdin), p1 == p2) ? EOF : *p1++
    #else
    #define get() getchar()
    #endif
    template<typename T> inline void read(T &t){
        T x = 0, f = 1;
        char c = getchar();
        while(!isdigit(c)){
            if(c == '-') f = -f;
            c = getchar();
        }
        while(isdigit(c)) x = x * 10 + c - '0', c = getchar();
        t = x * f;
    }
    template<typename T, typename ... Args> inline void read(T &t, Args&... args){
        read(t);
        read(args...);
    }
    template<typename T> void write(T t){
        if(t < 0) putchar('-'), t = -t;
        if(t >= 10) write(t / 10);
        putchar(t % 10 + '0');
    }
    template<typename T, typename ... Args> void write(T t, Args... args){
        write(t), putchar(' '), write(args...);
    }
    template<typename T> void writeln(T t){
        write(t);
        puts("");
    }
    template<typename T> void writes(T t){
        write(t), putchar(' ');
    }
};
using namespace fastio;
typedef long long ll;
typedef pair<int, int> PII;
const int N = 5e4 + 5, inf = 2147483647;
int n, m, cnt, w[N];
struct node{
	int fa, siz, son[2], num, val;
}a[N * 200];
struct SplayTree{
	#define ls a[x].son[0]
	#define rs a[x].son[1]
	int rt;
	inline int New(int val, int fa){a[++cnt] = node{fa, 1, {0, 0}, 1, val}; return cnt;};
	inline void update(int x){a[x].siz = a[ls].siz + a[rs].siz + a[x].num;}
	inline void rotate(int x){
		int y = a[x].fa, z = a[y].fa, k = a[y].son[1] == x;
		a[x].fa = z, a[z].son[a[z].son[1] == y] = x;
		a[y].son[k] = a[x].son[k ^ 1], a[a[x].son[k ^ 1]].fa = y;
		a[x].son[k ^ 1] = y, a[y].fa = x;
		update(y), update(x);
	}
	inline void splay(int x, int goal){
		while(a[x].fa != goal){
			int y = a[x].fa, z = a[y].fa;
			if(z != goal) rotate((a[z].son[1] == y) ^ (a[y].son[1] == x) ? x : y);
			rotate(x);
		}
		if(!goal) rt = x;
	}
	inline int getval(int k){
		int x = rt;
		while(x){
			if(k <= a[ls].siz) x = ls;
			else if(k <= a[ls].siz + a[x].num) return splay(x, 0), a[x].val;
			else k -= a[x].num + a[ls].siz, x = rs;
		}
	}
	inline int getrk(int val){
		int x = rt, res = 0;
		while(x){
			if(val < a[x].val) x = ls;
			else if(val == a[x].val) return splay(x, 0), res + a[ls].siz;
			else res += a[ls].siz + a[x].num, x = rs;
		}
	}
	inline void insert(int val){
		if(!rt){
			rt = New(val, 0);
			return;
		}
		int fa = 0, x = rt;
		for(; x; fa = x, x = a[x].son[val > a[x].val]){
			if(a[x].val == val){
				a[x].num++, splay(x, 0);
				return; 
			}
		}
		x = New(val, fa); if(fa) a[fa].son[val > a[fa].val] = x;
		splay(x, 0);
	}
	inline int getid(int val){
		int x = rt;
		for(; a[x].val != val && x; x = a[x].son[val > a[x].val]);
		splay(x, 0); return x;
	}
	inline int pre(int val){
		int x = getid(val);
		if(a[x].val < val) return x; x = a[x].son[0];
		while(a[x].son[1]) x = a[x].son[1];
		return x;
	}
	inline int nxt(int val){
		int x = getid(val);
		if(a[x].val > val) return x; x = a[x].son[1];	
		while(a[x].son[0]) x = a[x].son[0];
		return x;		
	}
	inline void del(int val){
		int l = pre(val), r = nxt(val);
		if(!l || !r){
			return;
		}
		splay(l, 0), splay(r, l);
		if(a[a[r].son[0]].num > 1) a[a[r].son[0]].num--, splay(a[r].son[0], 0);
		else a[a[r].son[0]].fa = a[r].son[0] = 0;
		update(r), update(l); 
	}
	#undef ls
	#undef rs
}tr[N << 2];
#define ls x << 1
#define rs x << 1 | 1
void build(int x, int l, int r){
	tr[x].insert(-inf), tr[x].insert(inf);
	_rep(i, l, r) tr[x].insert(w[i]);
	if(l == r) return;
	int mid = l + r >> 1;
	build(ls, l, mid), build(rs, mid + 1, r);
}
int qrk(int x, int l, int r, int L, int R, int val){
	if(l >= L && r <= R){
		tr[x].insert(val); 
		int res = tr[x].getrk(val);
		tr[x].del(val);
		return res;
	}
	int mid = l + r >> 1, res = 0;
	if(L <= mid) res += qrk(ls, l, mid, L, R, val);
	if(R > mid) res += qrk(rs, mid + 1, r, L, R, val);
	return res;
}
inline int qval(int L, int R, int k){
	int l = 0, r = 1e8, res = 0;
	while(l <= r){
		int mid = l + r >> 1;
		if(qrk(1, 1, n, L, R, mid) + 1 <= k){
			res = mid, l = mid + 1;
		}
		else r = mid - 1;
	}
	return res;
}
int qpre(int x, int l, int r, int L, int R, int val){
	if(l >= L && r <= R){
		tr[x].insert(val); 
		int res = a[tr[x].pre(val)].val; 
		tr[x].del(val);
		return res;
	}
	int mid = l + r >> 1, res = -inf;
	if(L <= mid) res = max(res, qpre(ls, l, mid, L, R, val));
	if(R > mid) res = max(res, qpre(rs, mid + 1, r, L, R, val));
	return res;
}
int qnxt(int x, int l, int r, int L, int R, int val){
	if(l >= L && r <= R){
		tr[x].insert(val); int res = a[tr[x].nxt(val)].val; tr[x].del(val);
		return res;
	}
	int mid = l + r >> 1, res = inf;
	if(L <= mid) res = min(res, qnxt(ls, l, mid, L, R, val));
	if(R > mid) res = min(res, qnxt(rs, mid + 1, r, L, R, val));
	return res;
}
void modify(int x, int l, int r, int p, int val){
	tr[x].del(w[p]);
	tr[x].insert(val);
	if(l == r){w[p] = val; return;}
	int mid = l + r >> 1;
	if(p <= mid) modify(ls, l, mid, p, val);
	else modify(rs, mid + 1, r, p, val);
}
int main(){
	read(n, m);
	_rep(i, 1, n) read(w[i]);
	build(1, 1, n);
	while(m--){
		int opt, l, r, k, p;
		read(opt);
		if(opt == 1){
			read(l, r, k);
			writeln(qrk(1, 1, n, l, r, k) + 1);
		}
		if(opt == 2){
			read(l, r, k);
			writeln(qval(l, r, k));
		}
		if(opt == 3){
			read(p, k);
			modify(1, 1, n, p, k);
		}
		if(opt == 4){
			read(l, r, k);
			writeln(qpre(1, 1, n, l, r, k));
		}
		if(opt == 5){
			read(l, r, k);

			writeln(qnxt(1, 1, n, l, r, k));
		}
	}
    return 0;
}
2023/5/17 17:02
加载中...