萌新求助 20pts
查看原帖
萌新求助 20pts
560516
喵仔牛奶楼主2023/2/4 19:27

如题,差错无果

测评记录: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
*/
2023/2/4 19:27
加载中...