30pts 悬赏关注
查看原帖
30pts 悬赏关注
365296
koobee楼主2023/8/26 20:45
#include<bits/stdc++.h>
using namespace std;
const int N = 5e5+5;
int g[N][21], f[N][21], n, d, a[N], flag[N], head[N], cnt, m;
struct node{
	int to, nxt;
} e[2*N];
void add(int x, int y){
	e[++cnt] = {y, head[x]}, head[x] = cnt;
}
void dfs(int u, int fa){
	if(flag[u] == 1) g[u][0] = f[u][0] = a[u];
	for(int i = 1; i <= d; i++) g[u][d] = a[u];
	for(int i = head[u]; i; i = e[i].nxt){
		int v = e[i].to;
		if(v == fa) continue;
		dfs(v, u);
		for(int j = d; j >= 0; j--){
			g[u][j] = min(g[u][j]+f[v][j], g[v][j+1]+f[u][j+1]);
			if(j < d) g[u][j] = min(g[u][j], g[u][j+1]);
		}
		f[u][0] = g[u][0];
		int mn = 2e9;
		for(int j = 1; j <= d; j++)
			mn = min(mn, f[v][j-1]), f[u][j] += mn;
		for(int j = 1; j <= d; j++)
			f[u][j] = min(f[u][j], f[u][j-1]);
	}
}
int main(){
//	freopen("julian.in","r",stdin);
//	freopen("julian.out","w",stdout);
	cin>>n>>d;
	for(int i = 1; i <= n; i++) cin>>a[i];
	cin>>m;
	for(int i = 1; i <= m; i++){
		int x;
		cin>>x;
		flag[x] = 1;
	}
	for(int i = 1; i < n; i++){
		int u, v;
		cin>>u>>v;
		add(u, v), add(v, u);
	}
	memset(g, 0x3f, sizeof(g));
	dfs(1, 0);
	cout<<g[1][0];
	return 0;
}
2023/8/26 20:45
加载中...