题解:P6798 「StOI-2」简单的树
前言
题解区怎么全是树剖啊,树剖多难写,我来写一发倍增的。
思路
题意很清楚了,直接说思路了。
本题最关键的性质:这个
由于
考虑修改
首先来处理
由于
接下来是
这个问题是平凡的,因为只要满足
既然
-
小于的
l 的部分(这里并没有严格定义等号的位置,理论上等号放在哪都可以)修改后的
val 一定是c' 。具体地,val 是l ,l + 1 ,l + 2 ……r ,直接等差数列求和计算答案即可。 -
在
l 和r 之间的部分这部分是最难受的,因为一开始,
val 是v' ,但在c' 增大的过程中,val 会变成c' 。我们考虑推一下这个贡献的表达式,具体地,val 是v' ,v' (一共有(v' - l) 个)……v' ,v' + 1 ,v' + 2 ……r 。
观察式子,我们需要维护的是
-
大于
r 的部分这一部分中,
c' 不可能成为val ,直接用v' 计算答案即可。
实现
就像前文说的,我们采用倍增的思想,在链上跳,得以快速找出那三个段点的位置,找到段点后,由于我们已经找到了每一段可能成为新的
通过记录
::::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;
}
::::