求助!树链剖分WA on #4
查看原帖
求助!树链剖分WA on #4
701254
Mu_leaf楼主2023/10/1 08:14
#include<bits/stdc++.h>
#define int long long
#define mid ((l+r)>>1)
#define bug printf("1\n");
using namespace std;
const int N=1e6+5;
struct edge{
	int u,v,w;
	bool tag;
}line[N];
struct node{
	int v,cost;
};
vector<node> V[N];
int f[N],val[N],a[N];
int n,m,tot;
bool cmp(edge a,edge b){
	return a.w<b.w;
}void init(){
	for(int i=1;i<=n;i++)f[i]=i;
}
int find(int x){
	if(f[x]==x) return x;
	return f[x]=find(f[x]);
}
void K(){
	sort(line+1,line+tot+1,cmp);
//	cout << line[1].u << "\n";
	for(int i=1;i<=tot;i++){
		int fu=find(line[i].u),fv=find(line[i].v);
		if(fu!=fv){
//			cout << 1 << "\n";
			f[fu]=fv;
			line[i].tag=1;
			V[line[i].u].push_back((node){line[i].v,line[i].w});
			V[line[i].v].push_back((node){line[i].u,line[i].w});
		}
	}
}int dep[N],siz[N],son[N],id[N],top[N],fa[N],cnt;
void dfs1(int x,int f){
	dep[x]=dep[f]+1;
	siz[x]=1;
	fa[x]=f;
	int fat=-1;
	for(int i=0;i<V[x].size();i++){
		int v=V[x][i].v;
		if(v==f) continue;
		dfs1(v,x);
		siz[x]+=siz[v];
		if(fat<siz[v]) fat=siz[v],son[x]=v;
	}
}
void dfs2(int x,int f,int tp){
	id[x]=++cnt;
	top[x]=tp;
	if(!son[x]) return;
	dfs2(son[x],x,tp);
	for(int i=0;i<V[x].size();i++){
		int v=V[x][i].v;
		int w=V[x][i].cost;
		if(v==f || v==son[x]) continue;
		dfs2(v,x,v);
	}
}
void dfs3(int x,int f){
	for(int i=0;i<V[x].size();i++){
		int v=V[x][i].v;
		int w=V[x][i].cost;
		if(v==f) continue;
		val[id[v]]=w;
		dfs3(v,x);
	}
}
void build(int x,int l,int r){
	if(l==r){a[x]=val[l];return;}
	build(x<<1,l,mid);
	build(x<<1|1,mid+1,r);
	a[x]=max(a[x<<1],a[x<<1|1]);
}
int query(int x,int l,int r,int lt,int rt){
	if(lt>rt) return 0;
	if(l>=lt && r<=rt) return a[x];
	int res=0;
	if(lt <= mid) res=max(res,query(x<<1,l,mid,lt,rt));
	if(rt >  mid) res=max(res,query(((x<<1)|1),mid+1,r,lt,rt));
	return res;
}
int answer(int u,int v){
	int ans=0;
	while(top[u]!=top[v]){
		if(dep[top[u]]<dep[top[v]]) swap(u,v);
		ans=max(ans,query(1,1,n,id[top[u]]+1,id[u]));
		u=fa[top[u]];
	}
	if(dep[u]>dep[v]) swap(u,v);
	ans=max(ans,query(1,1,n,id[u]+1,id[v]));
	return ans;
}
signed main(){
	cin >> n >> m;
	
	for(int i=1,u,v,w;i<=m;i++){
		cin >> u >> v >> w;
		line[++tot]=(edge){u,v,w,0};
	}
	init();
	K();
	dfs1(1,1);
	dfs2(1,1,1);
	dfs3(1,1);
	build(1,1,n);
	for(int i=1;i<=m;i++){
		if(line[i].tag) continue;
		cout << answer(line[i].u,line[i].v) << "\n";
	}
	return 0;
}

RT.

2023/10/1 08:14
加载中...