recursive DP
Dynamic programming 动态规划就是把过程中的算出来的数存在dict list array hash map lru_cache里。
这个存储方式也叫memoization,备忘录。
主要是用于 recursive function,比如斐波那契,0,1 然后第三个数是前两个数之和,第四个数是第三个和第四个之和,以此类推。fib(n)每次 都要重新算一下,就是 2^N 次,算量爆炸增长,所以如果存下来,比如放到一个list 或者 dict 里面,就是 O(n)。
recursive 往往是 top-down 的思维,从最终的大问题往回推,直到最小问题为止。
而 for loop 是 bottom-up 思维,从最小问题开始,不断 loop,直到算到最终大问题的结果。
但是 for loop 往往需要一个 regular DP space,往上推的时候是按照顺序的。
但是对于一些 graph,或者 DAG,Directed Acyclic Graph, 每个 node的相连往往是不规律的,有多有少,这个时候从最小问题往上 for loop 就非常难以实现了。
Top-down vs. bottom-up

添加图片注释,不超过 140 字(可选)
DAG
2025 AOC day11 的 puzzle 就是这样一个例子,想要从 you 到 out 有多少条路径

添加图片注释,不超过 140 字(可选)
g={'aaa': ['you', 'hhh'],
'you': ['bbb', 'ccc'],
'bbb': ['ddd', 'eee'],
'ccc': ['ddd', 'eee', 'fff'],
'ddd': ['ggg'],
'eee': ['out'],
'fff': ['out'],
'ggg': ['out'],
'hhh': ['ccc', 'fff', 'iii'],
'iii': ['out']}
定义函数
ways(node) = 从 node 出发最终能走到 out 的路径数
状态转移,State Transition
ways(you) = ways(bbb) + ways(ccc)
#状态转移到 bbb 和 ccc
ways(bbb)=ways(eee)+ways(ddd)
#以此类推,最终 case
ways(eee)=ways(out)
ways(out)=1
那就可以这样定义func
from functools import lru_cache
@lru_cache(None)
def ways(node):
if node=='out': return 1
return sum(ways(child) for child in g[node])
为什么不把 g 放在 arg 里呢,因为lru_cache会想要 hash arg,而 dict 无法被 hash
第二问长这样,从 svr 到 out,图中必须穿过 fft 和 dac,问有几条路

添加图片注释,不超过 140 字(可选)
那就这么想
The logic is:
ways(svr)=ways(aaa)+ways(bbb)
# aaa meet fft
ways(aaa)=ways(fft)
# count as normal
...
ways(out)=1
# for bbb
ways(bbb)=ways(tty)
ways(tty)=ways(ccc)
# suppose goes to ccc-ddd-hub-out,and does not meet fft nor dac
ways(ccc)=ways(hub)
...
ways(out) =0
然后
@lru_cache(None)
def ways(node,seen_fft=False,seen_dac=False):
if node=='fft': seen_fft=True
if node=='dac': seen_dac=True
if node=='out':
if seen_fft and seen_dac:
return 1
else: return 0
return sum(ways(child,seen_fft,seen_dac) for child in g[node])
有一个 m x n 的 grid,起点是左上角的 0,0,终点是 m-1,n-1,求问从 0,0 到 m-1,n-1 有多少条路线。
m=4,n=3
way(0,0)=way(0,1)+way(1,0)
way(0,1)=way(0,2)+way(1,1)
...
way(3,1)=way(3,2) # hit the wall, go right
way(3,2)=way(3,3)
way(3,4)=1 # hit final
func
from functools import cache
@cache
def way(x,y):
if x==m-1 and y==n-1: return 1
if x==m-1: return way(x,y+1)
if y==n-1: return way(x+1,y)
return way(x+1,y)+way(x,y+1)
m=4, n=3
way(0,0) # return 10
强盗robber
有一个强盗,在一个街list上抢,每个item表示这户人家有多少钱,不可以两个相邻的抢,否则会police。
a=[2,7,9,3,1] # return 12, as 2+9+1=12
from functools import cache
@cache
def m(idx):
if idx==0: return a[idx]
if idx==1: return max(a[:2])
return max(m(idx-1),m(idx-2)+a[idx])
爬楼梯,每次爬1或者2步,直到top,最少cost是多少
cost = [10,15,20] # return 15
@cache
def pay(idx):
if idx<=1: return 0
return min(pay(idx-1)+cost[idx-1],pay(idx-2)+cost[idx-2])