题解:洛谷 P3398 仓鼠找 sugar

【题目来源】

洛谷:P3398 仓鼠找 sugar - 洛谷

【题目描述】

小仓鼠的和他的基(mei)友(zi)sugar 住在地下洞穴中,每个节点的编号为 \(1\sim n\)。地下洞穴是一个树形结构。这一天小仓鼠打算从从他的卧室(\(a\))到餐厅(\(b\)),而他的基友同时要从他的卧室(\(c\))到图书馆(\(d\))。他们都会走最短路径。现在小仓鼠希望知道,有没有可能在某个地方,可以碰到他的基友?

小仓鼠那么弱,还要天天被 zzq 大爷虐,请你快来救救他吧!

【输入】

第一行两个正整数 \(n\)\(q\),表示这棵树节点的个数和询问的个数。

接下来 \(n-1\) 行,每行两个正整数 \(u\)\(v\),表示节点 \(u\) 到节点 \(v\) 之间有一条边。

接下来 \(q\) 行,每行四个正整数 \(a\)\(b\)\(c\)\(d\),表示节点编号,也就是一次询问,其意义如上。

【输出】

对于每个询问,如果有公共点,输出大写字母 Y;否则输出N

【输入样例】

5 5
2 5
4 2
1 3
1 4
5 1 5 1
2 2 1 4
4 1 3 4
3 1 1 5
3 5 1 4

【输出样例】

Y
N
Y
Y
Y

【核心思想】

  1. 问题分析:给定一棵树,\(q\) 次询问,每次询问两条路径 \((a,b)\)\((c,d)\) 是否有公共点。这是一个LCA + 路径相交判定问题,关键在于利用树上路径的距离性质,将路径相交问题转化为距离不等式。

  2. 算法选择

    • LCA 倍增预处理:BFS 计算每个节点的深度 dep\(2^k\) 级祖先 fa,支持 \(O(\log n)\) 的 LCA 查询
    • 距离公式dist(u,v) = dep[u] + dep[v] - 2*dep[lca(u,v)]
    • 路径相交判定定理:树上两条路径 \((a,b)\)\((c,d)\) 相交 \(\Leftrightarrow\) \(dist(a,b) + dist(c,d) \geq \max(dist(a,c)+dist(b,d), dist(a,d)+dist(b,c))\)
  3. 关键步骤

    • 建树与预处理:读入 \(n, q\)\(n-1\) 条边,BFS 计算深度和倍增数组
    • 处理每次询问(读入 \(a, b, c, d\)):
      • 计算 \(ab = dist(a,b)\), \(cd = dist(c,d)\)
      • 计算 \(ac = dist(a,c)\), \(bd = dist(b,d)\)
      • 计算 \(ad = dist(a,d)\), \(bc = dist(b,c)\)
      • \(ab + cd \geq ac + bd\)\(ab + cd \geq ad + bc\),输出 Y;否则输出 N
  4. 时间/空间复杂度

    • 时间复杂度:\(O(n \log n + q \log n)\),预处理 \(O(n \log n)\),每次询问 \(4\) 次 LCA 查询 \(O(\log n)\)
    • 空间复杂度:\(O(n \log n)\),倍增数组
  5. 树上路径相交判定的核心思想

    • 四边形不等式:树上四条路径的距离满足特定关系。若 \((a,b)\)\((c,d)\) 相交,则交点 \(p\) 同时在两条路径上,此时 \(ab + cd = (ap+pb) + (cp+pd)\),而 \(ac+bd = (ap+pc) + (bp+pd)\),两者关系由交点位置决定
    • LCA 降维:将树上路径长度转化为深度运算,避免实际遍历路径
    • 对称性利用:通过比较三组距离和的最大值,覆盖所有可能的相交情况
    • 适用于"树上多路径相交判定"问题,核心在于LCA距离公式和路径相交的代数判定

【算法标签】

普及+ #最近公共祖先

【代码详解】

