[NOIP2024] 树上查询

GitHub跳转原题关系图返回列表

把连续编号区间的 LCA 深度转成相邻 LCA 深度数组的区间最小值,再离线二分答案。

OJ: luogu

题目 ID: P11364

难度:省选/NOI-

标签:树形结构LCA二分线段树离线

日期: 2026-06-22 19:30

题意

给定一棵以 1 为根的树,节点编号为 1..n。节点深度定义为从根到该节点路径上的节点数量。

LCA*(l, r) 表示编号在 [l, r] 内所有节点的最近公共祖先。每次询问给出 l, r, k,要在 [l, r] 的所有长度至少为 k 的连续编号子区间 [l', r'] 中,求 dep(LCA*(l', r')) 的最大值。

如果 k = 1,可以只选一个节点,答案就是编号区间 [l, r] 内节点深度的最大值。

思路

先看一个可以直接验证想法的朴素解:

cpp
// brute.cpp:小数据暴力解,用来帮助理解题意并辅助对拍。
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 55;

int n, q;
vector<int> tree[MAXN];
int parent_node[MAXN], depth_node[MAXN];

void dfs_build(int u, int fa) {
    parent_node[u] = fa;
    depth_node[u] = (fa == 0 ? 1 : depth_node[fa] + 1);
    for (int i = 0; i < (int)tree[u].size(); i++) {
        int v = tree[u][i];
        if (v == fa) {
            continue;
        }
        dfs_build(v, u);
    }
}

int lca_pair(int u, int v) {
    while (depth_node[u] > depth_node[v]) {
        u = parent_node[u];
    }
    while (depth_node[v] > depth_node[u]) {
        v = parent_node[v];
    }
    while (u != v) {
        u = parent_node[u];
        v = parent_node[v];
    }
    return u;
}

int lca_interval(int l, int r) {
    int cur = l;
    for (int i = l + 1; i <= r; i++) {
        cur = lca_pair(cur, i);
    }
    return cur;
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    cin >> n;
    for (int i = 1; i < n; i++) {
        int u, v;
        cin >> u >> v;
        tree[u].push_back(v);
        tree[v].push_back(u);
    }

    dfs_build(1, 0);

    cin >> q;
    while (q--) {
        int l, r, k;
        cin >> l >> r >> k;

        int ans = 0;
        for (int left = l; left <= r; left++) {
            for (int right = left + k - 1; right <= r; right++) {
                int x = lca_interval(left, right);
                ans = max(ans, depth_node[x]);
            }
        }
        cout << ans << '\n';
    }

    return 0;
}

暴力做法会枚举每个询问里的所有连续子区间,再逐个求这些节点的 LCA。这个思路非常直观,但一个询问就可能有二次方个子区间,无法处理 5 * 10^5 的数据范围。

关键是把“连续编号节点的 LCA”变成一个数组问题。定义:

text
b[i] = dep(lca(i, i + 1))  (1 <= i < n)

对于长度大于 1 的连续编号区间 [L, R],有:

text
dep(LCA*(L, R)) = min(b[L], b[L + 1], ..., b[R - 1])

这张表展示节点区间和 b 数组区间之间的对应关系:

节点区间 对应的 b 数组位置 LCA 深度
[L, L] dep[L]
[L, L + 1] b[L] b[L]
[L, R] b[L..R-1] min(b[L..R-1])

表里最重要的是第三行:长度大于 1 的节点区间不再直接看所有节点,而是看相邻编号 LCA 深度数组的一段最小值。相邻位置能串起整个连续编号区间,所以只要每一对相邻节点在某个深度处仍有共同祖先,整段节点也在这个深度处有共同祖先。

于是当 k >= 2 时,询问 [l, r, k] 等价于:

text
在 b[l..r-1] 中找一个长度至少为 k-1 的连续子段,
最大化这个子段的最小值。

对答案深度 x 做判定:把所有满足 b[i] >= x 的位置激活。若 [l, r-1] 内存在连续至少 k-1 个激活位置,就说明可以选出一个长度至少为 k 的节点区间,其 LCA 深度至少为 x

这个判定对 x 单调,所以可以二分答案。为了同时处理所有询问,代码使用并行二分:每一轮把询问按当前二分中点分桶,然后按阈值从大到小扫描并激活位置。

线段树维护每个区间的四个量:

  • len:区间长度;
  • pref:最长激活前缀;
  • suf:最长激活后缀;
  • best:区间内最长连续激活段。

合并两个线段树节点时,跨过中点的连续激活段长度是左儿子的 suf + 右儿子的 pref。查询 [l, r-1] 后,如果 best >= k-1,当前二分值可行。

k = 1 的询问单独处理:预处理节点深度的 ST 表,直接查询编号区间 [l, r] 内最大深度。

代码

cpp
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 500005;
const int LOG = 20;

struct Edge {
    int to;
    int next;
};

