这题是不是数据有一些水
查看原帖
这题是不是数据有一些水
905677
whysoseriousQAQ楼主2023/8/28 00:59

刚学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;
    }
}
2023/8/28 00:59
加载中...