「CF 1458E」Range Diameter Sum

树论综合好题啊,\(O(N \log^2 N)\) 做法还是太难打了,听说有好打的 \(O(N \sqrt{N} \log N)\) 做法,瞅了瞅,确实没往不同的增量只有 \(O(\sqrt{N})\) 个去想,写写双 \(\log\) 做法吧。

拿到这道题第一感觉就是非常棘手,直径是树上的东西,但统计是按下标统计区间的直径,这两个东西拼不在一起,所以直径不太能用传统的 dp 或两遍 dfs 来求,所以联想到直径相关的树上圆理论。

对于一个点集的直径,了解过树上圆理论的就知道等价于找到能覆盖该点集的直径最小的圆,时间 \(O(\log N)\),这样我们就将一段区间的直径相关信息以一个二元组 \((u, r)\) 用少量元素表示了出来。接下来考虑怎么统计 \(O(N^2)\) 个区间的答案,考虑分治这种比较经典的统计方法,对于一个分治区间 \([L, R]\),中点为 \(mid\),枚举区间左端点 \(l\),考虑右端点 \(r\),对于形如 \([l, mid]\)\([mid + 1, r]\) 的信息我们可以 \(O(N \log^2 N)\) 预处理,然后 \([l, r]\) 就可以合并 \([l, mid]\)\([mid + 1, r]\) 两个圆来得到 \([l, r]\) 的圆计算答案,对于两个圆 \((x_1, r_1)\)\((x_2, r_2)\)(注意这里是半径)合并,分 \(3\) 种情况:

  • \((x_1, r_1)\) 包含 \((x_2, r_2)\),直径为 \(2 r_1\),这种情况容易发现随着左端点递减 \((x_1, r_1)\) 代表的圆不断变大,随着右端点递增 \((x_2, r_2)\) 代表的圆不断变小,所以对于每个左端点 \(l\),满足这种情况的右端点是一段前缀,双指针维护。

  • \((x_1, r_1)\)\((x_2, r_2)\) 包含,直径为 \(2 r_2\),做法同上,对于每个左端点 \(l\),满足这种情况的右端点是一段后缀。

  • \((x_1, r_1)\)\((x_2, r_2)\) 互不包含,在考虑完前两种情况后,对于递减的左端点 \(l\),满足这种情况的右端点就是一段不断右移的区间。考虑合并的新圆,新直径 \(d\)\(r_1 + r_2 + dis(x_1, x_2)\),新圆心 \(x\)\(x_1\)\(x_2\) 移动 \(\frac{d}{2} - r1\) 个单位得到的点,由于我们计算答案只关心直径为多少,所以观察直径,发现 \(r_1\)\(r_2\) 容易统计,问题在于 \(dis(x_1, x_2)\),等价于在树上动态激活或失效一个点,动态查询到某个点 \(u\) 的距离之和,点分树解决。

现在讨论点分树怎么解决提出的问题,我们对于每个结点 \(u\) 维护 \(sum_u, cnt_u, sum2_u\) 分别表示 \(u\) 点分树上子树中所有激活点在原树上到 \(u\) 距离之和,\(u\) 点分树上子树中激活点数量和 \(u\) 点分树上子树中所有激活点在原树上到 \(fa_u\) 距离之和,由于对于两点 \(x, y\),其点分树上 \(\text{lca} z\) 在原树上一定是 \(x, y\) 链上某点,所以 \(dis(x, y)\) 就等于 \(dis(x, z) + dis(z, y)\),所以对于修改 \(v\),我们只需修改其点分树上祖先的信息即可,然后对于查询 \(v\),我们先加上 \(sum_v\),然后假设现在遍历到其祖先 \(t\)\(t\) 可以等于 \(v\)),答案加上 \(sum_{fa_t} + dis(v, fa_t) \cdot cnt_{fa_t} - sum2_t - dis(v, t) \cdot cnt_t\) 即可,如果看不懂可以看小D

到这里就做完了,代码比较 shit,供读者参考。

