P3976 旅游 - 洛谷

题意

给定一棵 n 个点的树,每个点有一个权值(宝石价格)。有 q 次操作,每次给出 a,b,v,表示从 a 到 b 的路径上所有点的权值都增加 v,然后输出这条路径上能获得的最大利润。
利润定义为:在路径上按从 a 到 b 的顺序,选择一个点买入,再在之后的一个点卖出,卖价减买价的最大值;如果最大值小于 0,则输出 0。

思路和代码

暴力

第一个暴力的想法就是:取出路径上a到b上所有点的权值,然后扫描维护最小值、最大差值来求出答案
需要取出路径a到b的点,我们需要找到lca,又因为这个题目需要动态修改权值,所以我使用了树链剖分+线段树

具体实现步骤

  • 树链剖分预处理:依然使用树链剖分(HLD),将树映射到 DFS 序上,并实现 get_path(u, v) 函数,该函数能够把路径上的节点按实际行走顺序(从 u 到 v)收集到一个 vector<int> 中。

  • 线段树:这个版本中的线段树非常简单,只维护区间和懒标记,支持两种操作:

    1. 区间加(路径上所有节点加 v)。

    2. 单点查询(查询某个节点当前的权值)。

  • 每次询问的处理流程

    1. 首先对路径 [a, b] 执行区间加 v(更新价格)。

    2. 调用 get_path(a, b),获得按顺序排列的节点编号数组。

    3. 遍历这个数组,对于每个节点 x,通过线段树单点查询(或查询长度为 1 的区间)得到其当前价格 vv

    4. 维护一个变量 mi(遇到的最小价格)和 res(最大差值)。

    5. 对于每个 vvres = max(res, vv-mi),然后 mi = min(mi, vv)

    6. 最后输出 res

具体步骤拆解完了,那么难点就来到了a到b的路径如何取出 不会好难啊,我用倍增好了
ok,我现在来详细讲解一下这段 get_path 函数。

get_path函数

总体思路

我们要从 u 走到 v,需要将路径拆成若干段,每段都是一条重链上的连续区间。树剖标准做法是:

  • top[u] != top[v] 时,深度较深的链顶所在的链段需要先处理,然后将节点上跳到链顶的父亲,重复直到同一条链。
  • 最后处理同一条链上的剩余部分。

但是,标准的树剖 path 操作(例如区间加)并不关心节点顺序,只关心区间覆盖。而这里我们需要节点的顺序,因此必须考虑每段的方向(从上到下还是从下到上)以及拼接顺序。

变量说明
1
vector<int> pathl, pathr;
  • pathl:用于存放从起点 u 出发到 LCA(最近公共祖先)方向上的节点(包括 LCA),按正确的行走顺序累积。
  • pathr:用于存放从 LCA 到终点 v 方向上的节点(不包括 LCA,避免重复),但采用逆序(即最先加入的是靠近 v 的节点,最后加入的是靠近 LCA 的节点)。这样最终可以反转 pathr 得到正确顺序,然后接到 pathl 后面。
循环部分:while(top[u] != top[v])

每次循环,处理 top[u]top[v] 中深度更深的那个链段(因为深度更深意味着离 LCA 更远,应该先处理)。

分支 1:if(dep[top[u]] < dep[top[v]])
  • 含义:top[v]top[u] 更深,说明 v 所在的链顶更深,我们应该先处理 v 侧的链段。
  • 处理内容:
    1
    2
    3
    4
    for(int i = dfn[v]; i >= dfn[top[v]]; --i) {
    pathr.push_back(rnk[i]);
    }
    v = fa[top[v]];
    • 当前链段是从 top[v]v(深度递增),但实际行走方向是从 v 向上走到 top[v](从深到浅)。
    • 由于 DFS 序上 dfn[top[v]]dfn[v] 是顺序(从上到下),而我们需要的方向是倒序(从 vtop[v]),所以循环从 dfn[v] 递减到 dfn[top[v]],这样 pathr 中先加入的是 v,最后加入的是 top[v],符合从 v 向上走的方向。
    • 但是这个链段在最终路径中的位置是在 LCA 到 v 的那一侧,且顺序应该是从 LCA 到 v(从上到下)。由于 pathr 采用逆序,先加入的是靠近 v 的节点,后加入的是靠近 LCA 的。之后我们会整体反转 pathr,就能得到正确的从上到下的顺序。
    • 处理完后,v 跳到链顶的父亲(即进入更上层的链)。
