P17141 [NOI 2026] 传送

· · 题解

来个根号做法,感觉很显然啊。

思路:

显然,策略一定是随机一会儿然后再沿着最短路径走:

那么显然,对于 u \to v 的询问,一定是选择一个 d,那么 S = \{ i | dis(i, v) \le d\} 内的点是往 v 直接走,外面的点就一直传送,算一下期望:

\frac{n + \sum_{i \in S} dis(i, v)}{|S|}

于是有了一个 O(nq) 做法,将询问离线到每个 v 上,然后预处理 c_i 表示 dis(j, v) = i 的点的数量,就可以算答案了。

考虑优化,把上面式子用 c_i 表示一下:

ans(d) = \frac{n + \sum_{i = 0}^d c_i i}{\sum_{i = 0}^d c_i}

推一下 ans(d) \le ans(d + 1) 的条件:

\frac{n + \sum_{i = 0}^d c_i i}{\sum_{i = 0}^d c_i} \le \frac{n + \sum_{i = 0}^{d + 1} c_i i}{\sum_{i = 0}^{d + 1} c_i} n \sum_{i = 0}^{d + 1} c_i + \sum_{i = 0}^d \sum_{j = 0}^{d + 1} c_i c_j i \le n \sum_{i = 0}^d c_i + \sum_{i = 0}^{d + 1} \sum_{j = 0}^d c_i c_j i n c_{d + 1} + \sum_{i = 0}^d c_{d + 1} c_i i \le \sum_{j = 0}^d c_{d + 1} (d + 1) c_j n + \sum_{i = 0}^d c_i i \le \sum_{j = 0}^d (d + 1) c_j n + \sum_{i = 0}^d c_i (i - d - 1) \le 0

考虑 f(d) = n + \sum_{i = 1}^d c_i(i - d - 1),显然 f(d + 1) = f(d) - \sum_{i = 1}^{d + 1} c_i,即 f(x) 是一个单调递减的函数,所以 ans(x) 是一个单谷的函数,先减后增。

那么可以想到二分,然后只需要查询 |S| 以及 \sum_{i \in S} dis(i, v) 即可,直接点分树是 O(n \log^2 n) 的。

但是实际上不需要那么麻烦,我们都推出来 f(d + 1) = f(d) - \sum_{i = 1}^{d + 1} c_i 这个式子了,在 f(d) 第一次 \le 0 的时候显然就是最优决策,而 f 变化每次至少减去 d + 1,初始是 n,所以最优的 dO(\sqrt n) 级别的,精细是 \le \sqrt{2n} 的。

然后怎么做?你需要数一个点 \le O(\sqrt n) 邻域的信息,显然可以直接 dp,令 dp_{u, i} 表示 dis(u, v) = i 的点的数量,那么有转移:

dp_{u, i} = \sum_{(u, v) \in E} dp_{v, i - 1} - (deg_u - 1) dp_{u, i - 2}

对于 i \le 2 的情况要特殊处理;于是我们枚举 i,因为只根 i, i - 1, i - 2 有关,所以直接滚动即可;过程中记录 s_u 表示当前 \le idp_{u, i} 之和,以及 all_u 表示当前 \le i 的所有 s 的累加和,在 all_u \ge n 时到达了最优决策点。

注意询问的时候如果 u 本身就在 v 的最优 d 邻域内,那么答案直接就是 dis(u, v) 而不是上面那个式子;时间复杂度为 O(n \sqrt n),轻微卡常。

完整代码:

 #include<bits/stdc++.h>