点击查看代码
/*
address:https://codeforces.com/problemset/problem/1458/F
AC 2026/9/9 14:39
*/
#include<bits/stdc++.h>
using namespace std;
typedef long long LL;
const int N = 2e5 + 5;
int n;
vector<int>G[N];
int top[N], siz[N], dep[N], fa[N], son[N], rnk[N];
vector<int>num[N];
inline void dfs1(int u) {
    dep[u] = dep[fa[u]] + 1;
    siz[u] = 1;
    for (int v : G[u])
        if (v != fa[u]) {
            fa[v] = u;
            dfs1(v);
            siz[u] += siz[v];
            if (siz[v] > siz[son[u]]) son[u] = v;
        }
}
inline void dfs2(int u) {
    rnk[u] = num[top[u]].size();
    num[top[u]].push_back(u);
    if (son[u]) top[son[u]] = top[u], dfs2(son[u]);
    for (int v : G[u])
        if (v != fa[u] && v != son[u]) {
            top[v] = v;
            dfs2(v);
        }
}
inline int lca(int u, int v) {
    while (top[u] != top[v])
        if (dep[top[u]] > dep[top[v]]) u = fa[top[u]];
        else v = fa[top[v]];
    return dep[u] > dep[v] ? v : u;
}
inline int lift(int u, int k) {
    while (rnk[u] < k) k -= rnk[u] + 1, u = fa[top[u]];
    return num[top[u]][rnk[u] - k];
}
inline int dis(int u, int v) {
    int anc = lca(u, v);
    return dep[u] + dep[v] - (dep[anc] << 1);
}
struct circle {
    int u, r;
    circle() { u = r = 0; }
    circle(int u, int r) : u(u), r(r) {}
    friend inline int check(circle x, circle y) {
        int d = dis(x.u, y.u);
        return d + y.r <= x.r ? -1 : d + x.r <= y.r ? 1 : 0;
    }
    circle operator + (const circle& o)const {
        int op = check(*this, o);
        if (op == 1) return o;
        else if (op == -1) return *this;
        int anc = lca(u, o.u);
        circle ret(u, r + o.r + dep[u] + dep[o.u] - (dep[anc] << 1) >> 1);
        if (dep[u] - dep[anc] < ret.r - r) ret.u = lift(o.u, ret.r - o.r);
        else ret.u = lift(u, ret.r - r);
        return ret;
    }
}a[N];
int Siz[N], up[N];
int val[N][30];
bool vis[N];
inline int find(int u, int fa, int n) {
    for (int v : G[u])
        if (v != fa && !vis[v])
            if (Siz[v] > n >> 1) return find(v, u, n);
    return u;
}
inline void ergodic(int u, int fa) {
    Siz[u] = 1;
    for (int v : G[u])
        if (v != fa && !vis[v]) {
            ergodic(v, u);
            Siz[u] += Siz[v];
        }
}
int rt;
inline void dfs(int u, int fa) {
    ergodic(u, fa);
    u = find(u, fa, Siz[u]);
    if (!fa) rt = u;
    vis[u] = true;
    up[u] = fa;
    for (int v : G[u])
        if (!vis[v]) dfs(v, u);
}
LL sum[N], sum2[N];
int cnt[N];
inline void Insert(int u) {
    for (int i = 0, v = u;v;++i, v = up[v])
        ++cnt[v], sum[v] += val[u][i],
        sum2[v] += val[u][i + 1];
}
inline void Erase(int u) {
    for (int i = 0, v = u;v;++i, v = up[v])
        --cnt[v], sum[v] -= val[u][i],
        sum2[v] -= val[u][i + 1];
}
inline LL query(int u) {
    LL ret = sum[u];
    for (int i = 1, v = u;up[v];++i, v = up[v])
        ret += sum[up[v]] + 1ll * cnt[up[v]] * val[u][i],
        ret -= sum2[v] + 1ll * cnt[v] * val[u][i];
    return ret;
}
LL ans;
LL pre[N];
inline void solve(int l, int r) {
    if (l == r) return;
    int mid = l + r >> 1;
    solve(l, mid), solve(mid + 1, r);
    for (int i = l;i <= r;++i) a[i] = { i, 0 };
    for (int i = mid - 1;i >= l;--i) a[i] = a[i] + a[i + 1];
    for (int i = mid + 2;i <= r;++i) a[i] = a[i - 1] + a[i];
    pre[mid] = 0;
    for (int i = mid + 1;i <= r;++i) pre[i] = pre[i - 1] + a[i].r;
    int j = mid + 1, k = mid + 1;
    for (int i = mid;i >= l;--i) {
        while (k <= r && check(a[i], a[k]) != 1) Insert(a[k++].u);
        while (j <= r && check(a[i], a[j]) == -1) Erase(a[j++].u);
        ans += 1ll * a[i].r * (j - mid - 1) + pre[r] - pre[k - 1];
        ans += 1ll * a[i].r * (k - j) + pre[k - 1] - pre[j - 1] + query(a[i].u) >> 1;
    }
    while (j < k) Erase(a[j++].u);
}
int main() {
    scanf("%d", &n);
    for (int i = 1, cntn = n;i < n;++i) {
        int u, v;scanf("%d%d", &u, &v);
        ++cntn;
        G[u].push_back(cntn), G[cntn].push_back(u);
        G[v].push_back(cntn), G[cntn].push_back(v);
    }
    dfs1(1);
    top[1] = 1;
    dfs2(1);
    dfs(1, 0);
    for (int i = 1;i <= (n << 1) - 1;++i)
        for (int j = 0, u = i;u;++j, u = up[u]) val[i][j] = dis(i, u);
    solve(1, n);
    printf("%lld\n", ans);
    return 0;
}
posted @ 2026-09-10 09:01  keysky  阅读(3)  评论(0)    收藏  举报