马风良好,清楚症结所在(不足20行),玄2关求助
查看原帖
马风良好,清楚症结所在(不足20行),玄2关求助
302394
dingshengyang楼主2023/9/27 14:16

求调 rank1,rank2 两函数的实现正确性与使用正确性

#include<bits/stdc++.h>
using namespace std;
const int N = 1000005;

struct Splay{
	struct node{
		int v;int p,cnt,size,s[2];
		void init(int _v,int _p){
			v = _v,p = _p;
			cnt = size = 1;
		}
	}tr[N];
	int idx;
    int root;
    void pushup(int u){
		tr[u].size = tr[u].cnt + tr[tr[u].s[0]].size + tr[tr[u].s[1]].size;
	}
	void rotate(int x){
		int y = tr[x].p,z = tr[y].p;
		int k = tr[y].s[1] == x;
		tr[x].p = z,tr[z].s[tr[z].s[1] == y] = x;
		tr[tr[x].s[k^1]].p = y,tr[y].s[k] = tr[x].s[k^1];
		tr[y].p = x;tr[x].s[k^1] = y;
		pushup(y),pushup(x);
	}
	void splay(int x,int k){
		while(tr[x].p != k){
			int y = tr[x].p,z = tr[y].p;
			if(z != k){
				if((tr[z].s[1] == y) ^ (tr[y].s[1] == x)){
					rotate(x);
				}else rotate(y);
			}
			rotate(x);
		}
		if(!k)root = x;
	}
	public:
	void insert(int v){
		int u = root,p = 0;
		while(u){
			if(tr[u].v == v){
				tr[u].cnt ++;
				splay(u,0);
				return;
			}
			p = u,u = tr[u].s[tr[u].v<v];
		} 
		u = ++idx;
		tr[u].init(v,p);
		if(p){
			tr[p].s[tr[p].v<v] = u;
		}
		splay(u,0);
	}
	int pre(int v,bool ret_vaule = false){
		int u = root,p;
		while(u){
			if(tr[u].v<v)p = u,u = tr[u].s[1];
			else u = tr[u].s[0];
		}
		if(!ret_vaule)return p;
		return tr[p].v;
	}
	int succ(int v,bool ret_vaule = false){
		int u = root,p;
		while(u){
			if(tr[u].v>v)p = u,u = tr[u].s[0];
			else u = tr[u].s[1];
		}
		if(!ret_vaule)return p;
		return tr[p].v;
	}
	void erase(int x){
		int L = pre(x);
		int R = succ(x);
		splay(L,0);
		splay(R,L);
		if(tr[R].s[0] == 0)return;
		if(--tr[tr[R].s[0]].cnt == 0)tr[R].s[0] = 0;
		else splay(tr[R].s[0],0);
	}
	int FindKth(int k){
		int u = root;k ++;
		while(1){
			if(tr[tr[u].s[0]].size >= k)u = tr[u].s[0];
			else if(tr[tr[u].s[0]].size + tr[u].cnt >= k)return splay(u,0),tr[u].v;
			else k -= tr[tr[u].s[0]].size + tr[u].cnt,u = tr[u].s[1];
		}
	}
	int Find(int x){
		int u = root;
		while(u){
			if(tr[u].v == x)return u;
			u = tr[u].s[x>tr[u].v];
		}
		return 0;
	}
	int rank1(int x){//x < y how many numbers greater than me
		int pos = Find(x); 
        if(!pos)pos = pre(x);
		if(pos){
			splay(pos,0);
			return tr[tr[pos].s[1]].size - 1; 
		}
		return 0;
	}
	int rank2(int x){//x > y
		int pos = Find(x); 
        if(!pos)pos = succ(x);
		if(pos){
			splay(pos,0);
			return tr[tr[pos].s[0]].size - 1; 
		}
		return 0;
	}
	void out(int u){
		if(tr[u].s[0])out(tr[u].s[0]);
		printf("%d X %d\n",tr[u].v,tr[u].cnt);
		if(tr[u].s[1])out(tr[u].s[1]);
	}
	int size(){
		return tr[root].size-2;
	}
	Splay(){
		insert(-0x3f3f3f3f),insert(0x3f3f3f3f);
        // puts("done");
	}
}tr1,tr2;
//tr1:x < y
//tr2:x > y
vector<pair<int,int>> opt;
int rmed[N];
int tag;
int main(){
    int n;
    cin >> n;
    while(n --){
        string op;
        int a,b,c;
        cin >> op >> a;
        if(op == "Add"){
            cin >> b >> c;
            if(a > 0)tr2.insert(floor(double(c-b)/a)),
                opt.push_back({floor(double(c-b)/a),2});
            if(a < 0)tr1.insert(ceil(double(c-b)/a)),
                opt.push_back({ceil(double(c-b)/a),1});
            if(a == 0){
                if(b > c)tag ++;
                opt.push_back({b > c,0});
            }
        }else if(op == "Del"){
            a --;
            if(rmed[a])continue;
            rmed[a] = 1;
            if(opt[a].second == 2)tr2.erase(opt[a].first);
            if(opt[a].second == 1)tr1.erase(opt[a].first);
            if(opt[a].second == 0)tag -= opt[a].first;
        }else{
            cout << tag + tr1.rank1(a) + tr2.rank2(a) << endl; 
        }
		tr1.out(tr1.root);

		puts("\n\n\n");

		tr2.out(tr2.root);
    }
    return 0;
}
2023/9/27 14:16
加载中...