[Link] https://www.acmicpc.net/problem/1949 · 한국어 · 日本語
A tree has a population weight at each village. Choose villages with the maximum total population subject to the rule that no two adjacent villages are both chosen. The input contains N, then the N population values, followed by the N - 1 undirected roads.
Root the tree at village 1 iteratively. Store the order in which vertices are visited; processing that order in reverse guarantees every child is handled before its parent without recursion (and therefore avoids stack overflow on a long path). For each village u, maintain two values:
take[u]: the best total whenuis selected. Its children cannot be selected, sotake[u] = population[u] + sum(skip[child]).skip[u]: the best total whenuis not selected. Each child may be selected or skipped, soskip[u] = sum(max(take[child], skip[child])).
The answer is max(take[root], skip[root]). Since populations are positive, selecting the root alone is always a positive valid choice; the empty selection cannot improperly win. A single node therefore returns its population. In a path, neighboring choices compete through the two states, and in a star, selecting the center competes with selecting its leaves.
The adjacency lists, traversal order, and DP arrays each use O(N) space. Every vertex and edge is processed a constant number of times, for O(N) time. Use long for the DP totals so summed populations do not overflow an int.
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
import java.io.BufferedReader;
import java.io.IOException;
import java.io.InputStreamReader;
import java.util.ArrayList;
import java.util.StringTokenizer;
public class Main {
public static void main(String[] args) throws IOException {
BufferedReader input = new BufferedReader(new InputStreamReader(System.in));
int n = Integer.parseInt(input.readLine().trim());
long[] population = new long[n];
StringTokenizer values = new StringTokenizer(input.readLine());
for (int i = 0; i < n; i++) {
population[i] = Long.parseLong(values.nextToken());
}
ArrayList<Integer>[] graph = new ArrayList[n];
for (int i = 0; i < n; i++) {
graph[i] = new ArrayList<>();
}
for (int i = 0; i < n - 1; i++) {
StringTokenizer edge = new StringTokenizer(input.readLine());
int a = Integer.parseInt(edge.nextToken()) - 1;
int b = Integer.parseInt(edge.nextToken()) - 1;
graph[a].add(b);
graph[b].add(a);
}
int[] parent = new int[n];
int[] order = new int[n];
int size = 0;
order[size++] = 0;
parent[0] = -1;
for (int i = 0; i < size; i++) {
int node = order[i];
for (int neighbor : graph[node]) {
if (neighbor == parent[node]) {
continue;
}
parent[neighbor] = node;
order[size++] = neighbor;
}
}
long[] take = new long[n];
long[] skip = new long[n];
for (int i = n - 1; i >= 0; i--) {
int node = order[i];
take[node] = population[node];
for (int child : graph[node]) {
if (parent[child] == node) {
take[node] += skip[child];
skip[node] += Math.max(take[child], skip[child]);
}
}
}
System.out.println(Math.max(take[0], skip[0]));
}
}