
예를 들어 해당 트리에서 11과 2의 LCA는 7이다.
def __init__(self, edges: List[List[int]]):
self.graph = defaultdict(list)
self.SIZE = 0
for s, e in edges:
self.graph[s].append(e)
self.graph[e].append(s)
self.SIZE = max(self.SIZE, s, e)
self.depth = [0] * (self.SIZE + 1)
self.parent = [[0] * (self.SIZE + 1) for _ in range(LOG)]
self.build()
def build(self, root: int = 1):
visited = [False] * (self.SIZE + 1)
q = deque([(root, 1)])
# build depth
while q:
n, d = q.popleft()
self.depth[n] = d
visited[n] = True
for dst in self.graph[n]:
if visited[dst]:
continue
self.parent[0][dst] = n
q.append((dst, d + 1))
# build parent
for k in range(1, LOG):
for node in range(self.SIZE + 1):
if self.depth[node] == 0:
continue
mid = self.parent[k - 1][node]
self.parent[k][node] = self.parent[k - 1][mid]
if self.depth[a] < self.depth[b]:
a, b = b, a
diff = self.depth[a] - self.depth[b]
if diff > 0:
for level in range(LOG):
if diff & (1 << level):
a = self.parent[level][a]
for level in range(LOG - 1, -1, -1):
if self.parent[level][a] != self.parent[level][b]:
a = self.parent[level][a]
b = self.parent[level][b]
LOG = 21
class LCA:
def __init__(self, edges: List[List[int]]):
self.graph = defaultdict(list)
self.SIZE = 0
for s, e in edges:
self.graph[s].append(e)
self.graph[e].append(s)
self.SIZE = max(self.SIZE, s, e)
self.depth = [0] * (self.SIZE + 1)
self.parent = [[0] * (self.SIZE + 1) for _ in range(LOG)]
self.build()
def build(self, root: int = 1):
visited = [False] * (self.SIZE + 1)
q = deque([(root, 1)])
# build depth
while q:
n, d = q.popleft()
self.depth[n] = d
visited[n] = True
for dst in self.graph[n]:
if visited[dst]:
continue
self.parent[0][dst] = n
q.append((dst, d + 1))
# build parent
for k in range(1, LOG):
for node in range(self.SIZE + 1):
if self.depth[node] == 0:
continue
mid = self.parent[k - 1][node]
self.parent[k][node] = self.parent[k - 1][mid]
def query(self, a: int, b: int) -> int:
if self.depth[a] < self.depth[b]:
a, b = b, a
diff = self.depth[a] - self.depth[b]
if diff > 0:
for level in range(LOG):
if diff & (1 << level):
a = self.parent[level][a]
if a == b:
return a
for level in range(LOG - 1, -1, -1):
if self.parent[level][a] != self.parent[level][b]:
a = self.parent[level][a]
b = self.parent[level][b]
return self.parent[0][a]