Splay 有没有记忆代码的好方法?
  • 板块灌水区
  • 楼主不知名用户
  • 当前回复12
  • 已保存回复12
  • 发布时间2023/7/9 19:49
  • 上次更新2023/11/3 10:52:08
查看原帖
Splay 有没有记忆代码的好方法?
316358
不知名用户楼主2023/7/9 19:49

先声明一点:我确实理解 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;
}
2023/7/9 19:49
加载中...