【MX-J2-T3】Piggy and Trees

f(u,v,i) 为 i 到 u-v 路径的距离,闭式化简化得答案 = D*(n-2)/2,D 用边贡献 size*(n-size) 累加。

OJ: luogu

题目 ID: P10842

难度:普及+/提高-

标签:树形结构计数组合计数数学

日期: 2026-08-14 15:01

形式化题目

给定一棵 nn 个结点的树。对点对 (u,v)(u, v) 与点 ii,定义 f(u,v,i)f(u,v,i):在满足 dis(u,x)+dis(v,x)=dis(u,v)dis(u,x)+dis(v,x)=dis(u,v) 的点 xx 中,dis(x,i)dis(x,i) 的最小值。求

u<vif(u,v,i)mod(109+7)\sum_{u < v} \sum_{i} f(u,v,i) \bmod (10^9+7)

思路

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

cpp
/**
 * Author by Rainboy blog: https://rainboylv.com github: https://github.com/rainboylvx
 * rbook: -> https://rbook.roj.ac.cn  https://rbook2.roj.ac.cn
 * rainboy的学习导航网站: https://idx.roj.ac.cn
 * create_at: 2026-08-14 15:01
 * update_at: 2026-08-14 16:10
 */
// brute.cpp:小数据暴力解,按题意直接计算:
// 对每对 (u,v) 枚举所有点 x 找出满足 dis(u,x)+dis(v,x)=dis(u,v) 的集合,
// 再对每个 i 求集合中 dis(x,i) 的最小值并累加。用来帮助理解题意并辅助对拍。
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 55;

int n;
vector<int> g[MAXN];
int dist[MAXN][MAXN]; // dist[u][v]:u 到 v 的距离

// BFS 求从 start 出发到所有点的距离。
void bfs(int start) {
    queue<int> q;
    q.push(start);
    dist[start][start] = 0;
    while (!q.empty()) {
        int u = q.front();
        q.pop();
        for (int v : g[u]) {
            if (dist[start][v] == -1) {
                dist[start][v] = dist[start][u] + 1;
                q.push(v);
            }
        }
    }
}

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;
        g[u].push_back(v);
        g[v].push_back(u);
    }

    memset(dist, -1, sizeof(dist));
    for (int i = 1; i <= n; i++) bfs(i);

    long long ans = 0;
    for (int u = 1; u <= n; u++) {
        for (int v = u + 1; v <= n; v++) {
            for (int i = 1; i <= n; i++) {
                int best = INT_MAX;
                for (int x = 1; x <= n; x++) {
                    if (dist[u][x] + dist[v][x] == dist[u][v]) {
                        best = min(best, dist[x][i]);
                    }
                }
                ans += best;
            }
        }
    }
    cout << ans % (long long)(1e9 + 7) << '\n';

    return 0;
}

brute.cpp 按定义直接计算:对每对 (u,v)(u,v) 枚举所有点 xx 检查等式,得到候选集合,再对每个 ii 求集合中到 xx 的最小距离。三层枚举加距离查询是 O(n4)O(n^4)n=2×105n = 2 \times 10^5 完全不可行。

两步化简把问题变成一遍 DFS:

第一步,看清 f 是什么。树上三角不等式取等 dis(u,x)+dis(v,x)=dis(u,v)dis(u,x)+dis(v,x)=dis(u,v) 当且仅当 xxuuvv 的路径上。所以 f(u,v,i)f(u,v,i) 就是点 ii 到路径 uu-vv 的最小距离。设 wwii 到路径的垂足,则

f(u,v,i)=dis(i,u)+dis(i,v)dis(u,v)2f(u,v,i) = \frac{dis(i,u)+dis(i,v)-dis(u,v)}{2}

第二步,交换求和次序。记 D=u<vdis(u,v)D = \sum_{u<v} dis(u,v) 为所有点对距离和。把闭式代入三层和式:

u<vif(u,v,i)=12[(n1)i,jdis(i,j)nD]=D(n2)2\sum_{u<v}\sum_i f(u,v,i) = \frac{1}{2}\left[(n-1)\sum_{i,j}dis(i,j) - nD\right] = \frac{D(n-2)}{2}

最后,DD 可以用边贡献线性求出:一条边把树分成大小为 aabb 的两部分,恰好有 a×ba \times b 个点对穿过这条边,所以 D=a×bD = \sum_{\text{边}} a \times b。一遍 DFS 求子树大小即可。

下面这张表展示样例 1(星形,中心 1 连 2,3,4)每条边的贡献:

