求助最后一个测试点
查看原帖
求助最后一个测试点
905677
whysoseriousQAQ楼主2023/8/27 12:34

我在get_k的地方多进行了一次旋转,结果wa在最后一个点了,不知道为什么

#include <iostream>
#include <algorithm>

using namespace std;
const int N = 1e5 + 10, inf = 1e9;

struct node
{
    int s[2], p, v, sz;
    void init(int _v, int _p)
    {
        v = _v, p = _p;
        sz = 1;
    }
}tr[N];
int root, idx;
int L, R;
int delta;

void pushup(int x)
{
    tr[x].sz = tr[tr[x].s[0]].sz + tr[tr[x].s[1]].sz + 1;
}

void rotate(int x)
{
    int y = tr[x].p, z = tr[y].p;
    int k = tr[y].s[1] == x;
    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;
}

int insert(int v)
{
    int u = root, p = 0;
    while (u) p = u, u = tr[u].s[v > tr[u].v];
    u = ++ idx;
    if (p) tr[p].s[v > tr[p].v] = u;
    tr[u].init(v, p);
    splay(u, 0);
    return u;
}
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_suf(int v)
{
    find(v);
    int x = root;
    if (tr[x].v >= v) return x;
    x = tr[x].s[1];
    while (tr[x].s[0]) x = tr[x].s[0];
    splay(x, 0);
    return x;
}

void del(int v)
{
    int l = L, r = get_suf(v);
    splay(l, 0), splay(r, l);
    tr[r].s[0] = 0;
    splay(r, 0);
}

int get_k(int k)
{
    int u = root;
    while (u)
    {
        if (tr[tr[u].s[0]].sz >= k) u = tr[u].s[0];
        else if (tr[tr[u].s[0]].sz + 1 == k) break;
        else k -= tr[tr[u].s[0]].sz + 1, u = tr[u].s[1];
    }
    //splay(u, 0)为什么在此处多旋转一次就会wa最后一个点
    return tr[u].v;
}
int main()
{
    int n, m; cin >> n >> m;
    L = insert(-inf), R = insert(inf);
    int tot = 0;
    while (n -- )
    {
        char op[2];
        int k; cin >> op >> k;
        if (*op == 'I')
        {
            if (k >= m) k -= delta, insert(k), tot ++ ;
        }
        else if (*op == 'A') delta += k;
        else if (*op == 'S')
        {
            delta -= k;
            del(m - delta);
        }
        else
        {
            if (tr[root].sz - 2 < k) cout << -1 << endl;
            else cout << get_k(tr[root].sz - k) + delta << endl;
        }
    }

    cout << tot - (tr[root].sz - 2) << endl;
}
2023/8/27 12:34
加载中...