萌新求助
查看原帖
萌新求助
905677
whysoseriousQAQ楼主2023/8/31 01:04

这是在洛谷上可以ac的代码,让我很迷惑的是get_suc那里,如果我使用注销掉的代码代替现有的代码,在另一个oj上就会wa在最后一个点,但我实在不理解这俩写法有什么差别

#include <iostream>

using namespace std;
const int N = 5e4 + 10, inf = 2147483647;
struct node
{
    int s[2], v, p;
    int sz;
    void init(int _v, int _p)
    {
        v = _v, p = _p;
        sz = 1;
    }
}tr[N << 6];

struct Segment
{
    int l, r;
}seg[N << 2];

int root[N << 2], idx, w[N];

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, int b)
{
    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[b] = x;
}

void insert(int v, int b)
{
    int x = root[b], p = 0;
    while (x) p = x, x = tr[x].s[v > tr[x].v];
    x = ++ idx;
    if (p) tr[p].s[v > tr[p].v] = x;
    tr[x].init(v, p);
    splay(x, 0, b);
}

void find(int v, int b)
{
    int x = root[b];
    while (tr[x].s[v > tr[x].v] && tr[x].v != v) x = tr[x].s[v > tr[x].v];
    splay(x, 0, b);
}

// int get_sp(int v, bool op, int b)
// {
//     find(v, b);
//     int x = root[b];
//     if (!op && tr[x].v == v) goto c;
//     if (v < tr[x].v == op) return x;
//     c:
//     x = tr[x].s[op];
//     while (tr[x].s[op ^ 1]) x = tr[x].s[op ^ 1];
//     return x;
// }
int get_pre(int v, int b)
{
    find(v, b);
    int u = root[b], res = -inf;
    if (tr[u].v < v) return tr[u].v;
    u = tr[u].s[0];
    while (tr[u].s[0] && tr[u].v == v) u = tr[u].s[0];
    while (tr[u].s[1]) u = tr[u].s[1];
    return tr[u].v;
}

// void find(int v, int b)
// {
//     int x = root[b];
//     while (tr[x].s[v > tr[x].v] && tr[x].v != v) x = tr[x].s[v > tr[x].v];
//     splay(x, 0, b);
// }
int get_suc(int v, int b)
{
    // find(v, b);
    int u = root[b], res = inf;
    // if (tr[u].v > v) return tr[u].v;
    // u = tr[u].s[1];
    // while(tr[u].s[1] && tr[u].v == v) u = tr[u].s[1];
    // while (tr[u].s[0]) u = tr[u].s[0];
    // return tr[u].v;  
    while (u)
    {
        // if (tr[u].v > v) res = min(res, tr[u].v), u = tr[u].s[0];
        // else u = tr[u].s[1];
        if (tr[u].v > v) res = min(res, tr[u].v);
        u = tr[u].s[v >= tr[u].v];
    }
    return res;

}

void update(int p, int q, int b)
{
    int x = root[b];
    while (x)
    {
        if (tr[x].v == p) break;
        x = tr[x].s[p > tr[x].v];
    }
    splay(x, 0, b);
    int l = tr[x].s[0], r = tr[x].s[1];
    while (tr[l].s[1]) l = tr[l].s[1];
    while (tr[r].s[0]) r = tr[r].s[0];
    splay(l, 0, b), splay(r, l, b);
    tr[r].s[0] = 0;
    pushup(r), pushup(l);
    insert(q, b);
}

int get_rk(int v, int b)
{
    int u = root[b], res = 0;
    while (u)
    {
        if (tr[u].v < v) res += tr[tr[u].s[0]].sz + 1, u = tr[u].s[1];
        else u = tr[u].s[0];
    }
    return res;
}
void build(int id, int l, int r)
{
    seg[id] = {l, r};
    // root[id] = id;
    insert(-inf, id), insert(inf, id);
    for (int i = l; i <= r; i ++) insert(w[i], id);
    if (l == r) return;
    int mid = l + r >> 1;
    build(id << 1, l, mid), build(id << 1 | 1, mid + 1, r);
}

