P6798 题解
Antiphilia
·
·
题解
先想 f(a,0) 怎么快速求
接下来只需要计算 $f(a,i)-f(a,i-1)$,看出来这个差的意义其实是当 $c_a=i-1$ 时,$a$ 是多少个点的子树最大值。
如果原先这个点的子树最大点为 $a$ 且没有替换的,则 $mx=c_a,se<c_a$,新最大值可以为新 $a$ 或原次大。否则新最大值可以为新 $a$ 或原最大。
就有公式:
$$f(a,i)=f(a,0)+\sum_{j,se_j<c_a}\max(0, i-se_j)+\sum_{j,se_j \ge c_a}\max(0, i-mx_j)$$
其中 $j$ 为 $a$ 的根链上的点。
对于 $se_j$ 的条件,在树上仍然是一段连续区间,依旧倍增去做。
以计算式子第二项为例
$$
\begin{aligned}
\sum_{i=1}^m\sum_{j,se_j<c_a}\max(0, i-se_j) &= \sum_{j,se_j<c_a}\sum_{i=se_j}^mi-se_j \\
&= \sum_{j,se_j<c_a}\frac{(se_j+m)(m-se_j+1)}{2}-se_j(m-se_j+1) \\
&= \sum_{j,se_j<c_a}\frac{se_j(se_j+1)}{2}-(m+1)\sum_{j,se_j<c_a}se_j+\sum_{j,se_j<c_a}\frac{m(m+1)}{2}
\end{aligned}
$$
第三项同理。
时间 $\mathcal{O}((n+q)\log n)$,空间 $\mathcal{O}(n\log n)$。
```cpp
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int mod = 998244353;
const int MAXN = 5e5;
const int BUFSIZE = 1 << 20;
char ibuf[BUFSIZE], *p1 = ibuf, *p2 = ibuf;
char obuf[BUFSIZE], *p3 = obuf;
#define getchar() (p1 == p2 && (p2 = (p1 = ibuf) + fread(ibuf, 1, BUFSIZE, stdin), p1 == p2) ? EOF : *p1++)
#define putchar(x) (p3 - obuf < BUFSIZE ? (*p3++ = (x)) : (fwrite(obuf, 1, p3 - obuf, stdout), p3 = obuf, *p3++ = (x)))
inline void flush() {
if (p3 != obuf) fwrite(obuf, 1, p3 - obuf, stdout);
}
inline int read() {
int x = 0, w = 1;
char c = getchar();
while (c < '0' || c > '9') {
if (c == '-') w = -1;
c = getchar();
}
while (c >= '0' && c <= '9') {
x = (x << 3) + (x << 1) + (c ^ 48);
c = getchar();
}
return x * w;
}
inline void write(int x) {
if (x < 0) {
putchar('-');
x = -x;
}
char stk[20];
int top = 0;
do {
stk[top++] = x % 10 + '0';
x /= 10;
} while (x);
while (top--) putchar(stk[top]);
}
inline void read(char *s) {
char c;
while ((c = getchar()) != EOF && (c == ' ' || c == '\n' || c == '\r'));
if (c == EOF)
return ;
int cnt = 0;
s[cnt++] = c;
while ((c = getchar()) != EOF && c != ' ' && c != '\n' && c != '\r')
s[cnt++] = c;
s[cnt] = '\0';
}
inline void write(const char *s) {
for (int i = 0; s[i]; i++)
putchar(s[i]);
putchar('\n');
}
int n;
int q;
int opt;
int c[MAXN + 1];
struct edge {
int to, nxt;
} E[MAXN * 2 + 1];
int e;
int head[MAXN + 1];
void add(int u, int v) {
E[++e].to = v;
E[e].nxt = head[u];
head[u] = e;
}
int tot;
int mx[MAXN + 1];
int se[MAXN + 1];
int smx[MAXN + 1];
int sse[MAXN + 1];
int sse2[MAXN + 1];
int smx2[MAXN + 1];
int f0[MAXN + 1];
int fa[MAXN + 1][19];
void dfs(int u, int f) {
mx[u] = c[u];
for (int i = head[u], v; i; i = E[i].nxt) {
v = E[i].to;
if (v == f)
continue;
fa[v][0] = u;
for (int i = 1; i <= 18; i++)
fa[v][i] = fa[fa[v][i - 1]][i - 1];
dfs(v, u);
if (mx[v] > mx[u]) {
se[u] = max(mx[u], se[v]);
mx[u] = mx[v];
} else if (mx[v] < mx[u])
se[u] = max(se[u], mx[v]);
else
se[u] = mx[v];
}
tot += mx[u];
(tot >= mod) && (tot -= mod);
}
void dfs2(int u, int f) {
smx[u] = smx[f] + mx[u];
(smx[u] >= mod) && (smx[u] -= mod);
sse[u] = sse[f] + se[u];
(sse[u] >= mod) && (sse[u] -= mod);
sse2[u] = (sse2[f] + (ll)se[u] * (se[u] + 1) / 2) % mod;
smx2[u] = (smx2[f] + (ll)mx[u] * (mx[u] + 1) / 2) % mod;
for (int i = head[u], v; i; i = E[i].nxt) {
v = E[i].to;
if (v == f)
continue;
dfs2(v, u);
}
}
int slv(int n, int a) {
int rs = 0, cnt = 0, sum = 0, sum2 = 0, t = min(n, c[a] - 1), i, j;
if (se[a] <= t) {
i = a;
for (int j = 18; ~j; j--) {
if (fa[i][j] && se[fa[i][j]] <= t) {
cnt |= 1 << j;
sum += sse[i] - sse[fa[i][j]];
(sum >= mod) && (sum -= mod);
(sum < 0) && (sum += mod);
sum2 += sse2[i] - sse2[fa[i][j]];
(sum2 >= mod) && (sum2 -= mod);
(sum2 < 0) && (sum2 += mod);
i = fa[i][j];
}
}
cnt++;
sum += se[i];
(sum >= mod) && (sum -= mod);
sum2 = (sum2 + (ll)se[i] * (se[i] + 1) / 2) % mod;
rs = (rs + sum2 + ((ll)n * (n + 1) >> 1) % mod * cnt - (ll)(n + 1) * sum) % mod;
cnt = sum = sum2 = 0;
}
if (mx[a] <= n) {
i = a;
for (int j = 18; ~j; j--) {
if (fa[i][j] && mx[fa[i][j]] <= n) {
cnt |= 1 << j;
sum += smx[i] - smx[fa[i][j]];
(sum >= mod) && (sum -= mod);
(sum < 0) && (sum += mod);
sum2 += smx2[i] - smx2[fa[i][j]];
(sum2 >= mod) && (sum2 -= mod);
(sum2 < 0) && (sum2 += mod);
i = fa[i][j];
}
}
cnt++;
sum += mx[i];
(sum >= mod) && (sum -= mod);
sum2 = (sum2 + (ll)mx[i] * (mx[i] + 1) / 2) % mod;
rs = (rs + sum2 + ((ll)n * (n + 1) >> 1) % mod * cnt - (ll)(n + 1) * sum) % mod;
cnt = sum = sum2 = 0;
}
if (se[a] < c[a] && mx[a] <= n) {
i = a;
for (int j = 18; ~j; j--) {
if (fa[i][j] && se[fa[i][j]] < c[a] && mx[fa[i][j]] <= n) {
cnt |= 1 << j;
sum += smx[i] - smx[fa[i][j]];
(sum >= mod) && (sum -= mod);
(sum < 0) && (sum += mod);
sum2 += smx2[i] - smx2[fa[i][j]];
(sum2 >= mod) && (sum2 -= mod);
(sum2 < 0) && (sum2 += mod);
i = fa[i][j];
}
}
cnt++;
sum += mx[i];
(sum >= mod) && (sum -= mod);
sum2 = (sum2 + (ll)mx[i] * (mx[i] + 1) / 2) % mod;
rs = (rs - sum2 - ((ll)n * (n + 1) >> 1) % mod * cnt + (ll)(n + 1) * sum) % mod;
}
(rs < 0) && (rs += mod);
return rs;
}
int lastans;
int main () {
n = read(), q = read(), opt = read();
for (int i = 1; i <= n; i++)
c[i] = read();
for (int i = 1, u, v; i < n; i++) {
u = read(), v = read();
add(u, v);
add(v, u);
}
dfs(1, 0);
dfs2(1, 0);
for (int i = 1, j; i <= n; i++) {
f0[i] = tot;
if (se[i] < c[i]) {
j = i;
for (int k = 18; ~k; k--) {
if (fa[j][k] && se[fa[j][k]] < c[i]) {
f0[i] = (f0[i] + (ll)sse[j] - smx[j] - sse[fa[j][k]] + smx[fa[j][k]]) % mod;
j = fa[j][k];
}
}
f0[i] += se[j] - mx[j];
(f0[i] >= mod) && (f0[i] -= mod);
(f0[i] < 0) && (f0[i] += mod);
}
}
for (int i = 1, l, r, a; i <= q; i++) {
l = (read() + opt * lastans) % n + 1;
r = (read() + opt * lastans) % n + 1;
if (l > r)
swap(l, r);
a = (read() + opt * lastans) % n + 1;
lastans = 1ull * f0[a] * (r - l + 1) % mod;
lastans -= slv(l - 1, a);
(lastans < 0) && (lastans += mod);
lastans += slv(r, a);
(lastans >= mod) && (lastans -= mod);
write(lastans), putchar('\n');
}
flush();
return 0;
}
```