tarjan缩点+倍增lca+树上差分 30pts求助
查看原帖
tarjan缩点+倍增lca+树上差分 30pts求助
752257
Miangoa楼主2023/7/24 20:58
#include<bits/stdc++.h>
#define int long long
#define N 1000001
#define M 8000002

using namespace std;

int n,m,q,s,t,ans,sum[N];
int h[N],e[M][2],cnt,v[N];
int nh[N],ne[M][2],ncnt,nv[N];
int dfn[N],low[N],bhd,bel[N],siz;
bool vis[N],pas[M];
stack<int> sta;
int d[N],f[N][70];

inline int read() {
	int x=0;
	char ch=getchar();
	while(ch<'0'||ch>'9')
		ch=getchar();
	while(ch>='0'&&ch<='9')
		x=x*10+ch-'0',ch=getchar();
	return x;
}

void add(int x,int y) {
	e[++cnt][0]=y;
	e[cnt][1]=h[x];
	h[x]=cnt;
	e[++cnt][0]=x;
	e[cnt][1]=h[y];
	h[y]=cnt;
}

void nadd(int x,int y) {
	ne[++cnt][0]=y;
	ne[cnt][1]=nh[x];
	nh[x]=cnt;
}

void init() {
	cnt=1;
	n=read(),m=read();
	for(int i=1; i<=n; i++)
		v[i]=read();
	for(int i=1; i<=m; i++)
		add(read(),read());
	q=read();
}

void tarjan(int p,int fa) {
	dfn[p]=low[p]=++bhd;
	vis[p]=true;
	sta.push(p);
	for(int i=h[p]; i; i=e[i][1]) {
		if(e[i][0]==fa)
			continue;
		if(!dfn[e[i][0]])
			tarjan(e[i][0],p);
		if(vis[e[i][0]])
			low[p]=min(low[p],low[e[i][0]]);
	}
	if(dfn[p]==low[p]) {
		siz++;
		while(!sta.empty()) {
			int x=sta.top();
			sta.pop();
			vis[x]=false;
			bel[x]=siz;
			nv[bel[x]]+=v[x];
			if(x==p)
				break;
		}
	}
}

void dfs(int p) {
	for(int i=1; f[f[p][i-1]][i-1]; i++)
		f[p][i]=f[f[p][i-1]][i-1];
	for(int i=nh[p]; i; i=ne[i][1])
		if(!d[ne[i][0]]) {
			d[ne[i][0]]=d[p]+1;
			f[ne[i][0]][0]=p;
			dfs(ne[i][0]);
		}
}

int lca(int a,int b) {
	if(d[a]!=d[b]) {
		if(d[a]<d[b])
			swap(a,b);
		for(int i=50; i>=0; i--)
			if(f[a][i]&&d[f[a][i]]>=d[b])
				a=f[a][i];
	}
	if(a==b)
		return a;
	for(int i=50; i>=0; i--)
		if(f[a][i]!=f[b][i])
			a=f[a][i],b=f[b][i];
	return f[a][0];
}

void run(int p) {
	vis[p]=true;
	for(int i=nh[p]; i; i=ne[i][1])
		if(!vis[ne[i][0]]) {
			f[ne[i][0]][0]=p;
			dfs(ne[i][0]);
			sum[p]+=sum[ne[i][0]];
		}
}

signed main() {
	init();
	tarjan(1,0);
	for(int i=1; i<=n; i++)
		for(int j=h[i]; j; j=e[j][1])
			if(bel[i]!=bel[e[j][0]])
				nadd(bel[i],bel[e[j][0]]);
	d[1]=1;
	dfs(1);
	for(int i=1; i<=q; i++) {
		int x=read(),y=read();
		int c=lca(bel[x],bel[y]);
		sum[bel[x]]++,sum[bel[y]]++,sum[c]--,sum[f[c][0]]--;
	}
	memset(vis,0,sizeof(vis));
	run(1);
	for(int i=1; i<=siz; i++)
		if(sum[i])
			ans+=nv[i];
	printf("%d",ans);
}
2023/7/24 20:58
加载中...