求助平衡树全WA,de了一晚上,没看出哪里错了
查看原帖
求助平衡树全WA,de了一晚上,没看出哪里错了
244165
_121017_楼主2022/6/28 20:27
#include<bits/stdc++.h>
#define int long long
using namespace std;
const int inf=1145141919810;
const int mod=1e6;
const int N=3e6+5;
int n,root,node;
int ls[N],rs[N],num[N],cnt[N],size[N],fa[N];
void push_up(int p){
	size[p]=((cnt[p]+size[ls[p]])%mod+size[rs[p]])%mod;
}
bool lorr(int p){
	return (ls[fa[fa[p]]]==fa[p]&&ls[fa[p]]==p)||(rs[fa[fa[p]]]==fa[p]&&rs[fa[p]]==p);
}
void rotate(int p){
	int fath=fa[p]; int gfath=fa[fath]; fa[p]=gfath;
	if(ls[gfath]==fath) ls[gfath]=p;
	else rs[gfath]=p;
	if(ls[fath]==p) ls[fath]=rs[p],fa[rs[p]]=fath,rs[p]=fath,fa[fath]=p;
	else rs[fath]=ls[p],fa[ls[p]]=fath,ls[p]=fath,fa[fath]=p;
	push_up(fath),push_up(p);
}
void splay(int p,int q){
	if(abs(p)==inf||abs(q)==inf)return ;
	while(fa[p]!=q) 
		if(lorr(p)&&fa[fa[p]]!=q) rotate(fa[p]),rotate(p);
		else rotate(p);
	if(!q) root=p;
}
void setup(int &p){
	p=++node;
	ls[p]=rs[p]=num[p]=cnt[p]=size[p]=0;
}
void insert(int x){
	int p=root,fath=0;
	while(p){
		if(x==num[p]) break; fath=p;
		if(num[p]>x) p=ls[p];
		else p=rs[p]; 
	}
	if(!p) setup(p);
	fa[p]=fath,num[p]=x,cnt[p]++,size[p]++;
	if(num[fath]>x) ls[fath]=x;
	else rs[fath]=x;
	splay(p,0);
}
void setout(int p){
	cnt[p]=num[p]=size[p]=ls[p]=rs[p]=fa[p]=0;
	if(p==node) node--;
}
int lowerr(int x){
	int p=root,id=-inf;
	while(p){
		if(num[p]>=x) p=ls[p];
		else id=p,p=rs[p];
	}
	return id;
}
int lower(int x){
	int p=root,id=-inf;
	while(p){
		if(num[p]==x) break;
		if(num[p]>x) p=ls[p];
		else id=p,p=rs[p];
	}
	splay(id,0); return id;
}
int upperr(int x){
	int p=root,id=-inf;
	while(p){
		if(num[p]>x) id=p,p=ls[p];
		else p=rs[p];
	}
	return id;
}
int upper(int x){
	int p=root,id=-inf;
	while(p){
		if(num[p]==x) break;
		if(num[p]>x) id=p,p=ls[p];
		else p=rs[p];
	}
	splay(id,0); return id;
}
void destroy(int x){
//	cout<<num[lowerr(x)]<<" "<<num[upperr(x)]<<endl;
	splay(lowerr(x),0),splay(upperr(x),root);
	int p=ls[rs[root]];
	size[root]--,size[fa[p]]--; cnt[p]--; 
	if(cnt[p]<=0) ls[fa[p]]=0,setout(p);
}
void Ot(int p){
	if(!p) return ;
	Ot(ls[p]);
	printf("%lld ",num[p]);
	Ot(rs[p]);
}
signed main(){
	cin>>n;
	int ans=0; insert(-inf); insert(inf);
	for(int i=1,op,x;i<=n;i++){
		scanf("%lld%lld",&op,&x);
		if(op==0) insert(x);
		else{
			int a=lower(x);	if(abs(a)!=inf) a=num[a];
			int b=upper(x);	if(abs(b)!=inf) b=num[b];
//			cout<<a<<" "<<b<<"  ybsb"<<endl;
			if(a==-inf&&b==inf) continue;
			else if(abs(x-a)<=abs(x-b)) ans=(ans+abs(x-a))%mod,destroy(a);
			else ans=(ans+abs(x-b))%mod,destroy(b);		
		}
	}
	cout<<ans;
	return 0;
}
2022/6/28 20:27
加载中...