P3953 逛公园 - 洛谷

题意

给定一张 $N$ 个点 $M$ 条边的带权有向图(边权非负),起点为 $1$,终点为 $N$。设 $1 \to N$ 的最短路长度为 $d$。求总长度不超过 $d + K$ 的路径条数,答案对 $P$ 取模。若存在无穷多条满足条件的路径,输出 -1

思路

刚拿到题目时,因为看到了“0 权边”和“无限多条合法路线”,很容易被引导到一个非常直观但在本题中行不通的方向上。下面我先分享一下我最初踩坑的思路,然后再推导出正确的解法。

一开始的错误思路:Tarjan 缩点 + DAG 上的最短路

看到图中可能有 0 权边构成的环,我的第一反应是:用 Tarjan 算法求强连通分量,把 0 权环缩成一个点

  1. 用 Tarjan 算法跑出所有 SCC;

  2. 如果存在节点数 $>1$ 且边权和为 $0$ 的 SCC,就判定有 $0$ 权环,直接输出 -1

  3. 将缩点后的 DAG 进行拓扑排序或 DP。

为什么这个思路彻底错了?

  1. 粗暴判 -1 会导致严重误判:存在 0 权环就一定有无数条路线吗?并不是!如果这个 0 权环在一个根本无法到达终点 $N$ 的角落,或者到达这个 0 权环并走向终点的总距离早就超过了 $d+K$,那这个环对答案毫无影响。

  2. 缩点会丢失内部路径信息:题目要求的是所有长度不超过 $d+K$ 的路径。如果把一个 0 权环缩成一个点,环内不同的绕行方式、不同的进出点都不知道怎么处理,导致 DP 根本无从下手。。。

正确思路:反向最短路 + 记忆化搜索 DP

意识到 Tarjan 走不通后,我们得回到题目的核心提示:$K \le 50$

这个极小的数据范围疯狂暗示我们:这题应该把 $K$ 作为一个维度塞进 DP 状态里

我们真正要关心的是:当前走的路,比绝对的最短路“浪费”了多少距离?

1. 参照物:反向图最短路

为了知道“浪费”了多少,我们需要一个参照物。如果我们在节点 $u$,准备走向 $v$,我们需要知道 $u$ 到终点 $N$ 的最短距离 $dis[u]$,以及 $v$ 到终点 $N$ 的最短距离 $dis[v]$。

所以要在反向图上以终点 $N$ 为起点跑一次 Dijkstra,求出各节点到 $N$ 的最短距离 $dis[u]$。

2. 状态定义与转移

DP 状态定义:$f[u][k]$ 表示当前位于节点 $u$,还可以多消耗 $k$ 额外距离到达终点 $N$ 的合法方案数。

如果我们选择走边 $u \to v$(边权为 $w$),这条边会消耗掉多少额度?

  • 完美路线是从 $u$ 直接走最短路,代价是 $dis[u]$。

  • 实际路线是从 $u$ 走到 $v$,再从 $v$ 走最短路,代价是 $w + dis[v]$。

  • 额外消耗的距离就是这两者的差值:$\Delta k = (dis[v] + w) - dis[u]$。

所以,走到 $v$ 后,我们剩余的绕路额度就变成了 $nxtk = k - \Delta k$。

转移方程

$$f[u][k] = \sum_{u \to v} f[v][k - (dis[v] + w - dis[u])]$$

3. 如何简单的判断 0 权环

既然不用 Tarjan,怎么找 0 权环?利用 DFS 的调用栈

在递归搜索 dfs(u, k) 时,维护 vis[u][k] 标记状态是否在当前递归栈内:

  • 入栈时设 vis[u][k] = 1,出栈时回溯设 vis[u][k] = 0

  • 若搜索过程中再次遇到 vis[u][k] == 1 的状态,说明绕行一圈后消耗配额为 0(因为题目说每条边有一个非负权值),直接标记全局变量 ok = 1 并退出。

分块代码详解

1. 反向图 Dijkstra 求最短路