#include <bits/stdc++.h>
using namespace std;
const int N = 100005, M = N * 2;
int n, Q;  // n: 节点数,Q: 查询次数
int dep[N];  // 节点深度
int fa[N][20];  // 倍增数组,fa[i][j]表示节点i向上跳2^j步到达的祖先节点
queue<int> q;  // BFS队列
int h[N], e[M], ne[M], idx;  // 邻接表存储树

// 添加无向边
void add(int a, int b)
{
    e[idx] = b, ne[idx] = h[a], h[a] = idx++;
}

// 预处理节点深度和倍增数组
void bfs(int root)
{
    memset(dep, 0x3f, sizeof(dep));  // 初始化深度为无穷大
    dep[0] = 0, dep[root] = 1;  // 虚拟0节点深度为0,根节点深度为1
    q.push(root);

    while (!q.empty())
    {
        int t = q.front();
        q.pop();

        for (int i = h[t]; i != -1; i = ne[i])
        {
            int j = e[i];  // 子节点j
            if (dep[j] > dep[t] + 1)  // 如果j的深度大于t的深度+1(即j未访问)
            {
                dep[j] = dep[t] + 1;  // 计算j的深度
                q.push(j);  // j入队列
                fa[j][0] = t;  // j向上走2^0=1步即为父节点t

                for (int k = 1; k <= 19; k++)  // 递推计算fa[j][k]
                {
                    // j向上走2^k步等于j向上走2^(k-1)步后再走2^(k-1)步
                    fa[j][k] = fa[fa[j][k - 1]][k - 1];
                }
            }
        }
    }
}

// LCA倍增算法,计算节点x和节点y的最近公共祖先
int lca(int x, int y)
{
    if (dep[x] < dep[y])  // 保证x为深度较大的节点
        swap(x, y);

    // 步骤1:把节点x向上跳,直到与节点y深度相同
    for (int k = 19; k >= 0; k--)  // 从高位开始尝试跳
        if (dep[fa[x][k]] >= dep[y])
            x = fa[x][k];

    if (x == y)  // 如果跳完后x等于y,说明y是x的祖先
        return x;

    // 步骤2:两个节点同时向上跳,跳到公共祖先的下一层
    for (int k = 19; k >= 0; k--)
    {
        if (fa[x][k] != fa[y][k])  // 如果跳2^k步后祖先不同
        {
            x = fa[x][k];
            y = fa[y][k];
        }
    }
    return fa[x][0];  // 返回x的父节点,即x和y的最近公共祖先
}

// 计算节点a和节点b之间的距离
int dist(int a, int b)
{
    int p = lca(a, b);  // 求最近公共祖先
    return dep[a] + dep[b] - 2 * dep[p];  // 距离 = a深度 + b深度 - 2*lca深度
}

int main()
{
    memset(h, -1, sizeof(h));
    cin >> n >> Q;

    for (int i = 1; i < n; i++)  // 输入n-1条边
    {
        int u, v;
        cin >> u >> v;
        add(u, v), add(v, u);  // 添加无向边
    }

    bfs(1);  // 从节点1开始预处理深度和倍增数组

    while (Q--)
    {
        int a, b, c, d;
        cin >> a >> b >> c >> d;

        int ab = dist(a, b);  // 路径a-b的距离
        int cd = dist(c, d);  // 路径c-d的距离

        int ac = dist(a, c);  // 节点a到节点c的距离
        int bd = dist(b, d);  // 节点b到节点d的距离

        int ad = dist(a, d);  // 节点a到节点d的距离
        int bc = dist(b, c);  // 节点b到节点c的距离

        // 判断路径a-b和c-d是否有交点
        // 条件:路径a-b和c-d有交点当且仅当ab+cd >= max(ac+bd, ad+bc)
        if (ab + cd >= ac + bd && ab + cd >= ad + bc)
            cout << "Y" << endl;
        else
            cout << "N" << endl;
    }
    return 0;
}

【运行结果】

5 5
2 5
4 2
1 3
1 4
5 1 5 1
Y
2 2 1 4
N
4 1 3 4
Y
3 1 1 5
Y
3 5 1 4
Y
posted @ 2026-08-10 12:29  团爸讲算法  阅读(16)  评论(0)    收藏  举报