如题,差错无果
测评记录:https://www.luogu.com.cn/record/101452894
#include <bits/stdc++.h>
using namespace std;
typedef long long LL;
const LL N = 2e5+ 5, inf = 1e16;
LL n, q, k, u, v, cnt, lg[N], f[N][20], depth[N], a[N], c[N];
struct matrix {
LL a[3][3];
matrix() { memset(a, 0x3f, sizeof a); }
matrix operator * (const matrix& x) const {
matrix c;
// cout << '\n';
// print(), cout << "*\n", x.print(), cout << "=\n";
for (int i = 0; i < 3; i ++)
for (int k = 0; k < 3; k ++)
for (int j = 0; j < 3; j ++)
c.a[i][j] = min(c.a[i][j], a[i][k] + x.a[k][j]);
// c.print(), cout << '\n';
return c;
}
void print() const {
for (int i = 0; i < 3; i ++) {
for (int j = 0; j < 3; j ++)
cout << a[i][j] << ' ';
cout << '\n';
}
}
} mat[N][20], rev[N][20];
vector<int> G[N];
void dfs(int u, int fa) {
f[u][0] = fa, depth[u] = depth[fa] + 1;
for (int v : G[u]) {
if (v != fa) dfs(v, u);
c[u] = min(c[u], a[v]);
}
}
matrix LCA(int x, int y) {
matrix resx, resy;
bool flag = 0;
for (int i = 0; i < 3; i ++) resx.a[i][i] = resy.a[i][i] = 0;
if (depth[x] < depth[y]) swap(x, y), flag = 1;
while (depth[x] > depth[y]) {
if (!flag) resx = mat[x][lg[depth[x] - depth[y]]] * resx;
else resy = mat[x][lg[depth[x] - depth[y]]];
x = f[x][lg[depth[x] - depth[y]]];
}
if (x == y) return resy * mat[x][0] * resx;
for (int i = 19; i >= 0; i --)
if (f[x][i] != f[y][i]) {
resx = mat[x][i] * resx, resy = resy * rev[y][i];
x = f[x][i], y = f[y][i];
}
return resy * mat[y][0] * mat[f[x][0]][0] * mat[x][0] * resx;
}
LL solve(int u, int v) {
if (u == v) return a[u];
if (depth[u] < depth[v]) swap(u, v);
matrix res;
if (k == 1) res.a[0][0] = a[u], res.a[1][0] = 0, res.a[2][0] = 0;
if (k == 2) res.a[0][0] = a[u], res.a[1][0] = inf, res.a[2][0] = 0;
if (k == 3) res.a[0][0] = a[u], res.a[1][0] = inf, res.a[2][0] = inf;
matrix ans = LCA(f[u][0], v);
return (ans * res).a[0][0];
}
int main() {
memset(c, 0x3f, sizeof c);
cin >> n >> q >> k;
for (int i = 1; i <= n; i ++) cin >> a[i];
for (int i = 2; i <= n; i ++) lg[i] = lg[i >> 1] + 1;
for (int i = 1; i < n; i ++)
cin >> u >> v, G[u].push_back(v), G[v].push_back(u);
dfs(1, 0);
for (int i = 1; i <= n; i ++) {
if (k == 1) {
mat[i][0].a[0][0] = a[i], mat[i][0].a[0][1] = inf, mat[i][0].a[0][2] = inf;
mat[i][0].a[1][0] = inf, mat[i][0].a[1][1] = 0, mat[i][0].a[1][2] = inf;
mat[i][0].a[2][0] = inf, mat[i][0].a[2][1] = inf, mat[i][0].a[2][2] = 0;
}
if (k == 2) {
mat[i][0].a[0][0] = a[i], mat[i][0].a[0][1] = a[i], mat[i][0].a[0][2] = inf;
mat[i][0].a[1][0] = 0, mat[i][0].a[1][1] = inf, mat[i][0].a[1][2] = inf;
mat[i][0].a[2][0] = inf, mat[i][0].a[2][1] = inf, mat[i][0].a[2][2] = 0;
}
if (k == 3) {
mat[i][0].a[0][0] = a[i], mat[i][0].a[0][1] = a[i], mat[i][0].a[0][2] = a[i];
mat[i][0].a[1][0] = 0, mat[i][0].a[1][1] = c[i], mat[i][0].a[1][2] = a[i] + c[i];
mat[i][0].a[2][0] = inf, mat[i][0].a[2][1] = 0, mat[i][0].a[2][2] = inf;
}
rev[i][0] = mat[i][0];
}
for (int j = 1; j <= 19; j ++)
for (int i = 1; i <= n; i ++) {
mat[i][j] = mat[i][j - 1] * mat[f[i][j - 1]][j - 1];
rev[i][j] = rev[f[i][j - 1]][j - 1] * rev[i][j - 1];
f[i][j] = f[f[i][j - 1]][j - 1];
}
for (int i = 1; i <= q; i ++)
cin >> u >> v, cout << solve(u, v) << '\n';
return 0;
}
/*
7 3 2
1 2 3 4 5 6 7
1 2
1 3
2 4
2 5
3 6
3 7
*/