Centroid decomposition is a way to split a tree recursively so that every path in the tree is examined at exactly one place, and so that no vertex sits in more than about log2(n) of the pieces. It turns many questions about all paths, such as how many vertex pairs are at most K apart or how far the nearest marked vertex is from v, from quadratic brute force into O(n log n) or O(n log^2 n) work.

The idea: pick a vertex whose removal leaves pieces no bigger than half the tree, handle every path through it, delete it, and recurse. This page proves such a vertex exists, finds it without recursion, traces a small example by hand, and builds two tested Python programs: pair counting and a nearest-marked-vertex structure.

Background that helps: divide and conquer in general, and tree diameter for the single-path version of the same question.

Paths through a tree

Take a tree with n vertices and weighted edges. Some questions ask about one path, such as the distance between u and v, which lowest common ancestor techniques answer in O(log n). Others ask about all n(n-1)/2 paths at once: count pairs at distance at most K, or decide whether some path has length exactly L. Trying every pair is quadratic at best.

The useful observation is that if you fix one vertex c, every path in the tree either passes through c or lies entirely inside one of the pieces you get by deleting c. Paths through c are easy to enumerate from c: each is two half-paths hanging off c in different directions. Solve the through-c case with one traversal, then recurse. The choice of c matters: the end of a long path leaves a piece of size n-1 and the recursion goes n deep.

Centroids and why the recursion is shallow

A centroid is a vertex whose removal leaves every remaining component with at most n/2 vertices. Every tree has one, and the proof is also the algorithm. Start anywhere. If some neighbour's side holds more than n/2 vertices, step into it. The side you left holds fewer than n/2 vertices, so you never step back, and the walk must stop at a vertex with no heavy side: a centroid. A tree has one or two; either works.

In code you compute subtree sizes once from an arbitrary root of the current component. For a vertex v with subtree size size[v], the components after deleting v are each child subtree, of size size[child], plus the part above v, of size total - size[v]. v is a centroid when the largest of those is at most total/2. The test below writes this as 2 * biggest <= total to stay in integers.

Depth follows directly. Each recursive call works on a component at most half the size of its parent component, so after k levels a component has at most n/2^k vertices, and the recursion is at most floor(log2 n) + 1 levels deep. Each level touches every surviving vertex once, which is O(n) per level and O(n log n) overall. Linking each centroid to the centroids of its pieces gives the centroid tree.

Building the centroid tree without recursion

The implementation below builds the centroid tree for an adjacency list of (neighbour, weight) pairs. It is iterative: a recursive search on a 200,000-vertex path overflows Python's recursion limit, though the decomposition is only 18 levels deep. The removed array keeps each search inside its component.

def build_centroid_tree(n, adj):
    """Return (parent, level) of the centroid tree. adj[v] = list of (u, w)."""
    removed = [False] * n
    size = [0] * n
    cparent = [-1] * n
    level = [0] * n

    def find_centroid(root):
        # Iterative DFS: order[] lists the component, par[] its DFS parents.
        order, par = [root], {root: -1}
        for v in order:                      # order grows while we scan it
            for u, _ in adj[v]:
                if not removed[u] and u != par[v]:
                    par[u] = v
                    order.append(u)
        for v in reversed(order):            # children before parents
            size[v] = 1 + sum(size[u] for u, _ in adj[v]
                              if not removed[u] and u != par[v])
        total = len(order)
        for v in order:
            biggest = total - size[v]        # the part above v
            for u, _ in adj[v]:
                if not removed[u] and u != par[v]:
                    biggest = max(biggest, size[u])
            if 2 * biggest <= total:
                return v

    stack = [(0, -1, 0)]                     # (any vertex of component, centroid parent, depth)
    while stack:
        start, p, d = stack.pop()
        c = find_centroid(start)
        cparent[c], level[c] = p, d
        removed[c] = True
        for u, _ in adj[c]:
            if not removed[u]:
                stack.append((u, c, d + 1))
    return cparent, level

After c is removed, each live neighbour lies in a different piece, so one stack entry per neighbour starts one search per piece.

Worked example: nine vertices by hand

