3 条题解

  • 1
    @ 2026-8-20 15:08:33

    不知道有没有小馋猫需要树链剖分套线段树套vector维护凸包构式做法。

    先放一下:

    #include <bits/stdc++.h>
    using namespace std;
    
    typedef long long LL;
    const int N = 3e5 + 10;
    const LL INF = 4e18;
    
    int n, K;
    vector<int> G[N];
    int fa[N], dep[N], siz[N], son[N];
    int dfn[N], _dfn[N], top[N], tsp;
    
    struct ew {
        LL z, b;
    } a[N];
    
    LL bsum[N];   // 从根到每个点的 b 前缀和
    LL val[N];    // z_i^2
    
    // 树链剖分
    void dfs1(int x, int f) {
        fa[x] = f;
        dep[x] = dep[f] + 1;
        siz[x] = 1;
        son[x] = 0;
        for (int y : G[x]) {
            if (y == f) continue;
            dfs1(y, x);
            siz[x] += siz[y];
            if (siz[y] > siz[son[x]]) son[x] = y;
        }
    }
    
    void dfs2(int x, int tp) {
        dfn[x] = ++tsp;
        _dfn[tsp] = x;
        top[x] = tp;
    
        bsum[x] = bsum[fa[x]] + a[x].b;
        val[x] = a[x].z * a[x].z;
    
        if (son[x]) dfs2(son[x], tp);
        for (int y : G[x]) {
            if (y == fa[x] || y == son[x]) continue;
            dfs2(y, y);
        }
    }
    
    // 线段树维护凸包
    
    #define lc(p) (p << 1)
    #define rc(p) ((p << 1) | 1)
    struct Line {
        LL k, b;   // y = k * x + b
    };
    
    struct node {
        vector<Line> hull;
    } tr[N << 2];
    
    // 判断 l2 是否无用(从低到高的下凹凸包)
    bool jg(const Line& l1, const Line& l2, const Line& l3) {
        // l1.k < l2.k < l3.k
        // 若 (b2-b1)/(k1-k2) >= (b3-b2)/(k2-k3),则 l2 无用
        return (__int128)(l2.b - l1.b) * (l2.k - l3.k) >=
               (__int128)(l3.b - l2.b) * (l1.k - l2.k);
    }
    
    vector<Line> helpHull(vector<Line>& lines) {
        vector<Line> hull;
        for (auto& ln : lines) {
            while (hull.size() >= 2 && jg(hull[hull.size() - 2], hull.back(), ln))
                hull.pop_back();
            hull.push_back(ln);
        }
        return hull;
    }
    
    void build(int p, int l, int r) {
        if (l == r) {
            int t = _dfn[l];
            if (t != 1) {   // 非根节点才对应一条边
                tr[p].hull.push_back({-bsum[fa[t]], val[t]});
            }
            return;
        }
        int mid = (l + r) >> 1;
        build(lc(p), l, mid);
        build(rc(p), mid + 1, r);
    
        vector<Line> merged;
        merged.reserve(tr[lc(p)].hull.size() + tr[rc(p)].hull.size());
    
        merge(tr[lc(p)].hull.begin(), tr[lc(p)].hull.end(),
              tr[rc(p)].hull.begin(), tr[rc(p)].hull.end(),
              back_inserter(merged),
              [](const Line& a, const Line& b) { return a.k < b.k; });
    
        // 去重(斜率相同保留截距最大的)
        vector<Line> uniq;
        for (auto& ln : merged) {
            if (!uniq.empty() && uniq.back().k == ln.k) {
                if (ln.b > uniq.back().b) uniq.back().b = ln.b;
            } 
    		else {
                uniq.push_back(ln);
            }
        }
    
        tr[p].hull = helpHull(uniq);
    }
    
    LL eval(const Line& ln, LL x) {
        return ln.k * x + ln.b;
    }
    
    // 找 vector 内和 x 结合最大的 y 
    LL getMax(const vector<Line>& hull, LL x) {
        int l = 0, r = (int)hull.size() - 1;
        while (l < r) {
            int mid = (l + r) >> 1;
            if (eval(hull[mid], x) <= eval(hull[mid + 1], x))
                l = mid + 1;
            else
                r = mid;
        }
        return eval(hull[l], x);
    }
    
    LL query(int p, int l, int r, int ql, int qr, LL x) {
        if (ql <= l && r <= qr) {
            if (tr[p].hull.empty()) return -INF;
            return getMax(tr[p].hull, x);
        }
        int mid = (l + r) >> 1;
        LL res = -INF;
        if (ql <= mid) res = max(res, query(lc(p), l, mid, ql, qr, x));
        if (qr > mid) res = max(res, query(rc(p), mid + 1, r, ql, qr, x));
        return res;
    }
    
    // 查询路径上深度不小于 L 的 s
    LL query_path(int p, int L, LL x) {
        LL res = -INF;
        while (top[p] != 1) {
            int t = top[p];
            if (dep[p] < L) break;
    
            int l = dfn[t], r = dfn[p];
            if (dep[t] < L) {
                int offset = L - dep[t];
                l = dfn[t] + offset;
            }
            if (l <= r) {
                res = max(res, query(1, 1, n, l, r, x));
            }
            p = fa[t];
        }
    
        // 最后一条重链(以根为链头)
        if (p != 0 && dep[p] >= L) {
            int t = 1;
            int l = dfn[t], r = dfn[p];
            if (dep[t] < L) {
                int offset = L - dep[t];
                l = dfn[t] + offset;
            }
            if (l <= r) {
                res = max(res, query(1, 1, n, l, r, x));
            }
        }
        return res;
    }
    
    int main() {
        ios::sync_with_stdio(false);
        cin.tie(0);
    
        cin >> n >> K;
        for (int i = 1; i < n; ++i) {
            int p;
            cin >> p;
            G[p].push_back(i + 1);
        }
        for (int i = 1; i < n; ++i) cin >> a[i + 1].z;
        for (int i = 1; i < n; ++i) cin >> a[i + 1].b;
    
        if (n == 1) {
            cout << 0 << '\n';
            return 0;
        }
    
        dfs1(1, 0);
        tsp = 0;
        dfs2(1, 1);
        build(1, 1, n);
    
        LL ans = -INF;
    
        for (int x = 2; x <= n; x ++) {
            LL z_x = a[x].z;
            LL S_x = bsum[x];
    
            // 边数限制
    		// dep[x] - dep[s] + 1 <= K  =>  dep[s] >= dep[x] - K + 1
            int L = max(2, dep[x] - K + 1);
            if (L > dep[x]) continue;
    
            LL t = query_path(x, L, z_x);
            if (t == -INF) continue;
    
            LL sum = t + z_x * z_x + z_x * S_x;
            ans = max(ans, sum);
        }
    
        cout << ans << "\n";
        return 0;
    }
    
    

    信息

    ID
    12642
    时间
    1000ms
    内存
    512MiB
    难度
    10
    标签
    递交数
    10
    已通过
    2
    上传者