我们建反向图 radj,从终点 $N$ 开始跑最短路。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
void dij(int s){
for(int i=1;i<=n;++i) dis[i]=inf;
priority_queue<State> pq;
pq.push({s,0});
dis[s]=0;
while(!pq.empty()){
auto [u,d]=pq.top();
pq.pop();
if(dis[u]<d) continue;
for(auto [v,w]:radj[u]){
if(dis[v]>dis[u]+w){
dis[v]=dis[u]+w;
pq.push({v,dis[v]});
}
}
}
}

2. 记忆化搜索与环检测

在状态转移时,要过滤掉无法到达终点的节点 $v$(dis[v] == inf

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
int dfs(int u,int k){
if(vis[u][k]){
ok=1; // 在递归栈中遇到了相同的 (u, k),抓到合法路径上的 0 权环
return 0;
}
if(f[u][k]!=-1) return f[u][k]; // 记忆化

vis[u][k]=1; // 标记入栈
ll res=0;
if(u==n) res=1; // 抵达终点,算作一条合法路线(即使额度没用完也算,因为路径已经确定)

for(auto [v,w]:adj[u]){
if(dis[v]==inf) continue; // 无法到达终点的点直接过滤
int nxtk = k - (dis[v] + w - dis[u]); // 计算下一个k
if(nxtk>=0){
res = (res + dfs(v, nxtk)) % p;
if(ok) return 0; // 如果深层递归发现了 0 ,就直接返回
}
}
vis[u][k]=0; // 标记出栈
return f[u][k]=res;
}

3. 多组数据重置(血泪教训)

巨坑啊啊啊啊

1
2
3
4
5
6
7
8
9
10
11
void init(){
for(int i=1;i<=n;++i){
adj[i].clear();
radj[i].clear();
for(int j=0;j<=k;++j){
f[i][j]=-1;
vis[i][j]=0;
}
}
ok=0;
}

完整 AC 代码

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
#include<bits/stdc++.h>
using namespace std;
using ll=long long;
const int maxn=1e5+5;
const ll inf=0x3f3f3f3f3f3f3f3fLL;

struct Edge{ int v;ll w;};
vector<Edge> adj[maxn],radj[maxn];
ll f[maxn][55];
int vis[maxn][55];
int n,m,k,p;
struct State{
int u; ll d;
bool operator<(const State &o) const{
return d>o.d;
}
};
ll dis[maxn];
void dij(int s){
for(int i=1;i<=n;++i) dis[i]=inf;
priority_queue<State> pq;
pq.push({s,0});
dis[s]=0;
while(!pq.empty()){
auto [u,d]=pq.top();
pq.pop();
if(dis[u]<d) continue;
for(auto [v,w]:radj[u]){
if(dis[v]>dis[u]+w){
dis[v]=dis[u]+w;
pq.push({v,dis[v]});
}
}
}
}
int ok;
int dfs(int u,int k){
if(vis[u][k]){
ok=1;
return 0;
}
if(f[u][k]!=-1) return f[u][k];
vis[u][k]=1;
ll res=0;
if(u==n) res=1;
for(auto [v,w]:adj[u]){
if(dis[v]==inf) continue;
int nxtk=k-(dis[v]+w-dis[u]);
if(nxtk>=0){
res=(res+dfs(v,nxtk)%p)%p;
if(ok) return 0;
}
}
vis[u][k]=0;
return f[u][k]=res;
}
void init(){
for(int i=1;i<=n;++i){
adj[i].clear();
radj[i].clear();
for(int j=0;j<=k;++j){
f[i][j]=-1;
vis[i][j]=0;
}
}
ok=0;
}
void solve(){
cin>>n>>m>>k>>p;
init();
for(int i=0;i<m;++i){
int u,v;ll w;
cin>>u>>v>>w;
adj[u].push_back({v,w});
radj[v].push_back({u,w});
}
dij(n);
ll ans=dfs(1,k)%p;
if(ok){
cout<<"-1\n";
return ;
}
cout<<ans<<"\n";
}
int main(){
ios::sync_with_stdio(0);cin.tie(0);
int __T=1;
cin>>__T;
while(__T--) solve();
return 0;
}