One tree, its centroid tree: every path meets the shallowest centroid it touchesOriginal tree (9 vertices)012345678remove 3: pieces of size 3, 1 and 4, all at most 9/2Centroid tree: 3 levels, within the bound floor(log2 9) + 1 = 4314602578Level 0: centroid 3. Level 1: centroids 1, 4, 6 of the three pieces. Level 2: single vertices.The path 2-1-3-5 is handled at centroid 3; the path 7-6-8 never reaches 3 and is handled at centroid 6.
The nine-vertex example and its centroid tree. Yellow is the level-0 centroid, green the level-1 centroids.

Use the tree with edges 0-1, 1-2, 1-3, 3-4, 3-5, 5-6, 6-7 and 6-8, all of weight 1. Rooted at 0, the subtree sizes are 9 at 0, 8 at 1, 6 at 3 and 4 at 5. Vertex 1 fails the test: the piece through 3 has 6 vertices, more than 4.5. Vertex 3 passes: deleting it leaves {0, 1, 2} with 3 vertices, {4} with 1 and {5, 6, 7, 8} with 4. So 3 is the level-0 centroid.

Recursing, {0, 1, 2} has centroid 1, {4} is its own centroid, and {5, 6, 7, 8} has centroid 6, because deleting 6 leaves three single vertices. Everything left is a single vertex at level 2. Three levels, within the bound floor(log2 9) + 1 = 4.

Now count pairs at distance at most K = 2. At centroid 3, the distances to every vertex of the component are 0 for 3 itself, then 1, 2, 2 in the piece {1, 0, 2}, 1 in {4}, and 1, 2, 3, 3 in {5, 6, 7, 8}. Among those nine values, pairs summing to at most 2 are the six involving the zero and a value of at most 2, plus the three pairs of ones, so 9. Subtract pairs whose two ends sit in the same piece, because their real path does not go through 3. None qualify here, so 9 pairs pass through 3: (3,1), (3,0), (3,2), (3,4), (3,5), (3,6), (1,4), (1,5) and (4,5).

The pieces add 3 pairs in {0, 1, 2} and 6 in {5, 6, 7, 8}, for 18 in total, matching brute force: 8 edges plus 10 pairs at distance 2.

Counting pairs within distance K

The trace contains the whole pair-counting algorithm. At each centroid, collect the distance from the centroid to every vertex in its component, count pairs whose distances add up to at most K, then subtract the same count computed separately inside each piece. That subtraction is the step people get wrong. A same-piece pair's real path turns before reaching c, so it was counted with the wrong length; remove it here and the recursion counts it correctly.

def count_pairs_within(n, adj, K):
    """Number of unordered vertex pairs whose path length is <= K."""
    removed = [False] * n

    def distances(src, start_dist):
        out, stack = [], [(src, -1, start_dist)]
        while stack:
            v, p, d = stack.pop()
            out.append(d)
            for u, w in adj[v]:
                if u != p and not removed[u]:
                    stack.append((u, v, d + w))
        return out

    def pairs_le(ds):                        # pairs i<j in ds with ds[i]+ds[j] <= K
        ds.sort()
        i, j, cnt = 0, len(ds) - 1, 0
        while i < j:
            if ds[i] + ds[j] <= K:
                cnt += j - i
                i += 1
            else:
                j -= 1
        return cnt

    total = 0
    cparent, level = build_centroid_tree(n, adj)
    for c in sorted(range(n), key=lambda v: level[v]):   # shallow centroids first
        removed[c] = True                    # branches must not walk back through c
        everything = [0]                     # the centroid itself, distance 0
        for u, w in adj[c]:
            if not removed[u]:
                branch = distances(u, w)
                total -= pairs_le(branch)    # both ends in one branch: path skips c
                everything += branch
        total += pairs_le(everything)
    return total

Processing centroids in level order is enough to reproduce the recursion, because centroids on the same level live in disjoint components. Mark c removed before walking its branches: a first draft of this code did it afterwards, each branch walked back through c, and a randomised brute-force comparison caught it at once.

Cost: the distances at one level add up to O(n), and sorting them costs O(n log n) per level, so the total is O(n log^2 n). With small integer weights you can replace the sort with counting sort and get O(n log n).

The centroid tree as a query structure

The centroid tree is also a data structure. For any two vertices u and v, the path between them passes through their lowest common ancestor in the centroid tree, the first centroid that separated them. So dist(u, v) = dist(u, c) + dist(c, v) for that c, and every centroid ancestor of v gives an upper bound dist(v, c) + dist(c, x) for any x in c's component. Storing, for each vertex, its O(log n) centroid ancestors with exact distances answers a classic query: mark vertices red over time, and ask for the distance from v to the nearest red vertex.

