这是在洛谷上可以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;
}
}
}