Loading

[2026-2-21] 题解:P10773 [NOISG 2021 Qualification] Truck

矩阵维护树剖线段树

好像没人写线段树维护矩阵,那我来写一篇吧。

  • 优点:不需要脑子
  • 缺点:常数大

题目的操作是:在一个拥有权值 \(D_i\)\(T_i\) 的边上移动时,运送的价值 \(V\) 会先减少 \(T_i\),然后产生 \((V - T_i) \times D_i\) 的费用 \(C\)

线性递推式:

  • \(V_{new} = V - T_i\)
  • \(C_{new} = C + (V - T_i) \cdot D_i = C + V \cdot D_i - T_i \cdot D_i\)

考虑矩阵乘法表示上边的式子。我们构造一个 \(3 \times 1\) 的状态列向量 \(\begin{bmatrix} V \\ C \\ 1 \end{bmatrix}\),为了完成上述转移,我们可以推导出对应的 \(3 \times 3\) 转移矩阵:

\[\begin{bmatrix} 1 & 0 & -T_i \\ D_i & 1 & -T_i \cdot D_i \\ 0 & 0 & 1 \end{bmatrix} \times \begin{bmatrix} V \\ C \\ 1 \end{bmatrix} = \begin{bmatrix} V - T_i \\ C + V \cdot D_i - T_i \cdot D_i \\ 1 \end{bmatrix} \]

由于矩阵乘法满足结合律,我们自然就可以用线段树来维护区间的矩阵乘积了。

树上的路径查询 \((x, y)\) 通常会经过 LCA。

有一点需要注意,矩阵乘法不满足交换律,所以:

  1. \(x\) 往上走到 LCA:对应树剖中 dfn 逆序的矩阵连乘。
  2. 从 LCA 往下走到 \(y\):对应树剖中 dfn 正序的矩阵连乘。

在线段树中,我们需要同时维护两个方向的矩阵乘积:

  • m[k] 表示区间正序乘积(从左到右)。
  • mm[k] 表示区间逆序乘积(从右到左)。

另外,题目要求的是“到达终点时价值刚好为 \(G\)”。我们可以直接在现有的树剖里面顺手维护一下,求出路径上 \(\sum T_i\) 的值,从而反推算出起点的初始价值 \(V_0 = G + \sum T_i\)

常熟太大?

下面两个任选一个都能过:

1、少用模运算,多用减法。

2、像我一样,把矩阵乘法拆开。

struct Matrix {
    int m[4][4];
    void init(int x) {
        for(int i=1; i<=3; i++)
            for(int j=1; j<=3; j++)
                m[i][j] = 0; 
        for(int i=1; i<=3; i++) m[i][i] = x;
    }
};

Matrix operator *(const Matrix &x, const Matrix &y) {
    Matrix res;
    // 把矩阵乘法拆开,太弱智了,还是建议大家优化模运算
    if (y.m[3][3] == 1) {
        res.m[1][1] = res.m[2][2] = res.m[3][3] = 1;
        res.m[1][2] = res.m[3][1] = res.m[3][2] = 0;

        res.m[1][3] = (x.m[1][3] + y.m[1][3]) % P;
        res.m[2][1] = (x.m[2][1] + y.m[2][1]) % P;
        res.m[2][3] = (x.m[2][1] * y.m[1][3] % P + y.m[2][3] + x.m[2][3]) % P;
    } else {
        res.init(0); 
        res.m[3][1] = 1;
        res.m[1][1] = (y.m[1][1] + x.m[1][3]) % P;
        res.m[2][1] = (x.m[2][1] * y.m[1][1] % P + y.m[2][1] + x.m[2][3]) % P;
    }
    return res;
}

完整代码(丑陋)

#include <bits/stdc++.h>
#define int long long
#define endl "\n"

using namespace std;

const int N = 1e5 + 10, P = 1e9 + 7;
int n, g, f[N][30], dep[N], son[N], top[N], siz[N], dfn[N];
int tot, nfd[N], D[N], T[N], tt[N * 8];

