[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。
有一点需要注意,矩阵乘法不满足交换律,所以:
- 从 \(x\) 往上走到 LCA:对应树剖中
dfn逆序的矩阵连乘。 - 从 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;
}

浙公网安备 33010602011771号