「ZJOI2008」树的统计 - 树链剖分

题意

给出一棵树,每个点有一个可修改的点权,每次查询两点之间所有点的点权和或点权最大值。

思路

作为树链剖分线段树的复习题,难度不高但是细节很多

树链剖分的策略是轻重边路径剖分,这种策略可以保证整棵树上的轻边和链的数量都不超过 \(\log n\)

树链剖分精要:

  • 第一遍 dfs 求出 son fa dep 为重链划分做准备
  • 第一遍 dfs 求出 top tid pos 通过 dfn 序重新编号,使得同链上节点编号连续。该节点与重儿子相连,最终构成链
  • 通过在链上的跳跃来求出一段路径的值,dep 维护高度
注意:

1. `top` `fa` 存的是原来的节点编号

2. 数据结构中的节点使用重新的标号 `rid`

代码
#include <cstdio>

#define ll long long
#define int long long
const int maxn = 3e4 + 10, maxm = maxn << 1;
struct edge { int to, nxt; } e[maxm];
int w[maxn], head[maxn], n, a[maxn];
int son[maxn], fa[maxn], dep[maxn], sz[maxn];
int top[maxn], pos[maxn], tid[maxn];

inline void swap(int &x, int &y) {
  x ^= y, y ^= x, x ^= y;
}

inline int max(int x, int y) {
  return x > y ? x : y;
}

inline void addline(int from, int to) {
  static int ecnt = 0;
  e[++ecnt] = (edge){ to, head[from] }, head[from] = ecnt;
}

void dfs1(int x, int f, int d) {
  sz[x] = 1, fa[x] = f, dep[x] = d;
  for (int i = head[x]; i; i = e[i].nxt) {
    if (e[i].to == f) continue;
    dfs1(e[i].to, x, d + 1);
    sz[x] += sz[e[i].to];
    if (sz[e[i].to] > sz[son[x]]) son[x] = e[i].to;
  }
}

void dfs2(int x, int rt) {
  static int dfnt = 0;
  top[x] = rt, pos[tid[x] = ++ dfnt] = x; a[tid[x]] = w[x];
  if (son[x]) dfs2(son[x], rt);
  for (int i = head[x]; i; i = e[i].nxt) {
    if (e[i].to != fa[x] && e[i].to != son[x]) dfs2(e[i].to, e[i].to);
  }
}

int lca(int x, int y) {
  if (top[x] == top[y]) return dep[x] < dep[y] ? x : y;
  while (top[x] != top[y]) {
    if (dep[top[x]] < dep[top[y]]) swap(x, y);
    x = fa[top[x]];
  }
  return dep[x] < dep[y] ? x : y;
}

struct node { ll sum; int mx; } tree[maxn << 2];

inline void update(int x) {
  tree[x].sum = tree[x << 1].sum + tree[x << 1 | 1].sum;
  tree[x].mx = max(tree[x << 1].mx, tree[x << 1 | 1].mx);
}

inline void build(int k, int l, int r) {
  if (l == r) {
    tree[k] = (node) {a[l], a[l]};
    return;
  }
  int mid = (l + r) >> 1;
  build(k << 1, l, mid), build(k << 1 | 1, mid + 1, r);
  update(k);
}

inline void modify(int k, int l, int r, int p, int val) {
  if (l == r) {
    tree[k].mx = tree[k].sum = val;
    a[l] = val;
    return;
  }
  int mid = (l + r) >> 1;
  if (p <= mid) modify(k << 1, l, mid, p, val);
  else modify(k << 1 | 1, mid + 1, r, p, val);
  update(k);
}

int queryM(int k, int l, int r, int L, int R) {
  if (L > R) printf("ERROR");
  if (l == L && r == R) return tree[k].mx;
  int mid = (l + r) >> 1;
  if (R <= mid) return queryM(k << 1, l, mid, L, R);
  else if (L > mid) return queryM(k << 1 | 1, mid + 1, r, L, R);
  else return max(queryM(k << 1, l, mid, L, mid), queryM(k << 1 | 1, mid + 1, r, mid + 1, R));
}

ll queryS(int k, int l, int r, int L, int R) {
  if (L > R) printf("ERROR");
  if (l == L && r == R) return tree[k].sum;
  int mid = (l + r) >> 1;
  if (R <= mid) return queryS(k << 1, l, mid, L, R);
  else if (L > mid) return queryS(k << 1 | 1, mid + 1, r, L, R);
  else return queryS(k << 1, l, mid, L, mid) + queryS(k << 1 | 1, mid + 1, r, mid + 1, R);
}

int jumpM(int f, int t) {
  if (top[f] == top[t]) {
    return queryM(1, 1, n, tid[t], tid[f]);
  }
  int ans = -1e9;
  while (top[f] != top[t]) {
    ans = max(ans, queryM(1, 1, n, tid[top[f]], tid[f]));
    f = fa[top[f]];
  }
  ans = max(ans, queryM(1, 1, n, tid[t], tid[f]));
  return ans;
}

ll jumpS(int f, int t) {
  if (top[f] == top[t]) {
    return queryS(1, 1, n, tid[t], tid[f]);
  }
  ll ans = 0;
  while (top[f] != top[t]) {
    ans += queryS(1, 1, n, tid[top[f]], tid[f]);
    f = fa[top[f]];
  }
  ans += queryS(1, 1, n, tid[t], tid[f]);
  return ans;
}

int queryM(int x, int y) {
  int Lca = lca(x, y);
  int aa = jumpM(x, Lca), bb = jumpM(y, Lca);
  return max(aa, bb);
}

ll queryS(int x, int y) {
  int Lca = lca(x, y);
  ll aa = jumpS(x, Lca), bb = jumpS(y, Lca);
  return aa + bb - a[tid[Lca]];
}

signed main() {
  scanf("%lld", &n);
  for (int i = 1, a, b; i < n; ++ i) {
    scanf("%lld %lld", &a, &b);
    addline(a, b), addline(b, a);
  }
  for (int i = 1; i <= n; ++ i)
    scanf("%lld", &w[i]);
  dfs1(1, 1, 0); dfs2(1, 1);
  build(1, 1, n);
  int t, x, y; char s[10];
  scanf("%lld", &t);
  while (t--) {
    scanf("%s %lld %lld", s, &x, &y);
    if (s[0] == 'C') modify(1, 1, n, tid[x], y);
    else if (s[1] == 'M') printf("%lld\n", queryM(x, y));
    else printf("%lld\n", queryS(x, y));
  }
}
posted @ 2021-01-19 20:38  trswnca  阅读(79)  评论(0)    收藏  举报