萌新刚学dp求助 44分
查看原帖
萌新刚学dp求助 44分
503792
Svemit楼主2023/9/20 16:25

跟暴力一个分。。。也只过了暴力能过的点。。。。。。无语了

#include <bits/stdc++.h>
using namespace std;
typedef long long LL;
const int N = 2e5 + 5, INF = 0x3f3f3f3f;
const LL mod = 1e9 + 7;
int n, m;
LL p[N];
string type;
vector<int> e[N];
int fa[N][20];
LL f[N][2], g[N][2], w[N][20][2][2];
int dep[N];
void dfs1(int u, int fath) {
	fa[u][0] = fath;
	dep[u] = dep[fath] + 1;
	for(int i = 1; i <= 17; i ++) 
		fa[u][i] = fa[fa[u][i - 1]][i - 1];
	f[u][1] = p[u];
	for(auto v : e[u]) if(v != fath) {
		dfs1(v, u);
		f[u][0] += f[v][1];
		f[u][1] += min(f[v][0], f[v][1]);
	}
}
void dfs2(int u, int fath) {
	for(auto v : e[u]) if(v != fath) {
		w[v][0][0][1] = w[v][0][1][1] = f[u][1] - min(f[v][0], f[v][1]);
		w[v][0][1][0] = f[u][0] - f[v][1];
		for(int i = 1; i <= 17; i ++)
			for(int x = 0; x < 2; x ++)
				for(int y = 0; y < 2; y ++)
					for(int z = 0; z < 2; z ++)
						w[v][i][x][y] = min(w[v][i][x][y], w[v][i - 1][x][z] + w[fa[v][i - 1]][i - 1][z][y]);
		g[v][0] = g[u][1] + f[u][1] - min(f[v][0], f[v][1]);
		g[v][1] = min(g[v][0], g[u][0] + f[u][0] - f[v][1]);
		dfs2(v, u);
	}
}
LL solve(int u, int x, int v, int y) {
	if(dep[u] < dep[v]) swap(u, v), swap(x, y);
	if(!x && !y && fa[u][0] == v) return -1;
	array<LL, 2> su, sv, gu, gv;
	su[0] = su[1] = sv[0] = sv[1] = INF;
	su[x] = f[u][x], sv[y] = f[v][y];
	for(int i = 17; ~i; i --)
		if(dep[fa[u][i]] >= dep[v]) {
			gu[0] = gu[1] = INF;
			for(int j = 0; j < 2; j ++)
				for(int k = 0; k < 2; k ++)
					gu[k] = min(gu[k], su[j] + w[u][i][j][k]);
			su[0] = gu[0], su[1] = gu[1];
			u = fa[u][i];
		}
	if(u == v) return su[y] + g[v][y];
	for(int i = 17; ~i; i --)
		if(fa[u][i] != fa[v][i]) {
			gu[0] = gu[1] = gv[0] = gv[1] = INF;
			for(int j = 0; j < 2; j ++)
				for(int k = 0; k < 2; k ++) {
					gu[k] = min(gu[k], su[j] + w[u][i][j][k]);
					gv[k] = min(gv[k], sv[j] + w[v][i][j][k]);
				}
			su[0] = gu[0], su[1] = gu[1];
			sv[0] = gv[0], sv[1] = gv[1];
			u = fa[u][i], v = fa[v][i];
		}
	int lca = fa[u][0];
	LL res1 = f[lca][0] + g[lca][0] - f[u][1] - f[v][1] + su[1] + sv[1];
	LL res2 = f[lca][1] + g[lca][1] - min(f[u][0], f[u][1]) - min(f[v][0], f[v][1]) + min(su[0], su[1]) + min(sv[0], sv[1]);
	return min(res1, res2);
}
int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
	cin >> n >> m >> type;
	for(int i = 1; i <= n; i ++) {
		cin >> p[i];
	}
	for(int i = 1; i < n; i ++) {
		int u, v;
		cin >> u >> v;
		e[u].push_back(v);
		e[v].push_back(u);
	}
	memset(w, 0x3f, sizeof w);
	dfs1(1, 0), dfs2(1, 0);
	while(m --) {
		int u, x, v, y;
		cin >> u >> x >> v >> y;
		cout << solve(u, x, v, y) << '\n';
	}
    return 0;
}
2023/9/20 16:25
加载中...