n개의 노드로 구성된 n-ary 트리가 인접 리스트 형태의 2차원 리스트 tree로 주어지고, 각 노드의 색상 정보는 리스트 color에 담겨 있습니다. 트리의 루트는 tree[0]에 해당합니다.
문제 정의
i번째 노드의 특성은 다음과 같습니다.
- tree[i]: i번째 노드의 자식 노드와 부모 노드 정보
- color[i]: i번째 노드의 색상
어떤 노드 N을 루트로 하는 서브트리에 속한 모든 노드의 색상이 중복 없이 고유할 때, 이 노드 N을 '특수(special) 노드'라고 부릅니다. 즉, 주어진 트리에서 특수 노드가 총 몇 개인지 구하는 것이 이 문제의 목표입니다.
예시
tree = [
[1, 2],
[0],
[0, 3],
[2]
]colors = [1, 2, 1, 1]일 때 출력 결과는 2입니다.
실제로 하나씩 확인해 보면 다음과 같습니다.
- 노드 1(리프): 색상이 {2} 하나뿐이므로 특수 노드
- 노드 3(리프): 색상이 {1} 하나뿐이므로 특수 노드
- 노드 2: 서브트리 {2, 3}의 색상이 {1, 1}로 중복되므로 특수 노드가 아님
- 노드 0(루트): 전체 트리의 색상이 {1, 2, 1}로 중복되므로 특수 노드가 아님
따라서 특수 노드는 노드 1과 노드 3, 총 2개입니다.
해결 접근 방식
이 문제는 깊이 우선 탐색(DFS)과 집합(set)을 조합하면 효율적으로 해결할 수 있습니다. 핵심 아이디어는 각 노드의 서브트리에서 사용된 색상 집합을 재귀적으로 수집하고, 자식들의 색상 집합을 병합하는 과정에서 중복이 발견되면 해당 노드를 특수 노드에서 제외하는 것입니다.
전체 알고리즘의 흐름은 다음과 같습니다.
- result := 0으로 초기화합니다.
- dfs(0, -1)을 호출합니다.
- result를 반환합니다.
check_intersection() 함수
두 색상 집합 사이에 교집합이 존재하는지 확인하는 보조 함수입니다.
- colors의 길이가 child_colors보다 짧으면, colors의 각 원소 c에 대해 c가 child_colors에 존재하는지 검사하고, 존재하면 True를 반환합니다.
- 그렇지 않으면, child_colors의 각 원소 c에 대해 c가 colors에 존재하는지 검사하고, 존재하면 True를 반환합니다.
이처럼 더 작은 집합을 기준으로 순회하면 불필요한 비교를 줄여 성능을 높일 수 있습니다.
dfs() 함수
- colors := {color[node]}로 초기화합니다.
- tree[node]의 각 child에 대해 다음을 수행합니다.
- child가 prev(직전 노드)와 같지 않은 경우:
- child_colors := dfs(child, node)로 자식의 색상 집합을 얻습니다.
- colors와 child_colors가 모두 비어 있지 않으면:
- check_intersection(colors, child_colors)가 참이면 중복 색상이 있다는 의미이므로 colors := null로 설정합니다.
- 그렇지 않으면 두 집합을 병합합니다. 이때 더 작은 집합을 더 큰 집합 쪽에 합치는 small-to-large 방식으로 최적화합니다.
- 둘 중 하나라도 이미 null이라면 colors := null로 설정합니다.
- child가 prev(직전 노드)와 같지 않은 경우:
- 탐색이 끝난 후 colors가 null이 아니면 result를 1 증가시킵니다. 즉, 현재 노드는 특수 노드입니다.
- colors를 반환합니다.
구현 예제
다음 구현을 통해 더 잘 이해해 보겠습니다.
import collections
class Solution:
def solve(self, tree, color):
self.result = 0
def dfs(node, prev):
colors = {color[node]}
for child in tree[node]:
if child != prev:
child_colors = dfs(child, node)
if colors and child_colors:
if self.check_intersection(colors, child_colors):
colors = None
else:
if len(colors) < len(child_colors):
child_colors |= colors
colors = child_colors
else:
colors |= child_colors
else:
colors = None
if colors:
self.result += 1
return colors
dfs(0, -1)
return self.result
def check_intersection(self, colors, child_colors):
if len(colors) < len(child_colors):
for c in colors:
if c in child_colors:
return True
else:
for c in child_colors:
if c in colors:
return True
ob = Solution()
print(ob.solve([
[1, 2],
[0],
[0, 3],
[2]
], [1, 2, 1, 1]))입력
[
[1, 2],
[0],
[0, 3],
[2]
], [1, 2, 1, 1]출력
2