struct SegNode {
    int len;
    int pref;
    int suf;
    int best;
};

struct Query {
    int left;
    int right;
    int need;
    int id;
};

int n, q;
int head[MAXN], edge_cnt;
Edge edges[MAXN * 2];
int depth_node[MAXN];
int up[LOG][MAXN];
int adjacent_depth[MAXN];
int depth_log[MAXN];
int depth_st[LOG][MAXN];
int answer[MAXN];

vector<int> values;
int value_id[MAXN];
int position_order[MAXN];
vector<Query> queries;

SegNode seg[MAXN * 4];
int low_bound_idx[MAXN], high_bound_idx[MAXN], answer_idx[MAXN];
int bucket_head[MAXN], bucket_next[MAXN];
vector<int> used_bucket;

void add_edge(int u, int v) {
    edge_cnt++;
    edges[edge_cnt].to = v;
    edges[edge_cnt].next = head[u];
    head[u] = edge_cnt;
}

SegNode merge_node(SegNode a, SegNode b) {
    SegNode c;
    c.len = a.len + b.len;
    c.pref = a.pref;
    if (a.pref == a.len) {
        c.pref = a.len + b.pref;
    }
    c.suf = b.suf;
    if (b.suf == b.len) {
        c.suf = b.len + a.suf;
    }
    c.best = max(max(a.best, b.best), a.suf + b.pref);
    return c;
}

void build_seg(int p, int l, int r) {
    seg[p].len = r - l + 1;
    seg[p].pref = seg[p].suf = seg[p].best = 0;
    if (l == r) {
        return;
    }
    int mid = (l + r) >> 1;
    build_seg(p << 1, l, mid);
    build_seg(p << 1 | 1, mid + 1, r);
}

void reset_seg(int p, int l, int r) {
    seg[p].pref = seg[p].suf = seg[p].best = 0;
    if (l == r) {
        return;
    }
    int mid = (l + r) >> 1;
    reset_seg(p << 1, l, mid);
    reset_seg(p << 1 | 1, mid + 1, r);
}

void activate_position(int p, int l, int r, int pos) {
    if (l == r) {
        seg[p].pref = seg[p].suf = seg[p].best = 1;
        return;
    }
    int mid = (l + r) >> 1;
    if (pos <= mid) {
        activate_position(p << 1, l, mid, pos);
    } else {
        activate_position(p << 1 | 1, mid + 1, r, pos);
    }
    seg[p] = merge_node(seg[p << 1], seg[p << 1 | 1]);
}

SegNode query_seg(int p, int l, int r, int ql, int qr) {
    if (ql <= l && r <= qr) {
        return seg[p];
    }
    int mid = (l + r) >> 1;
    if (qr <= mid) {
        return query_seg(p << 1, l, mid, ql, qr);
    }
    if (ql > mid) {
        return query_seg(p << 1 | 1, mid + 1, r, ql, qr);
    }
    SegNode left_node = query_seg(p << 1, l, mid, ql, qr);
    SegNode right_node = query_seg(p << 1 | 1, mid + 1, r, ql, qr);
    return merge_node(left_node, right_node);
}

void build_lca() {
    vector<int> order;
    order.push_back(1);
    up[0][1] = 0;
    depth_node[1] = 1;

    for (int i = 0; i < (int)order.size(); i++) {
        int u = order[i];
        for (int e = head[u]; e != 0; e = edges[e].next) {
            int v = edges[e].to;
            if (v == up[0][u]) {
                continue;
            }
            up[0][v] = u;
            depth_node[v] = depth_node[u] + 1;
            order.push_back(v);
        }
    }

    for (int j = 1; j < LOG; j++) {
        for (int i = 1; i <= n; i++) {
            up[j][i] = up[j - 1][up[j - 1][i]];
        }
    }
}

int lca(int u, int v) {
    if (depth_node[u] < depth_node[v]) {
        swap(u, v);
    }
    int diff = depth_node[u] - depth_node[v];
    for (int j = 0; j < LOG; j++) {
        if (diff & (1 << j)) {
            u = up[j][u];
        }
    }
    if (u == v) {
        return u;
    }
    for (int j = LOG - 1; j >= 0; j--) {
        if (up[j][u] != up[j][v]) {
            u = up[j][u];
            v = up[j][v];
        }
    }
    return up[0][u];
}

void build_depth_rmq() {
    for (int i = 2; i <= n; i++) {
        depth_log[i] = depth_log[i >> 1] + 1;
    }
    for (int i = 1; i <= n; i++) {
        depth_st[0][i] = depth_node[i];
    }
    for (int j = 1; j < LOG; j++) {
        int len = 1 << j;
        for (int i = 1; i + len - 1 <= n; i++) {
            depth_st[j][i] = max(depth_st[j - 1][i], depth_st[j - 1][i + (len >> 1)]);
        }
    }
}

