Skip to content
最近公共祖先LCA

最近公共祖先LCA

树链剖分(HLD 求 LCA)

c++
struct HLD {
    int n;
    std::vector<std::vector<int>> adj;
    std::vector<int> siz, dep, top, son, parent;

    explicit HLD(int n)
        : n(n), adj(n + 1), siz(n + 1), dep(n + 1),
          top(n + 1), son(n + 1), parent(n + 1) {}

    void add(int u, int v) {
        adj[u].push_back(v);
        adj[v].push_back(u);
    }

    void dfs1(int u) {
        siz[u] = 1;
        dep[u] = dep[parent[u]] + 1;
        for (int v : adj[u]) {
            if (v == parent[u]) continue;
            parent[v] = u;
            dfs1(v);
            siz[u] += siz[v];
            if (siz[v] > siz[son[u]]) son[u] = v;
        }
    }

    void dfs2(int u, int up) {
        top[u] = up;
        if (son[u]) dfs2(son[u], up);
        for (int v : adj[u]) {
            if (v == parent[u] || v == son[u]) continue;
            dfs2(v, v);
        }
    }

    void work(int root = 1) {
        dfs1(root);
        dfs2(root, root);
    }

    int lca(int u, int v) const {
        while (top[u] != top[v]) {
            if (dep[top[u]] < dep[top[v]]) std::swap(u, v);
            u = parent[top[u]];
        }
        return dep[u] < dep[v] ? u : v;
    }

    int calc(int u, int v) const {
        return dep[u] + dep[v] - 2 * dep[lca(u, v)];
    }
};

树上倍增

c++
struct TreeLCA {
    int n;
    static constexpr int LOG = 20; // 支持 N <= 1,000,000
    std::vector<std::vector<int>> adj;
    std::vector<int> dep;
    std::vector<std::array<int, LOG>> fa;

    explicit TreeLCA(int n) : n(n), adj(n + 1), dep(n + 1), fa(n + 1) {}

    void add(int u, int v) {
        adj[u].push_back(v);
        adj[v].push_back(u);
    }

    void dfs(int u, int p) {
        fa[u][0] = p;
        dep[u] = dep[p] + 1;
        for (int i = 1; i < LOG; ++i) {
            fa[u][i] = fa[fa[u][i - 1]][i - 1];
        }
        for (int v : adj[u]) {
            if (v != p) dfs(v, u);
        }
    }

    void work(int root = 1) {
        dfs(root, 0);
    }

    int lca(int u, int v) const {
        if (dep[u] < dep[v]) std::swap(u, v);
        for (int i = LOG - 1; i >= 0; --i) {
            if (dep[u] - (1 << i) >= dep[v]) {
                u = fa[u][i];
            }
        }
        if (u == v) return u;
        for (int i = LOG - 1; i >= 0; --i) {
            if (fa[u][i] != fa[v][i]) {
                u = fa[u][i];
                v = fa[v][i];
            }
        }
        return fa[u][0];
    }

    int calc(int u, int v) const {
        return dep[u] + dep[v] - 2 * dep[lca(u, v)];
    }
};