#define lowbit(x) x & (-x)
#define ls(k) k << 1
#define rs(k) k << 1 | 1
#define fi first
#define se second
#define ctz(x) __builtin_ctz(x)
#define popcnt(x) __builtin_popcount(x)
#define open(s1, s2) freopen(s1, "r", stdin), freopen(s2, "w", stdout);
using namespace std;
typedef __int128 __;
typedef long double lb;
typedef double db;
typedef unsigned int uint;
typedef unsigned long long ull;
typedef long long ll;
const int N = 5e5 + 10, M = 1e6 + 10;
inline ll read(){
    ll x = 0, f = 1;
    char c = getchar();
    while(c < '0' || c > '9'){
        if(c == '-')
          f = -1;
        c = getchar();
    }
    while(c >= '0' && c <= '9'){
        x = (x << 1) + (x << 3) + (c ^ 48);
        c = getchar();
    }
    return x * f;
}
inline void write(ll x){
    if(x < 0){
        putchar('-');
        x = -x;
    }
    if(x > 9)
      write(x / 10);
    putchar(x % 10 + '0');
}
int c, n, m, cnt;
int du[N], head[N];
struct edge{
    int to, nxt;
}e[M];
inline void add(int u, int v){
    e[++cnt] = {v, head[u]};
    head[u] = cnt;
    e[++cnt] = {u, head[v]};
    head[v] = cnt;
    ++du[u], ++du[v];
}
namespace Tree{
    int siz[N], top[N], son[N], dep[N], fa[N];
    inline void init(){
        for(int i = 0; i < n; ++i)
          son[i] = n;
    }
    inline void dfs1(int u, int f){
        siz[u] = 1;
        for(int i = head[u]; i; i = e[i].nxt){
            int v = e[i].to;
            if(v == f)
              continue;
            fa[v] = u;
            dep[v] = dep[u] + 1;
            dfs1(v, u);
            siz[u] += siz[v];
            if(siz[v] > siz[son[u]])
              son[u] = v;
        }
    }
    inline void dfs2(int u, int k){
        top[u] = k;
        if(!son[u])
          return ;
        dfs2(son[u], k);
        for(int i = head[u]; i; i = e[i].nxt){
            int v = e[i].to;
            if(v == fa[u] || v == son[u])
              continue;
            dfs2(v, v);
        }
    }
    inline int LCA(int u, int v){
        while(top[u] != top[v]){
            if(dep[top[u]] < dep[top[v]])
              swap(u, v);
            u = fa[top[u]];
        }
        return dep[u] < dep[v] ? u : v;
    }
    inline int dis(int u, int v){
        return dep[u] + dep[v] - 2 * dep[LCA(u, v)];
    }
}
bool vis[N];
int mxd[N];
int dp[3][N], s[N], all[N];
vector<pair<int, int>> Q[N];
std::vector<std::pair<long long, int>> teleport(int _c, int _n, int _m, std::vector<int> u, std::vector<int> v, std::vector<int> x, std::vector<int> y){
    c = _c, n = _n, m = _m;
    for(int i = 0; i < n - 1; ++i)
      add(u[i], v[i]);
    Tree::init();
    Tree::dfs1(0, 0);
    Tree::dfs2(0, 0);
    for(int i = 0; i < m; ++i)
      Q[y[i]].push_back({x[i], i});
    int lim = min((int)sqrt(2 * n) + 1, n);
    // cerr << lim << '\n';
    for(int i = 0; i <= lim; ++i){
        int now = i % 3;
        if(!i){
            for(int u = 0; u < n; ++u){
                dp[now][u] = 1;
                s[u] += dp[now][u];
                all[u] += s[u];
            }
            continue;
        }
        if(i == 1){
            for(int u = 0; u < n; ++u){
                dp[now][u] = du[u];
                s[u] += dp[now][u];
                all[u] += s[u];
                if(all[u] >= n && !mxd[u])
                  mxd[u] = i, vis[u] = 1;
            }
            continue;
        }
        int pre = (i - 1) % 3, ppre = (i - 2) % 3;
        for(int u = 0; u < n; ++u)
          dp[now][u] = -(du[u] - (i != 2)) * dp[ppre][u];
        for(int u = 1; u < n; ++u)
          dp[now][u] += dp[pre][Tree::fa[u]], dp[now][Tree::fa[u]] += dp[pre][u];
        for(int u = 0; u < n; ++u){
            if(vis[u])
              continue;
            s[u] += dp[now][u], all[u] += s[u];
            if(all[u] >= n && !mxd[u])
              mxd[u] = i, vis[u] = 1;
        }
    }
    vector<pair<long long, int>> ans(m);
    for(int u = 0; u < n; ++u){
        if(Q[u].empty())
          continue;
        int d = mxd[u];
        int a = n + (d + 1) * s[u] - all[u], b = s[u];
        // cerr << u << ' ' << mxd[u] << ' ' << all[u] << ' ' << a << ' ' << b << '\n';
        int g = __gcd(a, b);
        a /= g, b /= g;
        for(auto t : Q[u]){
            int v = t.fi, id = t.se;
            int dis = Tree::dis(u, v);
            if(a < 1ll * b * dis)
              ans[id] = {a, b};
            else
              ans[id] = {dis, 1};
        }
    }
    return ans;
}
int main(){
    c = read(), n = read(), m = read();
    vector<int> u, v;
    for(int i = 0; i < n - 1; ++i){
        u.push_back(read());
        v.push_back(read());
    }
    vector<int> x, y;
    for(int i = 0; i < m; ++i){
        x.push_back(read());
        y.push_back(read());
    }
    vector<pair<long long, int>> ans = teleport(c, n, m, u, v, x, y);
    puts("7c9f2e4a61b8d305a4e6f93c0d12b7aa");
    for(auto t : ans){
        write(t.fi);
        putchar(' ');
        write(t.se);
        putchar('\n');
    }
    puts("7c9f2e4a61b8d305a4e6f93c0d12b7aa");
    return 0;
}