int query_depth_max(int l, int r) {
    int len = r - l + 1;
    int lg = depth_log[len];
    return max(depth_st[lg][l], depth_st[lg][r - (1 << lg) + 1]);
}

bool cmp_position_by_value(int a, int b) {
    if (value_id[a] != value_id[b]) {
        return value_id[a] > value_id[b];
    }
    return a < b;
}

void solve_parallel_binary_search() {
    int total_values = (int)values.size();
    int total_queries = (int)queries.size();
    if (total_queries == 0) {
        return;
    }

    for (int i = 0; i < total_queries; i++) {
        low_bound_idx[i] = 0;
        high_bound_idx[i] = total_values - 1;
        answer_idx[i] = 0;
    }

    int edge_count = n - 1;
    for (int i = 1; i <= edge_count; i++) {
        position_order[i] = i;
    }
    sort(position_order + 1, position_order + edge_count + 1, cmp_position_by_value);

    build_seg(1, 1, edge_count);

    bool changed = true;
    while (changed) {
        changed = false;
        used_bucket.clear();

        for (int i = 0; i < total_queries; i++) {
            if (low_bound_idx[i] <= high_bound_idx[i]) {
                changed = true;
                int mid = (low_bound_idx[i] + high_bound_idx[i]) >> 1;
                if (bucket_head[mid] == -1) {
                    used_bucket.push_back(mid);
                }
                bucket_next[i] = bucket_head[mid];
                bucket_head[mid] = i;
            }
        }

        if (!changed) {
            break;
        }

        reset_seg(1, 1, edge_count);
        int ptr = 1;

        for (int value_index = total_values - 1; value_index >= 0; value_index--) {
            while (ptr <= edge_count && value_id[position_order[ptr]] >= value_index) {
                activate_position(1, 1, edge_count, position_order[ptr]);
                ptr++;
            }

            for (int idx = bucket_head[value_index]; idx != -1; idx = bucket_next[idx]) {
                Query qu = queries[idx];
                SegNode res = query_seg(1, 1, edge_count, qu.left, qu.right);
                if (res.best >= qu.need) {
                    answer_idx[idx] = value_index;
                    low_bound_idx[idx] = value_index + 1;
                } else {
                    high_bound_idx[idx] = value_index - 1;
                }
            }
        }

        for (int i = 0; i < (int)used_bucket.size(); i++) {
            bucket_head[used_bucket[i]] = -1;
        }
    }

    for (int i = 0; i < total_queries; i++) {
        answer[queries[i].id] = values[answer_idx[i]];
    }
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    cin >> n;
    for (int i = 1; i < n; i++) {
        int u, v;
        cin >> u >> v;
        add_edge(u, v);
        add_edge(v, u);
    }

    build_lca();
    build_depth_rmq();

    if (n > 1) {
        values.reserve(n - 1);
        for (int i = 1; i < n; i++) {
            adjacent_depth[i] = depth_node[lca(i, i + 1)];
            values.push_back(adjacent_depth[i]);
        }
        sort(values.begin(), values.end());
        values.erase(unique(values.begin(), values.end()), values.end());

        for (int i = 1; i < n; i++) {
            value_id[i] = lower_bound(values.begin(), values.end(), adjacent_depth[i]) - values.begin();
        }
        for (int i = 0; i < (int)values.size(); i++) {
            bucket_head[i] = -1;
        }
    }

    cin >> q;
    for (int id = 1; id <= q; id++) {
        int l, r, k;
        cin >> l >> r >> k;
        if (k == 1) {
            answer[id] = query_depth_max(l, r);
        } else {
            Query qu;
            qu.left = l;
            qu.right = r - 1;
            qu.need = k - 1;
            qu.id = id;
            queries.push_back(qu);
        }
    }

    if (n > 1) {
        solve_parallel_binary_search();
    }

    for (int i = 1; i <= q; i++) {
        cout << answer[i] << '\n';
    }

    return 0;
}

复杂度

LCA 倍增预处理和深度 ST 表都是 O(nlogn)O(n log n)

并行二分有 O(logn)O(log n) 轮,每轮线段树激活位置和回答询问的总复杂度为 O((n+q)logn)O((n+q) \log n),所以总时间复杂度为:

O((n+q)log2n)O((n+q) \log^2 n)

空间复杂度为 O(nlogn+q)O(n log n + q)

总结

本题最关键的一步是发现连续编号区间 [L, R] 的 LCA 深度等于相邻数组 b[L..R-1] 的最小值。这样原来的树上区间 LCA 查询就转成了“区间中是否存在足够长的连续可行段”。

k = 1 要单独处理,因为单点区间没有相邻数组位置。k >= 2 时,用离线二分答案和线段树维护连续激活段,就可以批量回答所有询问。

一图流解析

这张图把本题的建模、关键转移、实现检查和训练方法压缩到一页,适合读完正文后复盘。

一图流解析