P1600 天天爱跑步 - 洛谷

题意

给定一棵 $n$ 个节点的树,有 $m$ 个玩家在树上跑步。每个玩家从起点 $S_i$ 沿着最短路径跑到终点 $T_i$,每秒跑过一条边。每个节点 $u$ 上有一个观察员,他只会在特定的时间 $w_u$ 进行观察。问每个观察员能看到多少个刚好在那一秒跑到该节点的玩家。

思路

如果我们遍历每一个玩家到每一个的时间,那必然超时,比如换个思路,从观察员入手,看看那个玩家可以到达的时间刚刚好为 $w_u$

本题的核心是:将任意一条树上路径 $S \to T$ 拆分为两段:上行路径 $S \to LCA$ 和 下行路径 $LCA \to T$。观察员位于节点 $u$,如果他能看到玩家,说明玩家到达 $u$ 的时间刚好为 $w_u$。

我们按上行和下行分类推导条件公式:

  1. 上行路径 ($S \to LCA$):

    点 $u$ 在 $S$ 到 $LCA$ 的路径上。玩家从 $S$ 走到 $u$ 的时间是 $dep[S] - dep[u]$。

    满足条件的等式:$dep[S] - dep[u] = w[u] \implies dep[S] = w[u] + dep[u]$

  2. 下行路径 ($LCA \to T$):

    点 $u$ 在 $LCA$ 到 $T$ 的路径上。玩家从 $S$ 走到 $u$ 的时间是 $S$ 到 $LCA$ 的距离加上 $LCA$ 到 $u$ 的距离,即 $(dep[S] - dep[LCA]) + (dep[u] - dep[LCA])$。

    满足条件的等式:$dep[S] + dep[u] - 2 \cdot dep[LCA] = w[u] \implies dep[S] - 2 \cdot dep[LCA] = w[u] - dep[u]$

推导出了这两个等式后,怎么快速统计答案呢?

既然等式的一边完全是和路径相关,另一边是只和当前$u$有关,想快速查有多少个值相等,最容易想到的自然是开个数组当桶来计数。

但这里有个麻烦点:这是一棵树,如果我们直接无脑往全局桶里扔数据,我们在算节点 $u$ 的时候,肯定会把别的树枝上八竿子打不着的路径也给算进去。为了只统计经过 $u$ 本身的路径,就得把DFS 遍历差分结合起来,玩一个小小的“障眼法”:

  • 进门先记账: 当我们刚 DFS 走到节点 $u$,还没进去搜它底下的节点时,先瞅一眼当前桶里需要匹配的位置(比如上行的 w[u] + dep[u])记下的数是多少,存作一个旧的值 $t_1$。

  • 放手去遍历: 然后正常去搜 $u$ 的所有子树。在遍历子树的过程中,只要遇到路径的起点,就老老实实往桶里 +1

  • 出门算差价: 等 $u$ 的所有子树都搜完了,准备往上回溯退出 $u$ 时,再看一眼桶里现在的值。此时桶里新多出来的数值(当前值减去快照 $t_1$),绝对且仅仅是 $u$ 的子树里刚塞进去的!这就完美剔除了树上其他分支的干扰。

利用进出子树前后的“时间差”,从一个全局杂乱的桶里,精准扣出了只属于当前节点的贡献。

分块代码

我结合我的代码来解释一下为什么这么写。

为了分别统计上行和下行,我开了两个桶:cnt[0] 处理上行,cnt[1] 处理下行。由于下行等式左侧 $dep[S] - 2 \cdot dep[LCA]$ 可能是负数,我们要统一加上了偏移量 maxn * 2 防止数组越界。

拆分路径与端点记录

对于输入的每一条路径,先求出 LCA。记录起点数 ss[s],下行终点 ulj[t],并将下行的特征值算好存入 path 结构体:

1
2
3
ss[s]++;
path.push_back({s, t, lca, dep[s] - 2 * dep[lca]});
ulj[t].push_back(i);

关键去重: 如果某条路径在 $LCA$ 处既满足上行条件,又满足下行条件(即 $dep[S] - dep[LCA] = w[LCA]$),它会在两个桶里各贡献一次,导致 $LCA$ 的答案多算 $1$。所以在读入时直接做前置抵消:

1
if(dep[lca] + w[lca] == dep[s]) ans[lca]--;

全局桶的树上差分统计

进入节点 $u$ 时,先记录桶里对应需要查询的值 t1t2(此时桶里的数据不包含 $u$ 的子树):

1
2
int t1 = cnt[0][w[u] + dep[u]];
int t2 = cnt[1][w[u] - dep[u] + maxn * 2];

递归遍历完所有的子树之后,把以当前节点为起点或终点的路径扔进桶里:

