Sum of Distances in Tree — Hard Problem & Solution

An undirected, connected tree has n nodes labelled 0 to n - 1 and the n - 1 edges in edges, each [ai, bi].

Problem statement

An undirected, connected tree has n nodes labelled 0 to n - 1 and the n - 1 edges in edges, each [ai, bi].

Return an array answer of length n where answer[i] is the sum of the distances (numbers of edges) between node i and every other node.

Example 1

Input: n = 5, edges = [[0,1],[1,2],[1,3],[3,4]]
Output: [8,5,8,6,9]
Explanation: From node 1 the distances are 1, 1, 1 and 2, which add up to 5.

Example 2

Input: n = 1, edges = []
Output: [0]

Example 3

Input: n = 3, edges = [[2,0],[2,1]]
Output: [3,3,2]

Constraints

  • 1 <= n <= 3 * 10^4
  • edges.length == n - 1
  • edges[i].length == 2
  • 0 <= ai, bi < n
  • ai != bi
  • the given input represents a valid tree

How to solve Sum of Distances in Tree

Re-rooting: compute the answer for the root directly, then shift the root across each edge, adjusting by how many nodes move closer and how many move farther.

Approach

  1. Root at node 0 and get a BFS order and parents.
  2. In reverse BFS order, accumulate count[u] (subtree size) and sub[u] (sum of distances from u to its subtree): count[p] += count[c], sub[p] += sub[c] + count[c].
  3. answer[0] = sub[0].
  4. In BFS order, for every child c of p: answer[c] = answer[p] − count[c] + (n − count[c]).

Why it works

Moving from p to its child c shortens the distance to each of the count[c] nodes in c's subtree by one and lengthens the distance to each of the other n − count[c] nodes by one; nothing else changes. Starting from the exact root value and applying this along every edge gives every node's exact sum.

Complexity

  • Time — O(n)
  • Space — O(n)

Pitfalls

  • Recursive DFS over a 3 · 10^4-node path can overflow the call stack; a BFS order avoids recursion.
  • sub[p] must add count[c] as well as sub[c] — every node of the child subtree is one edge farther from p.
  • n = 1 has no edges; the answer is [0].

Reference solution

Python

from typing import List

def sumOfDistancesInTree(n: int, edges: List[List[int]]) -> List[int]:
    adj = [[] for _ in range(n)]
    for a, b in edges:
        adj[a].append(b)
        adj[b].append(a)
    parent = [-1] * n
    seen = [False] * n
    seen[0] = True
    order = [0]
    for u in order:
        for v in adj[u]:
            if not seen[v]:
                seen[v] = True
                parent[v] = u
                order.append(v)
    count = [1] * n
    sub = [0] * n
    for u in reversed(order):
        p = parent[u]
        if p >= 0:
            count[p] += count[u]
            sub[p] += sub[u] + count[u]
    ans = [0] * n
    ans[0] = sub[0]
    for u in order:
        if u != 0:
            ans[u] = ans[parent[u]] - count[u] + (n - count[u])
    return ans

JavaScript

var sumOfDistancesInTree = function(n, edges) {
    var adj = [], parent = [], count = [], sub = [], ans = [];
    for (var i = 0; i < n; i++) { adj.push([]); parent.push(-2); count.push(1); sub.push(0); ans.push(0); }
    for (var e = 0; e < edges.length; e++) {
        adj[edges[e][0]].push(edges[e][1]);
        adj[edges[e][1]].push(edges[e][0]);
    }
    parent[0] = -1;
    var order = [0];
    for (var h = 0; h < order.length; h++) {
        var u = order[h];
        for (var j = 0; j < adj[u].length; j++) {
            var v = adj[u][j];
            if (parent[v] === -2) { parent[v] = u; order.push(v); }
        }
    }
    for (var t = n - 1; t > 0; t--) {
        var c = order[t], p = parent[c];
        count[p] += count[c];
        sub[p] += sub[c] + count[c];
    }
    ans[0] = sub[0];
    for (var k = 1; k < n; k++) {
        var w = order[k];
        ans[w] = ans[parent[w]] - count[w] + (n - count[w]);
    }
    return ans;
};

Also on the editorial tab: C, C#, C++, Go, Java, Kotlin, PHP, Ruby, Rust, Swift, TypeScript.

All 304 dynamic programming problems · the whole catalogue

Learn the technique: Dynamic Programming · Graph Data Structure