这份树剖求 LCA 为啥能过
查看原帖
这份树剖求 LCA 为啥能过
239895
Yusani_huh楼主2023/5/13 10:41

这是一份 能过掉这题 的代码:

#include<bits/stdc++.h>
using namespace std;
#define N 100103
#define LL long long
#define INF 0x3f3f3f3f
#define MOD 998244353
#define PII pair<int,int>
#define fi first
#define se second
int n,m,dfn[N],vn[N],st[N],idx,tot;
int fa[N],d[N],sz[N],hs[N],tp[N],cnt;
LL dp[N][2],k0[N][2],k1[N][2],bk[N][2],ans;
vector<int>g[N],vg[N];
bool ban[N];
PII e[53];
bool cmp(int a,int b)
	{return dfn[a]<dfn[b];}
void dfs(int u,int ft){
	dfn[u]=++tot,fa[u]=ft,
	d[u]=d[ft]+1;
	for(int v:g[u]){
		if(v==ft) continue;
		if(dfn[v]){
			ban[v]=ban[u]=true;
			if(dfn[u]<dfn[v])
				e[idx++]={u,v};
		}else{
			dfs(v,u),sz[u]+=sz[v];
			if(sz[v]>sz[hs[u]]) hs[u]=v;
		}
	}
}
void dfs2(int u,int top){
	tp[u]=top;
	if(hs[u]) dfs2(hs[u],top);
	for(int v:g[u])
		if(v!=fa[u]&&v!=hs[u]) dfs2(v,v);
}
int lca(int u,int v){
	while(tp[u]!=tp[v]){
		if(d[u]<d[v]) swap(u,v);
		u=fa[tp[u]];
	}return d[u]<d[v]?u:v;
}
void add(int u,int v){vg[u].push_back(v);}
void build(){
	int tt=0,top=0;
	for(int i=1;i<=n;++i)
		if(ban[i]) vn[++tt]=i;
	sort(vn+1,vn+tt+1,cmp);
	st[++top]=1;
	for(int i=1;i<=n;++i){
		if(vn[i]==1) continue;
		int lc=lca(st[top],vn[i]);
		if(lc!=st[top]){
			while(dfn[lc]<dfn[st[top-1]])
				add(st[top-1],st[top]),top--;
			if(dfn[lc]>dfn[st[top-1]])
				add(lc,st[top]),st[top]=lc;
			else add(lc,st[top]),top--;
		}
		st[++top]=vn[i];
	}
	for(int i=1;i<top;++i)
		add(st[i],st[i+1]);
}
void pre(int u,int ft){
	for(int v:g[u])
		if(v!=ft) pre(v,u),ban[u]=ban[u]||ban[v];
}
void getrev(int u,int ft){
	bk[u][0]=bk[u][1]=1;
	for(int v:g[u]){
		if(v==ft||ban[v]) continue;
		getrev(v,u);
		bk[u][0]=bk[u][0]*(bk[v][0]+bk[v][1])%MOD;
		bk[u][1]=bk[u][1]*bk[v][0]%MOD;
	}
}
void getcro(int u,int bn){
	dp[u][0]=dp[u][1]=1;
	for(int v:g[u]){
		if(v==fa[u]||v==bn) continue;
		getcro(v,u);
		dp[u][0]=dp[u][0]*(dp[v][0]+dp[v][1])%MOD;
		dp[u][1]=dp[u][1]*dp[v][0]%MOD;
	}
}
void getcot(int u,int ft){
	k0[u][0]=k0[u][1]=k1[u][0]=1;
	int nw=u;
	while(fa[nw]!=ft){
		getcro(fa[nw],nw),nw=fa[nw];
		LL r0=k0[u][0],r1=k0[u][1];
		k0[u][0]=(dp[nw][0]*r0+dp[nw][1]*k1[u][0])%MOD;
		k0[u][1]=(dp[nw][0]*r1+dp[nw][1]*k1[u][1])%MOD;
		k1[u][0]=dp[nw][0]*r0%MOD;
		k1[u][1]=dp[nw][0]*r1%MOD;
	}
}
void initdp(int u,int ft){
	getrev(u,ft),vn[++tot]=u;
	if(ft!=u) getcot(u,ft);
	for(int v:vg[u])
		if(v!=ft) initdp(v,u);
}
void DP(int u,int ft){
	for(int v:vg[u]){
		if(v==ft) continue;
		DP(v,u);
		dp[u][0]=dp[u][0]*(dp[v][0]*k0[v][0]%MOD+
						   dp[v][1]*k0[v][1]%MOD)%MOD;
		dp[u][1]=dp[u][1]*(dp[v][0]*k1[v][0]%MOD+
						   dp[v][1]*k1[v][1]%MOD)%MOD;
	}
}
int main(){
	scanf("%d%d",&n,&m);
	for(int i=1;i<=m;++i){
		int u,v;
		scanf("%d%d",&u,&v);
		g[u].push_back(v),
		g[v].push_back(u);
	}
	dfs(1,0);
	for(int i=1;i<=n;++i) g[i].clear();
	for(int i=1;i<=n;++i)
		if(fa[i]) g[fa[i]].push_back(i);
	dfs2(1,1),build(),pre(1,0);
	tot=0,initdp(1,1);
	int s=1<<idx;
	for(int i=0;i<s;++i){
		for(int j=1;j<=tot;++j)
			dp[vn[j]][0]=bk[vn[j]][0],
			dp[vn[j]][1]=bk[vn[j]][1];
		for(int j=0;j<idx;++j)
			if(i>>j&1) dp[e[j].fi][1]=0;
			else dp[e[j].fi][0]=0,dp[e[j].se][1]=0;
		DP(1,1),ans=(ans+dp[1][0]+dp[1][1])%MOD;
	}
	printf("%lld\n",ans);
	return 0;
}

然后我发现有什么地方不对劲:dfs 里面没有初始化 sz 数组。也就是说 sz 整个是 0。这就导致所有点的重儿子都是 0。于是树剖剖了个寂寞,我 lca 退化成了暴力。

然而这份代码提交记录在最优解第一页。

UPD:我把 lca 直接改成暴力 lca 发现答案不对了。现在我完全看不懂我自己写的东西了。

2023/5/13 10:41
加载中...