第 62 章 记忆化搜索与剪枝¶
配套例题:BISHI90 【模板】记忆化搜索、BISHI79 取数游戏、BISHI146 收集金币 来源:S3 day6《DP 入门》(rxz);S3 day2《链表 DLX 并查集》暴力搜索与剪枝 前置:60-DFS深度优先搜索、11-函数、14-标准库速查
搜索之所以慢,只有两个原因:
- 算重了——同一个子问题被反复求解 → 用记忆化解决;
- 算了没用的——明知不可能出解的分支还在往下走 → 用剪枝解决。
这一章讲的就是这两把刀。S3 day6 的《DP 入门》把记忆化的本质说得很清楚:
DP 的原理,可以理解为「丢弃冗余信息」。 暴力往往会枚举很多信息,但是未必所有的信息都对解决问题有帮助。 如果我们只考虑有用的信息,往往就可以大幅降低复杂度。
记忆化搜索 = DFS + 缓存 = 自顶向下的 DP。 它和递推(自底向上的 DP)算的是同一张表,只是填表顺序不同。
62.1 记忆化:把指数变成多项式¶
最小例子:斐波那契¶
# ❌ 朴素递归:O(φ^n),n = 40 就要跑好几秒
def fib(n):
if n <= 1:
return n # 边界:fib(0) = 0,fib(1) = 1
return fib(n - 1) + fib(n - 2) # 两棵子调用树几乎完全重叠,这才是慢的根源
调用树里 fib(35) 被算了成千上万次——这就是「冗余信息」。
加一行缓存,复杂度立刻变成 \(O(n)\):
from functools import lru_cache # LRU = Least Recently Used,最近最少使用
@lru_cache(maxsize=None) # 容量无上限:算过的结果一律留着,绝不淘汰
def fib(n):
if n <= 1:
return n
return fib(n - 1) + fib(n - 2) # 每个 n 只真正执行一次,其余调用都命中缓存
记忆化的三个前提¶
| 前提 | 含义 | 违反会怎样 |
|---|---|---|
| 函数是纯的 | 相同参数必须返回相同结果 | 缓存到错误的值 |
| 状态可哈希 | 参数能当 dict 的键 |
TypeError: unhashable type |
| 无后效性 | 子问题的解与「怎么走到这里」无关 | 结果错误 |
无后效性是 DP 和记忆化的共同前提(S3 day6 也强调了这一点)。 判断方法:把当前状态的所有参数写出来,问「知道这些就够了吗?」 如果还需要知道「之前走过哪些格子」,那状态就没编全,记忆化会算错。
记忆化 vs 递推¶
| 记忆化搜索(自顶向下) | 递推(自底向上) | |
|---|---|---|
| 写法 | 递归 + 缓存 | 循环填表 |
| 需要想转移顺序吗 | ❌ 不需要(这是最大优势) | ✅ 需要保证依赖先算 |
| 只算用得到的状态 | ✅ 是(状态空间稀疏时省一大截) | ❌ 全部都算 |
| Python 的递归深度 | ⚠️ 致命弱点 | ✅ 无问题 |
| Python 的速度 | 慢(函数调用 + 缓存查找) | 快 3–5 倍 |
| 滚动数组优化 | ❌ 做不了 | ✅ 可以 |
Python 选手的默认选择: 能写递推就写递推,只有在「转移顺序绕、状态空间稀疏、状态是奇怪的元组」 这三种情况下才用记忆化。这和 C++ 选手的习惯不一样—— C++ 里记忆化只慢一点点,Python 里差距是数倍且会爆栈。
62.2 functools.lru_cache 的用法¶
from functools import lru_cache
@lru_cache(maxsize=None) # maxsize=None -> 无上限,永不淘汰
def f(a, b, c):
...
Python 3.9 起还有一个等价的简写:
相关 API¶
| 用法 | 作用 |
|---|---|
@lru_cache(maxsize=None) |
不限容量,竞赛里几乎总是用这个 |
@lru_cache(maxsize=N) |
最多缓存 \(N\) 条,LRU 淘汰(会重算!) |
f.cache_clear() |
清空缓存——多组数据必须调用 |
f.cache_info() |
返回 (hits, misses, maxsize, currsize),调试神器 |
陷阱一:maxsize 不写 None 会退化¶
lru_cache 不带参数时 maxsize 默认是 128。
状态一多,缓存不停淘汰再重算,复杂度直接退回指数级,而且你看不出任何异常——
只是慢。竞赛里永远写 maxsize=None(或用 @cache)。
实测(\(500\times500\) 的二维递推,25 万个状态):
写法 耗时 @lru_cache(maxsize=None)0.100 s @cache0.099 s @lru_cache(maxsize=1<<20)0.110 s 有上限但够大时只慢 10%,但上限不够时是灾难。
陷阱二:参数必须可哈希¶
@lru_cache(maxsize=None)
def f(state):
...
f([1, 2, 3]) # ❌ TypeError: unhashable type: 'list'
f({1, 2}) # ❌ set 也不行
f((1, 2, 3)) # ✅ tuple 可以
f(frozenset({1,2})) # ✅ frozenset 可以
| 想传的东西 | 改成 |
|---|---|
list |
tuple(x) |
set |
frozenset(x) |
| 一组布尔标记 | 二进制 mask 整数(最快,见 46-位运算) |
| 二维网格 | 别传! 放外层作用域(闭包/全局),只传下标 |
⚠️ 千万不要把大数组当参数传给记忆化函数。 就算把它转成
tuple,每次调用都要对整个 tuple 求哈希, \(O(n)\) 的哈希 × 指数级的调用 = 灾难。 正确做法:不变的数据放外面,只把「变的那几个下标」当参数。
陷阱三:lru_cache 让递归深度翻倍¶
这是 Python 独有的、最隐蔽的坑。实测(sys.setrecursionlimit(1000)):
| 递归形式 | 实际可达深度 |
|---|---|
裸递归 f(n-1) |
996 |
@lru_cache 的 f(n-1) |
497 |
因为 lru_cache 的包装器本身也占一层栈帧计数。也就是说:
加了
lru_cache之后,你的有效递归深度只剩一半。
配合 60.4 的物理栈数据(Windows 主线程约 2500 层、 64 MB 线程栈约 7 万层),结论是:
| 递归深度 | 裸递归 | lru_cache 递归 |
|---|---|---|
| \(\le 450\) | ✅ 无需任何处理 | ✅ 无需任何处理 |
| \(\le 900\) | ✅ 无需任何处理 | ⚠️ 要 setrecursionlimit |
| \(\le 10^4\) | ⚠️ 开大线程栈 | ⚠️ 开大线程栈(栈耗用也翻倍) |
| \(\ge 10^5\) | ❌ 改迭代 | ❌ 改递推 |
陷阱四:多组数据不清缓存¶
@lru_cache(maxsize=None)
def solve(i, j):
return grid[i][j] + ... # ← 用到了外层的 grid
for _ in range(T):
grid = read_grid() # 换了一组数据
print(solve(0, 0)) # ❌ 拿到的是上一组的缓存!
# solve.cache_clear() # ← 必须加这一句
这是记忆化题最经典的 WA。两种解法:
- 每组数据后
solve.cache_clear(); - 把函数定义在循环内部(闭包),每组自带一个新缓存—— 更安全,但每组都要重新编译一次函数对象(\(T\) 很大时有开销)。
陷阱五:缓存的内存¶
lru_cache 的每条记录要存 key(元组)、value、链表节点,
一条约 200–300 字节。\(10^6\) 个状态就是 200–300 MB,很容易 MLE。
状态数超过 \(10^6\) 时,改用
list数组存记忆化表(每格 8 字节指针), 或者干脆改成递推 + 滚动数组。
62.3 手写记忆化 vs lru_cache¶
很多资料说「手写 dict 比 lru_cache 快」。在 CPython 里这是错的——
lru_cache 是 C 实现的(_functools),比 Python 层的 dict.get + 赋值更快。
实测(\(500 \times 500\) 的组合数递推,25 万个状态,CPython 3.9):
| 写法 | 耗时 | 相对 |
|---|---|---|
@lru_cache(maxsize=None) |
0.100 s | 1.0× |
@cache |
0.099 s | 1.0× |
手写 dict 记忆化 |
0.131 s | 1.3× 慢 |
| 自底向上递推 | 0.019 s | 5.3× 快 |
结论:
| 场景 | 选择 |
|---|---|
| 一般记忆化 | @lru_cache(maxsize=None),又快又短 |
| 状态是连续小整数、可以编码成一维下标 | list 数组手写记忆化(比 dict 快 2–3 倍) |
| 状态数 \(\ge 10^6\) 或递归深 | 改递推,快 5 倍且不爆栈 |
| 需要精细控制内存 | 手写 |
三种写法对照¶
from functools import lru_cache
# ---- 写法一:lru_cache(首选)----
# 三种写法算的是同一个东西:从 (0,0) 只往右/往下走到 (i,j) 的路径数
@lru_cache(maxsize=None)
def f1(i, j):
if i == 0 or j == 0:
return 1 # 贴着边走只有一条路,这是递归出口
return f1(i - 1, j) + f1(i, j - 1)
# ---- 写法二:手写 dict(状态稀疏 / 状态是奇怪的元组时用)----
memo = {}
def f2(i, j):
key = (i, j) # 元组可哈希,才能当 dict 的键
v = memo.get(key)
if v is None: # ★ 用 get + None 判断,比 `if key in memo` 少一次哈希
v = memo[key] = 1 if (i == 0 or j == 0) else f2(i - 1, j) + f2(i, j - 1)
return v # 赋值表达式顺手把结果写进缓存再返回
# ---- 写法三:list 数组(状态是连续小整数时最快)----
def f3_build(n, m):
"""状态编码成一维下标 i * (m+1) + j,用 list 当记忆化表。"""
memo = [-1] * ((n + 1) * (m + 1)) # -1 表示「没算过」;本题答案恒为正,不会混淆
def go(i, j):
k = i * (m + 1) + j # 每行 m+1 个格子,所以行首偏移是 i*(m+1)
v = memo[k]
if v < 0:
v = memo[k] = 1 if (i == 0 or j == 0) else go(i - 1, j) + go(i, j - 1)
return v
return go
写法二的
memo.get(key)比if key in memo: return memo[key]快: 后者要哈希两次。但要注意None不能是合法的返回值, 否则会把「已缓存的 None」误判成「没算过」。返回值可能是 0 时要小心:if v is None是对的,if not v是错的。
62.4 剪枝:让搜索树变小¶
剪枝的唯一目标:在还没往下走之前,就证明这个分支不可能产生(更优的)解。
S3 day2 讲 DLX 时给的暴力搜索流程,第 2 步就是一个标准的可行性剪枝:
- 任意选取一列,如果这一列上没有 1,则无解。
四类剪枝¶
| 类型 | 判断依据 | 例子 |
|---|---|---|
| 可行性剪枝 | 当前状态已经违反约束 → 直接返回 | 皇后互相攻击;已选的数超了容量 |
| 最优性剪枝 | 当前部分解 + 剩余最好情况 \(\le\) 已知答案 → 返回 | 「当前和 + 剩下所有正数 \(\le\) ans」 |
| 对称性剪枝 | 两个分支本质相同 → 只走一个 | 第一个皇后只枚举左半列;相同的物品只按一种顺序选 |
| 记忆化剪枝 | 这个状态算过了 → 查表 | 本章前半部分 |
还有两个不是「剪枝」但同样重要的技巧:
| 技巧 | 说明 |
|---|---|
| 优化搜索顺序 | 先搜分支少的(DLX 选 1 最少的列)、先搜大的(背包先放大物品) |
| 等效冗余消除 | 把「顺序无关」的枚举强制成一种顺序(组合枚举的 start 参数) |
最优性剪枝的模板¶
def dfs_best(i, cur, remain_sum):
"""搜索第 i 个物品,cur 是当前累积值,remain_sum 是 i..n-1 的总和上界。"""
global ans
if cur + remain_sum <= ans: # ★ 最优性剪枝:往后全拿也超不过已知答案
return # 取 <= 而不是 <,等于已知最优也没必要再搜
if i == n:
if cur > ans: # 物品用完,用这条完整方案更新答案
ans = cur
return
...
上界函数(
remain_sum)越紧,剪枝越狠。 常用的上界:后缀和、贪心解、放松约束后的最优解(这就是「分支限界法」)。 上界算得太慢会得不偿失——剪枝的收益必须大于它的计算成本。
对称性剪枝的模板¶
# N 皇后:第一行只枚举左半边,答案 * 2(n 为奇数时中间列单独算)
# 相同元素的排列去重:排序后,同一层里跳过与前一个相同且前一个没被用的
for i in range(n):
if used[i]:
continue # 这个位置的元素已经在路径上了
# a[i-1] 没被用 -> 本层是拿第二个相同元素打头,与「拿第一个打头」的分支完全同形;
# a[i-1] 已被用 -> 说明它在更浅的层被选走,此时 a[i] 是接着用的,合法
if i > 0 and a[i] == a[i - 1] and not used[i - 1]:
continue # ★ 保证相同元素按固定顺序被使用
...
Python 特有的「剪枝」:把内层循环下沉到 C¶
这是本教程反复出现的主题。搜索题里最典型的两处:
| 场景 | Python 层写法 | C 层写法 |
|---|---|---|
| 枚举全排列 | 手写回溯 | itertools.permutations |
| 枚举子集 | 手写回溯 | range(1 << n) + 位运算 |
| 求某行的和 | for 累加 |
sum(row) |
| 判某个候选集是否为空 | 遍历 | mask == 0(位运算) |
| 找可选项最少的那一列 | 遍历比较 | min(cols, key=len) |
62.5 迭代加深(IDDFS)¶
问题形态:答案的深度不大,但分支因子很大,BFS 会因为队列太宽而 MLE。
迭代加深(IDDFS,Iterative Deepening Depth-First Search) = 限定深度的 DFS,深度从 1 开始逐次加 1:
def iddfs(start, is_goal, neighbors, max_depth=50):
"""迭代加深:找最少步数。空间 O(d),时间与 BFS 同阶。"""
def dfs(u, depth, limit):
if is_goal(u):
return True # 先判目标再判深度,limit 步刚好到达也算数
if depth == limit:
return False # 本轮的预算用完,这条路暂时不再往下走
for v in neighbors(u):
if dfs(v, depth + 1, limit):
return True # 找到就层层向上返回,不必搜完剩余分支
return False
for limit in range(max_depth + 1): # ★ 深度上限逐次放宽
if dfs(start, 0, limit):
return limit # 第一次搜到即最少步数:更小的 limit 已失败过
return -1
「重复搜索不是浪费吗?」 不是。设分支因子为 \(b\), 则深度 \(\le d\) 的搜索树节点数是 \(O(b^d)\),而前面所有轮的总和是
也就是说,重复的部分只占最后一轮的 \(\frac{1}{b-1}\)。\(b \ge 3\) 时重复开销 \(\le 50\%\)。
| BFS | IDDFS | |
|---|---|---|
| 时间 | \(O(b^d)\) | \(O(b^d)\)(多一个小常数) |
| 空间 | \(O(b^d)\) ← 瓶颈 | \(O(d)\) |
| Python 递归深度 | 无问题 | ⚠️ 深度就是 \(d\),一般很小,安全 |
| 找到的解 | 最短 | 最短 |
IDA*(Iterative Deepening A-star)是 IDDFS 加上估价函数 \(h\):当 已走步数 + h(当前状态) > limit 时剪枝。
\(h\) 必须是乐观估计(不高于真实剩余步数),否则会找到非最优解。
典型题:八数码(\(h\) = 各数字到目标位置的曼哈顿距离之和)、魔方。
在 Python 里 IDDFS 的实际价值:主要是省内存。 时间上 Python 跑指数级搜索本来就吃力,\(b^d\) 超过 \(10^6\) 基本就没戏了。 所以 IDDFS 在 Python 里的适用区间是:答案步数 \(\le 15\)、分支因子 \(\le 4\) 左右。
62.6 例题¶
BISHI90 【模板】记忆化搜索(中等)¶
给定分段递归函数 $\(f(a,b,c) = \begin{cases} 1 & a \le 0 \text{ 或 } b \le 0 \text{ 或 } c \le 0 \\ f(a,b,c\!-\!1) + f(a,b\!-\!1,c\!-\!1) - f(a,b\!-\!1,c) & a < b \text{ 且 } b < c \\ f(a\!-\!1,b,c) + f(a\!-\!1,b\!-\!1,c) + f(a\!-\!1,b,c\!-\!1) - f(a\!-\!1,b\!-\!1,c\!-\!1) & \text{其他} \end{cases}\)$ \(T \le 10^3\) 组询问,\(1 \le a,b,c \le 100\),输出 \(f(a,b,c) \bmod (10^9+7)\)。 题面见 BISHI90 原题(牛客)。 题解见
solutions/BISHI90.py(已用官方样例验证)。
这题名字叫「记忆化搜索模板」,但在 Python 里正解偏偏不是记忆化搜索。 这个反差正是本章最有价值的地方。
朴素递归每次分裂成 3~4 个子问题、深度上百,是指数级,必须记忆化。 状态总数只有 \(101^3 \approx 1.03\times10^6\)。三条路:
| 做法 | 问题 |
|---|---|
lru_cache + 递归 |
递归深度可达 300(lru_cache 再翻倍到 600),且 \(10^6\) 个状态的缓存约 250 MB,又慢又险 |
手写 dict 记忆化 |
同样的内存问题 |
| 自底向上递推 | ✅ 完全绕开递归,还能用切片把内层循环下沉到 C |
递推顺序:所有转移的下标都不增(且至少有一维严格减小), 按 \(a\) 升序 → \(b\) 升序 → \(c\) 升序扫描时,用到的 \(f(a\!-\!1,*,*)\)(上一层已算完)、\(f(a,b\!-\!1,*)\)(同层前一个 \(b\))、 \(f(a,b,c\!-\!1)\)(同行前一个 \(c\))全都已经就绪。
更妙的是:「其他情况」分支只依赖上一层 \(a-1\),各个 \(c\) 之间互不依赖,
所以可以用 zip + 列表推导整行批量算出来,把内层循环压到 C 层。
import sys
MOD = 10 ** 9 + 7
N = 100
def build():
"""F[a][b][c] = f(a,b,c) mod MOD。一次性把整张表递推出来。"""
ones = [1] * (N + 1)
F = [[ones] * (N + 1)] # a = 0 层:base case 全为 1
for a in range(1, N + 1):
pa = F[a - 1] # 上一层,即所有 f(a-1, *, *)
layer = [ones] # b = 0 -> 全 1
for b in range(1, N + 1):
pb = pa[b] # f(a-1, b, *)
pbm = pa[b - 1] # f(a-1, b-1, *)
# 分支三覆盖的 c 范围:a >= b 时是全部;a < b 时是 c <= b
lim = N if a >= b else b
row = [1] # c = 0 -> base case
# ★ 批量算 c = 1..lim:只依赖上一层,用 zip 把循环丢到 C 层
# 四个切片依次是 f(a-1,b,c)、f(a-1,b-1,c)、f(a-1,b,c-1)、f(a-1,b-1,c-1):
# 后两个的切片整体左移一格,取到的正是 c-1
row.extend([(x + y + u - v) % MOD for x, y, u, v in
zip(pb[1:lim + 1], pbm[1:lim + 1], pb[0:lim], pbm[0:lim])])
if lim < N:
# a < b 且 c > b:分支二,c 方向串行依赖,只能逐个算
cur = layer[b - 1] # f(a, b-1, *)
prev = row[lim] # f(a, b, lim),串行递推的起点
for c in range(lim + 1, N + 1):
# 分支二:f(a,b,c) = f(a,b,c-1) + f(a,b-1,c-1) - f(a,b-1,c)
prev = (prev + cur[c - 1] - cur[c]) % MOD
row.append(prev) # 减法可能为负,% MOD 会自动转成非负
layer.append(row)
F.append(layer)
return F
def main():
F = build()
data = sys.stdin.buffer.read().split()
t = int(data[0])
out = []
p = 1
for _ in range(t):
a = int(data[p]); b = int(data[p + 1]); c = int(data[p + 2]); p += 3
out.append(str(F[a][b][c]))
sys.stdout.write("\n".join(out) + "\n")
main()
手动复核样例: \(f(1,1,1)\):\(a<b\) 不成立 → 分支三 \(= f(0,1,1)+f(0,0,1)+f(0,1,0)-f(0,0,0) = 1+1+1-1 = 2\) ✓ \(f(2,2,2) = f(1,2,2)+f(1,1,2)+f(1,2,1)-f(1,1,1) = 2+2+2-2 = 4\) ✓
四个坑:
a <= 0 or b <= 0 or c <= 0才是边界,不是== 0。 题目保证输入 \(\ge 1\),递推时只会减到 0,所以第 0 层全 1 即可;- 取模不要每一项都取,攒完一个表达式再
% MOD一次。 Python 大整数不会溢出,少一次取模就快一点; - 减法后可能为负,最后统一
% MOD会自动转正(Python 的%结果非负), 见 03-运算符与位运算; F = [[ones] * (N + 1)]里ones被共享了 101 次—— 只读的表可以共享,但绝不能就地修改。这是 05-列表 讲的浅拷贝陷阱的一个正面利用。
本题给出的判据:状态数 \(\ge 10^6\)、且转移的依赖方向规整(下标只减不增)时, 一律写递推,不写记忆化。递归在 Python 里的代价(栈 + 调用 + 缓存对象) 在这个量级上是无法接受的。
BISHI79 取数游戏(中等)¶
\(T \le 20\) 组,每组 \(N \times M \le 6 \times 6\) 的非负整数矩阵, 取出若干数使得任意两数不八连通相邻,求最大和。 题面见 BISHI79 原题(牛客)。 题解见
solutions/BISHI79.py(已用官方样例验证,用的是递推版)。
60 章 给了自底向上的递推版本。 这里给记忆化搜索版,用来直观展示「记忆化 = 搜索 + 缓存」:
状态设计:dfs(r, prev) = 「前 \(r\) 行已确定,第 \(r-1\) 行选的是 prev,
从第 \(r\) 行往下能取到的最大和」。
- 逐格搜索的状态是「已选了哪些格子」,指数爆炸;
- 逐行搜索的状态只有「上一行选了什么」,因为隔一行以上永远不相邻—— 这就是「丢弃冗余信息」。
import sys
from functools import lru_cache
def main():
data = sys.stdin.buffer.read().split()
p = 0
t = int(data[p]); p += 1
out = []
for _ in range(t):
n, m = int(data[p]), int(data[p + 1]); p += 2
rows = []
for _ in range(n):
rows.append([int(v) for v in data[p:p + m]])
p += m
full = 1 << m # m 位二进制,共 2^m 种「一行的选法」
masks = [s for s in range(full) if not (s & (s << 1))] # 行内不相邻
spread = [s | (s << 1) | (s >> 1) for s in range(full)] # 向左右各扩一位
val = [] # val[r][s] = 第 r 行选 s 的和
for r in range(n):
row = rows[r]
v = [0] * full
for s in masks: # 只有合法 mask 需要算权值,其余位置留 0
tot = 0
x = s
while x:
low = x & -x # lowbit:取出 x 最低位的那个 1
tot += row[low.bit_length() - 1] # 该 1 在第几位就加第几列的数
x ^= low # 抹掉这一位,循环次数 = 1 的个数
v[s] = tot
val.append(v)
# ★ 把 dfs 定义在循环内部:每组数据自带一个全新的缓存,
# 不会串到下一组去(等价于每组 cache_clear)
@lru_cache(maxsize=None)
def dfs(r, prev):
"""第 r 行开始往下取,上一行选的是 prev,返回能取到的最大和。"""
if r == n:
return 0 # 越过最后一行,后面再也取不到东西
best = 0 # 本行可以一个都不取,所以下界是 0
vr = val[r]
for s in masks:
if spread[prev] & s: # 扩位后仍相交 = 正上方或斜上方冲突
continue
cur = vr[s] + dfs(r + 1, s) # 本行的收益 + 往下的最优解
if cur > best:
best = cur
return best # 只依赖 (r, prev),所以缓存是安全的
out.append(str(dfs(0, 0)))
dfs.cache_clear() # 显式释放缓存内存
sys.stdout.write("\n".join(out) + "\n")
main()
状态数 \(= N \times 2^M \le 6 \times 64 = 384\),每个状态枚举 21 个合法 mask, 总量不到 \(10^4\),秒出。递归深度只有 \(N \le 6\),完全安全。
三个坑:
val表必须预处理。如果在dfs里现算 mask 的权值和, 同一个(r, s)会被算很多次——记忆化只缓存了函数返回值,没缓存中间计算;- 多组数据的缓存:这里把
dfs定义在for循环内部, 每组是一个全新的函数对象和全新的缓存。 如果把它提到外面,必须dfs.cache_clear(),否则第二组会读到第一组的答案; dfs(0, 0)的第二个参数 0 表示「第 \(-1\) 行是空集」,spread[0] = 0与任何 mask 都不冲突,正好当哨兵。
对照 60 章 的递推版: 两份代码算的是同一张表
dp[r][mask], 记忆化版不需要想「先算哪一行」(递归自动处理依赖), 递推版不需要担心递归深度且更快。 这题规模小,两种都行;规模一大,就只有递推能活。
BISHI146 收集金币(中等)¶
\(n, m \le 1000\) 的网格,\((i,j)\) 有 \(a_{i,j}\) 个金币。\(t\) 条信息 \(\{x,y,v\}\) 表示 \((x,y)\) 在第 \(v\) 回合永久变成墙(此前金币仍可收集)。 每回合先变墙,小 K 再移动;小 K 从 \((1,1)\) 出发,每回合只能向右或向下走一格。 求最多能收集多少金币。 题面见 BISHI146 原题(牛客)。
ℹ️ 本题的
solutions/题解文件尚未编写。下面的代码已由scripts/verify_docs.py用官方样例实测通过,但未在牛客提交。
这题的全部难点在把「时间」这一维消掉。
只能向右或向下走 ⟹ 走到 \((x,y)\) 的回合数是唯一确定的:
而 \((x,y)\) 在第 \(v\) 回合开始时变墙、小K 在同一回合之后才移动,所以:
格子 \((x,y)\) 可以踩 \(\iff v(x,y) > (x-1)+(y-1)\)(没给信息的格子 \(v = +\infty\))。
时间维就这样被消掉了——预处理时把每个格子标成「死 / 活」, 剩下的就是一道最裸的网格 DP。
import sys
def main():
data = sys.stdin.buffer.read().split()
p = 0
n = int(data[p]); m = int(data[p + 1]); p += 2
a = []
for _ in range(n):
a.append([int(x) for x in data[p:p + m]])
p += m
t = int(data[p]); p += 1
# dead[i][j] = 1 表示走到 (i,j) 的那一刻它已经是墙(0-indexed)
dead = [bytearray(m) for _ in range(n)]
for _ in range(t):
x = int(data[p]) - 1; y = int(data[p + 1]) - 1; v = int(data[p + 2]); p += 3
if v <= x + y: # ★ 到达时刻是 x+y,v <= x+y 就踩不到了
dead[x][y] = 1
NEG = -1 # -1 表示该格不可达
prev = [NEG] * m # 滚动数组:上一行的 f 值
if not dead[0][0]:
prev[0] = a[0][0] # 起点活着才有出发点
for j in range(1, m): # 第 0 行只能从左边过来
if not dead[0][j] and prev[j - 1] >= 0:
prev[j] = prev[j - 1] + a[0][j] # 左边不可达就整条右侧都断了
best = max(prev) # 答案取所有可达格的最大值,见下方坑 3
for i in range(1, n):
row = a[i]
drow = dead[i]
cur = [NEG] * m
left = NEG # cur[j-1],省一次列表索引
for j in range(m):
if drow[j]:
left = NEG
continue # 墙:cur[j] 保持 NEG,且左邻也失效
up = prev[j]
b = left if left > up else up # max(上, 左)
if b < 0:
left = NEG # 上、左都不可达,本格也到不了
else:
left = b + row[j] # 顺手把 cur[j] 记进 left,下一列直接用
cur[j] = left
prev = cur # 滚动:本行算完就顶替上一行
mx = max(cur)
if mx > best:
best = mx
sys.stdout.write("%d\n" % best)
main()
手动复核样例 1(\(3\times3\),信息 \((1,1,1),(1,2,1),(2,2,2)\),0-indexed 后是 \((0,0,v{=}1),(0,1,v{=}1),(1,1,v{=}2)\)):
| 格子 | \(x+y\) | \(v\) | 死? |
|---|---|---|---|
| \((0,0)\) | 0 | 1 | 活(\(1 > 0\)) |
| \((0,1)\) | 1 | 1 | 死(\(1 \le 1\)) |
| \((1,1)\) | 2 | 2 | 死(\(2 \le 2\)) |
于是唯一活路是 \((0,0)\to(1,0)\to(2,0)\to(2,1)\to(2,2)\), 金币 \(1+1+1+1+1 = 5\) ✓
四个坑:
- 「先变墙再移动」决定了判定是
v <= x + y而不是v < x + y。 差一个等号答案就错,样例 1 的 \((1,1)\) 正是卡在等号上; - 起点永远是活的:\(v \ge 1 > 0 = x+y\),所以题面那句 「起点在第一回合变墙视作不受影响」是自动满足的,不需要特判;
- 答案是所有可达格子的 \(\max\),不是 \(f(n-1,m-1)\)—— 小 K 可能被堵死在半路(样例 2 的答案就是起点的 1);
- 不可达要用一个「传染性」的标记。这里用 \(-1\),并且在
b < 0时把left也置回 \(-1\),保证不可达状态不会被当成 0 传下去。
Python 现实性评估(\(n = m = 1000\),\(10^6\) 个格子,时限「其他语言 2 秒」):
| 项 | 实测 |
|---|---|
读入 \(10^6\) 个 token + int() |
约 0.1 s |
| 主循环 \(10^6\) 次 | 约 0.15 s |
| 总计(随机数据) | 约 0.26 s |
宽裕。对照:改成 lru_cache 记忆化递归版,同样数据实测 0.76 s(慢 2.9 倍),
而且递归深度达 \(n+m = 2000\),必须 setrecursionlimit + 开 64 MB 线程栈
(见 60.4),否则在 Windows 上直接静默崩溃。
这题是本章的总结题: 它天然是一道「记忆化搜索」题(\(f(i,j)\) 依赖 \(f(i-1,j)\) 和 \(f(i,j-1)\)), 但在 Python 里,递推版又快 3 倍又不会爆栈。 记忆化是思考工具,递推是提交工具。
62.7 本章速查¶
| 要点 | 结论 |
|---|---|
| 记忆化的本质 | DFS + 缓存 = 自顶向下的 DP |
| 三个前提 | 函数纯 + 状态可哈希 + 无后效性 |
| 首选写法 | @lru_cache(maxsize=None) 或 @cache(3.9+) |
maxsize 默认值 |
128!不写 None 会不停淘汰重算 |
| 不可哈希参数 | list→tuple,set→frozenset,标记集→位 mask |
| 大数组当参数 | ❌ 绝对不要,放外层作用域只传下标 |
lru_cache 与递归深度 |
有效深度减半(包装器多占一层栈帧) |
| 多组数据 | 必须 cache_clear(),或把函数定义在循环内 |
| 缓存内存 | 每条约 200–300 字节,\(10^6\) 状态就 250 MB |
lru_cache vs 手写 dict |
lru_cache 快 1.3 倍(它是 C 实现) |
| 状态是连续小整数 | 用 list 数组手写记忆化,最快 |
| 递推 vs 记忆化 | 递推快 3–5 倍且不爆栈;Python 默认选递推 |
| 什么时候才用记忆化 | 转移顺序绕 / 状态稀疏 / 状态是奇怪元组 |
| 可行性剪枝 | 违反约束立刻返回 |
| 最优性剪枝 | 当前 + 剩余上界 <= ans 就返回;上界越紧越好 |
| 对称性剪枝 | 等价分支只走一个(皇后半列、相同元素固定顺序) |
| 搜索顺序 | 先搜分支最少的(DLX 的 \(S\) 启发式) |
| IDDFS | 时间同 BFS,空间 \(O(d)\);重复开销只占 \(\frac{1}{b-1}\) |
| IDA* | IDDFS + 乐观估价 \(h\);\(h\) 必须不超过真实剩余步数 |
| Python 下 IDDFS 的适用区 | 答案步数 \(\le 15\)、分支因子 \(\le 4\) |
| 症状 → 该用哪把刀 |
|---|
| 「同一个子问题算了很多遍」 |
| 「状态数不大但递归树巨大」 |
| 「怎么剪都还是慢,但状态数其实很少」 |
| 「状态数也很大」 |
| 「明知走不通还在往下走」 |
| 「已经不可能比现有答案更优」 |
| 「很多分支算出来是一样的」 |
| 「BFS 队列爆内存」 |
| 「递归深度 > 900」 |