先声明一点:我确实理解 Splay 的基本操作及原理(除了复杂度证明)。
我虽然看懂了老师的 Splay 代码,但每次写 Splay 时啥都想不起来,而且写的时候总是像背文章一样想“我现在应该写什么”,而不是顺理成章地想“我还需要写什么”。
极其希望能得到记忆的好方法。
附:老师给的模板代码(P3369)
#define MAXN 100005
#define ls(x) a[x].sons[0]
#define rs(x) a[x].sons[1]
#define fa(x) a[x].f
struct splay{
//分别为父节点编号,副本数量,关键码,子树大小,左右子树下标
int f,cnt,val,size,sons[2];
}a[MAXN];
int tot,root,INF=0x7fffffff;//tot为节点的个数
//获取一个节点是父节点的左子节点还是右子节点
int ident(int x){
return a[fa(x)].sons[0]==x?0:1;
}
//新建一个关键码为val的节点,父亲节点为f
int New(int val,int f){
a[++tot].val=val;
a[tot].cnt=a[tot].size=1;
a[tot].f=f;
a[tot].sons[0]=a[tot].sons[1]=0;
return tot;
}
//在旋转节点p之后,更新节点p维护的信息
void Update(int p){
a[p].size=a[ls(p)].size+a[rs(p)].size+a[p].cnt;
}
//把x节点连到fa节点的how子节点
void connect(int x,int fa,int how){
a[fa].sons[how]=x;
a[x].f=fa;
}
//单旋转函数
void rotate(int x){
int y=a[x].f,z=a[y].f;
int Yson=ident(x),Zson=ident(y);
connect(a[x].sons[Yson^1], y, Yson);
connect(y, x, Yson^1);
connect(x, z, Zson);
Update(y);
Update(x);
}
//把x旋转到to节点的位置
void splay(int x,int to){
to=fa(to);
while(fa(x)!=to){
int y=fa(x);
if(a[y].f==to) rotate(x);
else if(ident(x)==ident(y)){
rotate(y);
rotate(x);
}
else{
rotate(x);
rotate(x);
}
}
if(to==0){
root=x;
}
}
int find(int x){
int now=root;
while(1){
if(!now) return 0;
if(a[now].val==x){
splay(now,root);
return now;
}
int nxt=x<a[now].val?0:1;
now=a[now].sons[nxt];
}
}
int GetRankByVal(int val){
int p=find(val);
if(p==0){
return 0;
}
return a[ls(p)].size+1;
}
int GetValByRank(int p,int rank){
if(p==0) return INF;
if(a[ls(p)].size>=rank) return GetValByRank(ls(p),rank);
if(a[ls(p)].size+a[p].cnt>=rank){
splay(p,root);
return a[p].val;
}
return GetValByRank(rs(p),rank-a[ls(p)].size-a[p].cnt);
}
void Insert(int x){
//这里我们没有建立两个虚拟节点,所以在第一次插入的时候需要设置好根节点
if(root==0){
root=New(x,0);
return;
}
int now=root;
while(1){
a[now].size++;
if(a[now].val==x){
a[now].cnt++;
splay(now,root);
return;
}
int nxt=x<a[now].val?0:1;
if(!a[now].sons[nxt]){
int p=New(x,now);
a[now].sons[nxt]=p;
splay(p, root);
return;
}
now=a[now].sons[nxt];
}
}
int join(int r1,int r2){
//找到r1的最大元素
int maxson=r1;
while(rs(maxson)) maxson=rs(maxson);
splay(maxson, r1);
connect(r2,maxson,1);
Update(maxson);
return maxson;
}
void delet(int x){
int p=find(x);
if(p){
if(a[p].cnt>1){
a[p].cnt--;
a[p].size--;
return;
}
if(!a[p].sons[0] && !a[p].sons[1]){
root=0;
return;
}
if(!a[p].sons[0]){
root=a[p].sons[1];
a[root].f=0;
return;
}
int left=a[p].sons[0];
root=join(left,a[p].sons[1]);
a[root].f=0;
}
}
void split(int x,int &r1,int &r2){
int p=find(x);
if(p){
r1=ls(p);
r2=rs(p);
}
}
int GetPre(int x){
Insert(x);
int p=find(x);
int now=ls(p);
while(rs(now)) now=rs(now);
int ans=a[now].val;
delet(x);
return ans;
}
int GetNext(int x){
Insert(x);
int p=find(x);
int now=rs(p);
while(ls(now)) now=ls(now);
int ans=a[now].val;
delet(x);
return ans;
}