题解:P6798 「StOI-2」简单的树

· · 题解

前言

题解区怎么全是树剖啊,树剖多难写,我来写一发倍增的。

思路

题意很清楚了,直接说思路了。

本题最关键的性质:这个 val_i 自下往上看是单调不降的

由于 val 是针对子树内的情况,于是我们将问题抽象成对一条链进行修改。所以其实这条链以外的点是不受影响的,可以先累计这部分的答案。

考虑修改 c_u 后能成为最大值的元素,容易发现只能是去掉 c_u 最大值(后文记为 v')和修改后的 c_i(后文记为 c')。

首先来处理 v',容易发现有两种情况。

由于 val 单调不减的性质,又因为次大值严格不大于最大值,所以我们可以得出一个结论,一定存在一个点,使得在这个点下面的点使用的都是次大值,上面的点使用的都是最大值,于是我们不妨使用倍增寻找这个点。

接下来是 c',我们思考,什么时候这个 c' 可以当 val_i

这个问题是平凡的,因为只要满足 c' \ge v' 即可,但是,这里的 c' \in \left[ l, r \right],它是一个变量!

既然 c' 是变量,那我们就考虑寻找一些不变的量(这里指好维护的量),那么我们就想到了 val。以 lr 为段点给这一条链上的 val 分段。

