萌新刚学ddp 10 ^ -998244353 秒样例不过万紫千红
查看原帖
萌新刚学ddp 10 ^ -998244353 秒样例不过万紫千红
519384
Link_Cut_Y楼主2023/7/5 15:16

球球了,既然进来了就帮忙看看吧。

#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;
}

2023/7/5 15:16
加载中...