时限开到 3 s 就行了。
Code:
#include <bits/stdc++.h>
#define rep(i, s, t) for(int i=s; i<=t; ++i)
#define F first
#define S second
#define pii pair<int, int>
#define ll long long
#define debug(x) cout<<#x<<":"<<x<<endl;
const int N=50010, inf=2147483647;
using namespace std;
namespace treap
{
struct node
{
int x, y, siz, l, r;
} t[N<<6]; int tot;
int New(int x) {t[++tot]={x, rand(), 1}; return tot;}
void up(int p) {t[p].siz=t[t[p].l].siz+t[t[p].r].siz+1;}
void split(int p, int val, int &x, int &y)
{
if(!p) return x=y=0, void();
if(t[p].x<=val) split(t[p].r, val, t[x=p].r, y);
else split(t[p].l, val, x, t[y=p].l);
up(p);
}
int merge(int x, int y)
{
if(!x || !y) return x|y;
if(t[x].y<t[y].y)
{
t[x].r=merge(t[x].r, y);
return up(x), x;
}
t[y].l=merge(x, t[y].l);
return up(y), y;
}
void insert(int &rt, int val)
{
int x, y; split(rt, val, x, y);
rt=merge(merge(x, New(val)), y);
}
void remove(int &rt, int val)
{
int x, y, z;
split(rt, val, x, z);
split(x, val-1, x, y);
y=merge(t[y].l, t[y].r);
rt=merge(merge(x, y), z);
}
int Less(int &rt, int val)
{
int x, y, ans;
split(rt, val-1, x, y);
ans=t[x].siz;
return rt=merge(x, y), ans;
}
int nxt(int p, int val)
{
int ans=inf;
while(p)
if(t[p].x<=val) p=t[p].r;
else ans=t[p].x, p=t[p].l;
return ans;
}
int pre(int p, int val)
{
int ans=-inf;
while(p)
if(t[p].x>=val) p=t[p].l;
else ans=t[p].x, p=t[p].r;
return ans;
}
}
struct segnode
{
int l, r, rt;
} t[N<<2]; int n, m, a[N];
#define lc p<<1
#define rc p<<1|1
void build(int p, int l, int r)
{
t[p]={l, r};
rep(i, l, r) treap::insert(t[p].rt, a[i]);
if(l==r) return;
int m=l+r>>1;
build(lc, l, m), build(rc, m+1, r);
}
void modify(int p, int i, int x)
{
treap::remove(t[p].rt, a[i]);
treap::insert(t[p].rt, x);
if(t[p].l==t[p].r) return;
int m=t[p].l+t[p].r>>1;
if(i<=m) modify(lc, i, x); else modify(rc, i, x);
}
int pre(int p, int l, int r, int x)
{
if(l<=t[p].l && t[p].r<=r)
return treap::pre(t[p].rt, x);
int m=t[p].l+t[p].r>>1;
int res=-inf;
if(l<=m) res=max(res, pre(lc, l, r, x));
if(r>m) res=max(res, pre(rc, l, r, x));
return res;
}
int nxt(int p, int l, int r, int x)
{
if(l<=t[p].l && t[p].r<=r)
return treap::nxt(t[p].rt, x);
int m=t[p].l+t[p].r>>1;
int res=inf;
if(l<=m) res=min(res, nxt(lc, l, r, x));
if(r>m) res=min(res, nxt(rc, l, r, x));
return res;
}
int Less(int p, int l, int r, int x)
{
if(l<=t[p].l && t[p].r<=r)
return treap::Less(t[p].rt, x);
int m=t[p].l+t[p].r>>1;
int res=0;
if(l<=m) res+=Less(lc, l, r, x);
if(r>m) res+=Less(rc, l, r, x);
return res;
}
int kth(int LL, int RR, int k)
{
int l=0, r=1e8, ans;
while(l<=r)
{
ll mid=l+r>>1;
if(Less(1, LL, RR, mid)<k) ans=mid, l=mid+1;
else r=mid-1;
}
return ans;
}
int main()
{
scanf("%d%d", &n, &m);
rep(i, 1, n) scanf("%d", a+i);
build(1, 1, n);
rep(i, 1, m)
{
int o, l, r, x;
scanf("%d", &o);
if(o==3) scanf("%d%d", &l, &x);
else scanf("%d%d%d", &l, &r, &x);
if(o==1) printf("%d\n", Less(1, l, r, x)+1);
if(o==2) printf("%d\n", kth(l, r, x));
if(o==3) modify(1, l, x), a[l]=x; // 一定记得a[l]=x!
if(o==4) printf("%d\n", pre(1, l, r, x));
if(o==5) printf("%d\n", nxt(1, l, r, x));
}
return 0;
}