跳转至

第 62 章 记忆化搜索与剪枝

配套例题:BISHI90 【模板】记忆化搜索、BISHI79 取数游戏、BISHI146 收集金币 来源:S3 day6《DP 入门》(rxz);S3 day2《链表 DLX 并查集》暴力搜索与剪枝 前置60-DFS深度优先搜索11-函数14-标准库速查

搜索之所以慢,只有两个原因:

  1. 算重了——同一个子问题被反复求解 → 用记忆化解决;
  2. 算了没用的——明知不可能出解的分支还在往下走 → 用剪枝解决。

这一章讲的就是这两把刀。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 起还有一个等价的简写:

from functools import cache       # 3.9+,等价于 lru_cache(maxsize=None)


@cache
def f(a, b, c):
    ...

相关 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)           # ❌ 默认值就是 128!
def f(i, j): ...

lru_cache 不带参数时 maxsize 默认是 128。 状态一多,缓存不停淘汰再重算,复杂度直接退回指数级,而且你看不出任何异常—— 只是慢。竞赛里永远写 maxsize=None(或用 @cache)。

实测(\(500\times500\) 的二维递推,25 万个状态):

写法 耗时
@lru_cache(maxsize=None) 0.100 s
@cache 0.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_cachef(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。两种解法:

  1. 每组数据后 solve.cache_clear()
  2. 把函数定义在循环内部(闭包),每组自带一个新缓存—— 更安全,但每组都要重新编译一次函数对象(\(T\) 很大时有开销)。

陷阱五:缓存的内存

lru_cache 的每条记录要存 key(元组)、value、链表节点, 一条约 200–300 字节\(10^6\) 个状态就是 200–300 MB,很容易 MLE。

状态数超过 \(10^6\) 时,改用 list 数组存记忆化表(每格 8 字节指针), 或者干脆改成递推 + 滚动数组。


62.3 手写记忆化 vs lru_cache

很多资料说「手写 dictlru_cache 快」。在 CPython 里这是错的—— lru_cacheC 实现的(_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. 任意选取一列,如果这一列上没有 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)\),而前面所有轮的总和是

\[b^0 + b^1 + \cdots + b^{d-1} = \frac{b^d - 1}{b - 1} \approx \frac{b^d}{b-1}\]

也就是说,重复的部分只占最后一轮的 \(\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\)

四个坑

  1. a <= 0 or b <= 0 or c <= 0 才是边界,不是 == 0。 题目保证输入 \(\ge 1\),递推时只会减到 0,所以第 0 层全 1 即可;
  2. 取模不要每一项都取,攒完一个表达式再 % MOD 一次。 Python 大整数不会溢出,少一次取模就快一点;
  3. 减法后可能为负,最后统一 % MOD 会自动转正(Python 的 % 结果非负), 见 03-运算符与位运算
  4. 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\),完全安全

三个坑

  1. val 表必须预处理。如果在 dfs 里现算 mask 的权值和, 同一个 (r, s) 会被算很多次——记忆化只缓存了函数返回值,没缓存中间计算
  2. 多组数据的缓存:这里把 dfs 定义在 for 循环内部, 每组是一个全新的函数对象和全新的缓存。 如果把它提到外面,必须 dfs.cache_clear(),否则第二组会读到第一组的答案;
  3. 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)\) 的回合数是唯一确定的

\[t(x,y) = (x-1) + (y-1)\]

\((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\)

四个坑

  1. 「先变墙再移动」决定了判定是 v <= x + y 而不是 v < x + y。 差一个等号答案就错,样例 1 的 \((1,1)\) 正是卡在等号上;
  2. 起点永远是活的\(v \ge 1 > 0 = x+y\),所以题面那句 「起点在第一回合变墙视作不受影响」是自动满足的,不需要特判
  3. 答案是所有可达格子的 \(\max\),不是 \(f(n-1,m-1)\)—— 小 K 可能被堵死在半路(样例 2 的答案就是起点的 1);
  4. 不可达要用一个「传染性」的标记。这里用 \(-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 会不停淘汰重算
不可哈希参数 listtuplesetfrozenset,标记集→位 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」