void output(int u)//中序遍历输出
{
    if (tr[u].s[0]) output(tr[u].s[0]);
    if (tr[u].v != -inf && tr[u].v != inf) cout << tr[u].v << " ";
    if (tr[u].s[1]) output(tr[u].s[1]);
}

int query_rk(int id, int ql, int qr, int v)
{
    if (seg[id].l == ql && seg[id].r == qr) return get_rk(v, id) - 1;
    int mid = seg[id].l + seg[id].r >> 1;
    if (qr <= mid) return query_rk(id << 1, ql, qr, v);
    else if (ql > mid) return query_rk(id << 1 | 1, ql, qr, v);
    return query_rk(id << 1, ql, mid, v) + query_rk(id << 1 | 1, mid + 1, qr, v);
}

// int query_sp(int id, int ql, int qr, int v, bool op)
// {
//     if (seg[id].l == ql && seg[id].r == qr) return tr[get_sp(v, op, id)].v;
//     int mid = seg[id].l + seg[id].r >> 1;
//     if (qr <= mid) return query_sp(id << 1, ql, qr, v, op);
//     else if (ql > mid) return query_sp(id << 1 | 1, ql, qr, v, op);
//     else
//     {
//         if (!op) return max(query_sp(id << 1, ql, mid, v, op), query_sp(id << 1 | 1, mid + 1, qr, v, op));
//         else return min(query_sp(id << 1, ql, mid, v, op), query_sp(id << 1 | 1, mid + 1, qr, v, op));
//     }
// }
int query_pre(int id, int a, int b, int x)
{
    if (seg[id].l == a && seg[id].r == b) return get_pre(x, id);
    int mid = seg[id].l + seg[id].r >> 1;
    if (b <= mid) return query_pre(id << 1, a, b, x);
    else if (a > mid) return query_pre(id << 1 | 1, a, b, x);
    return max(query_pre(id << 1, a, mid, x), query_pre(id << 1 | 1, mid + 1, b, x));
}

int query_suc(int id, int a, int b, int x)
{
    if (seg[id].l == a && seg[id].r == b) return get_suc(x, id);
    int mid = seg[id].l + seg[id].r >> 1;
    if (b <= mid) return query_suc(id << 1, a, b, x);
    else if (a > mid) return query_suc(id << 1 | 1, a, b, x);
    return min(query_suc(id << 1, a, mid, x), query_suc(id << 1 | 1, mid + 1, b, x));
}

void change(int id, int pos, int v)
{
    update(w[pos], v, id);
    if (seg[id].l == seg[id].r) return;
    int mid = seg[id].l + seg[id].r >> 1;
    if (pos <= mid) change(id << 1, pos, v);
    else change(id << 1 | 1, pos, v);
}

int main()
{
    int n, m; cin >> n >> m;
    for (int i = 1; i <= n; i ++) cin >> w[i];
    build(1, 1, n);
    // output(root[1]);
    while (m --)
    {
        int op, l, r, pos, v; cin >> op;
        if (op != 3)
        {
            cin >> l >> r >> v;
            if (op == 1) cout << query_rk(1, l, r, v) + 1 << endl;
            else if (op == 2)
            {
                int L = 0, R = 1e8;
                while (L < R)
                {
                    int mid = L + R + 1 >> 1;
                    if (query_rk(1, l, r, mid) + 1 <= v) L = mid;
                    else R = mid - 1;
                }
                cout << L << endl;
            }
            else if (op == 4) cout << query_pre(1, l, r, v) << endl;
            else cout << query_suc(1, l, r, v) << endl;
            // else cout << query_sp(1, l, r, v, op - 4) << endl;
        }
        else 
        {
            cin >> pos >> v;
            change(1, pos, v);
            w[pos] = v;
        }
    }
}
2023/8/31 01:04
加载中...