S = (v' - l) \times v' + \dfrac{(v' + r)(r - v' + 1)}{2} = \dfrac{1}{2}(v'^2 + v' - 2v'l + r^2 + r)

观察式子,我们需要维护的是 v'^2v',也就是最大值,次大值的一次和二次和

实现

就像前文说的,我们采用倍增的思想,在链上跳,得以快速找出那三个段点的位置,找到段点后,由于我们已经找到了每一段可能成为新的 val 的元素的计算方式,直接计算即可。这里分类讨论比较繁琐,具体见代码。

通过记录

::::info[小清新代码~]

#include "bits/stdc++.h"
//#include "bits/extc++.h"
#define ll long long
#define int long long
#define eb emplace_back 
#define pii std::pair<int, int> 
#define mkp std::make_pair
#define smin(a, b) (a = min(a, b))
#define smax(a, b) (a = max(a, b))
#define fi first
#define se second
#define rep(i, j, k) for (int i = (j); i <= (k); ++i)
#define per(i, j, k) for (int i = (j); i >= (k); --i)
using namespace std;

const int N = 5e5 + 7, K = 23, P = 998244353;
const int INV2 = (P + 1) / 2;
int n, q, opt, lst = 0, c[N], mx[N], cm[N], fa[K][N], dep[N], pw[K], s[N], S = 0;
int sm[K][N], sm2[K][N], sc[K][N], sc2[K][N];
vector<int> E[N];

inline void dfs(int x, int fat) {
    mx[x] = c[x];
    for (auto to : E[x]) if (to != fat) {
        dep[to] = dep[x] + 1;
        dfs(to, x);
        fa[0][to] = x;
        if (mx[to] > mx[x]) {
            cm[x] = max(mx[x], cm[to]);
            mx[x] = mx[to];
        } else if (mx[to] > cm[x]) {
            cm[x] = mx[to];
        } if (cm[to] > cm[x]) {
            cm[x] = cm[to];
        }
    }
}
inline void dfss(int x, int fat) {
    sm[0][x] = mx[fa[0][x]];
    sm2[0][x] = mx[fa[0][x]] * mx[fa[0][x]] % P;
    sc[0][x] = cm[fa[0][x]];
    sc2[0][x] = cm[fa[0][x]] * cm[fa[0][x]] % P;
    for (auto to : E[x]) if (to != fat) {
        s[to] = s[x] + mx[to];
        dfss(to, x);
    }
}
inline void add(int &a, int b) {
    b %= P;
    (b += P) %= P;
    a += b;
    if (a >= P) a -= P;
}

inline pii getsc(int x, int y) {
    int d = dep[x] - dep[y], ret1 = cm[x], ret2 = cm[x] * cm[x] % P;
    per(i, 20, 0) if ((d >> i) & 1)
        add(ret1, sc[i][x]), add(ret2, sc2[i][x]), x = fa[i][x];
    return mkp(ret1, ret2);
}
inline pii getsm(int x, int y) {
    int d = dep[x] - dep[y], ret1 = mx[x], ret2 = mx[x] * mx[x] % P;
    per(i, 20, 0) if ((d >> i) & 1)
        add(ret1, sm[i][x]), add(ret2, sm2[i][x]), x = fa[i][x];
    return mkp(ret1, ret2);
}

signed main() {
    std::ios::sync_with_stdio(false);   
    std::cin.tie(0);
    std::cout.tie(0);

    cin >> n >> q >> opt;
    pw[0] = 1;
    rep(i, 1, 20) pw[i] = pw[i - 1] << 1;
    rep(i, 1, n) cin >> c[i];
    rep(i, 1, n - 1) {
        int u, v;
        cin >> u >> v;
        E[u].eb(v);
        E[v].eb(u);
    }

    dep[1] = 1, dep[n + 1] = INT_MAX;
    dfs(1, 1);
    rep(i, 1, n) S += mx[i];
    s[1] = mx[1];
    dfss(1, 1);
    mx[0] = LLONG_MAX;
    rep(j, 1, 20) rep(i, 1, n) {
        fa[j][i] = fa[j - 1][fa[j - 1][i]];
        sm[j][i] = (sm[j - 1][i] + sm[j - 1][fa[j - 1][i]]) % P;
        sm2[j][i] = (sm2[j - 1][i] + sm2[j - 1][fa[j - 1][i]]) % P;
        sc[j][i] = (sc[j - 1][i] + sc[j - 1][fa[j - 1][i]]) % P;
        sc2[j][i] = (sc2[j - 1][i] + sc2[j - 1][fa[j - 1][i]]) % P;
    } 

    rep(i, 1, q) {
        int l, r, x, ans = 0;
        cin >> l >> r >> x;
        l = (l + opt * lst) % n + 1;
        r = (r + opt * lst) % n + 1;
        x = (x + opt * lst) % n + 1;
        if (l > r) swap(l, r);

        add(ans, (S - s[x]) * (r - l + 1) % P);

        int x1 = x, x2 = x, x3 = x;
        int SC = 0, SC2 = 0, SM = 0, SM2 = 0;
        per(i, 20, 0) if (mx[fa[i][x1]] == c[x]) 
           x1 = fa[i][x1];
        if (mx[x1] != c[x]) x1 = n + 1;

        fa[0][n + 1] = x, dep[n + 1] = dep[x] + 1; //由于我们的 x2 x3 都是闭区间,所以设置一个虚拟节点来处理一些边界情况
        rep(i, 1, 20) fa[i][n + 1] = fa[i - 1][fa[i - 1][n + 1]];

        per(i, 20, 0) if (cm[fa[i][x2]] <= l && dep[fa[i][x2]] >= dep[x1]) 
            x2 = fa[i][x2];
        if (cm[x2] > l) x2 = n + 1;
        if (dep[x2] > dep[x1]) {
            per(i, 20, 0) if (cm[fa[i][x3]] <= r && dep[fa[i][x3]] >= dep[x1])
                x3 = fa[i][x3];
            if (cm[x3] > r) x3 = n + 1;
            if (dep[x3] == dep[x1]) {
                per(i, 20, 0) if (mx[fa[i][x3]] <= r)
                    x3 = fa[i][x3];
            }
            if ((mx[x3] == c[x] ? cm[x3] : mx[x3]) > r)
                x3 = n + 1;
        } else {
            per(i, 20, 0) if (mx[fa[i][x2]] <= l)
                x2 = fa[i][x2];
            if ((mx[x2] == c[x] ? cm[x2] : mx[x2]) > l)
                x2 = n + 1;
            x3 = x1;
            per(i, 20, 0) if (mx[fa[i][x3]] <= r)
                x3 = fa[i][x3];
            if ((mx[x3] == c[x] ? cm[x3] : mx[x3]) > r)
                x3 = n + 1;
        }

        // <- l   x~x2
        int sum;
        if (dep[x] >= dep[x2]) {
            sum = (l + r) * (r - l + 1) / 2;
            add(ans, sum * (dep[x] - dep[x2] + 1));
        }
        // l <- -> r   x2~x3

        if (dep[x2] > dep[x1] && dep[x1] > dep[x3]) {
            auto tmp = getsc(fa[0][x2], x1);
            SC = tmp.fi, SC2 = tmp.se;
            sum = SC2 + (1 - 2 * l) * SC % P + (r * r % P + r) % P * (dep[x2] - dep[x1]) % P;
            sum %= P;
            (sum *= INV2) %= P;
            add(ans, sum);

            tmp = getsm(fa[0][x1], x3);
            SM = tmp.fi, SM2 = tmp.se;
            sum = SM2 + (1 - 2 * l) * SM % P + (r * r % P + r) % P * (dep[x1] - dep[x3]) % P;
            sum %= P;
            (sum *= INV2) %= P;
            add(ans, sum);
        } else if (dep[x1] <= dep[x3] && dep[x2] > dep[x3]) {
            auto tmp = getsc(fa[0][x2], x3);
            SC = tmp.fi, SC2 = tmp.se;

            sum = SC2 + (1 - 2 * l) * SC % P + (r * r % P + r) % P * (dep[x2] - dep[x3]) % P;
            sum %= P;
            (sum *= INV2) %= P;
            add(ans, sum);
        } else if (dep[x2] > dep[x3]) {
            auto tmp = getsm(fa[0][x2], x3);
            SM = tmp.fi, SM2 = tmp.se;

            sum = SM2 + (1 - 2 * l) * SM % P + (r * r % P + r) % P * (dep[x2] - dep[x3]) % P;
            sum %= P;
            (sum *= INV2) %= P;
            add(ans, sum);
        }
        // r -> x3 ~ 1
        if (dep[x1] < dep[x3]) {
            auto tmp = getsc(fa[0][x3], x1);
            SC = tmp.fi;
            (SC *= (r - l + 1)) %= P;

            add(ans, SC);

            if (dep[x1] > 1) {
                tmp = getsm(fa[0][x1], 1);
                SM = tmp.fi;
                (SM *= (r - l + 1)) %= P;
                add(ans, SM);
            }
        } else if (dep[x3] > 1) {
            auto tmp = getsm(fa[0][x3], 1);
            SM = tmp.fi;    
            (SM *= (r - l + 1)) %= P;

            add(ans, SM);
        } 

        cout << (lst = ans) << '\n';
    }

    return 0;
}

::::