60分,错的询问好像都是3和4:
#include <iostream>
#include <cassert>
#include <climits>
#include <array>
struct Node {
int num = 0, cnt = 1, h = 1, siz = 1;
Node *ls = nullptr, *rs = nullptr;
void update();
// Balance defined as h(l) - h(r)
int balance();
};
constexpr int MAXN = 100'005;
std::array<Node, MAXN> nodes;
int nodes_cnt = 0;
int get_h(Node* p) {
return p ? p->h : 0;
}
int get_siz(Node* p) {
return p ? p->siz : 0;
}
void Node::update() {
siz = cnt + get_siz(ls) + get_siz(rs);
h = 1 + std::max(get_h(ls), get_h(rs));
}
int Node::balance() {
return get_h(ls) - get_h(rs);
}
// Handles the left left situation
void left_plus(Node*& curr) {
auto ls = curr->ls;
curr->ls = ls->rs;
ls->rs = curr;
curr->update();
ls->update();
curr = ls;
}
// Handles the right right situation
void right_plus(Node*& curr) {
auto rs = curr->rs;
curr->rs = rs->ls;
rs->ls = curr;
curr->update();
rs->update();
curr = rs;
}
// Handles left right
void left_right(Node*& curr) {
right_plus(curr->ls);
left_plus(curr);
}
// Handles right left
void right_left(Node*& curr) {
left_plus(curr->rs);
right_plus(curr);
}
// Maintain the AVL properties at the node and updates info
void maintain(Node*& curr) {
curr->update();
if (curr->balance() > 1) {
if (curr->balance() >= 0)
left_plus(curr);
else
left_right(curr);
} else if (curr->balance() < -1) {
if (curr->balance() <= 0)
right_plus(curr);
else
right_left(curr);
}
}
// Inserts number x into subtree
void insert(int x, Node*& curr) {
if (!curr) {
curr = &nodes[++nodes_cnt];
curr->num = x;
return;
}
if (x == curr->num) {
curr->cnt++;
curr->update();
return;
}
if (x < curr->num)
insert(x, curr->ls);
else if (x > curr->num)
insert(x, curr->rs);
maintain(curr);
}
// Finds the max element in subtree, returns a copy of the node.
Node find_max(Node* curr) {
while (curr->rs)
curr = curr->rs;
return *curr;
}
// Removes x from subtree mul times
void remove(int x, Node*& curr, int mul = 1) {
if (!curr)
return;
if (x == curr->num) {
curr->cnt -= mul;
if (curr->cnt) {
// No structural changes for curr
curr->update();
return;
}
// Now we know curr->cnt == 0
if (!curr->ls && !curr->rs)
// Leaf node
curr = nullptr;
else if (curr->ls && !curr->rs)
curr = curr->ls;
else if (!curr->ls && curr->rs)
curr = curr->rs;
else {
auto lmax = find_max(curr->ls);
curr->num = lmax.num;
curr->cnt = lmax.cnt;
remove(lmax.num, curr->ls, lmax.cnt);
}
if (curr)
maintain(curr);
return;
} else if (x < curr->num) {
remove(x, curr->ls);
} else {
remove(x, curr->rs);
}
maintain(curr);
}
// Get the cnt of numbers smaller than x
int small_cnt(int x, Node* curr) {
if (!curr)
return 0;
return curr->num < x ? get_siz(curr->ls) + curr->cnt + small_cnt(x, curr->rs) : small_cnt(x, curr->ls);
}
// Get by rank
int by_rank(int rnk, Node* curr) {
if (rnk <= get_siz(curr->ls))
return by_rank(rnk, curr->ls);
if (get_siz(curr->ls) + 1 <= rnk && rnk <= get_siz(curr->ls) + curr->cnt)
return curr->num;
if (rnk > get_siz(curr->ls) + curr->cnt)
return by_rank(rnk - get_siz(curr->ls) - curr->cnt, curr->rs);
}
// Get predecessor of x
int get_pred(int x, Node* curr) {
if (!curr)
return INT_MIN;
return curr->num >= x ? get_pred(x, curr->ls) : std::max(curr->num, get_pred(x, curr->rs));
}
// Get successor of x
int get_suc(int x, Node* curr) {
if (!curr)
return INT_MAX;
return curr->num <= x ? get_suc(x, curr->rs) : std::min(curr->num, get_suc(x, curr->ls));
}
int main() {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
std::cout.tie(nullptr);
int n;
Node* root = nullptr;
std::cin >> n;
while (n--) {
int op, x;
std::cin >> op >> x;
switch (op) {
case 1:
insert(x, root);
break;
case 2:
remove(x, root);
break;
case 3:
std::cout << small_cnt(x, root) + 1 << '\n';
break;
case 4:
std::cout << by_rank(x, root) << '\n';
break;
case 5:
std::cout << get_pred(x, root) << '\n';
break;
case 6:
std::cout << get_suc(x, root) << '\n';
break;
default:
return 2;
}
}
}