球球了,既然进来了就帮忙看看吧。
#include <iostream>
#include <cstring>
#include <cstdio>
#define rep(i, a, b) for (int i = (a); i <= (b); i ++ )
#define per(i, a, b) for (int i = (a); i >= (b); i -- )
using namespace std;
const int N = 100010;
const int INF = 0x3f3f3f3f;
int h[N], e[N << 1], ne[N << 1], val[N], idx;
int fa[N], sz[N], dep[N], son[N], top[N], id[N], cnt;
int f[N][2], g[N][2], End[N], n, m;
void add(int a, int b) {
e[ ++ idx] = b, ne[idx] = h[a], h[a] = idx;
}
struct Matrix {
int a[3][3];
Matrix() { memset(a, 0, sizeof a); }
void makeINF() { memset(a, -0x3f, sizeof a); }
}M[N];
Matrix operator * (Matrix A, Matrix B) {
Matrix ans; ans.makeINF();
rep(i, 0, 1) rep(j, 0, 1) rep(k, 0, 1)
ans.a[i][j] = max(ans.a[i][j], A.a[i][k] + B.a[k][j]);
return ans;
}
void dfs1(int u, int father) {
fa[u] = father, dep[u] = dep[fa[u]] + 1, sz[u] = 1;
for (int i = h[u]; i; i = ne[i]) {
if (e[i] == father) continue;
dfs1(e[i], u); sz[u] += sz[e[i]];
if (sz[son[u]] < sz[e[i]]) son[u] = e[i];
}
}
void dfs2(int u, int t) {
top[u] = t, id[u] = ++ cnt, End[t] = cnt;
if (son[u]) dfs2(son[u], t);
for (int i = h[u]; i; i = ne[i]) {
if (e[i] == fa[u] or e[i] == son[u]) continue;
dfs2(e[i], e[i]);
}
}
void dfs3(int u) {
g[u][1] = val[u];
for (int i = h[u]; i; i = ne[i]) {
int v = e[i];
if (v == fa[u] or v == son[u]) continue;
dfs3(v);
g[u][0] += max(f[v][0], f[v][1]);
g[u][1] += f[v][0];
}
f[u][0] += g[u][0];
f[u][1] += g[u][1];
if (!son[u]) return;
dfs3(son[u]);
f[u][0] += max(f[son[u]][1], f[son[u]][0]);
f[u][1] += f[son[u]][0];
}
struct node {
int l, r;
Matrix sum;
}tr[N << 2];
#define ls u << 1
#define rs u << 1 | 1
void pushup(int u) {
tr[u].sum = tr[ls].sum * tr[rs].sum;
}
void build(int u, int l, int r) {
tr[u] = {l, r};
if (l == r) {
tr[u].sum = M[r];
return;
}
int mid = l + r >> 1;
build(ls, l, mid), build(rs, mid + 1, r);
pushup(u);
}
void modify(int u, int x) {
if (tr[u].l == tr[u].r) {
tr[u].sum = M[x];
return;
}
int mid = tr[u].l + tr[u].r >> 1;
if (x <= mid) modify(ls, x);
else modify(rs, x);
pushup(u);
}
Matrix query(int u, int l, int r) {
if (tr[u].l >= l && tr[u].r <= r) return tr[u].sum;
int mid = tr[u].l + tr[u].r >> 1;
if (r <= mid) return query(ls, l, r);
if (l > mid) return query(rs, l, r);
return query(ls, l, r) * query(rs, l, r);
}
void solve(int u, int w) {
M[id[u]].a[1][0] -= val[u];
M[id[u]].a[1][0] += w;
val[u] = w;
while (u) {
Matrix last = query(1, id[top[u]], End[top[u]]);
modify(1, id[u]);
Matrix now = query(1, id[top[u]], End[top[u]]);
u = fa[top[u]];
M[id[u]].a[0][0] -= max(last.a[0][0], last.a[1][0]);
M[id[u]].a[0][0] += max(now.a[0][0], last.a[1][0]);
M[id[u]].a[0][1] = M[id[u]].a[0][0];
M[id[u]].a[1][0] -= last.a[0][0];
M[id[u]].a[1][0] += now.a[0][0];
}
}
int main() {
scanf("%d%d", &n, &m);
for (int i = 1; i <= n; i ++ )
scanf("%d", &val[i]);
for (int i = 1; i < n; i ++ ) {
int a, b; scanf("%d%d", &a, &b);
add(a, b); add(b, a);
}
dfs1(1, 0), dfs2(1, 1); dfs3(1);
for (int u = 1; u <= n; u ++ ) {
M[id[u]].a[0][0] = g[u][0];
M[id[u]].a[0][1] = g[u][0];
M[id[u]].a[1][0] = g[u][1];
M[id[u]].a[1][1] = -INF;
}
build(1, 1, n);
while (m -- ) {
int u, w;
scanf("%d%d", &u, &w);
solve(u, w);
Matrix ans = query(1, id[1], End[1]);
cout << max(ans.a[0][0], ans.a[1][0]) << endl;
}
return 0;
}