救救孩子 TLE90
查看原帖
救救孩子 TLE90
758679
phoenixzhan楼主2023/9/10 09:51

和题解一模一样

TLE

#include <bits/stdc++.h>
using namespace std;
#define pb push_back
#define pii pair<int, int>
#define mp make_pair
#define fi first
#define se second
#define deb(var) cerr << #var << '=' << var << "; "
// #define int long long
const int maxn = 1e6 + 5;
struct Mat {
	int n, m, w[2][2];
	Mat() {
		n = m = 0; w[0][0] = w[0][1] = w[1][0] = w[1][1] = -1e9;
	}
	Mat(int x) {
		n = m = x;
		for (int i = 0; i < n; i++)
			for (int j = 0; j < n; j++) w[i][j] = -(i != j) * 1e9;
	}
	Mat(int x, int y) {
		n = x, m = y; memset(w, 0, sizeof w);
	}
	int* operator [](int k) { return w[k]; }
	friend Mat operator * (Mat a, Mat b) {
		Mat c; c.n = a.n, c.m = b.m;
		for (int i = 0; i < c.n; i++)
			for (int j = 0; j < c.m; j++)
				for (int k = 0; k < a.m; k++)
					c.w[i][j] = max(c.w[i][j], a.w[i][k] + b.w[k][j]); return c;
	}
};
int n, q, a[maxn];
vector<int> g[maxn];
int siz[maxn], hvy[maxn], ba[maxn];
void init(int u, int fa) {
	siz[u] = 1;
	for (int i = 0; i < g[u].size(); i++) {
		int v = g[u][i]; if (v == fa) continue; ba[v] = u;
		init(v, u); siz[u] += siz[v]; if (siz[v] > siz[hvy[u]]) hvy[u] = v;
	}
}
int sum[maxn][2]; Mat f[maxn];   // 轻儿子
int dfn[maxn], tim, top[maxn], bot[maxn], Ref[maxn]; 
Mat getmat(int u) {
	Mat c; c.n = c.m = 2;
	c[0][0] = sum[u][0], c[0][1] = sum[u][0], c[1][0] = sum[u][1], c[1][1] = -1e9; return c;
} 
int tval[maxn], rt[maxn];   // 权重 
struct Segt {
	Mat w[maxn], pr[maxn]; int mid[maxn], ls[maxn], rs[maxn], tot;
	Segt() {
		for (int i = 0; i <= maxn; i++) w[i] = Mat(2); tot = 0;
	} 
	void pu(int u) {
		w[u] = w[ls[u]] * pr[u] * w[rs[u]];
	}
	void build(int &u, int l, int r) {
		if (l > r) return u = 0, void();
		u = ++tot;
		int sum = 0; for (int i = l; i <= r; i++) sum += tval[Ref[i]];
		int mnu = r, qwq = 0;
		for (int i = l; i <= r; i++) {
		    qwq += tval[Ref[i]];
			if (qwq * 2 > sum) {
			    mnu = i; break;
			}
		} 
		mid[u] = mnu;
		pr[u] = getmat(Ref[mnu]);
		build(ls[u], l, mnu - 1); 
		build(rs[u], mnu + 1, r); pu(u);
	}
	void upd(int u, int l, int r, int p, Mat x) {
		if (p == mid[u]) pr[u] = x;
		else if (p < mid[u]) upd(ls[u], l, mid[u] - 1, p, x); else upd(rs[u], mid[u] + 1, r, p, x); pu(u);
	} 
	Mat query(int u, int l, int r, int L, int R) { return w[u]; }
} seg;
void dfs(int u, int fa, int tp) {
	sum[u][1] = a[u];
	top[u] = tp; bot[top[u]] = u; dfn[u] = ++tim; Ref[tim] = u;
	if (hvy[u]) dfs(hvy[u], u, tp);
	tval[u] = 1;
	for (int i = 0; i < g[u].size(); i++) {
		int v = g[u][i]; if (v != fa && v != hvy[u]) {
			dfs(v, u, v); sum[u][0] += max(f[v][0][0], f[v][1][0]); sum[u][1] += f[v][0][0]; tval[u] += siz[v];
		}
	}
	if (tp == u) seg.build(rt[u], dfn[u], dfn[bot[u]]), f[u] = seg.query(rt[u], dfn[u], dfn[bot[u]], dfn[u], dfn[bot[u]]) * Mat(2, 1);
}
void upd(int u, int x) {
	sum[u][1] += x - a[u]; seg.upd(rt[top[u]], dfn[top[u]], dfn[bot[top[u]]], dfn[u], getmat(u)); a[u] = x;
	while (u) {
		int t = top[u];
		sum[ba[t]][0] -= max(f[t][0][0], f[t][1][0]); sum[ba[t]][1] -= f[t][0][0];
		f[t] = seg.query(rt[t], dfn[t], dfn[bot[t]], dfn[t], dfn[bot[t]]) * Mat(2, 1);
		sum[ba[t]][0] += max(f[t][0][0], f[t][1][0]); sum[ba[t]][1] += f[t][0][0];
		if (ba[t]) seg.upd(rt[top[ba[t]]], dfn[top[ba[t]]], dfn[bot[top[ba[t]]]], dfn[ba[t]], getmat(ba[t])); u = ba[t]; 
	}
}
int query() {
	return max(f[1][0][0], f[1][1][0]);
}
signed main() {
//	freopen("D:\\data.in","r",stdin);
//	freopen("D:\\t.out","w",stdout);
	ios::sync_with_stdio(0), cin.tie(0), cout.tie(0);
	cin >> n >> q;
	for (int i = 1; i <= n; i++) cin >> a[i];
	for (int i = 1, u, v; i < n; i++) cin >> u >> v, g[u].pb(v), g[v].pb(u);
	init(1, 0); dfs(1, 0, 1);
	int lst = 0;
	while (q--) {
		int x, y;
		cin >> x >> y; 
		x ^= lst;
		upd(x, y); lst = query(); cout << lst << "\n";
	} 
 	return 0;
}
2023/9/10 09:51
加载中...