P6798 题解

· · 题解

先想 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; } ```