分支 2:else(即 dep[top[u]] >= dep[top[v]]
  • 含义:top[u] 更深或相等,处理 u 侧的链段。
  • 处理内容:
    1
    2
    3
    4
    for(int i = dfn[u]; i >= dfn[top[u]]; --i) {
    pathl.push_back(rnk[i]);
    }
    u = fa[top[u]];
    • 当前链段是从 top[u]u(深度递增),但实际行走方向是从 u 向上走到 top[u](从深到浅)。
    • 同样从 dfn[u] 递减到 dfn[top[u]],这样先加入的是 u,最后加入的是 top[u],符合从 u 向上走的方向。
    • 这个链段在最终路径中位于前半部分(从起点 u 到 LCA),顺序正好是这种“从下到上”的顺序,所以直接放入 pathl 末尾是正确的(因为先处理的是靠近 u 的段,后处理的靠近 LCA 的段,而 pathl 按处理顺序累积,最终会得到 u → ... → LCA 的完整顺序)。
    • 处理完后,u 跳到链顶的父亲。

注意pathlpathr 的处理方式不同:pathl 是直接按正确顺序,pathr 是逆序(先加靠近终点的)。这是因为路径两段的方向不同。

循环结束后的处理

此时 uv 已经在同一条重链上,即 top[u] == top[v]

情况 1:if(dep[u] > dep[v])
  • 含义:uv 的下方,剩余路径是从 u 向上走到 v
  • 处理:
    1
    2
    3
    for(int i = dfn[u]; i >= dfn[v]; --i) {
    pathl.push_back(rnk[i]);
    }
    • uv 是向上走,所以从 dfn[u] 递减到 dfn[v],顺序正确,直接追加到 pathl 末尾。
情况 2:else(即 dep[u] <= dep[v]
  • 含义:vu 的下方或相等,剩余路径是从 u 向下走到 v
  • 处理:
    1
    2
    3
    for(int i = dfn[u]; i <= dfn[v]; ++i) {
    pathl.push_back(rnk[i]);
    }
    • uv 是向下走,DFS 序递增,直接顺序加入 pathl 末尾。

注意:这里没有将剩余部分放入 pathr,而是直接放入 pathl,是因为此时 uv 已经同链,并且如果 uv 上方,这部分属于从 LCA 到 v 的下行段,但实际路径顺序是从 u(此时 u 就是 LCA 或上方)到 v,它紧接在 pathl 之后。如果 uv 下方,则这部分属于从 u 向上到 LCA 的上行段,自然应该追加到 pathl 后面。

最终合并
1
2
3
reverse(pathr.begin(), pathr.end());
pathl.insert(pathl.end(), pathr.begin(), pathr.end());
return pathl;
  • 由于 pathr 中存放的是从 v 到 LCA 的方向(先靠近 v,后靠近 LCA),我们需要反转它,得到从 LCA 到 v 的正确顺序。
  • 然后插入到 pathl 末尾,整个路径就是 u → ... → LCA → ... → v 的顺序,正好是从起点到终点的行走路径。

为什么 pathr 要逆序?
因为我们在处理 v 侧的链段时,是从下往上处理的(先处理靠近 v 的链段,再处理更上层的)。但是最终需要从上往下(LCA → v),所以将这些片段按处理顺序累积是反的,最后整体反转即可。

pathl 则不需要反转,因为我们处理 u 侧时是从下往上,但最终就是从 u 到 LCA 的顺序,所以处理顺序正好匹配。


正确性示例

假设树:1-2-3-4(链),u=4, v=1(向上走)。

  • 初始 top[4]=1, top[1]=1,同链,dep[4]>dep[1],执行 for(i=dfn[4]; i>=dfn[1]; --i),得到 [4,3,2,1],正确。

假设树分叉:u 在左子树,v 在右子树。

  • 处理 u 侧时,从下到上加入 pathl,得到 u → ... → LCA 的上半段。
  • 处理 v 侧时,从下到上加入 pathr,得到 v → ... → LCA 的反序,反转后得到 LCA → ... → v
  • 拼接后得到 u → ... → LCA → ... → v,正确。

最后代码实现

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int maxn=5e4+5;
const ll inf=4e18;
#define ls (u<<1)
#define rs (u<<1 | 1)
struct SegTree{
struct Node{
int l,r;
ll sum,tag;
}tr[maxn<<2];
int len(int u){
return tr[u].r-tr[u].l+1;
}
void push_up(int u){
tr[u].sum=tr[ls].sum+tr[rs].sum;
}
void apply(int u,ll val){
tr[u].sum+=val*len(u);
tr[u].tag+=val;
}
void push_down(int u){
if(!tr[u].tag) return ;
apply(ls,tr[u].tag);
apply(rs,tr[u].tag);
tr[u].tag=0;
}
void build(int u,int l,int r,ll val[]){
tr[u]={l,r,0,0};
if(l==r)
{
tr[u].sum=val[l];
return ;
}
int mid=(l+r)>>1;
build(ls,l,mid,val);
build(rs,mid+1,r,val);
push_up(u);
}
void upd(int u,int l,int r,ll val){
if(l<=tr[u].l&&tr[u].r<=r){
apply(u,val);
return ;
}
push_down(u);
int mid=(tr[u].l+tr[u].r)>>1;
if(l<=mid) upd(ls,l,r,val);
if(mid<r) upd(rs,l,r,val);
push_up(u);
}
ll ask(int u,int l,int r){
if(l<=tr[u].l&&tr[u].r<=r) return tr[u].sum;
push_down(u);
int mid=(tr[u].l+tr[u].r)>>1;
ll res=0;
if(l<=mid) res+=ask(ls,l,r);
if(mid<r) res+=ask(rs,l,r);
return res;
}
}seg;
#undef ls
#undef rs
namespace HLD{
vector<int> adj[maxn];
int fa[maxn],dep[maxn],siz[maxn],son[maxn];
int top[maxn],dfn[maxn],rnk[maxn],tim;
void init(int n){
tim=0;
for(int i=0;i<=n;++i){
adj[i].clear();
fa[i]=dep[i]=siz[i]=0;
son[i]=-1;
top[i]=dfn[i]=rnk[i]=0;
}
}
void add_edge(int u,int v){
adj[v].push_back(u);
adj[u].push_back(v);
}
void dfs1(int u,int p){
fa[u]=p;
dep[u]=(p==-1?0:dep[p]+1);
siz[u]=1,son[u]=-1;
for(int v:adj[u]){
if(v==p) continue;
dfs1(v,u);
siz[u]+=siz[v];
if(son[u]==-1||siz[v]>siz[son[u]]){
son[u]=v;
}
}
}
void dfs2(int u,int t){
top[u]=t,dfn[u]=++tim,rnk[tim]=u;
if(son[u]!=-1) dfs2(son[u],t);
for(int v:adj[u]){
if(v==fa[u] || v==son[u]) continue;
dfs2(v,v);
}
}
void build(int root){
tim=0;
dfs1(root,-1);
dfs2(root,root);
}
template<class F>
void path(int u,int v,F work){
while(top[u]!=top[v]){
if(dep[top[u]]<dep[top[v]]) swap(u,v);
work(dfn[top[u]],dfn[u]);
u=fa[top[u]];
}
if(dep[u]>dep[v]) swap(u,v);
work(dfn[u],dfn[v]);
}
pair<int,int> subtree(int u){
return {dfn[u],dfn[u]+siz[u]-1};
}
vector<int> get_path(int u,int v){
vector<int> pathl,pathr;
while(top[u]!=top[v]){
if(dep[top[u]]<dep[top[v]]) {
for(int i=dfn[v];i>=dfn[top[v]];--i){
pathr.push_back(rnk[i]);
}
v=fa[top[v]];
}else{
for(int i=dfn[u];i>=dfn[top[u]];--i){
pathl.push_back(rnk[i]);
}
u=fa[top[u]];
}
}
if(dep[u]>dep[v]){
for(int i=dfn[u];i>=dfn[v];--i) {
pathl.push_back(rnk[i]);
}
}
else{
for(int i=dfn[u];i<=dfn[v];++i){
pathl.push_back(rnk[i]);
}
}
reverse(pathr.begin(),pathr.end());
pathl.insert(pathl.end(),pathr.begin(),pathr.end());
return pathl;
}
}
ll val[maxn],seq[maxn];
void upd_path(int u,int v,ll x){
HLD::path(u,v,[&](int l,int r){
seg.upd(1,l,r,x);
});
}
int n,m;
void solve()
{
cin>>n;
for(int i=1;i<=n;++i) cin>>val[i];
for(int i=1;i<n;++i){
int u,v;
cin>>u>>v;
HLD::add_edge(u,v);
}
HLD::build(1);
for(int u=1;u<=n;++u){
seq[HLD::dfn[u]]=val[u];
}
seg.build(1,1,n,seq);
cin>>m;
while(m--){
int u,v;ll w;
cin>>u>>v>>w;
upd_path(u,v,w);
vector<int> ans=HLD::get_path(u,v);
ll mi=inf,res=0;
for(int i=0;i<(int)ans.size();++i){
ll vv=seg.ask(1,HLD::dfn[ans[i]],HLD::dfn[ans[i]]);
res=max(vv-mi,res);
mi=min(vv,mi);
}
cout<<res<<"\n";
}
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int ___T=1;
//cin>>___T;
while(___T--) solve();
return 0;
}

复杂度

  • 优点:思路极其简单,非常符合直觉。

  • 缺点:每次询问,get_path 和遍历数组的时间复杂度都是 O(路径长度),而树上的路径长度最坏是 O(n)。单次查询就需要 O(n) 的时间,总复杂度 O(nq)。交上去30分

什么!?看了你题解半天,结果是T??
别急别急!来看我的第二版!!

正解

把每个点都取出来在遍历时间复杂度太高了,我们既然已经使用了线段树,为什么不直接在线段树里面维护答案呢?
所以说这一版不再使用 get_path 提取所有节点,而是将路径分解为区间后直接利用线段树的区间信息进行合并,避免了遍历每个节点。但路径分解逻辑与 get_path 非常相似,都使用 top 跳转,只是合并的内容变成了 Info 结构体,并且需要处理方向翻转。所以说!理解这个暴力版的 get_path 非常有助于理解这版的方向处理。方向处理也是本题最大的难点。

思路详解

一、线段树节点信息

在线段树中,每个节点(对应一段连续的 DFS 序区间)维护一个 Info,包含:

1
2
3
4
struct Info {
ll mx, mi; // 区间最大值、最小值
ll lmx, rmx; // 从左到右(lmx)和从右到左(rmx)的最大差价
};
  • mi:区间最小值。
  • mx:区间最大值。
  • lmx:在区间内按从左到右的顺序,卖出价 − 买入价的最大值(可以为负,但最终答案取 max(0, lmx))。
  • rmx:在区间内按从右到左的顺序,卖出价 − 买入价的最大值。

对于叶子节点(单个元素 (x)),mi = mx = xlmx = rmx = 0(无法交易)。

二:区间合并(merge 函数)

合并左区间 (A) 和右区间 (B)(顺序为 (A) 在前,(B) 在后)得到 (C):

1
2
3
4
5
6
7
8
Info merge(const Info &a, const Info &b) {
Info res;
res.mi = min(a.mi, b.mi);
res.mx = max(a.mx, b.mx);
res.lmx = max(max(a.lmx, b.lmx), b.mx - a.mi);
res.rmx = max(max(a.rmx, b.rmx), a.mx - b.mi);
return res;
}
  • lmx 的第三项 b.mx - a.mi 表示:在左区间买入(选最小值),在右区间卖出(选最大值)——符合从左到右的顺序。
  • rmx 的第三项 a.mx - b.mi 表示:在右区间买入(选最小值),在左区间卖出(选最大值)——符合从右到左的顺序。

这种合并满足结合律,所以我们可以将多个区间按顺序合并,得到整体信息。

三、区间整体加与懒标记

当整个区间每个点都增加 (v) 时:

  • mimx 分别增加 (v)。
  • lmxrmx 不变(因为差值不变)。

因此在线段树的 apply 函数中:

1
2
3
4
5
void apply(int u, ll val) {
tr[u].info.mx += val;
tr[u].info.mi += val;
tr[u].tag += val;
}

并且维护 tag 用于下传。

四、路径加操作(upd_path

利用 HLD::path 模板函数,将路径 分解成若干连续区间,每个区间调用线段树的 upd 进行区间加:

1
2
3
4
5
void upd_path(int u, int v, ll x) {
HLD::path(u, v, [&](int l, int r) {
seg.upd(1, l, r, x);
});
}

HLD::path 的标准实现保证了覆盖路径上的所有点且不重复。

五、查询路径最大利润(get_ans

我们需要得到路径上按顺序的 lmx。但树剖分解得到的是多个区间,它们的方向可能与路径方向不一致。因此我们采用两个累积变量:

  • resl:存放从起点 u 向上到 lca 的部分,按正确顺序累积。
  • resr:存放从 lca 向下到终点 v 的部分,但采用逆序(即将新得到的段放在 resr 的前面),以便最后合并。

初始化:resl = resr = {-inf, inf, 0, 0}(空信息)。

循环跳转阶段(top[u] != top[v]
1
2
3
4
5
6
7
8
9
10
11
12
13
14
while (top[u] != top[v]) {
if (dep[top[u]] < dep[top[v]]) {
// 处理 v 侧
SegTree::Info a = seg.ask(1, dfn[top[v]], dfn[v]);
resr = seg.merge(a, resr);
v = fa[top[v]];
} else {
// 处理 u 侧
SegTree::Info a = seg.ask(1, dfn[top[u]], dfn[u]);
swap(a.lmx, a.rmx); // 翻转方向
resl = seg.merge(resl, a);
u = fa[top[u]];
}
}
  • dep[top[u]] < dep[top[v]] 时,说明 top[v] 更深,先处理 v 所在的链段。

    • 该链段在 DFS 序上是 top[v] → v(从上到下),而实际路径方向是从 v 向上走到 top[v](从下到上),但最终这部分在整条路径中位于 LCA 到 v 的方向,即从上到下
    • 我们查询得到的信息是 top[v] → v 的顺序,正好是最终需要的顺序(从上到下),所以不需要翻转,直接作为整体放在 resr前面(因为 resr 是逆序累积,先处理的靠近 v,后处理的靠近 LCA,最终 resr 会通过多次 merge(a, resr) 变成 LCA → ... → v 的正确顺序)。
    • 然后 v 跳到链顶的父亲,继续处理上层。
  • 否则(dep[top[u]] >= dep[top[v]],处理 u 侧的链段。

    • 该链段在 DFS 序上是 top[u] → u(从上到下),但实际路径方向是 u 向上走到 top[u](从下到上),而最终这部分位于从起点 u 到 LCA 的方向,即从下到上
    • 查询得到的是 top[u] → u 的顺序,我们需要的是 u → top[u],因此必须翻转方向,即交换 lmxrmx
    • 翻转后的信息按正确顺序(从 utop[u])合并到 resl 的末尾。
    • 然后 u 跳到链顶的父亲。
同链剩余部分

循环结束后,uv 在同一条重链上。

1
2
3
4
5
6
7
8
if (dep[u] > dep[v]) {
SegTree::Info a = seg.ask(1, dfn[v], dfn[u]);
swap(a.lmx, a.rmx); // 方向:u -> v(向上),翻转
resl = seg.merge(resl, a);
} else {
SegTree::Info a = seg.ask(1, dfn[u], dfn[v]);
resr = seg.merge(a, resr); // 方向:u -> v(向下),不翻转
}
  • 如果 dep[u] > dep[v],剩余路径从 u 向上到 v,方向是从下到上,因此查询 [dfn[v], dfn[u]](顺序 v → u),需要翻转后合并到 resl
  • 否则(dep[u] <= dep[v]),剩余路径从 u 向下到 v,方向是从上到下,查询 [dfn[u], dfn[v]](顺序 u → v),直接合并到 resr 前面。
最终合并
1
return seg.merge(resl, resr).lmx;

resl 已正确累积了 u → ... → LCAresr 已正确累积了 LCA → ... → v(因为逆序累积,多次 merge(a, resr) 使得最终顺序是从上到下)。合并两者得到整个路径的 Info,其 lmx 即为最大利润(若为负,输出 0)。


为什么 resr 采用逆序累积?

假设 LCA 下方有两条链段需要合并,从下往上处理时,先遇到靠近 v 的段 (S_1),再遇到靠近 LCA 的段 (S_2)。我们希望最终顺序是 (S_2) 在前,(S_1) 在后。逆序累积通过 resr = merge(S_new, resr) 实现:第一次 resr = merge(S1, empty) = S1,第二次 resr = merge(S2, S1),正好是 (S2) 在前,(S1) 在后,顺序正确。


最后代码实现

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int maxn=5e4+5;
const ll inf=4e18;
#define ls (u<<1)
#define rs (u<<1 | 1)
struct SegTree{
struct Info{
ll mx,mi;
ll lmx,rmx;
};
struct Node{
int l,r;
ll tag;
Info info;
}tr[maxn<<2];

Info merge(const Info &a,const Info &b){
Info res;
res.mi=min(a.mi,b.mi);
res.mx=max(a.mx,b.mx);
res.lmx=max(max(a.lmx,b.lmx),b.mx-a.mi);
res.rmx=max(max(a.rmx,b.rmx),a.mx-b.mi);
return res;
}
int len(int u){
return tr[u].r-tr[u].l+1;
}
void push_up(int u){
tr[u].info=merge(tr[ls].info,tr[rs].info);
}
void apply(int u,ll val){
tr[u].info.mx+=val;
tr[u].info.mi+=val;
tr[u].tag+=val;
}
void push_down(int u){
if(!tr[u].tag) return ;
apply(ls,tr[u].tag);
apply(rs,tr[u].tag);
tr[u].tag=0;
}
void build(int u,int l,int r,ll val[]){
tr[u].l=l,tr[u].r=r,tr[u].tag=0;
if(l==r)
{
tr[u].info={val[l],val[l],0,0};
return ;
}
int mid=(l+r)>>1;
build(ls,l,mid,val);
build(rs,mid+1,r,val);
push_up(u);
}
void upd(int u,int l,int r,ll val){
if(l<=tr[u].l&&tr[u].r<=r){
apply(u,val);
return ;
}
push_down(u);
int mid=(tr[u].l+tr[u].r)>>1;
if(l<=mid) upd(ls,l,r,val);
if(mid<r) upd(rs,l,r,val);
push_up(u);
}
Info ask(int u,int l,int r){

if(l<=tr[u].l&&tr[u].r<=r) return tr[u].info;
push_down(u);
int mid=(tr[u].l+tr[u].r)>>1;
if(r<=mid) return ask(ls,l,r);
if(mid<l) return ask(rs,l,r);
return merge(ask(ls,l,r),ask(rs,l,r));
}
}seg;
#undef ls
#undef rs
namespace HLD{
vector<int> adj[maxn];
int fa[maxn],dep[maxn],siz[maxn],son[maxn];
int top[maxn],dfn[maxn],rnk[maxn],tim;
void init(int n){
tim=0;
for(int i=0;i<=n;++i){
adj[i].clear();
fa[i]=dep[i]=siz[i]=0;
son[i]=-1;
top[i]=dfn[i]=rnk[i]=0;
}
}
void add_edge(int u,int v){
adj[v].push_back(u);
adj[u].push_back(v);
}
void dfs1(int u,int p){
fa[u]=p;
dep[u]=(p==-1?0:dep[p]+1);
siz[u]=1,son[u]=-1;
for(int v:adj[u]){
if(v==p) continue;
dfs1(v,u);
siz[u]+=siz[v];
if(son[u]==-1||siz[v]>siz[son[u]]){
son[u]=v;
}
}
}
void dfs2(int u,int t){
top[u]=t,dfn[u]=++tim,rnk[tim]=u;
if(son[u]!=-1) dfs2(son[u],t);
for(int v:adj[u]){
if(v==fa[u] || v==son[u]) continue;
dfs2(v,v);
}
}
void build(int root){
tim=0;
dfs1(root,-1);
dfs2(root,root);
}
template<class F>
void path(int u,int v,F work){
while(top[u]!=top[v]){
if(dep[top[u]]<dep[top[v]]) swap(u,v);
work(dfn[top[u]],dfn[u]);
u=fa[top[u]];
}
if(dep[u]>dep[v]) swap(u,v);
work(dfn[u],dfn[v]);
}
pair<int,int> subtree(int u){
return {dfn[u],dfn[u]+siz[u]-1};
}
ll get_ans(int u,int v){
SegTree::Info resl={-inf,inf,0,0};
SegTree::Info resr={-inf,inf,0,0};
while(top[u]!=top[v]){
if(dep[top[u]]<dep[top[v]]){
SegTree::Info a=seg.ask(1,dfn[top[v]],dfn[v]);
resr=seg.merge(a,resr);
v=fa[top[v]];
}
else {
SegTree::Info a=seg.ask(1,dfn[top[u]],dfn[u]);
swap(a.lmx,a.rmx);
resl=seg.merge(resl,a);
u=fa[top[u]];
}
}
if(dep[u]>dep[v]){
SegTree::Info a=seg.ask(1,dfn[v],dfn[u]);
swap(a.lmx,a.rmx);
resl=seg.merge(resl,a);
}
else resr=seg.merge(seg.ask(1,dfn[u],dfn[v]),resr);
return seg.merge(resl,resr).lmx;
}
}
ll val[maxn],seq[maxn];
void upd_path(int u,int v,ll x){
HLD::path(u,v,[&](int l,int r){
seg.upd(1,l,r,x);
});
}
int n,m;
void solve()
{
cin>>n;
for(int i=1;i<=n;++i) cin>>val[i];
for(int i=1;i<n;++i){
int u,v;
cin>>u>>v;
HLD::add_edge(u,v);
}
HLD::build(1);
for(int u=1;u<=n;++u){
seq[HLD::dfn[u]]=val[u];
}
seg.build(1,1,n,seq);
cin>>m;
while(m--){
int u,v;ll w;
cin>>u>>v>>w;
upd_path(u,v,w);
cout<<HLD::get_ans(u,v)<<"\n";
}
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int ___T=1;
//cin>>___T;
while(___T--) solve();
return 0;
}

复杂度

  • 树链剖分预处理:(O(n))。
  • 线段树建树:(O(n))。
  • 每次操作(路径加 + 查询):
    • upd_path 执行树剖路径分解,每条链调用 seg.upd 区间加,复杂度 (O(\log^2 n))。
    • get_ans 同样分解路径,每条链调用 seg.ask 区间查询,复杂度 (O(\log^2 n))。
  • 总复杂度 (O((n+q)\log^2 n)),足以通过。

区别总结

维度 第一版(暴力) 第二版(正解)
信息粒度 每个点单独处理 维护区间整体信息(最值、双向差值)
路径合并方式 提取所有点,线性扫描 分解为区间,合并 Info 结构体
单次时间复杂度 (O(\text{路径长度})) (O(\log^2 n))
是否需要单点查询 需要,用于获得每个点当前权值 不需要,直接利用区间信息合并
实现复杂度 中等,需仔细处理方向

后记

什么?写完感觉还想写题目练练手?
P2486 染色 - 洛谷
写!这里就不一步一步写思路了,和前面基本上一样,难点主要也是merge函数和路径方向

代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int maxn=1e5+5,inf=1e9;
#define ls (u<<1)
#define rs (u<<1 | 1)
struct SegTree{
struct Info{
ll lc,rc,sum;
};
struct Node{
int l,r;
int tag;
Info info;
}tr[maxn<<2];
Info merge(Info a,Info b){
Info res;
if(a.sum==0) res=b;
else if(b.sum==0) res=a;
else {
res={a.lc,b.rc,a.sum+b.sum};
if(a.rc==b.lc) res.sum--;
}
return res;
}
int len(int u){
return tr[u].r-tr[u].l+1;
}
void push_up(int u){
tr[u].info=merge(tr[ls].info,tr[rs].info);
}
void apply(int u,ll val){
tr[u].tag=val;
tr[u].info={val,val,1};
}
void push_down(int u){
if(!tr[u].tag) return ;
apply(ls,tr[u].tag);
apply(rs,tr[u].tag);
tr[u].tag=0;
}
void build(int u,int l,int r ,ll val[]){
tr[u]={l,r,0};
if(l==r){
tr[u].info={val[l],val[l],1};
return ;
}
int mid=(l+r)>>1;
build(ls,l,mid,val);
build(rs,mid+1,r,val);
push_up(u);
}
void upd(int u,int l,int r,ll val){
if(l<=tr[u].l&&tr[u].r<=r){
apply(u,val);
return ;
}
push_down(u);
int mid=(tr[u].l+tr[u].r)>>1;
if(l<=mid) upd(ls,l,r,val);
if(mid<r) upd(rs,l,r,val);
push_up(u);
}
Info ask(int u,int l,int r){
if(l<=tr[u].l&&tr[u].r<=r) return tr[u].info;
push_down(u);
int mid=(tr[u].l+tr[u].r)>>1;
Info res={0,0,0};
if(l<=mid) res=merge(res,ask(ls,l,r));
if(mid<r) res=merge(res,ask(rs,l,r));
return res;
}
}seg;
#undef ls
#undef rs
namespace HLD{
vector<int> adj[maxn];
int up[maxn],dep[maxn],siz[maxn],son[maxn];
int top[maxn],dfn[maxn],rnk[maxn],tim;
void init(int n){
tim=0;
for(int i=0;i<=n;++i){
adj[i].clear();
up[i]=dep[i]=siz[i]=0;
son[i]=-1;
top[i]=dfn[i]=rnk[i]=0;
}
}
void dfs1(int u,int p){
up[u]=p;
dep[u]=(p==-1?0:dep[p]+1);
siz[u]=1,son[u]=-1;
for(int v:adj[u]){
if(v==p) continue;
dfs1(v,u);
siz[u]+=siz[v];
if(son[u]==-1||siz[v]>siz[son[u]]){
son[u]=v;
}
}
}
void dfs2(int u,int t){
top[u]=t,dfn[u]=++tim,rnk[tim]=u;
if(son[u]!=-1) dfs2(son[u],t);
for(int v:adj[u]){
if(v==up[u]||v==son[u]) continue;
dfs2(v,v);
}
}
void build(int root){
tim=0;
dfs1(root,-1);
dfs2(root,root);
}
int get_lca(int u,int v){
while(top[u]!=top[v]){
if(dep[top[u]]<dep[top[v]]) swap(u,v);
u=up[top[u]];
}
return dep[u]<dep[v]?u:v;
}
ll ask_path(int u,int v){
SegTree::Info resl={0,0,0},resr={0,0,0};
while(top[u]!=top[v]){
if(dep[top[u]]<dep[top[v]]){
SegTree::Info aa=seg.ask(1,dfn[top[v]],dfn[v]);
resr=seg.merge(aa,resr);
v=up[top[v]];
} else{
SegTree::Info aa=seg.ask(1,dfn[top[u]],dfn[u]);
swap(aa.lc,aa.rc);
resl=seg.merge(resl,aa);
u=up[top[u]];
}
}
if(dfn[u]>dfn[v]){
SegTree::Info aa=seg.ask(1,dfn[v],dfn[u]);
swap(aa.lc,aa.rc);
resl=seg.merge(resl,aa);
} else{
SegTree::Info aa=seg.ask(1,dfn[u],dfn[v]);
resl=seg.merge(resl,aa);
}
return seg.merge(resl,resr).sum;

}
template<class F>
void path( int u, int v, F work ) {
while( top[u] != top[v] ) {
if( dep[top[u]] < dep[top[v]] ) swap( u, v );
work( dfn[top[u]], dfn[u] );
u = up[top[u]];
}
if( dep[u] > dep[v] ) swap( u, v );
work( dfn[u], dfn[v] );
}
}
void upd_path(int u,int v,ll x){
HLD::path(u,v,[&](int l,int r){
seg.upd(1,l,r,x);
});
}
ll a[maxn],seq[maxn];
int n,m;
void solve()
{
cin>>n>>m;
for(int i=1;i<=n;++i) cin>>a[i];
for(int i=1;i<n;++i){
int u,v;
cin>>u>>v;
HLD::adj[u].push_back(v);
HLD::adj[v].push_back(u);
}
HLD::build(1);
for(int u=1;u<=n;++u){
seq[HLD::dfn[u]]=a[u];
}
seg.build(1,1,n,seq);
while(m--){
char op;
int a,b;ll c;
cin>>op>>a>>b;
if(op=='Q'){
cout<<HLD::ask_path(a,b)<<"\n";
}else {
cin>>c;
upd_path(a,b,c);
}
}
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int ___T=1;
//cin>>___T;
while(___T--) solve();
return 0;
}