大分类讨论选手wa on #17,求助
查看原帖
大分类讨论选手wa on #17,求助
307940
aaaaaaaawsl楼主2022/11/2 19:40

思路是n遍dij跑出全源最短路,然后枚举看看是否能合并

#include<iostream>
#include<cstdio>
#include<cstring>
#include<algorithm>
#include<queue>
#define int long long
#define max(a,b) ((a)<(b)?(b):(a))

using namespace std;

inline int read(){
	register int x = 0, f = 1; register char ch = getchar();
	for(; ch > '9' || ch < '0'; ch = getchar()) if(ch == '-') f = -1;
	for(; ch >= '0' && ch <= '9'; ch = getchar()) x = (x << 1) + (x << 3) + (ch ^ '0');
	return x * f;
}

const int N = 2600;
const int M = 1e4 + 10;
int val[N];
int n, m, k, ans;
int dis[N][N];
int vis[N];
struct node{
	int pos, val;
	bool operator < (const node &a) const{
		return a.val < val;
	}
};
priority_queue<node> q;

int head[N], Next[M << 1], e[M << 1], idx;
inline void add(int a, int b){
	e[++ idx] = b; Next[idx] = head[a]; head[a] = idx;
}

void dij(int u){
	memset(vis, 0, sizeof vis);
	q.push(node{u, 0});
	while(q.size()){
		node now = q.top();
		q.pop();
		if(vis[now.pos]) continue;
		vis[now.pos] = 1;
		for(int i = head[now.pos]; i ; i = Next[i]){
			int j = e[i];
			if(dis[u][j] > dis[u][now.pos] + 1){
				dis[u][j] = dis[u][now.pos] + 1;
				if(!vis[j]) q.push(node{j, dis[u][j]});
			}
		}
	}
}

struct Node{
	int mx1 = -1919810;
	int mx2 = -1919810;
	int mx3 = -1919810;
	int ft, se, tr;
}f[N];

signed main(){
	n = read(); m = read(); k = read();
	k ++;
	for(int i = 2; i <= n; ++ i){
		val[i] = read();
	}
	for(int i = 1; i <= m; ++ i){
		int a = read(), b = read();
		add(a, b);
		add(b, a);
	}
	
	memset(dis, 0x3f, sizeof dis);
	for(int i = 1; i <= n; ++ i) dis[i][i] = 0;
	for(int i = 1; i <= n; ++ i) dij(i);
	
	for(int i = 2; i <= n; ++ i){ // 1 - j - i;
		for(int j = 2; j <= n; ++ j){
			if(dis[1][j] > k || dis[j][i] > k || i == j) continue;
			if(f[i].mx1 <= val[j] + val[i]){
				f[i].mx3 = f[i].mx2;
				f[i].mx2 = f[i].mx1;
				f[i].mx1 = val[j] + val[i];
				f[i].tr = f[i].se;
				f[i].se = f[i].ft;
				f[i].ft = j;
			}
			else if(f[i].mx2 <= val[j] + val[i]){
				f[i].mx3 = f[i].mx2;
				f[i].mx2 = val[j] + val[i];
				f[i].tr = f[i].se;
				f[i].se = j;
			}
			else if(f[i].mx3 <= val[j] + val[i]){
				f[i].mx3 = val[j] + val[i];
				f[i].tr = j;
			}
		}
	}
	for(int i = 2; i <= n; ++ i){
		for(int j = 2; j < i; ++ j){
			if(dis[i][j] > k) continue;
			if(f[i].ft != j){
				if(f[i].ft != f[j].ft && f[j].ft != i){
					ans = max(ans,f[i].mx1+f[j].mx1) ;
				}
				if(f[i].ft != f[j].se && f[j].se != i){
					ans = max(ans,f[i].mx1+f[j].mx2) ;
				}
				if(f[i].ft != f[j].tr && f[j].tr != i){
					ans = max(ans,f[i].mx1+f[j].mx3) ;
				}
			}
			if(f[i].se != j){
				if(f[i].se != f[j].ft && f[j].ft != i){
					ans = max(ans,f[i].mx2+f[j].mx1) ;
				}
				if(f[i].se != f[j].se && f[j].se != i){
					ans = max(ans,f[i].mx2+f[j].mx2) ;
				}
				if(f[i].se != f[j].tr && f[j].tr != i){
					ans = max(ans,f[i].mx2+f[j].mx3) ;
				}
			}
			if(f[i].tr != j){
				if(f[i].tr != f[j].ft && f[j].ft != i){
					ans = max(ans,f[i].mx3+f[j].mx1) ;
				}
				if(f[i].tr != f[j].se && f[j].se != i){
					ans = max(ans,f[i].mx3+f[j].mx2) ;
				}
				if(f[i].tr != f[j].tr && f[j].tr != i){
					ans = max(ans,f[i].mx3+f[j].mx3) ;
				}
			}
		}
	}
	printf("%lld", ans);
}
2022/11/2 19:40
加载中...