import java.util.*;
public class Main {
public static void main(String[] args) {
Scanner input = new Scanner(System.in);
int n = input.nextInt();
int[] happy = new int[n + 1];
List<Integer>[] tree = new ArrayList[n + 1];
Arrays.setAll(tree, e -> new ArrayList<Integer>());
for (int i = 1; i <= n; i++) {
happy[i] = input.nextInt();
}
boolean[] vis = new boolean[n + 1];
for (int i = 1; i < n; i++) {
int x = input.nextInt();
int y = input.nextInt();
tree[y].add(x);
vis[x] = true;
}
int root = 0;
for (int i = 1; i <= n; i++) {
if (!vis[i]) {
root = i;
break;
}
}
int[][] dp = new int[n + 1][2];
dfs(root, tree, happy, dp);
System.out.println(Math.max(dp[root][0], dp[root][1]));
}
private static void dfs(int node, List<Integer>[] tree, int[] happy, int[][] dp) {
dp[node][1] = happy[node];
for (Integer next : tree[node]) {
dfs(next, tree, happy, dp);
dp[node][0] += Math.max(dp[next][0], dp[next][1]);
dp[node][1] += dp[next][0];
}
}
}