WA 68pts Treap求调!!(貌似前驱和后继错了)
查看原帖
WA 68pts Treap求调!!(貌似前驱和后继错了)
519573
Daniel_yao楼主2023/7/17 11:43
#include <bits/stdc++.h>
#define int long long
#define H 19260817
#define rint register int
#define For(i,l,r) for(rint i=l;i<=r;++i)
#define FOR(i,r,l) for(rint i=r;i>=l;--i)
#define MOD 1000003
#define mod 1000000007
#define inf 1e9

using namespace std;

inline int read() {
  rint x=0,f=1;char ch=getchar();
  while(ch<'0'||ch>'9'){if (ch=='-') f=-1;ch=getchar();}
  while(ch>='0'&&ch<='9'){x=(x<<1)+(x<<3)+(ch^48);ch=getchar();}
  return x*f;
}

void print(int x){
  if(x<0){putchar('-');x=-x;}
  if(x>9){print(x/10);putchar(x%10+'0');}
  else putchar(x+'0');
  return;
}

const int N = 2e6 + 10;

struct Node {
  int l, r, cnt, v, val, size;
} t[N];

int n, m, op, root, idx, last, ans;

int New(int v) {
  t[++idx].v = v;
  t[idx].val = rand();
  t[idx].cnt = t[idx].size = 1;
  return idx;
}

void pushup(int p) {
  t[p].size = t[t[p].l].size + t[t[p].r].size + t[p].cnt;
}

void build() {
  New(-inf), New(inf);
  root = 1, t[1].l = 2;
  pushup(root);
} 

void zig(int &p) {
  int q = t[p].l;
  t[p].l = t[q].r, t[q].r = p, p = q;
  pushup(p), pushup(t[p].r);
}

void zag(int &p) {
  int q = t[p].r;
  t[p].r = t[q].l, t[q].l = p, p = q;
  pushup(p), pushup(t[p].l);
}

void insert(int &p, int v) {
  if(!p) {p = New(v);}
  else if(t[p].v == v) t[p].cnt++;
  else if(t[p].v > v) {
    insert(t[p].l, v);
    if(t[p].val < t[t[p].l].val) zig(p);
  } else {
    insert(t[p].r, v);
    if(t[p].val < t[t[p].r].val) zag(p);
  }
  pushup(p);
}

void remove(int &p, int v) {
  if(!p) return ;
  if(t[p].v == v) {
    if(t[p].cnt > 1) t[p].cnt--;
    else if(t[p].l || t[p].r){
      if(!t[p].r || t[t[p].l].val > t[t[p].r].val) {
        zig(p);
        remove(t[p].r, v);
      } else {
        zag(p);
        remove(t[p].l, v);
      }
    } else p = 0;
  }
  if(t[p].v > v) remove(t[p].l, v);
  else remove(t[p].r, v);
  pushup(p);
}

int rk(int p, int v) {
  if(!p) return 0;
  if(t[p].v == v) return t[t[p].l].size + 1;
  if(t[p].v > v) return rk(t[p].l, v);
  return t[t[p].l].size + t[p].cnt + rk(t[p].r, v);
}

int find(int p, int rk) {
  if(!p) return inf;
  if(t[t[p].l].size >= rk) return find(t[p].l, rk);
  if(t[t[p].l].size + t[p].cnt >= rk) return t[p].v;
  return find(t[p].r, rk - t[t[p].l].size - t[p].cnt);
}

signed main() {
  srand(time(0)); 
  n = read();
  while(n--) {
    op = read(); int x = read();
    if(op == 1) insert(root, x);
    if(op == 2) remove(root, x);
    if(op == 3) printf("%lld\n", rk(root, x));
    if(op == 4) printf("%lld\n", find(root, x));
    if(op == 5) printf("%lld\n", find(root, rk(root, x) - 1));
    if(op == 6) printf("%lld\n", find(root, rk(root, x + 1)));
  }
  return 0;
}

2023/7/17 11:43
加载中...