1
2
3
4
5
cnt[0][dep[u]] += ss[u];
for(int id : ulj[u]) {
int val = path[id].dis + maxn * 2;
cnt[1][val]++;
}

此时桶内增加的数量,就是 $u$ 的子树对 $u$ 产生的有效路径数!所以直接用当前桶的值减去之前记录的 t1t2 即可:

1
ans[u] += (cnt[0][w[u] + dep[u]] - t1) + (cnt[1][w[u] - dep[u] + maxn * 2] - t2);

撤销 LCA 处的贡献

树上差分最重要的一步是:路径在 $LCA$ 处就拐弯或者停止了,它不应该继续向上对其祖先产生贡献。

全局桶里存的本质上是“玩家的数量”。一条特定的路径就代表着一个在树上跑步的玩家。当程序的 DFS 走到玩家路径的下端点时,我们往桶里投入了这 $1$ 个玩家的特征数据(即执行了 +1)。既然当初只放进去了 $1$ 个标记,当回溯到最高点(LCA)需要作废这条路线时,自然也只需要从存钱罐里把这 $1$ 个标记拿出来即可,所以是减一。

因为我的逻辑顺序是 先计算 ans[u],后撤销,这意味着 $LCA$ 节点在撤销前,已经把这条路径统计进自己的答案里了。所以无论上行还是下行,撤销动作必须统一挂在 lca

1
2
3
// 挂载撤销任务
lcadown[lca].push_back(i);
lcaup[lca].push_back(i);

在 DFS 的最后,执行减操作,保证 $u$ 的父节点回溯时不会看到这些已经结束的路径:

1
2
for(int id : lcaup[u]) cnt[0][dep[path[id].s]]--;
for(int id : lcadown[u]) cnt[1][path[id].dis + maxn * 2]--;

完整代码

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
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int maxn=3e5+5;
struct lj{ int s,t,lca; ll dis;};
vector<lj> path;
vector<int> adj[maxn];
ll w[maxn],ans[maxn];
ll cnt[2][maxn<<2];//记录数量
int ss[maxn];//起点的数量
vector<int> ulj[maxn];//以u为终点的路径的id
vector<int> lcaup[maxn];//在lca处需撤销上面贡献的路径的id
vector<int> lcadown[maxn];//在lca点处需撤销下面贡献的路径的id
int up[maxn][25],dep[maxn];
void dfs(int u,int p){
dep[u]=dep[p]+1;
up[u][0]=p;
for(int i=1;i<20;++i){
up[u][i]=up[up[u][i-1]][i-1];
}
for(int v:adj[u]){
if(v==p) continue;
dfs(v,u);
}
}
int getlca(int u,int v){
if(dep[u]<dep[v]) swap(v,u);
int diff(dep[u]-dep[v]);
for(int i=19;i>=0;--i){
if((diff>>i)&1) u=up[u][i];
}
if(u==v) return u;
for(int i=19;i>=0;--i){
if(up[u][i]==up[v][i]) continue;
u=up[u][i],v=up[v][i];
}
return up[u][0];
}
void dfs2(int u,int p){
int t1=cnt[0][w[u]+dep[u]];
int t2=cnt[1][w[u]-dep[u]+maxn*2];
for(int v:adj[u]){
if(v==p) continue;
dfs2(v,u);
}
cnt[0][dep[u]]+=ss[u];
for(int id:ulj[u]){
int val=path[id].dis+maxn*2;
cnt[1][val]++;
}
ans[u]+=(cnt[0][w[u]+dep[u]]-t1)+(cnt[1][w[u]-dep[u]+maxn*2]-t2);
for(int id:lcaup[u]){
cnt[0][dep[path[id].s]]--;
}
for(int id:lcadown[u]){
int val=path[id].dis+maxn*2;
cnt[1][val]--;
}
}
int n,m;
void solve()
{
cin>>n>>m;
for(int i=1;i<n;++i){
int u,v;
cin>>u>>v;
adj[u].push_back(v);
adj[v].push_back(u);
}
dep[1]=1;
dfs(1,0);
for(int i=1;i<=n;++i) cin>>w[i];
for(int i=0;i<m;++i){
int s,t;
cin>>s>>t;
int lca=getlca(s,t);
ss[s]++;
path.push_back({s,t,lca,dep[s]-2*dep[lca]});
ulj[t].push_back(i);
lcadown[lca].push_back(i);
lcaup[lca].push_back(i);
if(dep[lca]+w[lca]==dep[s]) ans[lca]--;
}
dfs2(1,0);
for(int i=1;i<=n;++i){
cout<<ans[i]<<" ";
}
cout<<"\n";
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int ___T=1;
//cin>>___T;
while(___T--) solve();
return 0;
}