求助
查看原帖
求助
341373
Autofreeze楼主2022/4/27 23:50

有人和我一样 wa40 过吗,只有前两个点和最后两个点能过,是什么原因呢qwq

我的代码,用的是替罪羊树

#include<bits/stdc++.h>
#define N 3001001
#define MAX 2001
using namespace std;
typedef int ll;
typedef long double db;
const ll inf=1e9;
inline void read(ll &ret)
{
	ret=0;char c=getchar();bool pd=false;
	while(!isdigit(c)){pd|=c=='-';c=getchar();}
	while(isdigit(c)){ret=(ret<<1)+(ret<<3)+(c&15);c=getchar();}
	ret=pd?-ret:ret;
	return;
}
char s[N];
ll q,cnt,n,root;
#define ls(x) a[x].son[0]
#define rs(x) a[x].son[1]
#define fa(x) a[x].fa
struct node
{
	ll son[2],siz,fa;
	db l,r;
}a[801001];
inline void dfss(ll pos)
{
	if(!pos)
		return;
	dfss(ls(pos));
	printf("%d %d %d %d %d %.5lf %.5lf\n",pos,ls(pos),rs(pos),fa(pos),a[pos].siz,a[pos].l,a[pos].r);
	dfss(rs(pos));
	return;
}
vector<ll>lis;
inline void dfs(ll pos)
{
	if(!pos)
		return;
	dfs(ls(pos));
	lis.push_back(pos);
	ll tmp=rs(pos);
	ls(pos)=rs(pos)=a[pos].siz=fa(pos)=a[pos].l=a[pos].r=0;
	dfs(tmp);
	return;
}
inline void update(ll pos)
{
	a[pos].siz=1+a[ls(pos)].siz+a[rs(pos)].siz;
	return;
}
inline ll rebuild(ll l,ll r,db L,db R)
{
	if(l>r)
		return 0;
	ll mid=l+r>>1;
	ll pos=lis[mid];
	a[pos].siz=1;
	a[pos].l=L;
	a[pos].r=R;
	ls(pos)=rebuild(l,mid-1,L,(a[pos].l+a[pos].r)/2);
	rs(pos)=rebuild(mid+1,r,(a[pos].l+a[pos].r)/2,R);
	if(ls(pos))
		fa(ls(pos))=pos;
	if(rs(pos))
		fa(rs(pos))=pos;
	update(pos);
	return pos;
}
inline ll insert(ll now,db l,db r,ll pos)
{
	if(!now)
	{
		a[pos].l=l;
		a[pos].r=r;
		a[pos].siz=1;
		return pos;
	}
	if(s[pos]<s[now]||(s[pos]==s[now]&&(a[pos-1].l+a[pos-1].r)/2<(a[now-1].l+a[now-1].r)/2))
	{
		ls(now)=insert(ls(now),a[now].l,(a[now].l+a[now].r)/2,pos);
		update(now);
		fa(ls(now))=now;
		return now;
	}
	else
	{
		rs(now)=insert(rs(now),(a[now].l+a[now].r)/2,a[now].r,pos);
		update(now);
		fa(rs(now))=now;
		return now;
	}
}
inline ll check(ll now,ll pos)
{
	if(max(a[ls(now)].siz,a[rs(now)].siz)*4>a[now].siz*3)
		return now;
	if(now==pos)
		return 0;
	if(s[pos]<s[now]||(s[pos]==s[now]&&(a[pos-1].l+a[pos-1].r)/2<(a[now-1].l+a[now-1].r)/2))
		return check(ls(now),pos);
	else
		return check(rs(now),pos);
}
inline void ins(ll pos)
{
	if(!root)
		root=insert(root,0,1,pos);
	else
		root=insert(root,a[root].l,a[root].r,pos);
	ll tmp=check(root,pos);
	if(tmp)
	{
		lis.clear();
		ll now=fa(tmp);
		if(!now)
		{
			dfs(tmp);
			root=rebuild(0,(int)lis.size()-1,0,1);
		}
		else
		{
			if(rs(now)==tmp)
			{
				dfs(tmp);
				rs(now)=rebuild(0,(int)lis.size()-1,(a[now].l+a[now].r)/2,a[now].r);
				fa(rs(now))=now;
			}
			else
			{
				dfs(tmp);
				ls(now)=rebuild(0,(int)lis.size()-1,a[now].l,(a[now].l+a[now].r)/2);
				fa(ls(now))=now;
			}
		}
		
	}
	return;
}
char op[N];
ll mask;
char ss[N];
inline ll findpre(ll pos)
{
	pos=ls(pos);
	while(rs(pos))
		pos=rs(pos);
	return pos;
}
ll lens;
inline bool cmp(ll x)
{
	for(int i=0;i<lens;i++,x--)
	{
		if(x<0)
			return false;
		if(ss[i]>s[x])
			return false;
		else if(ss[i]<s[x])
			return true;
	}
	return true;
}
inline ll del(ll pos,ll d)
{
	if(pos==d)
	{
		if(!ls(pos)||!rs(pos))
		{
			ll x=ls(pos),y=rs(pos);
			fa(pos)=ls(pos)=rs(pos)=a[pos].siz=a[pos].l=a[pos].r=0;
			return x+y;
		}
		ll tmp=ls(pos),las=pos;
		while(rs(tmp))
		{
			a[tmp].siz--;
			las=tmp;
			tmp=rs(tmp);
		} 
		if(las==pos)
		{
			fa(rs(pos))=ls(pos);
			rs(ls(pos))=rs(pos);
			fa(pos)=ls(pos)=rs(pos)=a[pos].siz=a[pos].l=a[pos].r=0;
			update(tmp);
			return tmp;
		}
		else
		{
			fa(ls(tmp))=las;
			rs(las)=ls(tmp);
			ls(tmp)=ls(pos);
			rs(tmp)=rs(pos);
			a[tmp].l=a[pos].l;
			a[tmp].r=a[pos].r;
			fa(pos)=ls(pos)=rs(pos)=a[pos].siz=a[pos].l=a[pos].r=0;
			update(tmp);
			return tmp;
		}
	}
	if(s[d]<s[pos]||(s[d]==s[pos]&&(a[d-1].l+a[d-1].r)/2<(a[pos-1].l+a[pos-1].r)/2))
	{
		ls(pos)=del(ls(pos),d);
		fa(ls(pos))=pos;
		update(pos);
		return pos;
	}
	else
	{
		rs(pos)=del(rs(pos),d);
		fa(rs(pos))=pos;
		update(pos);
		return pos;
	}
}
inline ll findrank(ll pos)
{
	if(!pos)
		return 0;
	if(cmp(pos))
		return findrank(ls(pos));
	else
		return findrank(rs(pos))+a[ls(pos)].siz+1;
}
signed main()
{
	read(q);
	scanf("%s",s+1);
	n=strlen(s+1);
	for(int i=1;i<=n;i++)
		ins(i);
	for(int i=1;i<=q;i++)
	{
		scanf("%s",op+1);
		if(op[1]=='A')
		{
			scanf("%s",ss);
			ll len=strlen(ss);
			for(int j=0;j<len;j++)
			{
				mask=(mask*131+j)%len;
				swap(ss[mask],ss[j]);
			}
			for(int j=0;j<len;j++)
			{
				s[++n]=ss[j];
				ins(n);
			}
		}
		else if(op[1]=='D')
		{
			ll num;
			read(num);
			for(int j=1;j<=num;j++)
				root=del(root,n--);
		}
		else
		{
			scanf("%s",ss);
			ll len=strlen(ss);
			for(int j=0;j<len;j++)
			{
				mask=(mask*131+j)%len;
				swap(ss[mask],ss[j]);
			}
			reverse(ss,ss+len);
			ss[len]='Z'+1;
			lens=len+1;
			ll ans=findrank(root);
			ss[len-1]--;
			ans-=findrank(root);
			mask^=ans;
			printf("%d\n",ans);
		}
	}
	exit(0);
}
2022/4/27 23:50
加载中...