class NearestMarked:
    """Mark vertices red; query distance from v to the closest red vertex."""
    def __init__(self, n, adj):
        self.cparent, level = build_centroid_tree(n, adj)
        self.anc = [[] for _ in range(n)]    # anc[v] = [(centroid, dist(v, centroid)), ...]
        removed = [False] * n
        for c in sorted(range(n), key=lambda v: level[v]):
            removed[c] = True                # mark first, then walk the component
            stack = [(c, -1, 0)]
            while stack:
                v, p, d = stack.pop()
                self.anc[v].append((c, d))
                for u, w in adj[v]:
                    if u != p and not removed[u]:
                        stack.append((u, v, d + w))
        self.best = [float("inf")] * n       # best[c] = closest red vertex in c's component

    def mark(self, v):
        for c, d in self.anc[v]:
            self.best[c] = min(self.best[c], d)

    def query(self, v):
        return min(self.best[c] + d for c, d in self.anc[v])

Both mark and query touch at most floor(log2 n) + 1 ancestors, so each costs O(log n). The answer is exact: v and the nearest red vertex r were separated by some centroid c on both ancestor lists, and there best[c] + dist(v, c) equals dist(v, r). Memory is O(n log n). Unmarking needs a multiset per centroid instead of a minimum.

Both programs were checked against brute force on 400 random trees, and the builder produced 18 levels on a 200,000-vertex path, exactly floor(log2 200000) + 1.

Cost in practice

OperationTimeMemory
Build the centroid treeO(n log n)O(n)
Count pairs with distance at most KO(n log^2 n) with sortingO(n)
Ancestor distance listsO(n log n)O(n log n)
Nearest-marked mark or queryO(log n) eachO(n) for best[]

At most 20 levels for a million vertices, but each level re-traverses most of the tree through pointer-chasing adjacency lists. Store the graph in flat arrays in compiled languages, and time the build early in Python.

Failure modes

  • Searching through removed vertices. Every traversal must skip removed vertices, including the centroid you are currently processing. Forgetting it silently merges pieces.
  • Stale subtree sizes. Sizes must be recomputed inside each component. Reusing sizes from the original rooting picks vertices that are not centroids and the depth bound disappears.
  • Missing the same-piece subtraction, or subtracting with the wrong offset. The branch distances must include the edge from c to the branch root, exactly as they appear in the combined list.
  • Recursion depth. The decomposition is shallow but a recursive size computation is not; long paths crash recursive implementations.
  • Counting ordered pairs or self-pairs by accident. Decide up front whether (u, v) and (v, u) are one answer and whether u = v counts, then test a two-vertex tree.
  • Using it for path updates. Centroid decomposition answers questions about distances from a vertex. Adding a value to every edge on the path from u to v is a different problem.

Trade-offs against other tree techniques

TechniqueBest atWeak at
Centroid decompositionAll-pairs path counts, distance-to-a-set queriesPath updates, dynamic trees
Heavy-light decompositionPath sums and updates between two given verticesAggregates over all paths
Euler tourSubtree queries and updatesArbitrary paths
Small-to-large mergingSubtree statistics, simpler codePaths that go up and back down
LCA with binary liftingSingle distance queriesCounting over many pairs

A useful rule: if the query names two specific vertices, start with LCA or heavy-light decomposition. If it asks about every path, or about the distance from a vertex to the nearest member of a changing set, centroid decomposition is usually the intended tool.

What to do next

  • Type in build_centroid_tree and check that the 200,000-vertex path gives 18 levels.
  • Write an O(n^2) brute force first, then compare it with count_pairs_within on random small trees.
  • Reproduce the nine-vertex trace by hand, then by printing the distance lists at each centroid.
  • Extend NearestMarked to support unmarking with a per-centroid heap or sorted multiset.
  • Solve one exact-length problem: does any path have length exactly L? Use a set of distances per centroid.
  • Port the builder to flat arrays in C++ or Java and measure the build on a million-vertex random tree.
Key takeaway: Pick a vertex whose removal leaves pieces of at most half the tree, handle every path through it, remove it and recurse: that gives at most floor(log2 n) + 1 levels and lets all-paths questions run in O(n log n) to O(n log^2 n). Find centroids with fresh subtree sizes per component, skip removed vertices in every traversal, subtract same-piece pairs, keep the code iterative, and test against brute force on random trees before trusting it.