오늘은 트리의 지름이라는 문제를 풀어보았다. 처음에 풀었던 방식으로는 메모리 초과가 발생해서 결국은 다른 분들의 해설을 참고해서 풀었다. 참고한 방식이 좋은 문제라서 가져왔다.
트리(tree)는 사이클이 없는 무방향 그래프이다. 트리에서는 어떤 두 노드를 선택해도 둘 사이에 경로가 항상 하나만 존재하게 된다. 트리에서 어떤 두 노드를 선택해서 양쪽으로 쫙 당길 때, 가장 길게 늘어나는 경우가 있을 것이다. 이럴 때 트리의 모든 노드들은 이 두 노드를 지름의 끝 점으로 하는 원 안에 들어가게 된다.
이런 두 노드 사이의 경로의 길이를 트리의 지름이라고 한다. 정확히 정의하자면 트리에 존재하는 모든 경로들 중에서 가장 긴 것의 길이를 말한다.
입력으로 루트가 있는 트리를 가중치가 있는 간선들로 줄 때, 트리의 지름을 구해서 출력하는 프로그램을 작성하시오. 아래와 같은 트리가 주어진다면 트리의 지름은 45가 된다.
트리의 노드는 1부터 n까지 번호가 매겨져 있다.
파일의 첫 번째 줄은 노드의 개수 n(1 ≤ n ≤ 10,000)이다. 둘째 줄부터 n-1개의 줄에 각 간선에 대한 정보가 들어온다. 간선에 대한 정보는 세 개의 정수로 이루어져 있다. 첫 번째 정수는 간선이 연결하는 두 노드 중 부모 노드의 번호를 나타내고, 두 번째 정수는 자식 노드를, 세 번째 정수는 간선의 가중치를 나타낸다. 간선에 대한 정보는 부모 노드의 번호가 작은 것이 먼저 입력되고, 부모 노드의 번호가 같으면 자식 노드의 번호가 작은 것이 먼저 입력된다. 루트 노드의 번호는 항상 1이라고 가정하며, 간선의 가중치는 100보다 크지 않은 양의 정수이다.
첫째 줄에 트리의 지름을 출력한다.
처음 풀었을 때에는 모든 노드를 짝지어서 dfs를 돌린 다음 최대 가중치를 갱신해 주는 방법으로 구현하였는데 메모리 초과가 발생하였다.
아래의 코드처럼 작성하게 된다면 visited 배열이 수천만 번 생성이 되므로 메모리 제한이 128MB인 이 문제에서는 메모리 초과가 발생되었던 것이다. 그래서 이 코드로는 해결이 불가했던 것이다.(GC가 기존 visited 배열을 지워주지만 너무 많이 생성하면 GC의 제거보다 생성이 더 많아서 결국은 GC가 따라가지 못하여 메모리 초과가 발생한다.)
import java.io.BufferedReader;
import java.io.IOException;
import java.io.InputStreamReader;
import java.util.ArrayList;
import java.util.StringTokenizer;
class Main {
static int N;
static ArrayList<ArrayList<int[]>> graph = new ArrayList<>();
static boolean[] visited;
static int result = 0;
public static void main(String[] args) throws IOException {
BufferedReader br = new BufferedReader(new InputStreamReader(System.in));
N = Integer.parseInt(br.readLine());
// 그래프 생성
for (int i = 0; i < N + 1; i++) {
graph.add(new ArrayList<>());
}
// 간선 연결 - 가중치랑 같이
for (int i = 0; i < N - 1; i++) {
StringTokenizer link = new StringTokenizer(br.readLine());
int s = Integer.parseInt(link.nextToken());
int e = Integer.parseInt(link.nextToken());
int weight = Integer.parseInt(link.nextToken());
graph.get(s).add(new int[] {e, weight});
graph.get(e).add(new int[] {s, weight});
}
// 2중 for문으로 dfs 돌리기
for (int i = 1; i < N; i++) {
for (int j = i + 1; j < N + 1; j++) {
// 방문 리스트 생성
visited = new boolean[N + 1];
dfs(i, j, 0);
}
}
System.out.println(result);
}
private static void dfs(int start, int goal, int count) {
visited[start] = true;
if (start == goal) {
result = Math.max(result, count);
return;
}
for (int[] next : graph.get(start)) {
int node = next[0];
int weight = next[1];
if (!visited[node]) {
dfs(node, goal, count + weight);
}
}
}
}
그래서 다른 분들의 해설을 참조해 보니 처음 루트 노드에서 dfs 탐색을 통해 가장 가중치가 큰 노드를 찾은 다음 그 노드에서 다시 한번 dfs 탐색을 이용하여 가장 가중치가 큰 경우를 찾는 방법을 사용하면 시간과 메모리 측면에서 매우 효율적인 코드를 작성할 수 있다는 것을 알게 되었다.
아래는 정답 코드이다.
import java.io.BufferedReader;
import java.io.IOException;
import java.io.InputStreamReader;
import java.util.ArrayList;
import java.util.StringTokenizer;
class Main {
static int N;
static ArrayList<ArrayList<int[]>> graph = new ArrayList<>();
static boolean[] visited;
static int maxNode = 0;
static int result = 0;
public static void main(String[] args) throws IOException {
BufferedReader br = new BufferedReader(new InputStreamReader(System.in));
N = Integer.parseInt(br.readLine());
// 그래프 생성
for (int i = 0; i < N + 1; i++) {
graph.add(new ArrayList<>());
}
// 간선 연결 - 가중치랑 같이
for (int i = 0; i < N - 1; i++) {
StringTokenizer link = new StringTokenizer(br.readLine());
int s = Integer.parseInt(link.nextToken());
int e = Integer.parseInt(link.nextToken());
int weight = Integer.parseInt(link.nextToken());
graph.get(s).add(new int[] {e, weight});
graph.get(e).add(new int[] {s, weight});
}
// 루트 노드에서 가장 큰 가중치를 가진 maxNode를 찾은 후 maxNode에서 가중치가 가장 큰 노드를 찾기
visited = new boolean[N + 1];
dfs(1, 0);
visited = new boolean[N + 1];
dfs(maxNode, 0);
System.out.println(result);
}
private static void dfs(int node, int count) {
visited[node] = true;
if (count > result) {
result = Math.max(result, count);
maxNode = node;
}
for (int[] next : graph.get(node)) {
int n = next[0];
int weight = next[1];
if (!visited[n]) {
dfs(n,count + weight);
}
}
}
}