1
# Runtime: 1783 ms (Top 30.42%) | Memory: 59.4 MB (Top 95.18%)
2
from typing import List
3

4
ROOT_PARENT = -1
5

6

7
class Solution:
8
def sumOfDistancesInTree(self, n: int, edges: List[List[int]]) -> List[int]:
9
"""
10
@see https://leetcode.com/problems/sum-of-distances-in-tree/discuss/130583/C%2B%2BJavaPython-Pre-order-and-Post-order-DFS-O(N)
11
:param n:
12
:param edges:
13
:return:
14
"""
15
g = self.create_undirected_graph(
16
edges, n
17
) # as mentioned in the problem, this graph can be converted into tree
18

19
root = 0 # can be taken to any node between 0 and n - 1 (both exclusive)
20

21
# considering "root" as starting node, we create a tree.
22
# Now defining,
23
# tree_nodes[i] = number of nodes in the tree rooted at node i
24
# distances[i] = sum of distances of all nodes from ith node to all the
25
# other nodes of the tree
26
tree_nodes, distances = [0] * n, [0] * n
27

28
def postorder(rt: int, parent: int):
29
"""
30
updating tree_nodes and distances from children of rt. To update them, we must know their
31
values at children. And that is why post order traversal is used
32

33
After the traversal is done,
34
tree_nodes[rt] = all the nodes in tree rooted at rt
35
distances[rt] = sum of distances from rt to all the nodes of tree rooted at rt
36

37
:param rt:
38
:param parent:
39
"""
40
tree_nodes[rt] = 1
41

42
for c in g[rt]:
43
if c != parent:
44
postorder(c, rt)
45

46
# adding number of nodes in subtree rooted at c to tree rooted at rt
47
tree_nodes[rt] += tree_nodes[c]
48

49
# moving to rt from c will increase distances by nodes in tree rooted at c
50
distances[rt] += distances[c] + tree_nodes[c]
51

52
def preorder(rt: int, parent: int):
53
"""
54
we start with "root" and update its children.
55
distances[root] = sum of distances between root and all the other nodes in tree.
56

57
In this function, we calculate distances[c] with the help of distances[root] and
58
that is why preorder traversal is required.
59
:param rt:
60
:param parent:
61
:return:
62
"""
63
for c in g[rt]:
64
if c != parent:
65
distances[c] = (
66
n - tree_nodes[c]
67
) + ( # rt -> c increase this much distance
68
distances[rt] - tree_nodes[c]
69
) # rt -> c decrease this much distance
70
preorder(c, rt)
71

72
postorder(root, ROOT_PARENT)
73
preorder(root, ROOT_PARENT)
74

75
return distances
76

77
@staticmethod
78
def create_undirected_graph(edges: List[List[int]], n: int):
79
"""
80
:param edges:
81
:param n:
82
:return: graph from edges. Note that this undirect graph is a tree. (Any node can be
83
picked as root node)
84
"""
85
g = [[] for _ in range(n)]
86

87
for u, v in edges:
88
g[u].append(v)
89
g[v].append(u)
90

91
return g

0

WPM •0 •0

100%

ACC •0 •0

0s

TIME •0