struct node1 {
    int v, d, t;
};
vector<node1> e[N];

struct Matrix {
    int m[4][4];
    void init(int x) {
        for (int i = 0; i <= 3; i++)
            for (int j = 0; j <= 3; j++)
                m[i][j] = 0;
        for (int i = 1; i <= 3; i++)
            m[i][i] = x;
    }
} m[N * 8], mm[N * 8];

Matrix operator *(const Matrix &x, const Matrix &y) {
    Matrix res;
    if (y.m[3][3] == 1) {
        res.m[1][1] = res.m[2][2] = res.m[3][3] = 1;
        res.m[1][2] = res.m[3][1] = res.m[3][2] = 0;
        res.m[1][3] = (x.m[1][3] + y.m[1][3]) % P;
        res.m[2][1] = (x.m[2][1] + y.m[2][1]) % P;
        res.m[2][3] = (x.m[2][1] * y.m[1][3] % P + y.m[2][3] + x.m[2][3]) % P;
    } else {
        res.init(0);
        res.m[3][1] = 1;
        res.m[1][1] = (y.m[1][1] + x.m[1][3]) % P;
        res.m[2][1] = (x.m[2][1] * y.m[1][1] % P + y.m[2][1] + x.m[2][3]) % P;
    }
    return res;
}

Matrix operator +(Matrix x, Matrix y) {
    Matrix res;
    res.init(0);
    for (int i = 1; i <= 3; i++)
        for (int j = 1; j <= 3; j++)
            res.m[i][j] = (x.m[i][j] + y.m[i][j]) % P;
    return res;
}

struct node {
    int l, r;
} t[N * 8];

void pushup(int k) {
    m[k] = m[k * 2 + 1] * m[k * 2];
    mm[k] = mm[k * 2] * mm[k * 2 + 1];
    tt[k] = ((tt[k * 2] + tt[k * 2 + 1]) % P + P) % P;
}

void build(int k, int l, int r) {
    m[k].init(1);
    t[k].l = l, t[k].r = r;
    if (l == r) {
        m[k].m[1][3] = ((-T[nfd[l]]) % P + P) % P;
        m[k].m[2][1] = (D[nfd[l]] % P + P) % P;
        m[k].m[2][3] = ((-T[nfd[l]] * D[nfd[l]]) % P + P) % P;
        mm[k] = m[k];
        tt[k] = T[nfd[l]];
        return;
    }
    int mid = (l + r) >> 1;
    build(k * 2, l, mid);
    build(k * 2 + 1, mid + 1, r);
    pushup(k);
}

void modify(int k, int x) {
    if (t[k].l > x || t[k].r < x) return;
    if (t[k].l == t[k].r) {
        m[k].m[1][3] = ((-T[nfd[x]]) % P + P) % P;
        m[k].m[2][1] = (D[nfd[x]] % P + P) % P;
        m[k].m[2][3] = ((-T[nfd[x]] * D[nfd[x]]) % P + P) % P;
        mm[k] = m[k];
        tt[k] = T[nfd[x]];
        return;
    }
    modify(k * 2, x);
    modify(k * 2 + 1, x);
    pushup(k);
}

Matrix query(int k, int l, int r) {
    Matrix res;
    res.init(1);
    if (t[k].l > r || t[k].r < l) return res;
    if (l <= t[k].l && t[k].r <= r) return m[k];
    res = query(k * 2 + 1, l, r) * query(k * 2, l, r);
    return res;
}

Matrix query1(int k, int l, int r) {
    Matrix res;
    res.init(1);
    if (t[k].l > r || t[k].r < l) return res;
    if (l <= t[k].l && t[k].r <= r) return mm[k];
    res = query1(k * 2, l, r) * query1(k * 2 + 1, l, r);
    return res;
}

int querytt(int k, int l, int r) {
    if (t[k].l > r || t[k].r < l) return 0;
    if (l <= t[k].l && t[k].r <= r) return tt[k] % P;
    return ((querytt(k * 2, l, r) + querytt(k * 2 + 1, l, r)) % P + P) % P;
}

