fhq-treap50分RE求助
查看原帖
fhq-treap50分RE求助
469345
Sherlock___Holmes楼主2022/10/7 20:53

MnZn照着oi-wiki学的,所以用的是指针请见谅,另外码风丑请见谅

#include <cstdio>
#include <cstdlib>
#include <tuple>
#define int long long
#define re register
#define gc() getchar()
const int inf = 1e13;
const int mod = 1e6;
inline int read(){
	re int x = 0;re char c = gc();
	while (c < '0' || c > '9') c = gc();
	while (c >= '0' && c <= '9') x = (x << 1) + (x << 3) + (c ^ 48) , c = gc();
	return x;
}
struct Node{
	Node *ch[2];
	int val , prio , siz , cnt;
	Node (int val) : val(val) , siz(1) , cnt(1){
		ch[0] = ch[1] = nullptr;
		prio = rand();
	} Node (Node *_node){
		val = _node -> val , siz = _node -> siz , cnt = _node -> cnt , prio = _node -> prio;
	} inline void upd_siz(){
		siz = cnt;
		if (ch[0] != nullptr) siz += ch[0] -> siz;
		if (ch[1] != nullptr) siz += ch[1] -> siz;
	}
};
struct fhq_treap{
	Node *root;
	std::pair <Node * , Node *> split(Node *cur , re int key){
		if (cur == nullptr) return {nullptr , nullptr};
		if (cur -> val <= key){
			auto temp = split(cur -> ch[1] , key);
			cur -> ch[1] = temp.first;
			cur -> upd_siz();
			return {cur , temp.second};
		} else{
			auto temp = split(cur -> ch[0] , key);
			cur -> ch[0] = temp.second;
			cur -> upd_siz();
			return {temp.first , cur};
		}
	} std::tuple <Node * , Node * , Node *> split_by_rk(Node *cur , re int rk){
		if (cur == nullptr) return {nullptr , nullptr , nullptr};
		re int ls_size = cur -> ch[0] == nullptr ? 0 : cur -> ch[0] -> siz;
		if (rk <= ls_size){
			Node *l , *mid , *r;
			std::tie(l , mid , r) = split_by_rk(cur -> ch[0] , rk);
			cur -> ch[0] = r;
			cur -> upd_siz();
			return {l , mid , cur};
		} else if (rk <= ls_size + cur -> cnt){
			Node *lt = cur -> ch[0];
			Node *rt = cur -> ch[1];
			cur -> ch[0] = cur -> ch[1] = nullptr;
			return {lt , cur , rt};
		} else{
			Node *l , *mid , *r;
			std::tie(l , mid , r) = split_by_rk(cur -> ch[1] , rk - ls_size - cur -> cnt);
			cur -> ch[1] = l;
			cur -> upd_siz();
			return {cur , mid , r};
		}
	} Node* merge(Node* u , Node* v){
		if (u == nullptr && v == nullptr) return nullptr;
		if (u == nullptr && v != nullptr) return v;
		if (u != nullptr && v == nullptr) return u;
		if (u -> prio < v -> prio){
			u -> ch[1] = merge(u -> ch[1] , v);
			u -> upd_siz();
			return u;
		} else{
			v -> ch[0] = merge(v -> ch[0] , u);
			v -> upd_siz();
			return v;
		}
	} inline void insert(re int val){
		auto temp = split(root , val);
		auto l_tr = split(temp.first , val - 1);
		Node *new_node;
		if (l_tr.second == nullptr) new_node = new Node(val);
		else{
			++ l_tr.second -> cnt;
			l_tr.second -> upd_siz();
		} Node *l_tr_combined = merge(l_tr.first , l_tr.second == nullptr ? new_node : l_tr.second);
		root = merge(l_tr_combined , temp.second);
	} inline void del(re int val){
		auto temp = split(root , val);
		auto l_tr = split(temp.first , val - 1);
		if (l_tr.second -> cnt > 1){
			-- l_tr.second -> cnt;
			l_tr.second -> upd_siz();
			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);
	} inline int qrank_by_val(Node *cur , re int val){
		auto temp = split(cur , val - 1);
		re int ret = (temp.first == nullptr ? 0 : temp.first -> siz) + 1;
		root = merge(temp.first , temp.second);
		return ret;
	} inline int qval_by_rank(Node *cur , re int rk){
		Node *l , *mid , *r;
		std::tie(l , mid , r) = split_by_rk(cur , rk);
		re int ret = mid -> val;
		root = merge(merge(l , mid) , r);
		return ret;
	} inline int qprev(re int val){
		auto temp = split(root , val - 1);
		re int ret;
		if (temp.first == nullptr) ret = -inf;
		else ret = qval_by_rank(temp.first , temp.first -> siz);
		root = merge(temp.first , temp.second);
		return ret;
	} inline int qnext(re int val){
		auto temp = split(root , val);
		re int ret;
		if (temp.second == nullptr) ret = inf;
		else ret = qval_by_rank(temp.second , 1);
		root = merge(temp.first , temp.second);
		return ret;
	}
}tr;
signed main(){
	srand(time_t(0));
	re int n = read() , cnt = 0 , ans = 0;
	while (n --){
		re int a = read() , b = read();
		if (cnt == 0) tr.insert(b);
		else if (cnt > 0){
			if (a == 1) tr.insert(b);
			else{
				re int pick1 = tr.qprev(b) , pick2 = tr.qnext(b);
				if (pick2 - b < b - pick1){
					if (pick2 != inf)tr.del(pick2);
					(ans += pick2 - b) %= mod;
				} else{
					if (pick1 != -inf)tr.del(pick1);
					(ans += b - pick1) %= mod;
				}
			}
		} else{
			if (a == 0) tr.insert(b);
			else{
				re int pick1 = tr.qprev(b) , pick2 = tr.qnext(b);
				if (pick2 - b < b - pick1){
					tr.del(pick2);
					(ans += pick2 - b) %= mod;
				} else{
					tr.del(pick1);
					(ans += b - pick1) %= mod;
				}
			}
		}
		cnt += (a == 1 ? 1 : -1);
	}
	printf("%lld" , ans);
	return 0;
}
2022/10/7 20:53
加载中...