정점이 n개인 트리가 있다고 가정해 보겠습니다. 각 정점에는 1부터 n까지 번호가 붙어 있고, 루트 정점의 번호는 1입니다. 또한 모든 정점은 고유한 가중치 wi를 가지고 있습니다.
이때 n×n 크기의 행렬 A를 다음과 같이 정의할 수 있습니다. 행렬의 (x, y) 성분은 A(x, y) = Wf(x, y)인데, 여기서 f(x, y)는 정점 x와 y의 최소 공통 조상(LCA)입니다. 즉, 두 정점의 가장 가까운 공통 조상의 가중치가 행렬 성분이 되는 구조입니다. 이 글에서는 이렇게 만들어진 행렬 A의 행렬식(determinant)을 구하는 방법을 다룹니다. 트리의 간선 정보, 각 정점의 가중치, 전체 정점 수가 입력으로 주어집니다.
예시 입력과 출력
예를 들어 input_array = [[1, 2], [1, 3], [1, 4], [1, 5]], weights = [1, 2, 3, 4, 5], vertices = 5가 입력으로 주어진다면, 결과는 24가 됩니다.
이때 만들어지는 행렬 A는 다음과 같습니다.
| 1 | 1 | 1 | 1 | 1 |
| 1 | 2 | 1 | 1 | 1 |
| 1 | 1 | 3 | 1 | 1 |
| 1 | 1 | 1 | 4 | 1 |
| 1 | 1 | 1 | 1 | 5 |
이 행렬의 행렬식은 24입니다.
해결 접근 방법
이 문제는 깊이 우선 탐색(DFS)을 활용해 효율적으로 해결할 수 있습니다. 핵심 아이디어는 각 정점을 방문할 때마다 '현재 정점의 가중치에서 부모 정점의 가중치를 뺀 값'을 누적으로 곱하는 것입니다. 단계별 과정은 다음과 같습니다.
- w := 각 정점의 가중치와 인접 정점 목록을 담는 빈 리스트를 준비합니다.
- i를 0부터 vertices까지 반복하며, w에 weights[i]와 새 리스트를 추가합니다.
- enumerate(input_array)로 간선 정보를 순회하면서 다음을 수행합니다.
- p := item[0], q := item[1]
- w[p - 1][1] 끝에 q - 1을 추가합니다.
- w[q - 1][1] 끝에 p - 1을 추가합니다. (양방향 인접 리스트 구성)
- 행렬식 값 det := 1로 초기화합니다.
- stack := (0, 0) 튜플을 포함하는 스택을 만듭니다. (루트 정점, 부모 가중치 0)
- 스택이 비어 있지 않은 동안 다음을 반복합니다.
- i, parent_weight := 스택의 최상위 요소를 꺼냅니다.
- det := (det * (w[i][0] - parent_weight)) mod (10^9 + 7)
- w[i][1]의 각 자식 정점 t에 대해 (t, w[i][0]) 튜플을 스택에 추가합니다.
- 동시에 각 자식 t의 인접 리스트 w[t][1]에서 i를 제거하여, 트리의 위쪽 방향으로 되돌아가는 탐색을 차단합니다.
- 탐색이 끝나면 det를 반환합니다.
왜 이 공식이 성립할까?
수학적으로 잘 알려진 성질에 따르면, A(x, y) = W(LCA(x, y)) 형태의 행렬은 그 행렬식이 다음과 같이 매우 단순한 형태로 계산됩니다.
det(A) = ∏ (자식 정점의 가중치 − 부모 정점의 가중치)
위 예시에 적용해 보면 (1−0) × (2−1) × (3−1) × (4−1) × (5−1) = 1 × 1 × 2 × 3 × 4 = 24로 실제 결과와 일치합니다. 일반적인 n×n 행렬의 행렬식을 직접 계산하려면 상당한 연산량이 필요하지만, 이 성질을 활용하면 n×n 행렬을 명시적으로 만들지 않고도 트리를 한 번 순회하는 것만으로 답을 구할 수 있습니다. 결과값이 지수적으로 커질 수 있으므로 10^9 + 7로 나눈 나머지를 취하는 점도 참고하세요.
구현 예제
다음 구현을 통해 더 잘 이해해 보겠습니다.
def solve(input_array, weights, vertices): w = [[weights[i],[]] for i in range(vertices)] for i, item in enumerate(input_array): p,q = item[0], item[1] w[p - 1][1].append(q - 1) w[q - 1][1].append(p - 1) det = 1 stack = [(0,0)] while stack: i, parent_weight = stack.pop() det = (det * (w[i][0] - parent_weight)) % (10**9 + 7) stack += [(t,w[i][0]) for t in w[i][1]] for t in w[i][1]: w[t][1].remove(i) return det print(solve([[1, 2], [1, 3], [1, 4], [1, 5]], [1, 2, 3, 4, 5], 5))
입력
[[1, 2], [1, 3], [1, 4], [1, 5]], [1, 2, 3, 4, 5], 5
출력
24