求助 DFS剪枝部分分
查看原帖
求助 DFS剪枝部分分
247202
OnlyExtreme楼主2022/10/29 22:50

应该能有30~50,结果全 wa

#include <bits/stdc++.h>
using namespace std;

int n, m, k;
int score[3000];
vector<int> g[3000];

namespace k0 {
	long long ans = -1;
	bool vis[3000];
	int maxn[3000][110];
	
	void init() {
		memset(vis, 0, sizeof(vis));
		for(int i=1; i<=n; i++) {
			for(int j=1; j<=k; j++) {
				maxn[i][j] = -1;
			}
		}
	}
	
	void dfs(int cur, int fa, int cnt, int sc) {
		if(cnt == 5) {
			if(cur == 1) {
				ans = max(ans, 1ll*sc);
//				printf(" %d %d %d %d\n", cur, fa, cnt, sc);
			}
			return;
		}
		sc += score[cur];
		if(sc < maxn[cur][cnt]) return;
		maxn[cur][cnt] = sc;
		for(int i=0; i<g[cur].size(); i++) {
			int nxt = g[cur][i];
			if(nxt == fa) continue;
			if(vis[nxt]) continue;
			vis[nxt] = true;
			dfs(nxt, cur, cnt+1, sc);
			vis[nxt] = false;
		}
	}
	
	void solve() {
		init();
		dfs(1, 1, 0, 0);
		printf("%lld\n", ans);
		return;
	}
}

namespace sol {
	vector<int> reach[3000];
	long long ans = -1;
	bool vis[3000];
	int maxn[3000][110];
	
	void init() {
		memset(vis, 0, sizeof(vis));
		for(int i=1; i<=n; i++) {
			for(int j=1; j<=k; j++) {
				maxn[i][j] = -1;
			}
		}
	}
	
	void pre(int cur, int fa, int rt, int depth) {
		if(depth == k+2) return;
		bool f[3000];
		memset(f, 0, sizeof(f));
		if(cur != rt) reach[rt].push_back(cur);
		for(int i=0; i<g[cur].size(); i++) {
			int nxt = g[cur][i];
			if(nxt == fa) continue;
			if(f[nxt]) continue;
			f[nxt] = true;
			pre(nxt, cur, rt, depth+1);
		}
	}
	
	void dfs(int cur, int fa, int cnt, int sc) {
//		printf(" cur=%d cnt=%d sc=%d max=%d\n", cur, cnt, sc, maxn[cur][cnt]);
		if(cnt == 5) {
			if(cur == 1) {
				ans = max(ans, 1ll*sc);
//				printf(" %d %d %d %d\n", cur, fa, cnt, sc);
			}
			return;
		}
		sc += score[cur];
		if(sc < maxn[cur][cnt]) return;
		maxn[cur][cnt] = sc;
		for(int i=0; i<reach[cur].size(); i++) {
			int nxt = reach[cur][i];
			if(nxt == fa) continue;
			if(nxt == 1 && cnt < 4) continue;
			if(vis[nxt]) continue;
			vis[nxt] = true;
			dfs(nxt, cur, cnt+1, sc);
			vis[nxt] = false;
		}
	}
	
	void solve() {
		for(int i=1; i<=n; i++) {
			pre(i, i, i, 0);
		}
//		for(int i=1; i<=n; i++) {
//			for(int j=0; j<reach[i].size(); j++) {
//				printf("%d ", reach[i][j]);
//			}
//			printf("\n");
//		}
		init();
		dfs(1, 1, 0, 0);
		printf("%lld\n", ans);
		return;
	}
}


int main() {
//	freopen("holiday.in", "r", stdin);
//	freopen("holiday.out", "w", stdout);
	scanf("%d %d %d", &n, &m, &k);
	for(int i=1; i<n; i++)
		scanf("%d", &score[i+1]);
	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);
	}
//	if(k == 0) {
//		k0::solve();
//	}
	sol::solve();
	return 0;
}
2022/10/29 22:50
加载中...