BaekJoon 13511번 : 트리와 쿼리 2 (python)

owei·2024년 5월 23일

백준

목록 보기
62/62

📝 BaekJoon 13511번 : 트리와 쿼리 2 (P3 28.565%)


🔎 트리와 쿼리 2 문제


📌 아이디어

가중치를 더한 LCA알고리즘을 이용한 문제이며 Cost는 LCA알고리즘으로 간단하게 구현하고 k번째 정수를 구할 때 경우에 따라 분기를 해주어야 한다.


💭 풀이

  • 첫번째로 u에서 v로 가는 경로의 비용을 출력하기 위해서는 cost배열을 sparse_table을 이용해 정의해주고 이를 통해 cost를 구해주어야 한다.
  • u와 v의 공통 root를 구하여 u와 v에서 공통 root까지의 cost를 각각 구하여 합쳐주면 u에서 v까지 가는 경로의 비용을 구할 수 있다.
  • 두 번째로 u에서 v로 가는 경로에 k번째 정점을 출력하는 부분은 u와 v의 공통 root까지의 길이를 구한 다음 해당 k가 k<=depth_u + 1인지 아닌지를 나눠주어야 한다.
  • 만약 k<=depth_u + 1이라면 k번째 정점은 u에서 공통 root까지 가는길에 있는 root이기에 lca알고리즘을 이용하여 u에서 k-1번째의 root를 구하면 된다.
  • 만약 k > depth_u + 1인 경우라면 v에서 공통 root가는 길에 있기 때문에 v에서
    depth_v - (k-(depth_u + 1)) 위의 root를 구하면 되게 된다.

💻 코드

from collections import deque
import sys
input = sys.stdin.readline

n = int(input())
edge = [[] for _ in range(n+1)]

for _ in range(n-1) :
    a, b, c = map(int,input().split())
    edge[a].append((b,c))
    edge[b].append((a,c))

root = [0]*(n+1)
level = [0]*(n+1)

q = deque()
q.append((1,0))
check = [False]*(n+1)
check[1] = True
while q :
    x, count = q.popleft()
    level[x] = count

    for i in edge[x] :
        if not check[i[0]] :
            check[i[0]] = True
            root[i[0]] = (x, i[1])
            q.append((i[0],count+1))

sparse_table = [[0]*21 for _ in range(n+1)]
cost_table = [[0]*21 for _ in range(n+1)]

for i in range(2, n+1) :
    sparse_table[i][0] = root[i][0]
    cost_table[i][0] = root[i][1]

for i in range(1, 21) :
    for j in range(1, n+1) :
        sparse_table[j][i] = sparse_table[sparse_table[j][i-1]][i-1]
        cost_table[j][i] = cost_table[j][i-1] + cost_table[sparse_table[j][i-1]][i-1]

def kth_ancestor(u, k):
    for i in range(21):
        if (k >> i) & 1:
            u = sparse_table[u][i]

    return u

def lca(u,v) :
    if level[u] > level[v] :
        u, v = v, u

    diff = level[v] - level[u]
    for i in range(21):
        if (diff >> i) & 1:
            v = sparse_table[v][i]

    if u == v:
        return u

    for i in range(20, -1,-1):
        if sparse_table[u][i] != sparse_table[v][i]:
            u = sparse_table[u][i]
            v = sparse_table[v][i]

    return root[u][0]


m = int(input())
for _ in range(m) :
    s = list(map(int,input().split()))
    if s[0] == 1:
        u, v = s[1],s[2]
        cost = 0
        lca_node = lca(u,v)

        diff_u = level[u] - level[lca_node]
        diff_v = level[v] - level[lca_node]

        for i in range(21):
            if (diff_u >> i) & 1:
                cost += cost_table[u][i]
                u = sparse_table[u][i]

        for i in range(21):
            if (diff_v >> i) & 1:
                cost += cost_table[v][i]
                v = sparse_table[v][i]

        print(cost)

    else :
        u, v, k = s[1],s[2],s[3]
        lca_node = lca(u,v)
        depth_u = level[u] - level[lca_node]
        depth_v = level[v] - level[lca_node]
        
        if k <= depth_u + 1 :
            print(kth_ancestor(s[1], k-1))

        else :
            print(kth_ancestor(s[2], depth_v - (k-(depth_u + 1))))

profile
owei

0개의 댓글