Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Heavy-light decomposition: sum weights along a tree path

Last updated: 4 Oct 20269 min read
tutorial
IntermediateBy AITrove Editorial

Heavy-light decomposition assigns each tree node to a chain formed by repeatedly choosing a largest child subtree. Every move between chains crosses a light edge, and a root-to-leaf path crosses O(log N) such edges. Nodes receive contiguous positions within their chains. A Fenwick tree over those positions stores node weights, so a path query sums chain ranges until both endpoints share one chain. This example includes both endpoints and counts the common ancestor exactly once. The constructor requires one connected acyclic tree with distinct depot IDs; weights may change, but edges and node IDs are fixed after construction.

Operational case

The depot tree connects D-19 to D-26 and D-47; D-26 connects to D-52 and D-61; D-47 connects to D-83. With weights 7, 11, 13, 17, 19, and 23 respectively, the path from D-52 to D-83 sums 17+11+7+13+23=71. Replacing D-19's weight with 29 changes that path sum to 93. The path from D-61 to D-52 stays 47 because it does not pass through D-19. A path between any two vertices is unique only because the input is a tree. For a general road network, choosing one spanning tree would discard alternative routes and would answer a different question.

Working Python program

python
"""Heavy-light decomposition with a Fenwick tree for node-weight path sums."""


class Fenwick:
    def __init__(self, values: list[int]):
        self.tree = [0] * (len(values) + 1)
        for position, value in enumerate(values):
            self.add(position, value)

    def add(self, position: int, difference: int) -> None:
        position += 1
        while position < len(self.tree):
            self.tree[position] += difference
            position += position & -position

    def prefix(self, stop: int) -> int:
        result = 0
        while stop:
            result += self.tree[stop]
            stop -= stop & -stop
        return result

    def between(self, start: int, stop: int) -> int:
        return self.prefix(stop) - self.prefix(start)


class DepotPathSums:
    def __init__(self, depot_ids: list[str], links: list[tuple[str, str]], weights: dict[str, int]):
        if len(depot_ids) != len(set(depot_ids)) or not depot_ids or set(weights) != set(depot_ids):
            raise ValueError("depots must be distinct, nonempty, and weighted")
        if len(links) != len(depot_ids) - 1:
            raise ValueError("a tree needs one fewer link than vertices")
        self.index = {depot_id: position for position, depot_id in enumerate(depot_ids)}
        count = len(depot_ids)
        neighbors = [[] for _ in depot_ids]
        for first, second in links:
            left, right = self.index[first], self.index[second]
            if left == right:
                raise ValueError("self-link")
            neighbors[left].append(right)
            neighbors[right].append(left)
        self.parent = [-1] * count
        self.depth = [0] * count
        traversal = [0]
        seen = {0}
        for current in traversal:
            for neighbor in neighbors[current]:
                if neighbor == self.parent[current]:
                    continue
                if neighbor in seen:
                    raise ValueError("links contain a cycle")
                seen.add(neighbor)
                self.parent[neighbor] = current
                self.depth[neighbor] = self.depth[current] + 1
                traversal.append(neighbor)
        if len(seen) != count:
            raise ValueError("links are disconnected")
        subtree_size = [1] * count
        heavy = [-1] * count
        for current in reversed(traversal):
            for neighbor in neighbors[current]:
                if self.parent[neighbor] == current:
                    subtree_size[current] += subtree_size[neighbor]
                    if heavy[current] < 0 or subtree_size[neighbor] > subtree_size[heavy[current]]:
                        heavy[current] = neighbor
        self.head = [0] * count
        self.position = [0] * count
        next_position = 0
        pending = [(0, 0)]
        while pending:
            current, chain_head = pending.pop()
            while current >= 0:
                self.head[current] = chain_head
                self.position[current] = next_position
                next_position += 1
                for neighbor in neighbors[current]:
                    if self.parent[neighbor] == current and neighbor != heavy[current]:
                        pending.append((neighbor, neighbor))
                current = heavy[current]
        self.weights = [weights[depot_id] for depot_id in depot_ids]
        flattened = [0] * count
        for vertex in range(count):
            flattened[self.position[vertex]] = self.weights[vertex]
        self.sums = Fenwick(flattened)

    def update(self, depot_id: str, weight: int) -> None:
        vertex = self.index[depot_id]
        self.sums.add(self.position[vertex], weight - self.weights[vertex])
        self.weights[vertex] = weight

    def path_sum(self, first: str, second: str) -> int:
        left, right = self.index[first], self.index[second]
        total = 0
        while self.head[left] != self.head[right]:
            if self.depth[self.head[left]] < self.depth[self.head[right]]:
                left, right = right, left
            chain_head = self.head[left]
            total += self.sums.between(self.position[chain_head], self.position[left] + 1)
            left = self.parent[chain_head]
        if self.depth[left] > self.depth[right]:
            left, right = right, left
        return total + self.sums.between(self.position[left], self.position[right] + 1)


depots = ["D-19", "D-26", "D-47", "D-52", "D-61", "D-83"]
links = [("D-19", "D-26"), ("D-19", "D-47"), ("D-26", "D-52"), ("D-26", "D-61"), ("D-47", "D-83")]
weights = {"D-19": 7, "D-26": 11, "D-47": 13, "D-52": 17, "D-61": 19, "D-83": 23}
path_index = DepotPathSums(depots, links, weights)
print(path_index.path_sum("D-52", "D-83"))
path_index.update("D-19", 29)
print(path_index.path_sum("D-52", "D-83"))
print(path_index.path_sum("D-61", "D-52"))

Output

Output
71
93
47

Time, space, and tradeoff

For N nodes, adjacency, parent, depth, subtree size, chain head, position, and Fenwick arrays use O(N) space. This straightforward constructor inserts each flattened weight into Fenwick separately, so preprocessing costs O(N log N) time. A node-weight update costs O(log N). A path query crosses O(log N) chains and performs an O(log N) Fenwick range sum for each, giving O(log² N) time. An ordinary parent walk is simpler and costs O(path length) per query. This implementation uses iterative topology passes to avoid recursive DFS depth limits, but it is not a fully online tree structure: linking or cutting an edge requires rebuilding.

Common Mistakes

  • Do not count a chain boundary node twice when climbing toward the common ancestor.
  • Do not mistake node weights for edge weights without changing the flattened contract.
  • Do not accept a cyclic or disconnected input as one tree.
  • Do not claim fast link and cut operations from a fixed-topology decomposition.

Connected lessons

Test this invariant in the incident flag and path audit, then answer the contract quiz.

Binary lifting: ancestors and common managers adds a related operation contract.

Link-cut forests: change tree edges and sum a path adds a related operation contract.

Centroid decomposition: nearest marked depot on a fixed tree adds a distinct structure contract to compare.

data structures
range-query-structures
Storage details