刚学splay的蒟蒻,在get_rk中尝试了两种不同的写法,代码中被注销的是正确写法,可是错误的写法依然AC了这道题,hack方式很简单,对于
4
1 10
1 30
1 20
3 21
而言,应该输出3,可是下面的代码却输出的4,我想应该是哪里出了问题
#include <iostream>
using namespace std;
const int N = 1e5 + 10, inf = 1e9;
struct node
{
int s[2], v, p, sz, cnt;
void init(int _v, int _p)
{
v = _v, p = _p;
sz = 1;
cnt = 1;
}
}tr[N];
int root, idx;
void pushup(int x)
{
tr[x].sz = tr[tr[x].s[0]].sz + tr[tr[x].s[1]].sz + tr[x].cnt;
}
void rotate(int x)
{
int y = tr[x].p, z = tr[y].p;
int k = tr[y].s[1] == x;//k为0说明是左子树
tr[z].s[tr[z].s[1] == y] = x, tr[x].p = z;
tr[y].s[k] = tr[x].s[k ^ 1], tr[tr[x].s[k ^ 1]].p = y;
tr[x].s[k ^ 1] = y, tr[y].p = x;
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;
}
void find(int v)
{
int x = root;
while (tr[x].s[v > tr[x].v] && tr[x].v != v) x = tr[x].s[v > tr[x].v];
splay(x, 0);
}
int get_sp(int v, bool flag)//suf是true,pre是false
{
find(v);
int x = root;
if (!flag && tr[x].v == v) goto c;
if ((tr[x].v > v) == flag) return x;
c:
x = tr[x].s[flag];
while (tr[x].s[flag ^ 1]) x = tr[x].s[flag ^ 1];
splay(x, 0);
return x;
}
void del(int v)
{
int l = get_sp(v, false), r = get_sp(v, true);
splay(l, 0), splay(r, l);
int x = tr[r].s[0];
if (tr[x].cnt > 1)
{
tr[x].cnt --;
splay(x, 0);
}
else
{
tr[r].s[0] = 0;
splay(r, 0);
}
}
void insert(int v)
{
int x = root, p = 0;
while (x && tr[x].v != v) p = x, x = tr[x].s[v > tr[x].v];
if (x) tr[x].cnt ++;
else
{
x = ++ idx;
tr[p].s[v > tr[p].v] = x;
tr[x].init(v, p);
}
splay(x, 0);
}
int get_rk(int v)//找到某个数前面有多少数
{
// insert(v);
find(v);
// cout << tr[root].v << endl;
int res = tr[tr[root].s[0]].sz;
if (tr[root].v != v) {
// insert(v);
// res = tr[tr[root].s[0]].sz;
// del(v);
// return res;
return res + 1;
}
// else res = tr[tr[root].s[0]].sz;
return res;
}
int get_k(int k)
{
int x = root;
while (true)
{
int y = tr[x].s[0];
if (tr[y].sz + tr[x].cnt < k)
{
k = k - tr[y].sz - tr[x].cnt;
x = tr[x].s[1];
}
else if (tr[y].sz >= k) x = y;
else break;
}
return tr[x].v;
}
int main()
{
insert(-inf), insert(inf);
int n; cin >> n;
for (int i = 0; i < n; i ++)
{
int op, x; cin >> op >> x;
if (op == 1) insert(x);
if (op == 2) del(x);
if (op == 3) cout << get_rk(x) << endl;
if (op == 4) cout << get_k(x + 1) << endl;
if (op == 5) cout << tr[get_sp(x, false)].v << endl;
if (op == 6) cout << tr[get_sp(x, true)].v << endl;
}
}