【DS 求 debug】萌新 Splay WA 样例求调
查看原帖
【DS 求 debug】萌新 Splay WA 样例求调
356003
Moeebius楼主2022/8/14 21:03

RT,照着 OI-wiki\text{OI-wiki} 写的,代码用 clang-format格式化过了,马蜂应该还可以吧 qwq

#include <bits/stdc++.h>
using namespace std;

#define il inline
#define mkp make_pair
#define pii pair<int, int>
#define lll __int128
#define ll long long
#define For(i, j, k) for (int i = (j); i <= (k); ++i)
#define ForDown(i, j, k) for (int i = (j); i >= (k); --i)
#define pb push_back
#define init(filename)                                                         \
  freopen(filename ".in", "r", stdin);                                         \
  freopen(filename ".out", "w", stdout)
template <typename T> il void read(T &x) {
  x = 0;
  int f = 1;
  char c = getchar();
  while (!isdigit(c)) {
    if (c == '-')
      f = -1;
    c = getchar();
  }
  while (isdigit(c)) {
    x = x * 10 + c - '0';
    c = getchar();
  }
  x *= f;
}
template <typename T, typename... Args> il void read(T &x, Args &...y) {
  read(x);
  read(y...);
}

class Splay {
public:
#define MAXN 200005
  int sz[MAXN], ch[MAXN][2], fa[MAXN], val[MAXN], cnt[MAXN], trash[MAXN], rt,
      tot, top;
  Splay() {
    rt = tot = top = 0;
    sz[0] = ch[0][0] = ch[0][1] = fa[0] = cnt[0] = val[0] = 0;
  }
  il void pushUp(int p) { sz[p] = sz[ch[p][0]] + sz[ch[p][1]] + cnt[p]; }
  il int get(int p) { return p == ch[fa[p]][1]; }
  il int create() { return top ? trash[top--] : ++tot; }
  il void del(int p) {
    trash[++top] = p;
    sz[p] = ch[p][0] = ch[p][1] = fa[p] = cnt[p] = val[p] = 0;
  }

  il void rotate(int x) {
    int y = fa[x], z = fa[y], f = get(x);
    ch[y][f] = ch[x][f ^ 1];
    if (ch[x][f ^ 1])
      fa[ch[x][f ^ 1]] = y;
    ch[x][f ^ 1] = y;
    fa[y] = x, fa[x] = z;
    if (z)
      ch[z][y == ch[z][1]] = x;
    pushUp(y), pushUp(x);
  }
  il void splay(int x) {
    for (int f = fa[x]; f = fa[x], f; rotate(x)) {
      if (fa[f]) {
        rotate(get(x) == get(f) ? f : x);
      }
    }
    rt = x;
  }
  il void insert(int x) {
    if (!rt) {
      rt = create();
      cnt[rt]++;
      val[rt] = x;
      pushUp(x);
      return;
    }
    int cur = rt, f = 0;
    while (1) {
      if (val[cur] == x) {
        cnt[cur]++;
        pushUp(x), pushUp(f);
        splay(cur);
        return;
      }
      f = cur, cur = ch[f][x > val[f]];
      if (!cur) {
        cur = ch[f][x > val[f]] = create();
        fa[cur] = f;
        val[cur] = x;
        cnt[cur]++;
        pushUp(cur), pushUp(f);
        splay(f);
        return;
      }
    }
  }
  il int __pre_node() {
    int cur = ch[rt][0];
    while (ch[cur][1])
      cur = ch[cur][1];
    splay(cur);
    return cur;
  }
  il void erase(int x) {
    rank(x);
    if (cnt[rt] > 1) {
      cnt[rt]--;
      return;
    }
    if (!ch[rt][0] && !ch[rt][1]) {
      del(rt);
      return;
    }
    if (!ch[rt][0]) {
      int cur = rt;
      rt = ch[rt][1];
      fa[rt] = 0;
      del(cur);
      return;
    }
    if (!ch[rt][1]) {
      int cur = rt;
      rt = ch[rt][0];
      fa[rt] = 0;
      del(cur);
      return;
    }
    int cur = rt;
    int p = __pre_node();
    fa[ch[cur][1]] = p;
    ch[p][1] = ch[cur][1];
    rt = p;
    del(cur);
    pushUp(rt);
  }
  il int rank(int x) {
    int res = 0, cur = rt;
    while (1) {
      if (x < val[cur]) {
        cur = ch[cur][0];
      } else {
        res += sz[ch[cur][0]];
        if (x == val[cur]) {
          splay(cur);
          return res;
        }
        res += cnt[cur];
        cur = ch[cur][1];
      }
    }
  }
  il int kth(int k) {
    int cur = rt;
    while (1) {
      if (sz[ch[cur][0]] >= k) {
        cur = ch[cur][0];
      } else {
        k -= sz[ch[cur][0]] + cnt[cur];
        if (k <= 0) {
          splay(cur);
          return val[cur];
        }
        cur = ch[cur][1];
      }
    }
  }
  il int pre(int x) {
    insert(x);
    int p = __pre_node();
    int ans = val[p];
    erase(x);
    return ans;
  }
  il int nxt(int x) {
    insert(x);
    int cur = ch[rt][1];
    while (ch[cur][0])
      cur = ch[cur][0];
    splay(cur);
    int ans = val[cur];
    erase(x);
    return ans;
  }
};

Splay BT;
int n;

signed main() {
  read(n);
  while (n--) {
    int op, x;
    read(op, x);
    switch (op) {
    case 1:
      BT.insert(x);
      break;
    case 2:
      BT.erase(x);
      break;
    case 3:
      printf("%d\n", BT.rank(x));
      break;
    case 4:
      printf("%d\n", BT.kth(x));
      break;
    case 5:
      printf("%d\n", BT.pre(x));
      break;
    case 6:
      printf("%d\n", BT.nxt(x));
      break;
    }
  }
  return 0;
}
2022/8/14 21:03
加载中...