CP Notebook

← all categories

Heavy Light Decomposition 312fcf6e

Heavy-light decomposition: flattens a tree into O(\log n) contiguous ranges per root-to-node path. Pair with any segment tree (segment-tree-lazy.h, segment-tree-iterative.h, ...) keyed by pos[]; forPath(u, v, f) calls f(l, r) for each maximal contiguous [l, r] range (in pos[] order, both inclusive) along the path u-v; forSubtree(u) gives the single [l, r] covering u's subtree. Recursive dfsSize/ dfsDecompose; may need a larger stack for very deep/skewed trees.

Time: O(n) build; a path decomposes into O(\log n) ranges. tested

content/graphs/heavy-light-decomposition.h

struct HLD {
    int n, timer = 0;
    vector<vector<int>> adj;
    vector<int> par, dep, heavy, head, pos, sz;

    HLD(int n, vector<vector<int>> &adj, int root = 0)
        : n(n), adj(adj), par(n, -1), dep(n), heavy(n, -1), head(n), pos(n), sz(n) {
        dfsSize(root, -1);
        head[root] = root;
        dfsDecompose(root, -1);
    }

    int dfsSize(int u, int p) {
        par[u] = p; sz[u] = 1;
        int maxSz = 0;
        for (int v : adj[u]) {
            if (v == p) continue;
            dep[v] = dep[u] + 1;
            sz[u] += dfsSize(v, u);
            if (sz[v] > maxSz) { maxSz = sz[v]; heavy[u] = v; }
        }
        return sz[u];
    }

    void dfsDecompose(int u, int p) {
        pos[u] = timer++;
        if (heavy[u] != -1) {
            head[heavy[u]] = head[u];
            dfsDecompose(heavy[u], u);
        }
        for (int v : adj[u]) {
            if (v == p || v == heavy[u]) continue;
            head[v] = v;
            dfsDecompose(v, u);
        }
    }

    void forPath(int u, int v, function<void(int, int)> f) {
        while (head[u] != head[v]) {
            if (dep[head[u]] < dep[head[v]]) swap(u, v);
            f(pos[head[u]], pos[u]);
            u = par[head[u]];
        }
        if (dep[u] > dep[v]) swap(u, v);
        f(pos[u], pos[v]);
    }

    pair<int, int> forSubtree(int u) { return {pos[u], pos[u] + sz[u] - 1}; }

    int lca(int u, int v) {
        while (head[u] != head[v]) {
            if (dep[head[u]] < dep[head[v]]) swap(u, v);
            u = par[head[u]];
        }
        return dep[u] < dep[v] ? u : v;
    }
};