Use this to learn the idea, then write your own version.
123 4class UnionFind(object): 5 def __init__(self, n):6 self.set = range(n)7 self.rank = [0]*n8 9 def find_set(self, x):10 stk = []11 while self.set[x] != x: 12 stk.append(x)13 x = self.set[x]14 while stk:15 self.set[stk.pop()] = x16 return x17 18 def union_set(self, x, y):19 x, y = self.find_set(x), self.find_set(y)20 if x == y:21 return False22 if self.rank[x] > self.rank[y]: 23 x, y = y, x24 self.set[x] = self.set[y]25 if self.rank[x] == self.rank[y]:26 self.rank[y] += 127 return True28 29 30def binary_search(left, right, check):31 while left <= right:32 mid = left+(right-left)233 if check(mid):34 right = mid-135 else:36 left = mid+137 return left38 39 4041class Solution(object):42 def findMedian(self, n, edges, queries):43 """44 :type n: int45 :type edges: List[List[int]]46 :type queries: List[List[int]]47 :rtype: List[int]48 """49 def iter_dfs():50 lookup = [False]*len(adj)51 lookup2 = [[] for _ in xrange(len(adj))]52 for i, q in enumerate(queries):53 for x in q:54 lookup2[x].append(i)55 uf = UnionFind(len(adj))56 ancestor = range(len(adj))57 depth = [0]*len(adj)58 dist = [0]*len(adj)59 lca = [0]*len(queries)60 result = [0]*len(queries)61 stk = [(1, (0,))]62 while stk:63 step, args = stk.pop()64 if step == 1:65 u = args[0]66 for i in lookup2[u]:67 if queries[i][0] == queries[i][1]:68 lca[i] = u69 continue70 result[i] += dist[u]71 for x in queries[i]:72 if lookup[x]:73 lca[i] = ancestor[uf.find_set(x)]74 result[i] -= 2*dist[lca[i]]75 lookup[u] = True76 stk.append((2, (u, 0)))77 elif step == 2:78 u, i = args79 if i == len(adj[u]):80 continue81 v, w = adj[u][i]82 stk.append((2, (u, i+1)))83 if lookup[v]:84 continue85 dist[v] = dist[u]+w86 depth[v] = depth[u]+187 stk.append((3, (v, u)))88 stk.append((1, (v, u)))89 elif step == 3:90 v, u = args91 uf.union_set(v, u)92 ancestor[uf.find_set(u)] = u 93 return result, lca, dist, depth94 95 def iter_dfs2():96 lookup3 = [[] for _ in xrange(len(adj))]97 for i, (u, v) in enumerate(queries):98 if 2*(dist[u]-dist[lca[i]]) >= result[i]:99 lookup3[u].append((i, 0))100 else:101 lookup3[v].append((i, 1))102 result2 = [0]*len(queries)103 path = []104 stk = [(1, (0,))]105 while stk:106 step, args = stk.pop()107 if step == 1:108 u = args[0]109 path.append(u)110 for i, t in lookup3[u]:111 d = depth[u]-depth[lca[i]]112 if t == 0:113 j = binary_search(0, d, lambda x: 2*(dist[u]-dist[path[-(x+1)]]) >= result[i])114 result2[i] = path[-(j+1)]115 else:116 l = dist[queries[i][0]]-dist[lca[i]]117 j = binary_search(0, d-1, lambda x: 2*(l+(dist[path[-((d-1)+1)+x]]-dist[lca[i]])) >= result[i])118 result2[i] = path[-((d-1)+1)+j]119 stk.append((3, None))120 stk.append((2, (u, 0)))121 elif step == 2:122 u, i = args123 if i == len(adj[u]):124 continue125 v, w = adj[u][i]126 stk.append((2, (u, i+1)))127 if len(path) >= 2 and path[-2] == v:128 continue129 dist[v] = dist[u]+w130 depth[v] = depth[u]+1131 stk.append((1, (v, u)))132 elif step == 3:133 path.pop()134 return result2135 136 adj = [[] for _ in xrange(len(edges)+1)]137 for u, v, w in edges:138 adj[u].append((v, w))139 adj[v].append((u, w))140 result, lca, dist, depth = iter_dfs()141 return iter_dfs2()142 143 144145146147class Solution2(object):148 def findMedian(self, n, edges, queries):149 """150 :type n: int151 :type edges: List[List[int]]152 :type queries: List[List[int]]153 :rtype: List[int]154 """155 def dfs(u):156 for i in lookup2[u]:157 if queries[i][0] == queries[i][1]:158 lca[i] = u159 continue160 result[i] += dist[u]161 for x in queries[i]:162 if lookup[x]:163 lca[i] = ancestor[uf.find_set(x)]164 result[i] -= 2*dist[lca[i]]165 lookup[u] = True166 for v, w in adj[u]:167 if lookup[v]:168 continue169 dist[v] = dist[u]+w170 depth[v] = depth[u]+1171 dfs(v)172 uf.union_set(v, u)173 ancestor[uf.find_set(u)] = u174 175 def dfs2(u):176 path.append(u)177 for i, t in lookup3[u]:178 d = depth[u]-depth[lca[i]]179 if t == 0:180 j = binary_search(0, d, lambda x: 2*(dist[u]-dist[path[-(x+1)]]) >= result[i])181 result2[i] = path[-(j+1)]182 else:183 l = dist[queries[i][0]]-dist[lca[i]]184 j = binary_search(0, d-1, lambda x: 2*(l+(dist[path[-((d-1)+1)+x]]-dist[lca[i]])) >= result[i])185 result2[i] = path[-((d-1)+1)+j]186 for v, w in adj[u]:187 if len(path) >= 2 and path[-2] == v:188 continue189 dfs2(v)190 path.pop()191 192 adj = [[] for _ in xrange(len(edges)+1)]193 for u, v, w in edges:194 adj[u].append((v, w))195 adj[v].append((u, w))196 lookup = [False]*len(adj)197 lookup2 = [[] for _ in xrange(len(adj))]198 for i, q in enumerate(queries):199 for x in q:200 lookup2[x].append(i)201 uf = UnionFind(len(adj))202 ancestor = range(len(adj))203 dist = [0]*len(adj)204 depth = [0]*len(adj)205 result = [0]*len(queries)206 lca = [-1]*len(queries)207 dfs(0)208 result2 = [0]*len(queries)209 lookup3 = [[] for _ in xrange(len(adj))]210 for i, (u, v) in enumerate(queries):211 if 2*(dist[u]-dist[lca[i]]) >= result[i]:212 lookup3[u].append((i, 0))213 else:214 lookup3[v].append((i, 1))215 path = []216 dfs2(0)217 return result2218