一侧点数 另一侧点数 贡献
(1,2) 1 3 3
(1,3) 1 3 3
(1,4) 1 3 3

观察要点:每条边都连接中心与一个叶子,贡献都是 1×3=31 \times 3 = 3,总和 D=9D = 9;代入公式 9×(42)/2=99 \times (4-2)/2 = 9,与样例答案一致。星形里任意两点的路径都经过中心,f 的求和直观上与距离和成比例。

代码

cpp
/**
 * Author by Rainboy blog: https://rainboylv.com github: https://github.com/rainboylvx
 * rbook: -> https://rbook.roj.ac.cn  https://rbook2.roj.ac.cn
 * rainboy的学习导航网站: https://idx.roj.ac.cn
 * create_at: 2026-08-14 15:01
 * update_at: 2026-08-14 16:10
 */
// P10842 【MX-J2-T3】Piggy and Trees
// 结论:答案 = D * (n-2) / 2,其中 D 是所有点对距离之和。
// D = sum(每条边 size * (n - size)),一遍 DFS 求子树大小即可。
// 推导:dis(u,x)+dis(v,x)=dis(u,v) 当且仅当 x 在 u-v 路径上,
// f(u,v,i) = i 到 u-v 路径的距离 = (dis(i,u)+dis(i,v)-dis(u,v))/2。
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 200005;
const long long MOD = 1e9 + 7;

int n;
vector<int> g[MAXN]; // 树的邻接表
int sz[MAXN];        // sz[u]:以 u 为根的子树大小
long long D;         // 所有点对距离之和

// 求子树大小(dfs 遍历,仿 rbook 模板 dfs-traversal)。
void dfs(int u, int fa) {
    sz[u] = 1;
    for (int v : g[u]) {
        if (v == fa) continue;
        dfs(v, u);
        sz[u] += sz[v];
        // 边 (u, v) 把树分成大小为 sz[v] 和 n - sz[v] 的两部分
        D = (D + (long long)sz[v] * (n - sz[v])) % MOD;
    }
}

// 快速幂:求 a^b mod MOD。
long long quick_pow(long long a, long long b) {
    long long res = 1;
    while (b > 0) {
        if (b & 1) res = res * a % MOD;
        a = a * a % MOD;
        b >>= 1;
    }
    return res;
}

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;
        g[u].push_back(v);
        g[v].push_back(u);
    }

    dfs(1, 0);

    // 答案 = D * (n-2) / 2 = D * (n-2) * inv(2) mod MOD
    long long inv2 = quick_pow(2, MOD - 2);
    long long ans = D * ((n - 2 + MOD) % MOD) % MOD * inv2 % MOD;
    cout << ans << '\n';

    return 0;
}

复杂度

  • 时间:一遍 DFS 求子树大小并累计边贡献,O(n)O(n)
  • 空间:邻接表 O(n)O(n)

总结

这道题的价值在把三层求和整体化简:先识别出 ff 的几何意义(点到路径的距离)并写成距离的闭式,再交换求和次序,把"每对点 × 每个点"的 O(n3)O(n^3) 结构坍缩成只依赖 DD 的公式;而 DD 又通过边贡献 size×(nsize)size \times (n-size) 一遍 DFS 求出。这类"点对路径相关求和"的题目,优先尝试:距离闭式 → 求和交换 → 边贡献,往往能直接降到 O(n)O(n)

图示解析

这张 ASCII 图展示整道题的解题路线:

text
题意:sum_{u<v} sum_i f(u,v,i),f 是"满足距离等式的 x 中到 i 最近的距离"
        |
        | 暴力:O(n^4) 按定义枚举 x(brute.cpp)
        v
关键观察 1:dis(u,x)+dis(v,x)=dis(u,v) ⟺ x 在 u-v 路径上
   f(u,v,i) = 点 i 到路径 u-v 的最小距离
        |
        | 设垂足 w,代入三条距离关系
        v
关键观察 2:f = (dis(i,u)+dis(i,v)-dis(u,v))/2
        |
        | 交换求和次序
        v
化简:答案 = D*(n-2)/2,D = sum_{u<v} dis(u,v)
        |
        | 边贡献:每条边贡献 size*(n-size)
        v
实现:一遍 DFS 求子树大小,O(n)
   答案 = D * (n-2) * inv(2) mod MOD

图中主线是"识别几何意义 → 距离闭式 → 求和化简 → 边贡献"。真正要掌握的是:三层求和不要硬算,先找每一项的闭式,再整体交换求和次序。