rt,有没有人救一下
屎山代码
#include<bits/stdc++.h>
using namespace std;
int read()
{
int r=0,f=1;
char c=getchar();
while(!isdigit(c))
{
if(c=='-') f=0;
c=getchar();
}
while(isdigit(c))
{
r=(r<<1)+(r<<3)+c-48;
c=getchar();
}
return f?r:-r;
}
const int N=5e5+5,INF=1e8+5;
int n,m,a[N],tot,rt[N];
struct node
{
int ls,rs,siz;
}tr[40000005];
void pushup(int p)
{
tr[p].siz=tr[tr[p].ls].siz+tr[tr[p].rs].siz;
}
void update(int &p,int l,int r,int L,int val)
{
if(!p) p=++tot;
if(l==r)
{
tr[p].siz+=val;
return ;
}
int mid=(l+r)>>1;
if(L<=mid) update(tr[p].ls,l,mid,L,val);
else update(tr[p].rs,mid+1,r,L,val);
pushup(p);
}
int lowbit(int x){return x&(-x);}
void ADD(int x,int k,int val)
{
while(x<=n)
{
update(rt[x],0,INF,k,val);
x+=lowbit(x);
}
}
int q1[30],q2[30],cnt1,cnt2;
void getq1(int x)
{
cnt1=0;
while(x)
{
q1[++cnt1]=rt[x];
x-=lowbit(x);
}
}
void getq2(int x)
{
cnt2=0;
while(x)
{
q2[++cnt2]=rt[x];
x-=lowbit(x);
}
}
int query(int l,int r,int k)
{
if(l==r) return l;
int sum1=0,sum2=0,ssiz,mid=(l+r)>>1;
for(int i=1;i<=cnt2;i++)
sum2+=tr[tr[q2[i]].ls].siz;
for(int i=1;i<=cnt1;i++)
sum1+=tr[tr[q1[i]].ls].siz;
ssiz=sum2-sum1;
if(ssiz>=k)
{
for(int i=1;i<=cnt2;i++)
q2[i]=tr[q2[i]].ls;
for(int i=1;i<=cnt1;i++)
q1[i]=tr[q1[i]].ls;
return query(l,mid,k);
}
else
{
for(int i=1;i<=cnt2;i++)
q2[i]=tr[q2[i]].rs;
for(int i=1;i<=cnt1;i++)
q1[i]=tr[q1[i]].rs;
return query(mid+1,r,k-ssiz);
}
}
int getkth(int l,int r,int val)
{
if(l==r) return 0;
int mid=(l+r)>>1;
int sum=0;
for(int i=1;i<=cnt2;i++)
sum+=tr[tr[q2[i]].ls].siz;
for(int i=1;i<=cnt1;i++)
sum-=tr[tr[q1[i]].ls].siz;
if(val<=mid)
{
// return get
for(int i=1;i<=cnt2;i++)
q2[i]=tr[q2[i]].ls;
for(int i=1;i<=cnt1;i++)
q1[i]=tr[q1[i]].ls;
return getkth(l,mid,val);
}
else
{
for(int i=1;i<=cnt2;i++)
q2[i]=tr[q2[i]].rs;
for(int i=1;i<=cnt1;i++)
q1[i]=tr[q1[i]].rs;
return sum+getkth(mid+1,r,val);
}
}
// int rq1[30],rq2[30];
int getmin(int l,int r)
{
int sum=0;
for(int i=1;i<=cnt2;i++) sum+=tr[q2[i]].siz;
for(int i=1;i<=cnt1;i++) sum-=tr[q1[i]].siz;
if(sum==0) return 2147483647;
if(l==r) return l;
int mid=(l+r)>>1;
sum=0;
for(int i=1;i<=cnt2;i++) sum+=tr[tr[q2[i]].ls].siz;
for(int i=1;i<=cnt1;i++) sum-=tr[tr[q1[i]].ls].siz;
if(sum>0)
{
for(int i=1;i<=cnt2;i++) q2[i]=tr[q2[i]].ls;
for(int i=1;i<=cnt1;i++) q1[i]=tr[q1[i]].ls;
return getmin(l,mid);
}
else
{
for(int i=1;i<=cnt2;i++) q2[i]=tr[q2[i]].rs;
for(int i=1;i<=cnt1;i++) q1[i]=tr[q1[i]].rs;
return getmin(mid+1,r);
}
}
int getmax(int l,int r)
{
int sum=0;
for(int i=1;i<=cnt2;i++) sum+=tr[q2[i]].siz;
for(int i=1;i<=cnt1;i++) sum-=tr[q1[i]].siz;
if(sum==0) return -2147483647;
if(l==r) return l;
int mid=(l+r)>>1;
sum=0;
for(int i=1;i<=cnt2;i++) sum+=tr[tr[q2[i]].rs].siz;
for(int i=1;i<=cnt1;i++) sum-=tr[tr[q1[i]].rs].siz;
if(sum>0)
{
for(int i=1;i<=cnt2;i++) q2[i]=tr[q2[i]].rs;
for(int i=1;i<=cnt1;i++) q1[i]=tr[q1[i]].rs;
return getmax(mid+1,r);
}
else
{
for(int i=1;i<=cnt2;i++) q2[i]=tr[q2[i]].ls;
for(int i=1;i<=cnt1;i++) q1[i]=tr[q1[i]].ls;
return getmax(l,mid);
}
}
int getpre(int l,int r,int val)
{
int sum=0;
for(int i=1;i<=cnt2;i++) sum+=tr[q2[i]].siz;
for(int i=1;i<=cnt1;i++) sum-=tr[q1[i]].siz;
if(sum==0) return -2147483647;
if(l==r) return l;
int mid=(l+r)>>1;
if(val<=mid)
{
for(int i=1;i<=cnt2;i++) q2[i]=tr[q2[i]].ls;
for(int i=1;i<=cnt1;i++) q1[i]=tr[q1[i]].ls;
return getpre(l,mid,val);
}
else
{
int rq1[30],rq2[30];
for(int i=1;i<=cnt1;i++) rq1[i]=q1[i];
for(int i=1;i<=cnt2;i++) rq2[i]=q2[i];
for(int i=1;i<=cnt2;i++) q2[i]=tr[q2[i]].rs;
for(int i=1;i<=cnt1;i++) q1[i]=tr[q1[i]].rs;
int ans=getpre(mid+1,r,val);
if(ans!=-2147483647) return ans;
for(int i=1;i<=cnt1;i++) q1[i]=rq1[i];
for(int i=1;i<=cnt2;i++) q2[i]=rq2[i];
for(int i=1;i<=cnt2;i++) q2[i]=tr[q2[i]].ls;
for(int i=1;i<=cnt1;i++) q1[i]=tr[q1[i]].ls;
return getmax(l,mid);
}
}
int getnxt(int l,int r,int val)
{
int sum=0;
for(int i=1;i<=cnt2;i++) sum+=tr[q2[i]].siz;
for(int i=1;i<=cnt1;i++) sum-=tr[q1[i]].siz;
if(sum==0) return 2147483647;
if(l==r) return l;
int mid=(l+r)>>1;
if(val>mid)
{
for(int i=1;i<=cnt2;i++) q2[i]=tr[q2[i]].rs;
for(int i=1;i<=cnt1;i++) q1[i]=tr[q1[i]].rs;
return getnxt(mid+1,r,val);
}
else
{
int rq1[30],rq2[30];
for(int i=1;i<=cnt1;i++) rq1[i]=q1[i];
for(int i=1;i<=cnt2;i++) rq2[i]=q2[i];
for(int i=1;i<=cnt2;i++) q2[i]=tr[q2[i]].ls;
for(int i=1;i<=cnt1;i++) q1[i]=tr[q1[i]].ls;
int ans=getnxt(l,mid,val);
if(ans!=2147483647) return ans;
for(int i=1;i<=cnt1;i++) q1[i]=rq1[i];
for(int i=1;i<=cnt2;i++) q2[i]=rq2[i];
for(int i=1;i<=cnt2;i++) q2[i]=tr[q2[i]].rs;
for(int i=1;i<=cnt1;i++) q1[i]=tr[q1[i]].rs;
return getmin(mid+1,r);
}
}
int main()
{
n=read(); m=read();
for(int i=1;i<=n;i++)
{
a[i]=read();
ADD(i,a[i],1);
}
int c;
int x,y,z;
while(m--)
{
cin>>c>>x>>y;
if(c==1)
{
cin>>z;
getq1(x-1);
getq2(y);
printf("%d\n",getkth(0,INF,z)+1);
}
else if(c==2)
{
cin>>z;
getq1(x-1);
getq2(y);
cout<<query(0,INF,z)<<endl;
}
else if(c==3)
{
ADD(x,a[x],-1);
a[x]=y;
ADD(x,a[x],1);
}
else if(c==4)
{
cin>>z;
getq1(x-1);
getq2(y);
printf("%d\n",getpre(0,INF,z));
}
else if(c==5)
{
cin>>z;
getq1(x-1);
getq2(y);
printf("%d\n",getnxt(0,INF,z));
}
}
}