void dfs1(int x, int fa) {
    f[x][0] = fa;
    dep[x] = dep[fa] + 1;
    siz[x] = 1;
    for (auto i : e[x]) {
        int v = i.v;
        if (v == fa) continue;
        D[v] = i.d;
        T[v] = i.t;
        dfs1(v, x);
        siz[x] += siz[v];
        if (siz[v] > siz[son[x]]) son[x] = v;
    }
}

void dfs2(int x, int tp) {
    dfn[x] = ++tot;
    nfd[tot] = x;
    top[x] = tp;
    if (son[x]) dfs2(son[x], tp);
    for (auto i : e[x]) {
        int v = i.v;
        if (v == f[x][0] || v == son[x]) continue;
        dfs2(v, v);
    }
}

void init() {
    for (int i = 1; i <= 29; i++)
        for (int j = 1; j <= n; j++)
            f[j][i] = f[f[j][i - 1]][i - 1];
}

int lca(int x, int y) {
    if (dep[x] < dep[y]) swap(x, y);
    for (int i = 29; i >= 0; i--)
        if (dep[f[x][i]] >= dep[y])
            x = f[x][i];
    if (x == y) return x;
    for (int i = 29; i >= 0; i--)
        if (f[x][i] != f[y][i])
            x = f[x][i], y = f[y][i];
    return f[x][0];
}

int jump(int x, int d) {
    for (int i = 29; i >= 0; i--)
        if (dep[f[x][i]] >= d)
            x = f[x][i];
    return x;
}

Matrix queryline1(int x, int ff) {
    Matrix res;
    res.init(1);
    if (x == ff) return res;
    ff = jump(x, dep[ff] + 1);
    while (top[x] != top[ff]) {
        res = res * query(1, dfn[top[x]], dfn[x]);
        x = f[top[x]][0];
    }
    res = res * query(1, dfn[ff], dfn[x]);
    return res;
}

Matrix queryline2(int x, int ff) {
    Matrix res;
    res.init(1);
    if (x == ff) return res;
    ff = jump(x, dep[ff] + 1);
    while (top[x] != top[ff]) {
        res = query1(1, dfn[top[x]], dfn[x]) * res;
        x = f[top[x]][0];
    }
    res = query1(1, dfn[ff], dfn[x]) * res;
    return res;
}

int queryline(int x, int y) {
    int res = 0;
    while (top[x] != top[y]) {
        if (dep[top[x]] < dep[top[y]]) swap(x, y);
        res = ((res + querytt(1, dfn[top[x]], dfn[x])) % P + P) % P;
        x = f[top[x]][0];
    }
    if (dep[x] > dep[y]) swap(x, y);
    res = ((res + querytt(1, dfn[x], dfn[y])) % P + P) % P;
    return res;
}

signed main() {
    ios::sync_with_stdio(false);
    cin.tie(0);
    
    cin >> n >> g;
    for (int i = 1; i <= n - 1; i++) {
        int x, y, d, t;
        cin >> x >> y >> d >> t;
        e[x].push_back({y, d, t});
        e[y].push_back({x, d, t});
    }
    
    dfs1(1, 0);
    dfs2(1, 1);
    init();
    build(1, 1, n);

    int q;
    cin >> q;
    while (q--) {
        int op;
        cin >> op;
        if (op == 0) {
            int x, y, z;
            cin >> x >> y >> z;
            if (f[x][0] == y) swap(x, y);
            T[y] = z % P;
            modify(1, dfn[y]);
        } else {
            int x, y;
            cin >> x >> y;
            Matrix res;
            res.init(0);
            int o = lca(x, y);
            int G = ((g + queryline(x, y) - T[o]) % P + P) % P;
            res.m[1][1] = G;
            res.m[3][1] = 1;
            
            res = queryline1(y, o) * (queryline2(x, o) * res);
            cout << res.m[2][1] % P << endl;
        }
    }
    return 0;
}

posted @ 2026-02-21 10:29  TommyJin  阅读(14)  评论(0)    收藏  举报