Computer >> 컴퓨터 >  >> 프로그래밍 >> Python

Python으로 트리의 특수 노드(Special Node) 찾기


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)을 조합하면 효율적으로 해결할 수 있습니다. 핵심 아이디어는 각 노드의 서브트리에서 사용된 색상 집합을 재귀적으로 수집하고, 자식들의 색상 집합을 병합하는 과정에서 중복이 발견되면 해당 노드를 특수 노드에서 제외하는 것입니다.

전체 알고리즘의 흐름은 다음과 같습니다.

  1. result := 0으로 초기화합니다.
  2. dfs(0, -1)을 호출합니다.
  3. 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로 설정합니다.
  • 탐색이 끝난 후 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