提供一个树链剖分做法
查看原帖
提供一个树链剖分做法
130183
liyishui2003楼主2023/9/12 23:19

其实本质上和树状数组一样啦,只是我想的时候没有想到树状数组。

考虑判断串x在串y中出现的次数,等价为在fail树上x有多少个节点属于y。对trie图(也即fail树)进行树链剖分,然后再dfs一遍trie树,对于当前遍历到的节点u,在fail树上进行一个u到根节点0的路径+1。

统计答案的话就先把询问离线,dfs遍历到该点时,拎出该点对应的所有x,求一个x当前的点权即可。

#include<bits/stdc++.h>
using namespace std;
const int maxn=1e5+5;
int m;
vector<pair<int,int>>q[maxn];
struct node{
	int son[26],flag,fail;
}trie[maxn];
int n,idx,fail[maxn],id[maxn],fa[maxn];
vector<int>inv[maxn],to[maxn],v[maxn];// to:fail树的连边关系,v:trie图的连边关系
void getfail(){
	queue<int>q;
	for(int i=0;i<26;i++) {
		if(trie[0].son[i]){
			fail[trie[0].son[i]]=0;
			to[0].push_back(trie[0].son[i]);
			q.push(trie[0].son[i]);
		}    
	}
	while(!q.empty()){
		int u=q.front();
		q.pop();
		for(int i=0;i<26;i++){
			int v=trie[u].son[i];
			if(v){
				fail[v]=trie[fail[u]].son[i];
				to[trie[fail[u]].son[i]].push_back(v);
				q.push(v);
			}
			else trie[u].son[i]=trie[fail[u]].son[i];
		}
	}
}
int a[maxn<<2];// a 为该节点所管辖的区间和
void build(int rt,int l,int r){
	if(l==r){
		a[rt]=0;return;
	}
	int mid=(l+r)>>1;
	build(rt<<1,l,mid);
	build(rt<<1|1,mid+1,r);
	a[rt]=a[rt<<1]+a[rt<<1|1];
}

int top[maxn],siz[maxn],father[maxn],dep[maxn],son[maxn];
void dfs1(int u,int f){
	siz[u]=1;
	father[u]=f;
	dep[u]=dep[f]+1;
	int maxson=-1;
	for(auto to:to[u]){
		if(to==f) continue;
		dfs1(to,u);
		siz[u]+=siz[to];
		if(siz[to]>maxson) son[u]=to,maxson=siz[to];
	}
}
int dfs_order=-1;
int dfn[maxn];
void dfs2(int u,int topf){
	dfs_order++;// dfs order
	dfn[u]=dfs_order;
	top[u]=topf;
	if(!son[u]) return;
	dfs2(son[u],topf);
	for(auto to:to[u]){
		if(to==son[u]||to==fa[u]) continue;
		dfs2(to,to);
	}
}

int lazy[maxn<<2];
void pushdown(int rt,int l,int r){
	if(lazy[rt]){
		int mid=(l+r)>>1;
		lazy[rt<<1]+=lazy[rt];
		lazy[rt<<1|1]+=lazy[rt];
		a[rt<<1]+=lazy[rt]*(mid-l+1);
		a[rt<<1|1]+=lazy[rt]*(r-mid);
		lazy[rt]=0;
	}	
}
void add(int rt,int l,int r,int ql,int qr,int k){
	if(ql<=l&&r<=qr){
		a[rt]+=(r-l+1)*k;
		lazy[rt]+=k;
		return;
	}
	int mid=(l+r)>>1;
	if(lazy[rt]) pushdown(rt,l,r);
	if(ql<=mid) add(rt<<1,l,mid,ql,qr,k);
	if(mid<qr) add(rt<<1|1,mid+1,r,ql,qr,k); 
	a[rt]=a[rt<<1]+a[rt<<1|1];
}
void addpath(int x,int y,int k){
	while(top[x]!=top[y]){
		if(dep[top[x]]<dep[top[y]]) swap(x,y);
		add(1,0,n,dfn[top[x]],dfn[x],k);
		x=father[top[x]];
	}
	if(dep[x]>dep[y]) swap(x,y);
	add(1,0,n,dfn[x],dfn[y],k);
}
int query(int rt,int l,int r,int ql,int qr){
	if(ql<=l&&r<=qr){
		return a[rt];
	}
	pushdown(rt,l,r);
	int mid=(l+r)>>1;
	int ans=0;
	if(mid>=ql) ans+=query(rt<<1,l,mid,ql,qr);
	if(mid<qr) ans+=query(rt<<1|1,mid+1,r,ql,qr); 
	return ans; 
}
int qpath(int x,int y){
	int ans=0;
	while(top[x]!=top[y]){
		if(dep[top[x]]<dep[top[y]]) swap(x,y);
		ans+=query(1,0,n,dfn[top[x]],dfn[x]);
		x=father[top[x]];
	}
	if(dep[x]>dep[y]) swap(x,y);
	ans+=query(1,0,n,dfn[x],dfn[y]);
	return ans;
}

//	inv[u].push_back(cnt);
int ans[maxn];
void dfs(int u,int fa){
	
	addpath(u,0,1);
	// 把该点对应的,y拎出来,_go为编号
	for(auto _go:inv[u]){
		for(auto go:q[_go]){// go为该编号所对应的询问
			int qid=go.first;
			int x=go.second;
			x=id[x];
			ans[qid]=qpath(x,x);
		}
	}
	for(auto to:v[u]){
		if(to==fa) continue;
	    dfs(to,u);
	}
	addpath(u,0,-1);
}
int main(){
	
	ios_base::sync_with_stdio(false);
	cin.tie(0);
	cout.tie(0);
	
	//freopen("lys.in","r",stdin);
	
	string str;
	cin>>str;
	cin>>m;
	for(int i=1;i<=m;i++){
		int x,y;
		cin>>x>>y;
		q[y].push_back({i,x});
	}
	
	int l=str.length();
	int u=0,cnt=0;
	for(int i=0;i<l;i++){
		if(str[i]=='P'){
			cnt++;
			id[cnt]=u;
			inv[u].push_back(cnt);
		    trie[u].flag++;
		}
		else if(str[i]=='B'){
			u=fa[u];
		}
		else {
			int V=str[i]-'a';
			if(!trie[u].son[V]) trie[u].son[V]=++idx;
			fa[trie[u].son[V]]=u;
			v[u].push_back(trie[u].son[V]);
			u=trie[u].son[V];
		}
	}	
	n=idx;
	getfail();
	
	// 对fail树进行一个树链剖分
	dfs1(0,-1);
	dfs2(0,0);
	build(1,0,idx);
	// 遍历trie树
	dfs(0,-1);
	for(int i=1;i<=m;i++) cout<<ans[i]<<endl;
	
}
2023/9/12 23:19
加载中...