跟暴力一个分。。。也只过了暴力能过的点。。。。。。无语了
#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;
}