Lab 4:序列、树递归与树ok 10 项通过
从「把列表变成另一个列表」,到「把一棵树变成另一棵树」——这份 lab 把递归从数字搬到了数据结构上。
0. 这份作业在练什么
Lab 4 是整门课的一个分水岭。在它之前,你写的递归几乎都是「拿一个数字,把它变小一点,再递归」——
fact(n) 调 fact(n-1),sum_digits(n) 调 sum_digits(n // 10)。
这类递归有一个很舒服的性质:每次只有一条路往下走,展开出来是一根直线。
Lab 4 之后,递归开始长出分叉。一个 balanced(t) 的调用会派生出「每个分支各调一次」,
一个 num_trees(n) 的调用会派生出「每种拆分方式各调两次」。
调用结构不再是一根直线,而是一棵树——这就是树递归(tree recursion)这个名字的由来。
注意:树递归指的是调用结构像树,它和处理树这种数据是两回事,只是这份 lab 把两件事放在一起练了。
整份作业分成三块,难度是递进的:
| 模块 | 题目 | 真正在练什么 |
|---|---|---|
| 序列 | Q1 my_map、Q2 my_filter、Q3 my_reduce | 列表推导式的两种形态;以及「把函数当参数传」这件事的实感 |
| 数据抽象 | Q4 distance、Q5 closer_city、Q6 抽象屏障检查 | 只靠构造器 / 选择器写程序,让实现换掉之后代码还能跑 |
| 树 | Q7 WWPD、Q8 sum_tree / balanced、Q9 num_trees、Q10 only_paths | 树的数据抽象;「对每个分支递归,再把结果合起来」这个万能套路 |
- 列表推导式(list comprehension)有两种用法:
[f(x) for x in s]做映射,[x for x in s if p(x)]做筛选。Q1 和 Q2 就是这两句话本身。 - 高阶函数(higher-order function):
fn、pred、combiner都是参数位置上的函数。 你不需要知道它们是什么,只需要按约定的元数(arity)去调用。 - 抽象屏障(abstraction barrier):只准写
get_lat(city),不准写city[1]。 Q6 会真的把底层实现从列表偷换成字典来验你有没有作弊。 - 树的数据抽象:
tree(label, branches)构造,label(t)/branches(t)选择,is_leaf(t)判叶。不许写t[0]、t == x、x in t。 - 树的递归模板:
处理(t) = 用 label(t) 和 [处理(b) for b in branches(t)] 拼出答案。 叶子(branches(t)为空)通常不需要单独写 base case,因为空列表推导式自然返回[]。
做之前该掌握什么
三件事。第一,你要能读懂 lambda x: x * x 这样的表达式,知道它求值出来是一个函数对象,
而不是一个数。第二,你要习惯「递归调用返回的是答案,不是过程」——
写 sum_tree(b) 的时候就假定它已经正确返回了分支 b 的总和,不要在脑子里再展开一层。
第三,你要接受「列表推导式里的 for 每轮都真的执行一遍表达式」,
所以 [print(x) for x in s] 会真的打印,同时返回一堆 None——Q1 的第三个 doctest 就是考这个。
本仓库 labs/lab04/ 下的代码在 python3 ok --local 下的结果是
10 test cases passed! No cases failed.,
--score 给出的分项为 my_map / my_filter / my_reduce / distance / closer_city /
check_city_abstraction / sum_tree / balanced 各 1.0,Trees(WWPD)不计分。
两道选做题 num_trees、only_paths 也已写出并通过其 doctest
(python3 -m doctest lab04.py 报告 78 passed and 0 failed)。
本页贴出的每一段代码都逐字来自该文件。
1. Q1 my_map
题目要什么
写一个函数 my_map(fn, seq):把单参数函数 fn 作用到序列 seq 的每个元素上,
返回一个列表,里面装着这些结果,顺序与 seq 一致。
额外限制:函数体只能写一行(官方另有一个 my_map_syntax_check 测试会用 ast 解析你的源码,
确认函数体只有一条 Expr(docstring)加一条 Return)。
>>> my_map(lambda x: x*x, [1, 2, 3])
[1, 4, 9]
>>> my_map(lambda x: abs(x), [1, -1, 5, 3, 0])
[1, 1, 5, 3, 0]
>>> my_map(lambda x: print(x), ['cs61a', 'summer', '2023'])
cs61a
summer
2023
[None, None, None]
前两个 doctest 平平无奇。第三个 doctest 才是这题的真正考点:
传进去的 fn 是 lambda x: print(x)。调用它会做两件事——
把 x 打印到屏幕上(副作用),然后返回 None(print 的返回值)。
所以期望输出里先出现三行文字,再出现 [None, None, None]。
这告诉你:你必须真的对每个元素调用一次 fn,而且要把它的返回值——哪怕是 None——原样收进列表。
任何「偷偷跳过 None」的写法都会挂。
怎么想到的
拆成两个问题:怎么「对每个元素做一件事」,怎么「把结果收成列表」。
如果还没学列表推导式,你会写出这版:
def my_map(fn, seq): # 这不是最终答案
result = []
for element in seq:
result.append(fn(element))
return result
这段是对的,语义也完全清楚:建一个空列表,遍历,每次把 fn(element) 追加进去。
但它有四行函数体,过不了 my_map_syntax_check。
题目提示用列表推导式,而列表推导式恰恰就是这个循环的字面翻译:
result = []
for element in seq: → [ fn(element) for element in seq ]
result.append(fn(element)) └──收什么──┘ └────遍历什么────┘
return result
推导式的求值过程:先求 seq,得到一个序列;然后建一个新的空列表;
对序列中每个元素,把它绑定到名字 element 上,求值一次前面的表达式 fn(element),
把结果追加到新列表;遍历结束后,整个推导式的值就是这个新列表。
这里有个容易走的弯路:有人会想「Python 不是内置了 map 吗,直接 return map(fn, seq)」。
不行。Python 3 的内置 map 返回的是一个惰性的 map 对象,不是列表,
my_map(lambda x: x*x, [1, 2, 3]) 会显示成 <map object at 0x...> 而不是 [1, 4, 9]。
题面也明确说了「In Python, the map and filter built-ins have slightly different behavior than
the my_map and my_filter functions we are defining here」。
就算写 return list(map(fn, seq)) 能过 doctest,也完全绕过了这题想让你练的东西。
把「fn 是个不知道内容的函数」这件事当成不需要关心的事。
你只知道一条约定:fn 接一个参数。那你就写 fn(element),剩下的交给调用者。
这就是高阶函数的全部心智负担——比想象中少得多。
代码
def my_map(fn, seq):
"""Applies fn onto each element in seq and returns a list.
>>> my_map(lambda x: x*x, [1, 2, 3])
[1, 4, 9]
>>> my_map(lambda x: abs(x), [1, -1, 5, 3, 0])
[1, 1, 5, 3, 0]
>>> my_map(lambda x: print(x), ['cs61a', 'summer', '2023'])
cs61a
summer
2023
[None, None, None]
"""
return [fn(element) for element in seq]
逐处解释:
return—— 必须是return而不是print。列表推导式求值出一个值, 不return的话函数返回None,doctest 会显示Nothing而不是[1, 4, 9]。fn(element)—— 这是「收什么」的表达式,写在for前面。 注意不能写成fn(那会得到三个函数对象),也不能写成element(那就是原样复制一份 seq)。for element in seq——element是推导式内部的名字。 它不会泄漏到函数的其他地方:列表推导式在 Python 3 里有自己的作用域(每次迭代在一个隐含的帧里绑定element), 所以推导式结束后element这个名字在函数里并不存在。- 没有
if子句 —— 这一点和 Q2 形成对照:映射要保留全部元素,一个都不能少。
验证
手动追踪第三个 doctest,my_map(lambda x: print(x), ['cs61a', 'summer', '2023']):
fn → 那个 lambda 函数对象,seq → ['cs61a', 'summer', '2023']。[]。element = 'cs61a',求值 fn('cs61a')。进入 lambda 体,执行 print('cs61a'):屏幕出现 cs61a,print 返回 None;lambda 把这个 None 作为自己的返回值。新列表变成 [None]。element = 'summer',同理,屏幕出现 summer,列表变成 [None, None]。element = '2023',屏幕出现 2023,列表变成 [None, None, None]。[None, None, None],被 return 出去。None 的返回值,把它的 repr 打印出来:[None, None, None]。于是屏幕上的完整内容正是 doctest 里写的四行:三行打印 + 一行返回值。
注意顺序——三行打印全部发生在 函数返回之前,因为推导式是立即(eager)求值的,
遍历完才轮到 return。这一点和后面学的生成器表达式 (fn(x) for x in seq) 截然不同,
那种写法在你迭代它之前一个 print 都不会执行。
- 写成
[fn for element in seq]:漏了括号,收的是函数对象本身。 第一个 doctest 会得到[<function <lambda> at 0x...>, ..., ...],三个一模一样的东西。 - 用
seq.append(...)就地修改原列表:这会一边遍历一边往里加元素, 在 Python 里对列表这么干会导致无限循环(每加一个元素,遍历的终点就往后挪一格)。 题目要的是返回一个新列表,不要碰原来的。 - 先写
result = []再return:语义正确,但python3 ok -q my_map_syntax_check会失败,因为[type(x).__name__ for x in ast.parse(...).body[0].body]得到的是['Expr', 'Assign', 'For', 'Return']而不是要求的['Expr', 'Return']。 这个检查读的是你的源代码,所以「能跑对」不等于「能过」。
2. Q2 my_filter
题目要什么
my_filter(pred, seq):pred 是一个谓词函数(predicate)——
接一个参数、返回 True 或 False 的函数。
返回一个新列表,只保留那些让 pred 为真的元素,顺序不变。同样只准写一行函数体。
>>> my_filter(lambda x: x % 2 == 0, [1, 2, 3, 4])
[2, 4]
>>> my_filter(lambda x: (x + 5) % 3 == 0, [1, 2, 3, 4, 5])
[1, 4]
>>> my_filter(lambda x: print(x), [1, 2, 3, 4, 5])
1
2
3
4
5
[]
>>> my_filter(lambda x: max(5, x) == 5, [1, 2, 3, 4, 5, 6, 7])
[1, 2, 3, 4, 5]
四个 doctest 里,第三个又是陷阱题,而且这次的陷阱和 Q1 不一样。
lambda x: print(x) 返回的是 None。None 不是 True,
但它也不是 False——它是一个假值(falsy value)。
Python 的 if 判断的不是「等不等于 True」,而是「真值性(truthiness)」,
而 None、0、''、[]、{} 都是假值。
所以五个元素全都被过滤掉,结果是空列表 []——但五次 print 都实实在在发生了。
这个 doctest 的教学意图很明确:提醒你 pred 的返回值是被当条件用的,
不是被当值收进列表的。这正是 filter 和 map 的根本区别。
第四个 doctest lambda x: max(5, x) == 5 值得多看一眼:
max(5, x) 在 x <= 5 时是 5,在 x > 5 时是 x。
所以这个谓词等价于「x 小于等于 5」,于是保留 1 到 5,丢掉 6 和 7。
它是在提醒你:谓词可以是任意复杂的表达式,你不需要看懂它就能写对 my_filter。
怎么想到的
如果你已经写完了 Q1,这题的思维路径是「找出哪里不一样」。
先写出循环版:
def my_filter(pred, seq): # 这不是最终答案
result = []
for element in seq:
if pred(element):
result.append(element)
return result
把它和 Q1 的循环版并排看,差别只有两点:append 的是 element 而不是 fn(element);
多了一层 if。列表推导式恰好为这两点各准备了一个位置:
收什么(for 之前) | 遍历什么 | 要不要(for 之后) | |
|---|---|---|---|
my_map | fn(element) ← 变换 | for element in seq | 无 —— 全都要 |
my_filter | element ← 原样 | for element in seq | if pred(element) ← 筛选 |
列表推导式的完整形态是 [表达式 for 名字 in 序列 if 条件]。
表达式管「变成什么」,if 管「留不留」,两者互相独立。
Q1 只用了前者,Q2 只用了后者,两个都用就能一句话写出「把所有偶数平方」这种复合操作。
这里有个真会绊人的弯路:不少人第一次会把 if 写到前面去,写成
[element if pred(element) for element in seq]。这是语法错误,
Python 会报 SyntaxError: expected 'else' after 'if' expression。
原因是 for 之前的位置只能放表达式,而 A if C else B
(条件表达式)必须带 else——它总要求值出一个值来,没有「什么都不产生」这个选项。
真正的「跳过」只能由 for 之后的过滤子句完成。
代码
def my_filter(pred, seq):
"""Keeps elements in seq only if they satisfy pred.
>>> my_filter(lambda x: x % 2 == 0, [1, 2, 3, 4]) # new list has only even-valued elements
[2, 4]
>>> my_filter(lambda x: (x + 5) % 3 == 0, [1, 2, 3, 4, 5])
[1, 4]
>>> my_filter(lambda x: print(x), [1, 2, 3, 4, 5])
1
2
3
4
5
[]
>>> my_filter(lambda x: max(5, x) == 5, [1, 2, 3, 4, 5, 6, 7])
[1, 2, 3, 4, 5]
"""
return [element for element in seq if pred(element)]
element(第一个)—— 原样收进去。不是pred(element); 如果写成那样,第一个 doctest 会得到[True, True]而不是[2, 4]。if pred(element)—— 注意不要写成if pred(element) == True。 两个原因:一是啰嗦;二是它会改变语义——如果某个谓词返回1(而不是True),if 1为真但1 == True恰好也为真,看起来没事; 可要是返回的是非空字符串'yes',if 'yes'为真而'yes' == True为假,行为就分岔了。 直接用真值性是唯一正确的写法。- 求值次数 ——
pred(element)每个元素恰好被调用一次。 这在谓词有副作用(比如第三个 doctest 的print)时是可观察的:屏幕上恰好五行。
验证
追踪第二个 doctest,my_filter(lambda x: (x + 5) % 3 == 0, [1, 2, 3, 4, 5]):
element | element + 5 | % 3 | pred 结果 | 新列表 |
|---|---|---|---|---|
| 1 | 6 | 0 | True | [1] |
| 2 | 7 | 1 | False | [1] |
| 3 | 8 | 2 | False | [1] |
| 4 | 9 | 0 | True | [1, 4] |
| 5 | 10 | 1 | False | [1, 4] |
返回 [1, 4],与 doctest 一致。
再追踪第三个:pred 是 lambda x: print(x)。
element = 1 时求值 pred(1) —— 屏幕出现 1,返回值 None;
if None 为假,所以 1 不进列表。2 到 5 同理。
遍历完屏幕上有五行数字,新列表始终是空的,返回 []。
- 以为
None会报错:if None:完全合法,只是走 else 分支。 真正会报错的是把None当数用,比如None + 1→TypeError: unsupported operand type(s) for +: 'NoneType' and 'int'。 - 把过滤条件写死,比如
if element % 2 == 0。这样只能过第一个 doctest, 后面三个全挂。pred是参数,不是常量。 - 调用两次
pred,比如写成[element for element in seq if pred(element) and pred(element)](有人为了「保险」)。 第三个 doctest 会打印十行而不是五行,直接暴露。
3. Q3 my_reduce
题目要什么
my_reduce(combiner, seq):combiner 是一个两参数函数,
seq 保证非空。要把整个序列「折叠」成一个值。
>>> my_reduce(lambda x, y: x + y, [1, 2, 3, 4]) # 1 + 2 + 3 + 4
10
>>> my_reduce(lambda x, y: x * y, [1, 2, 3, 4]) # 1 * 2 * 3 * 4
24
>>> my_reduce(lambda x, y: x * y, [4])
4
>>> my_reduce(lambda x, y: x + 2 * y, [1, 2, 3]) # (1 + 2 * 2) + 2 * 3
11
前两个 doctest 有点误导性:加法和乘法都满足结合律和交换律,
所以不管你从左往右折还是从右往左折,结果都一样。
第四个 doctest 才把语义钉死:注释写的是 (1 + 2 * 2) + 2 * 3,
括号明确告诉你要从左往右折——先把 1 和 2 合成 5,再把 5 和 3 合成 11。
如果你从右往左折,得到的是 1 + 2 * (2 + 2 * 3) = 1 + 16 = 17,就错了。
第三个 doctest my_reduce(lambda x, y: x * y, [4]) → 4 是边界情况:
序列只有一个元素时,combiner 一次都不会被调用,直接返回那个元素本身。
这条约束会直接影响你的写法——它排除了「从某个初始值开始」的方案。
怎么想到的
先想清楚一件事:这题跟 Q1、Q2 有本质区别。 map 和 filter 是「每个元素独立处理」,元素之间不通信; reduce 是「结果要一路累积下来」,第 k 步的输入依赖第 k-1 步的输出。 所以列表推导式在这里帮不上忙——推导式没有办法让一轮看到上一轮的结果。
弯路一:想套初始值。 很自然的想法是仿照「求和从 0 开始、求积从 1 开始」写成:
result = 0 # 行不通
for element in seq:
result = combiner(result, element)
这在 combiner 是加法时对,是乘法时就错了(0 乘任何数都是 0)。
根本问题是:不同的 combiner 有不同的单位元(identity),
而你作为 my_reduce 的作者根本不知道 combiner 是什么。
求和的单位元是 0,求积是 1,字符串拼接是 '',而
lambda x, y: x + 2 * y 压根没有单位元。所以这条路死了。
弯路二:想不出初始值,那就不要初始值。 既然题目保证 seq 非空,
那就拿第一个元素当起点。这一步是本题的关键转折:
初始值不是从外面给的,是从数据里取的。
result = seq[0],然后从第二个元素开始往里折。
「非空」这个前提正是为了让 seq[0] 合法而写的——
如果 seq 可以为空,seq[0] 会抛 IndexError: list index out of range,
而且此时根本没有正确答案可返回。题目提前替你排除了这种情况。
接下来只剩一个技术问题:怎么「从第二个开始遍历」。有两个办法:
| 写法 | 含义 | 评价 |
|---|---|---|
for element in seq[1:] | 切片,得到去掉首元素的新序列 | 读起来最直白,本文采用 |
for i in range(1, len(seq)) | 按下标走,循环体里用 seq[i] | 也对,但多一层间接 |
切片 seq[1:] 会复制一份(对列表来说是浅拷贝),对本题的规模无所谓,换来的是可读性。
顺带一提,切片对空列表是安全的:[4][1:] 得到 [],循环一轮都不跑——
这正好让第三个 doctest 自动正确,不需要单独写 if len(seq) == 1。
还有一条递归的思路,也完全可行:
def my_reduce(combiner, seq): # 另一种写法,本仓库没用
if len(seq) == 1:
return seq[0]
return combiner(my_reduce(combiner, seq[:-1]), seq[-1])
注意这个递归版必须从尾部剥(seq[:-1] 和 seq[-1]),
才能保持「左折」的语义。如果写成 combiner(seq[0], my_reduce(combiner, seq[1:]))
就变成右折了,第四个 doctest 会输出 17。这个细节很能说明「结合方向」不是可以随便选的。
代码
def my_reduce(combiner, seq):
"""Combines elements in seq using combiner.
seq will have at least one element.
>>> my_reduce(lambda x, y: x + y, [1, 2, 3, 4]) # 1 + 2 + 3 + 4
10
>>> my_reduce(lambda x, y: x * y, [1, 2, 3, 4]) # 1 * 2 * 3 * 4
24
>>> my_reduce(lambda x, y: x * y, [4])
4
>>> my_reduce(lambda x, y: x + 2 * y, [1, 2, 3]) # (1 + 2 * 2) + 2 * 3
11
"""
# Start from the first element, then fold in the rest one at a time.
result = seq[0]
for element in seq[1:]:
result = combiner(result, element)
return result
result = seq[0]—— 起点取自数据本身,绕开了「不知道单位元」的死结。seq[1:]—— 跳过已经用掉的首元素。写成seq会导致首元素被折进去两次, 第一个 doctest 会得到11(多加了一个 1)。result = combiner(result, element)—— 参数顺序是 「累积值在前,新元素在后」。这条决定了折叠方向。 若写成combiner(element, result),第四个 doctest 会算成(2 + 2*1) = 4,再(3 + 2*4) = 11——巧的是也等于 11, 但第二步的中间值完全不同,换个 combiner 就会露馅(比如减法)。return result在循环外面 —— 这是极常见的错误位置。 缩进进循环里的话,第一轮就返回了,[1,2,3,4]求和会得到3。
验证
追踪第四个 doctest:my_reduce(lambda x, y: x + 2 * y, [1, 2, 3])。
result = seq[0] = 1
seq[1:] = [2, 3]
第 1 轮 element = 2
result = combiner(1, 2) = 1 + 2*2 = 1 + 4 = 5
第 2 轮 element = 3
result = combiner(5, 3) = 5 + 2*3 = 5 + 6 = 11
循环结束 return 11
和 docstring 注释 (1 + 2 * 2) + 2 * 3 完全对应:括号里那部分就是第 1 轮的 result。
再看第三个:my_reduce(lambda x, y: x * y, [4])。
result = seq[0] = 4;seq[1:] 是 [],for 循环体一次都不执行;
直接 return 4。combiner 一次都没被调用——这正是期望行为。
result = 0起手:第一个 doctest 侥幸通过(0 是加法单位元), 第二个立刻返回0。这个 bug 的可怕之处在于它「部分正确」,容易让人以为只是小问题。- 忘了
[1:]:my_reduce(lambda x, y: x * y, [1,2,3,4])仍然返回 24(因为首元素恰好是 1),但把序列换成[2,3]就会得到 12 而不是 6。 用第四个 doctest 一测就现形:会算成((1 + 2*1) + 2*2) + 2*3 = 13。 - 试图用列表推导式一行写完:本题没有要求一行,
题面里 Q3 的提示框写的是
"*** YOUR CODE HERE ***"而非「Use only a single line」。 硬凑推导式只会写出看不懂的东西。
4. Q4 distance
题目要什么
课程给了一个「城市」的数据抽象(data abstraction):
| 角色 | 函数 | 作用 |
|---|---|---|
| 构造器 constructor | make_city(name, lat, lon) | 把名字、纬度、经度打包成一个 city |
| 选择器 selector | get_name(city) | 取出名字 |
| 选择器 selector | get_lat(city) | 取出纬度 |
| 选择器 selector | get_lon(city) | 取出经度 |
要写 distance(city_a, city_b),返回两座城市坐标之间的欧几里得距离,
即 $\sqrt{(x_1-x_2)^2 + (y_1-y_2)^2}$。sqrt 已经从 math 导入好了。
>>> city_a = make_city('city_a', 0, 1)
>>> city_b = make_city('city_b', 0, 2)
>>> distance(city_a, city_b)
1.0
>>> city_c = make_city('city_c', 6.5, 12)
>>> city_d = make_city('city_d', 2.5, 15)
>>> distance(city_c, city_d)
5.0
注意期望值是 1.0 和 5.0 而不是 1 和 5:
sqrt 永远返回浮点数,所以哪怕结果是整数,也会显示成小数。
如果你自作聪明加个 int(...),doctest 会因为显示成 1 而失败。
这道题数学部分几乎没有难度,难度全在「怎么拿到坐标」。 这就是它真正的考点。
怎么想到的
公式里的 x1, y1, x2, y2 得从两个 city 对象里挖出来。
问题是——city 到底长什么样?
如果你翻到 lab04.py 的下半部分,会看到 make_city 的实现里赫然写着
return [name, lat, lon]。于是「聪明」的写法呼之欲出:
def distance(city_a, city_b): # 能过前两个 doctest,但会挂 Q6
return sqrt((city_a[1] - city_b[1]) ** 2 + (city_a[2] - city_b[2]) ** 2)
这段代码确实能通过 Q4 的 doctest。它的问题是: 你把「city 是一个三元素列表、纬度在下标 1」这个实现细节写进了自己的代码。 一旦有人把 city 改成字典存储,你的代码就会炸——而 Q6 干的就是这件事。
正确的心态是:假装你根本没看过 make_city 的实现。
你只知道题面上写的那四行接口说明。要纬度就调 get_lat,要经度就调 get_lon,
至于它们内部是查列表还是查字典,不关你事。
抽象屏障的意义不是「让代码更好看」,而是让两部分代码可以独立演化。
屏障上方(distance、closer_city)只依赖接口的行为;
屏障下方(make_city 等)可以随意改表示。
只要没人跨越屏障,改下面不需要动上面。这条原则在 61A 后半程(数据抽象、面向对象、解释器)会反复出现。
确定用选择器之后,还有个小选择:是把四个差值直接塞进一个大表达式,还是先起名字?
return sqrt((get_lat(city_a) - get_lat(city_b)) ** 2 + (get_lon(city_a) - get_lon(city_b)) ** 2)
这行是对的,但括号多到一眼数不清,很容易把 **2 写到括号外面变成
get_lat(city_a) - get_lat(city_b) ** 2——那是先平方再相减,结果完全不同,而且不会报错。
起两个中间名字能显著降低这种风险。
代码
from math import sqrt
def distance(city_a, city_b):
"""
Returns the distance between city_a and city_b according to their
coordinates.
>>> city_a = make_city('city_a', 0, 1)
>>> city_b = make_city('city_b', 0, 2)
>>> distance(city_a, city_b)
1.0
>>> city_c = make_city('city_c', 6.5, 12)
>>> city_d = make_city('city_d', 2.5, 15)
>>> distance(city_c, city_d)
5.0
"""
lat_difference = get_lat(city_a) - get_lat(city_b)
lon_difference = get_lon(city_a) - get_lon(city_b)
return sqrt(lat_difference ** 2 + lon_difference ** 2)
get_lat(city_a) - get_lat(city_b)—— 两次调用同一个选择器,作用在不同的 city 上。 差值的正负无所谓,因为马上要平方。lat_difference/lon_difference—— 局部名字,只是为了让最后一行读得懂。 它们和抽象屏障无关,纯粹是可读性。lat_difference ** 2——**的优先级高于+, 所以a ** 2 + b ** 2就是「两个平方相加」,不需要额外括号。 但注意**的优先级也高于一元负号:-2 ** 2是-4不是4。 这里因为先算好了差值再平方,不受影响。- 没有出现任何方括号下标 —— 这是这段代码最重要的性质。
验证
追踪第二个 doctest:city_c = make_city('city_c', 6.5, 12),
city_d = make_city('city_d', 2.5, 15)。
get_lat(city_c) = 6.5 get_lat(city_d) = 2.5 get_lon(city_c) = 12 get_lon(city_d) = 15 lat_difference = 6.5 - 2.5 = 4.0 lon_difference = 12 - 15 = -3 sqrt(4.0 ** 2 + (-3) ** 2) = sqrt(16.0 + 9) = sqrt(25.0) = 5.0
返回 5.0,与 doctest 一致。这是个 3-4-5 直角三角形,出题人特意挑的整数结果,
方便你一眼看出对错。注意 lon_difference 是负数 -3,
平方之后变成 9——这就是为什么公式里要平方而不是取差:距离没有方向。
- 用
city_a[1]取纬度:Q4 会通过,Q6 的check_city_abstraction会失败并报KeyError: 1—— 因为那时 city 已经变成字典,键是"lat"而不是1。详见下一节。 - 用
abs()代替平方:sqrt(abs(dx) + abs(dy))在第一个 doctest 里得到sqrt(0 + 1) = 1.0,居然过了; 第二个得到sqrt(4 + 3) ≈ 2.6458,才暴露。这是「用一个 doctest 就以为写对了」的经典教训。 - 把纬度经度搞混:因为要平方求和,一致地搞混其实不影响结果。
但如果只混一半(比如
get_lat(city_a) - get_lon(city_b)),结果就错了。
5. Q5 closer_city
题目要什么
closer_city(lat, lon, city_a, city_b):给一个坐标 (lat, lon) 和两座城市,
返回离这个坐标更近的那座城市的名字(字符串,不是 city 对象)。
如果两者一样远,规定算 city_b 更近。
题面还明确限定了可用的工具:只能用 get_name、get_lat、get_lon、
make_city 和你刚写的 distance。
>>> berkeley = make_city('Berkeley', 37.87, 112.26)
>>> stanford = make_city('Stanford', 34.05, 118.25)
>>> closer_city(38.33, 121.44, berkeley, stanford)
'Stanford'
>>> bucharest = make_city('Bucharest', 44.43, 26.10)
>>> vienna = make_city('Vienna', 48.20, 16.37)
>>> closer_city(41.29, 174.78, bucharest, vienna)
'Bucharest'
「平局算 city_b」这条约束看起来无关紧要(浮点数很少精确相等),
但它决定了你的比较该用 < 还是 <=,是个真实的边界条件。
怎么想到的
第一反应:我需要「坐标到 city_a 的距离」和「坐标到 city_b 的距离」,比一下谁小。
但立刻卡住了——distance 接的是两个 city,不是一个 city 和一对坐标。
我手上的 lat, lon 是两个裸的数字,喂不进去。
有三条路可以走:
| 方案 | 做法 | 问题 |
|---|---|---|
| A. 重写距离公式 | 在 closer_city 里再写一遍 sqrt((lat - get_lat(city_a))**2 + ...) | 能跑,但把同一段数学抄了两遍。distance 白写了 |
B. 改 distance 的签名 | 让它收四个数字 | 会破坏 Q4 的 doctest,而且题面没让你改 |
| C. 把坐标包装成一个 city | make_city('here', lat, lon),然后正常调 distance | 没有问题——这是正解 |
题面把 make_city 列进「你可以用的东西」里,就是在给这个提示。
当你手上的数据形状和现成函数的接口对不上时,先想想能不能把数据「造成」那个形状,
而不是去改函数。 这个动作在工程里叫 adapter(适配),在 61A 里以后会以
「用构造器补齐一个中间值」的形式反复出现。
顺带一提,这里的 'here' 这个名字是随便起的——它永远不会被 get_name 读出来,
只是构造器要求有个位置参数。
解决了距离怎么算,剩下的就是比较和返回什么。两个坑:
坑一:返回名字还是返回 city? doctest 期望 'Stanford'(带引号的字符串),
所以必须 return get_name(city_a) 而不是 return city_a。
后者在当前实现下会显示成 ['Stanford', 34.05, 118.25]。
坑二:平局怎么办? 规定平局算 city_b 近。
所以判断「a 更近」的条件必须是严格小于 <:
只有 a 严格更近时才返回 a,其余情况(b 更近、或者一样近)都返回 b。
如果写成 <=,平局时会返回 a,违反规定。
代码
def closer_city(lat, lon, city_a, city_b):
"""
Returns the name of either city_a or city_b, whichever is closest to
coordinate (lat, lon). If the two cities are the same distance away
from the coordinate, consider city_b to be the closer city.
>>> berkeley = make_city('Berkeley', 37.87, 112.26)
>>> stanford = make_city('Stanford', 34.05, 118.25)
>>> closer_city(38.33, 121.44, berkeley, stanford)
'Stanford'
>>> bucharest = make_city('Bucharest', 44.43, 26.10)
>>> vienna = make_city('Vienna', 48.20, 16.37)
>>> closer_city(41.29, 174.78, bucharest, vienna)
'Bucharest'
"""
# Build a city for the given coordinate so we can reuse `distance`.
here = make_city('here', lat, lon)
if distance(here, city_a) < distance(here, city_b):
return get_name(city_a)
return get_name(city_b)
here = make_city('here', lat, lon)—— 适配那一步。 只调用一次,之后两次distance复用同一个对象。<而不是<=—— 直接对应「平局算 b」。- 没有
else——if分支里已经return了, 控制流不可能落到后面去;所以第二个return写在函数体顶层即可。 写else: return ...也完全正确,只是多一层缩进。 get_name(...)—— 又一次只用选择器。整个函数里没有一个下标。
验证
追踪第一个 doctest:closer_city(38.33, 121.44, berkeley, stanford)。
berkeley 的坐标是 (37.87, 112.26),stanford 是 (34.05, 118.25)。
here 的坐标 = (38.33, 121.44)
distance(here, berkeley):
lat_difference = 38.33 - 37.87 = 0.46
lon_difference = 121.44 - 112.26 = 9.18
sqrt(0.46**2 + 9.18**2) = sqrt(0.2116 + 84.2724)
= sqrt(84.484) ≈ 9.1915
distance(here, stanford):
lat_difference = 38.33 - 34.05 = 4.28
lon_difference = 121.44 - 118.25 = 3.19
sqrt(4.28**2 + 3.19**2) = sqrt(18.3184 + 10.1761)
= sqrt(28.4945) ≈ 5.3381
9.1915 < 5.3381 ? → False
所以不进 if,执行 return get_name(city_b) = 'Stanford'
返回 'Stanford',与 doctest 一致。可以看到主要差距来自经度:
Berkeley 的经度 112.26 离 121.44 差了 9 度多,而 Stanford 的 118.25 只差 3 度多。
第二个 doctest 同理:目标点经度 174.78 离得很远,
Bucharest 经度 26.10(差 148.68),Vienna 经度 16.37(差 158.41),
纬度差分别是 3.14 和 6.91。Bucharest 的两项都更小,所以更近,返回 'Bucharest'。
这里 if 条件为真,走的是第一个 return。
- 返回 city 对象而不是名字:doctest 报
Expected: 'Stanford' Got: ['Stanford', 34.05, 118.25]。看到这个报错就知道少了个get_name。 - 直接把
(lat, lon)元组传给distance:distance((lat, lon), city_a)会在get_lat里执行city[1], 拿到的是lon,get_lon执行city[2]直接IndexError: tuple index out of range。 - 用
min硬凑:min(city_a, city_b, key=...)在这里 一是没学到,二是min在平局时返回第一个参数,恰好和题目要求相反。
6. Q6 抽象屏障检查 check_city_abstraction
题目要什么
这题没有代码要写——如果 Q4、Q5 写对了的话。 它是一个自动化的「你有没有偷看实现」检测器。运行:
python3 ok -q check_city_abstraction
它会把 city 的底层表示从列表偷偷换成字典,然后把 Q4、Q5 的 doctest 重跑一遍。
如果你的代码只用了构造器和选择器,换表示对你毫无影响,测试通过;
如果你用了 city[1] 之类的下标,代码立刻崩掉。
它是怎么做到的
秘密全在 lab04.py 底部那段「Treat all the following code as being behind an
abstraction layer」的代码里。构造器和选择器各有两套实现,
由一个全局开关 change_abstraction.changed 决定走哪一套:
def make_city(name, lat, lon):
if change_abstraction.changed:
return {"name" : name, "lat" : lat, "lon" : lon}
else:
return [name, lat, lon]
def get_lat(city):
if change_abstraction.changed:
return city["lat"]
else:
return city[1]
(get_name、get_lon 同构,分别取 "name" / city[0]
和 "lon" / city[2]。)
而 change_abstraction 是一个函数,它给自己挂了一个属性:
def change_abstraction(change):
"""
For testing purposes.
>>> change_abstraction(True)
>>> change_abstraction.changed
True
"""
change_abstraction.changed = change
change_abstraction.changed = False
Python 里函数也是对象,对象可以带属性。change_abstraction.changed = False
这行在模块加载时执行,给函数对象贴上一个标记;
之后 make_city 每次被调用都去读这个标记。
效果相当于一个「不会被别人的局部变量遮住」的全局开关。
这在 61A 后面讲「函数是一等公民」时会正式展开——现在你只要知道它是合法的就够了。
整个检查的 doctest 长这样(就是 check_city_abstraction 的 docstring):
>>> change_abstraction(True)
>>> city_a = make_city('city_a', 0, 1)
>>> city_b = make_city('city_b', 0, 2)
>>> distance(city_a, city_b)
1.0
...
>>> closer_city(41.29, 174.78, bucharest, vienna)
'Bucharest'
>>> change_abstraction(False)
头尾各一句:先打开开关,跑一遍前两题的全部 doctest,最后关掉开关 (不关的话后面的测试都会在字典模式下跑,互相干扰)。
如果你违反了屏障,会发生什么
假设你的 distance 写的是 city_a[1] - city_b[1]。
开关打开后,make_city 返回的是字典 {"name": 'city_a', "lat": 0, "lon": 1}。
于是 city_a[1] 变成「在字典里查键 1」——
KeyError: 1
因为这个字典的键是 "name"、"lat"、"lon" 三个字符串,
根本没有 1 这个键。注意它不是 TypeError:
字典的 [] 语法本身是合法的,只是键不存在。这个报错很有辨识度——
一看到 KeyError: 1 或 KeyError: 2,就知道是哪一行在用下标取坐标。
验证
我们的实现从头到尾只用了 get_lat、get_lon、get_name、make_city。
在字典模式下追踪一遍 distance(city_a, city_b)(city_a 坐标 (0, 1),
city_b 坐标 (0, 2)):
city_a = {"name": 'city_a', "lat": 0, "lon": 1}
city_b = {"name": 'city_b', "lat": 0, "lon": 2}
get_lat(city_a) → changed 为真 → city_a["lat"] → 0
get_lat(city_b) → city_b["lat"] → 0
lat_difference = 0 - 0 = 0
get_lon(city_a) → city_a["lon"] → 1
get_lon(city_b) → city_b["lon"] → 2
lon_difference = 1 - 2 = -1
sqrt(0 ** 2 + (-1) ** 2) = sqrt(1) = 1.0
结果 1.0,和列表模式下算出来的一模一样。
这正是抽象屏障承诺的东西:换实现不影响使用者。
把上面这段和 Q4 验证那段并排看——除了「选择器内部做了什么」那一列不同,
其余每一步的数值完全相同。distance 的作者永远不需要知道差别在哪。
本仓库 ok --score 中 check_city_abstraction: 1.0/1,即通过。
这题真正想让你记住的不是「别写下标」这条禁令,而是为什么: 在真实项目里,「换掉底层表示」是极常见的需求(列表换成字典、字典换成类、内存换成数据库)。 如果十几个函数都直接摸了底层,你就得改十几处,还得指望自己没漏; 如果它们都只走选择器,你只改选择器那三行。抽象屏障买的是未来的修改成本。
- 「我 Q4 Q5 都过了,Q6 不用管」:Q6 是单独计分的一题
(
check_city_abstraction: 1.0/1),过不了就丢分。 - 为了过 Q6,把
get_lat也改了:那是屏障下方的代码, 文件里明写「you shouldn't need to look at it」。改它等于改考题。 - 在
closer_city里直接比较 city 对象,比如if city_a == city_b。列表可以比,字典也可以比,但语义是「内容是否相同」, 和「哪个更近」毫无关系,而且同样属于跨越屏障——接口里没有「比较两个 city」这个操作。
7. Q7 WWPD:Trees
在做树的编程题之前,先把树的数据抽象摸清楚。这一节是
What Would Python Display(Python 会显示什么)概念题,
本仓库对应文件是 labs/lab04/tests/wwpd.py,答案已解锁为明文。
运行方式是 python3 ok -q wwpd -u。它不计分,但这是整份 lab 里性价比最高的部分——
把这几行搞明白,后面三道编程题会顺得多。
先把接口写清楚
lab04.py 底部给出的树抽象是这样实现的(同样在屏障下方,看一眼是为了理解,写题时别用):
def tree(label, branches=[]):
"""Construct a tree with the given label value and a list of branches."""
for branch in branches:
assert is_tree(branch), 'branches must be trees'
return [label] + list(branches)
def label(tree):
return tree[0]
def branches(tree):
return tree[1:]
def is_leaf(tree):
return not branches(tree)
三条必须记住的事实:
- 树被表示成一个列表:第 0 个元素是标签,从第 1 个起全是分支。
所以
branches(t)返回的是一个列表,里面每个元素本身又是一棵树。 branches参数必须是一个列表。tree(1, tree(2))是错的,tree(1, [tree(2)])才对。因为构造器会for branch in branches逐个断言。- 叶子(leaf)就是没有分支的树。
tree(5)用了默认参数branches=[], 得到[5],branches([5])是[],is_leaf返回True。 注意叶子仍然是一棵合法的树,不是「光秃秃的一个数」。
第一组
>>> from lab04 import *
>>> t = tree(1, tree(2))
______
>>> t = tree(1, [tree(2)])
______
>>> label(t)
______
>>> label(branches(t)[0])
______
>>> x = branches(t)
>>> len(x)
______
>>> is_leaf(x[0])
______
>>> branch = x[0]
>>> label(t) + label(branch)
______
>>> len(branches(branch))
______
答案依次是:Error、Nothing、1、2、1、
True、3、0。逐条讲为什么。
t = tree(1, tree(2)) → Error。
内层 tree(2) 先求值,得到 [2],这是一棵合法的树。
然后它被当作整个 branches 列表传给外层。
外层执行 for branch in branches,即 for branch in [2],
所以 branch 被绑成整数 2。接着 is_tree(2):
type(2) != list 为真,直接返回 False,
于是 assert 失败,抛出
AssertionError: branches must be trees。
这就是构造器里那句 assert 存在的全部意义——它把「忘了套方括号」这个错误
在构造的瞬间就抓住,而不是等到你几十行之后取分支时才莫名其妙地崩。t = tree(1, [tree(2)]) → Nothing。
这次分支列表是 [[2]],唯一的元素 [2] 通过了 is_tree,
构造成功,t 是 [1, [2]]。
但整行是一条赋值语句(assignment statement),不是表达式。
语句不产生值,交互式解释器什么也不显示。
这就是为什么答案是 Nothing 而不是 [1, [2]]。
顺带说明第 1 行为什么会 Error 而不是 Nothing:赋值语句要先求值右边,
求值过程中就炸了,压根轮不到赋值。label(t) → 1。
t 是 [1, [2]],label 取 t[0],即 1。label(branches(t)[0]) → 2。
从里往外算:branches(t) 是 t[1:],即 [[2]]——
一个装着一棵树的列表。
[0] 取出第一棵(也是唯一一棵)分支,得到 [2]。
label([2]) 取 [2][0],即 2。
这里最容易混的是「分支列表」和「一个分支」的区别:
branches(t) 是列表,branches(t)[0] 才是树。len(x) → 1。
x = branches(t) = [[2]],长度 1,因为 t 只有一个分支。
注意别把它和 len(t) 混了:len([1, [2]]) 也是 2,
因为标签也占一格。len(branches(t)) 才是「有几个分支」。is_leaf(x[0]) → True。
x[0] 是 [2]。branches([2]) 是 [2][1:],即 []。
not [] 为 True(空列表是假值)。所以它是叶子。label(t) + label(branch) → 3。
branch = x[0] = [2]。label(t) = 1,label(branch) = 2,
1 + 2 = 3。这行想强调的是:标签是普通的值,可以直接参与算术;
而树本身不行——t + branch 会得到 [1, [2], 2](列表拼接),
是个毫无意义的东西,而且跨越了抽象屏障。len(branches(branch)) → 0。
branch 是叶子,分支列表为空,长度 0。和第 6 行是同一件事的两种问法。第二组
>>> from lab04 import *
>>> b1 = tree(5, [tree(6), tree(7)])
>>> b2 = tree(8, [tree(9, [tree(10)])])
>>> t = tree(11, [b1, b2])
>>> for b in branches(t):
... print(label(b))
______
>>> for b in branches(t):
... print(is_leaf(branches(b)[0]))
...
______
>>> [label(b) + 100 for b in branches(t)]
______
>>> [label(b) * label(branches(b)[0]) for b in branches(t)]
______
先把树画出来。这一步别省,画完之后四问几乎不用想:
11 <- t
/ \
5 8 <- b1, b2
/ \ \
6 7 9
\
10
底层表示:
b1 = [5, [6], [7]]
b2 = [8, [9, [10]]]
t = [11, [5, [6], [7]], [8, [9, [10]]]]5 换行 8。
branches(t) 是 [b1, b2]。循环两轮:
b = b1 时 label(b1) = 5,打印 5;
b = b2 时 label(b2) = 8,打印 8。
for 语句本身不产生值,所以除了这两行打印之外没有别的输出。print(is_leaf(branches(b)[0])) → True 换行 False。
这行套了三层,从里往外拆:
·
b = b1 = [5, [6], [7]]:branches(b1) 是 [[6], [7]],
[0] 取出 [6]——也就是标签为 6 的那个叶子。
is_leaf([6]) 为 True。
·
b = b2 = [8, [9, [10]]]:branches(b2) 是 [[9, [10]]],
[0] 取出 [9, [10]]——标签为 9、还带着一个孩子 10 的那棵子树。
is_leaf 为 False。
翻成人话就是:「t 的每个孙辈里最左边那个,是不是叶子?」 b1 的最左孙是 6(是叶子),b2 的最左孙是 9(不是,它下面还有 10)。
[label(b) + 100 for b in branches(t)] → [105, 108]。
这就是 Q1 的 my_map,只不过映射的是树的分支:
5 + 100 = 105,8 + 100 = 108。
这次是表达式不是语句,所以解释器把返回的列表显示出来。
「对每个分支做点什么,收成一个列表」——记住这个形状,Q8 马上就要用。[label(b) * label(branches(b)[0]) for b in branches(t)] → [30, 72]。
「每个分支的标签,乘以它第一个孩子的标签」:
·
b1:label(b1) = 5,branches(b1)[0] 是 [6],
标签 6,5 * 6 = 30。
·
b2:label(b2) = 8,branches(b2)[0] 是 [9, [10]],
标签 9,8 * 9 = 72。
结果
[30, 72]。注意 [10] 从头到尾没被碰过——
label 只看根,不管下面还挂着什么。处理树表达式的固定手法:从最里层的括号往外读,每一步问自己「现在手上这个东西是一棵树,还是一个树的列表,还是一个标签值?」
| 表达式 | 结果的类型 | 能对它做什么 |
|---|---|---|
t | 树 | label(t)、branches(t)、is_leaf(t) |
branches(t) | 树的列表 | len(...)、[...][i]、for b in ... |
branches(t)[0] | 树 | 又回到第一行,可以继续递归 |
label(t) | 普通值(这里是数) | 加减乘除、比较、打印 |
- 忘了给单个分支套方括号:
tree(1, tree(2))→AssertionError: branches must be trees。第一组第一行专门考这个, 说明这是历年最高频的错误。 - 把
branches(t)当成一棵树用:label(branches(t))会取[[2]][0],得到[2]—— 不报错,但返回的是一棵树而不是一个标签,后面拿去做加法才会崩,很难查。 - 分不清
Nothing和None:赋值语句、for语句、返回None的函数调用,在交互式解释器里都什么也不显示, WWPD 里一律写Nothing。只有当None出现在一个容器里 (比如 Q1 的[None, None, None])才会被打印出来。 - 以为
tree(9, [tree(10)])里的 9 是叶子: 叶子的定义是「没有分支」,9 有一个分支 10,所以不是。 第二组第 2 行的False就是在考这个。
8. Q8 sum_tree 与 balanced
8.1 sum_tree:题目要什么
返回树 t 中所有标签的总和——包括根、所有内部节点、所有叶子。
>>> t = tree(4, [tree(2, [tree(3)]), tree(6)])
>>> sum_tree(t)
15
这棵树是:
4
/ \
2 6
/
3
4 + 2 + 3 + 6 = 158.1 怎么想到的
「所有标签的总和」= 「根的标签」+「每个分支的总和之和」。 这句中文本身就是答案,问题只在于怎么把它写成 Python。
先确认这个拆分是真的正确,而不只是听起来顺: 一棵树的每个节点,要么是根,要么落在某一个分支里,而且只落在一个分支里 (树没有环、没有共享节点)。所以「根 + 各分支」不重不漏地覆盖了全部节点。这就是递归的合法性依据。
接下来是一个新手常卡的地方:base case 在哪? 按照写数字递归的习惯,你会想先写:
def sum_tree(t): # 可以,但不必要
if is_leaf(t):
return label(t)
return label(t) + sum([sum_tree(b) for b in branches(t)])
这是对的。但仔细看第二行和第三行:当 t 是叶子时,
branches(t) 是 [],列表推导式产出 [],
sum([]) 是 0,所以第三行返回 label(t) + 0 = label(t)——
和第二行完全一样。base case 是多余的。
树的递归常常不需要显式的 base case,因为「叶子」这个终止条件
已经藏在「branches(t) 是空列表 → 推导式产出空列表 → 循环零轮」里了。
这和数字递归很不一样:fact(n) 不写 if n == 0 会无限递归下去,
但 sum_tree 不写 if is_leaf(t) 照样会停。
什么时候还是要写 base case? 当叶子的处理方式和内部节点确实不同时——
Q10 的 only_paths 就是那种情况。
另一个选择是用 sum 还是自己累加。sum 是内置函数,
接一个数字序列返回总和,空序列返回 0。这里用它最直接。
不用它的话要写成:
total = label(t)
for b in branches(t):
total = total + sum_tree(b)
return total
三行变一行,语义完全一样,但题面的 Challenge 明确说「Solve both of these parts with just 1 line of code each」, 所以我们用推导式版本。
8.1 代码
def sum_tree(t):
"""Add all elements in a tree.
>>> t = tree(4, [tree(2, [tree(3)]), tree(6)])
>>> sum_tree(t)
15
"""
return label(t) + sum([sum_tree(b) for b in branches(t)])
label(t)—— 当前节点自己贡献的那一份。一定要加上, 漏了它的话所有内部节点都不算数,例子里会得到 9 而不是 15。[sum_tree(b) for b in branches(t)]—— 对每个分支递归。 这就是 Q7 第二组第 3 行那个形状([对 b 做点什么 for b in branches(t)]), 只不过「做的事」是递归调用自己。写的时候不要在脑子里展开它, 就当sum_tree(b)已经正确返回了分支 b 的总和。sum(...)—— 把「每个分支的总和」这个列表加起来。sum([])为 0,这一点让叶子自动正确。- 为什么调用的是
sum_tree(b)而不是label(b)? 因为分支下面可能还有子孙。写label(b)只会加上直接孩子,例子里得到4+2+6=12, 漏掉了孙子 3。
8.1 验证
把 sum_tree(tree(4, [tree(2, [tree(3)]), tree(6)])) 真的展开:
sum_tree([4, [2, [3]], [6]])
label = 4
branches = [ [2,[3]] , [6] ]
→ 需要 sum_tree([2,[3]]) 和 sum_tree([6])
sum_tree([2, [3]])
label = 2
branches = [ [3] ]
→ 需要 sum_tree([3])
sum_tree([3])
label = 3
branches = [] <- 叶子
推导式 → []
sum([]) → 0
返回 3 + 0 = 3 <- 触底,开始回代
回到 sum_tree([2,[3]]):
sum([3]) = 3
返回 2 + 3 = 5
sum_tree([6])
label = 6
branches = [] <- 叶子
返回 6 + 0 = 6
回到最外层:
推导式的值 = [5, 6]
sum([5, 6]) = 11
返回 4 + 11 = 15
得到 15,与 doctest 一致。注意回代的顺序:
最深的 sum_tree([3]) 先返回 3,它的父亲才能算出 5,
最后根才能算出 15。递归调用的展开是自顶向下的,答案的产生是自底向上的。
8.2 balanced:题目要什么
判断一棵树是否「平衡」。定义有两条,缺一不可:
t的每个分支的总和(sum_tree)都相等;t的每个分支自己也是平衡的(递归定义)。
>>> t = tree(1, [tree(3), tree(1, [tree(2)]), tree(1, [tree(1), tree(1)])])
>>> balanced(t)
True
>>> t = tree(1, [t, tree(1)])
>>> balanced(t)
False
>>> t = tree(1, [tree(4), tree(1, [tree(2), tree(1)]), tree(1, [tree(3)])])
>>> balanced(t)
False
三个例子分别在考三件不同的事,值得先画出来:
例 1(True) 例 3(False)
1 1
/ | \ / | \
3 1 1 4 1 1
| / \ / \ \
2 1 1 2 1 3
各分支和:3, 3, 3 各分支和:4, 4, 4 <- 顶层看起来平衡!
且每个分支自己平衡 但分支 [1,[2],[1]] 的两个孩子和是 2 和 1
→ 不相等 → 那个分支自己不平衡 → 整体 False例 3 是全题的关键:如果你只检查了「顶层的分支和是否相等」,
会错误地返回 True。定义的第二条(分支自己也要平衡)不是废话。
例 2 则是另一个方向的考查:把例 1 那棵和为 7 的树(1+3+1+2+1+1+1 = 10,
准确说 sum_tree 为 10)和一个孤零零的 tree(1) 并列做分支,
两个分支的和一个是 10 一个是 1,第一条就不满足。
第二个例子写的是 t = tree(1, [t, tree(1)])——右边的 t 是旧的那棵。
先求值右边(用旧 t 建一棵新树),再把名字 t 重新绑到新树上。
所以这里并没有产生「自己包含自己」的循环结构,只是旧树成了新树的一个分支。
第三个例子又整个重新赋值,和前两个无关。
8.2 怎么想到的
定义已经是递归的了,几乎是直译。难点在两处。
难点一:怎么表达「所有分支的和都相等」?
你不能写 sum_tree(b1) == sum_tree(b2) == sum_tree(b3),因为分支个数不定。
标准手法是挑一个基准,让所有人跟它比。挑谁都行,最方便的是第一个分支
branches(t)[0]。「所有人都等于第一个」和「所有人两两相等」是等价的。
这里有个隐患:如果 t 是叶子,branches(t) 是空列表,
branches(t)[0] 会抛 IndexError。会不会崩?
不会——因为这个表达式写在推导式的循环体内,
而叶子的 branches(t) 为空,循环一轮都不执行,
那行代码根本不会被求值。这是本题最微妙的一点。
[... branches(t)[0] ... for b in branches(t)] 里,
branches(t)[0] 看起来危险,实际安全:循环零轮时,循环体里的任何表达式都不会被求值。
换句话说,「branches(t) 非空」这个前提被 for 子句自动保证了。
如果你把它挪到推导式外面(比如 first = branches(t)[0] 写成单独一行),
叶子就会立刻 IndexError: list index out of range。位置决定生死。
难点二:怎么把两个条件合起来?
定义说「每个分支都要满足:和等于基准 并且 自己平衡」。
「每个都要」用 all——它接一个序列,全为真才返回 True,
空序列返回 True(这正好让叶子平凡地平衡,符合直觉:叶子没有分支,无所谓平不平衡)。
「并且」用 and。于是一句话:
all([条件A(b) and 条件B(b) for b in branches(t)])
我一开始想过分开写两个 all:
return (all([sum_tree(b) == sum_tree(branches(t)[0]) for b in branches(t)])
and all([balanced(b) for b in branches(t)]))
这也是对的,逻辑等价,只是遍历了两遍分支。合成一个推导式更紧凑, 也更贴近「每个分支都要同时满足两件事」的原始表述。
8.2 代码
def balanced(t):
"""Checks if each branch has same sum of all elements and
if each branch is balanced.
>>> t = tree(1, [tree(3), tree(1, [tree(2)]), tree(1, [tree(1), tree(1)])])
>>> balanced(t)
True
>>> t = tree(1, [t, tree(1)])
>>> balanced(t)
False
>>> t = tree(1, [tree(4), tree(1, [tree(2), tree(1)]), tree(1, [tree(3)])])
>>> balanced(t)
False
"""
# Every branch must have the same total sum, and each branch must itself
# be balanced. A leaf has no branches, so it is trivially balanced.
return all([sum_tree(b) == sum_tree(branches(t)[0]) and balanced(b)
for b in branches(t)])
sum_tree(b) == sum_tree(branches(t)[0])—— 条件一。 每轮都重算一次基准,效率不高(可以先存起来),但保证了它只在循环体内被求值,从而叶子安全。and balanced(b)—— 条件二,递归。 注意是balanced(b)不是balanced(t);后者会无限递归, 报RecursionError: maximum recursion depth exceeded。and的短路(short-circuit):如果和不相等, 右边的balanced(b)根本不会被调用。这不影响正确性,只是省了点计算。all([...])—— 全真才真。all([])为True,叶子直接返回True。- 为什么不检查
t自己的标签?定义里就没有这一条。 平衡只约束分支之间,根的标签是多少无所谓—— 例 1 和例 3 的根都是 1,结果却不同。
8.2 验证
追踪第三个 doctest(那个「顶层看起来平衡但其实不平衡」的例子):
t = tree(1, [tree(4), tree(1, [tree(2), tree(1)]), tree(1, [tree(3)])])
记三个分支为
B1 = [4]
B2 = [1, [2], [1]]
B3 = [1, [3]]
第一步:算基准 sum_tree(branches(t)[0]) = sum_tree(B1) = 4
轮 1 b = B1
sum_tree(B1) = 4 == 4 ✓
balanced(B1):B1 是叶子,branches 为 [],all([]) = True ✓
→ True
轮 2 b = B2
sum_tree(B2) = 1 + 2 + 1 = 4 == 4 ✓
balanced(B2):
B2 的分支是 [2] 和 [1],基准 = sum_tree([2]) = 2
子轮 1 b = [2]:2 == 2 ✓ 且 balanced([2]) = True → True
子轮 2 b = [1]:sum_tree([1]) = 1,1 == 2 ? ✗ → False
all([True, False]) = False
→ True and False = False <- 就在这里挂掉
轮 3 b = B3
sum_tree(B3) = 1 + 3 = 4 == 4 ✓
balanced(B3):只有一个分支 [3],基准就是它自己,3 == 3 ✓
且 balanced([3]) = True → all([True]) = True
→ True
最终 all([True, False, True]) = False
返回 False,与 doctest 一致。
可以清楚看到:三个分支的和确实都是 4,第一条完全满足;
是第二条(B2 自己不平衡,它的两个孩子和分别是 2 和 1)把整体判成了 False。
如果你的实现漏了 and balanced(b),这个 doctest 就会返回 True 而失败。
再快速验一下第一个(True):三个分支
[3]、[1,[2]]、[1,[1],[1]] 的 sum_tree
分别是 3、3、3,相等;
[3] 是叶子平衡;[1,[2]] 只有一个分支,自比自等且叶子平衡;
[1,[1],[1]] 两个分支和都是 1,且都是叶子。全部为真,返回 True。
- 只比和,不递归:
all([sum_tree(b) == sum_tree(branches(t)[0]) for b in branches(t)])。 前两个 doctest 都过,第三个返回True而挂。这是本题设计出来专门抓的错。 - 把基准提到推导式外面:写成
first = sum_tree(branches(t)[0]) return all([sum_tree(b) == first and balanced(b) for b in branches(t)])
一碰到叶子(递归总会碰到叶子)就IndexError: list index out of range。要这么写就必须先加if is_leaf(t): return True。 - 用
label代替sum_tree: 「分支的总和」不是「分支的标签」。例 1 里三个分支标签是 3、1、1,会立刻判成False。 - 写
any而不是all: 「每个分支都要」是all。any只要有一个满足就为真, 例 3 会返回True。
9. Q9(选做)num_trees
题目要什么
满二叉树(full binary tree)的定义:每个节点要么有 2 个分支,要么有 0 个分支, 不允许只有 1 个。问:恰好有 n 个叶子的满二叉树,有多少种不同的形状?
注意这里数的是形状,节点上没有标签,所以「不同」指的是结构不同。
>>> num_trees(1)
1
>>> num_trees(2)
1
>>> num_trees(3)
2
>>> num_trees(8)
429
先把小情况画出来,这是唯一能建立直觉的办法:
n = 1 只有一种:一个孤零零的节点,它自己就是叶子
*
n = 2 只有一种:根 + 两个叶子
*
/ \
* *
n = 3 两种:左边挂一棵 2 叶子树,或者右边挂
* *
/ \ / \
* * * *
/ \ / \
* * * *为什么 n = 3 时不能是「根有三个孩子」?因为满二叉树规定分支数只能是 0 或 2。 为什么不能是「根 - 一个孩子 - 两个叶子」?因为那个中间节点只有 1 个分支,非法。
怎么想到的
第一反应往往是错的:想去「构造出所有树再数一遍」。 那太重了,而且 n = 8 时有 429 棵,手写生成逻辑很容易漏。 计数问题的正确姿势是:找一个递推关系,让「n 的答案」由「更小的答案」拼出来。
题面的 Hint 直接给了突破口:一棵满二叉树可以由「一个根 + 两棵更小的满二叉树」构成。
如果左子树有 a 个叶子、右子树有 b 个叶子,那么整棵树有 a + b 个叶子
(根不是叶子,因为它有 2 个分支)。
于是问题变成:把 n 拆成两个正整数之和 n = a + b,
每种拆法贡献 num_trees(a) × num_trees(b) 种形状。
左子树有 num_trees(a) 种长法,右子树有 num_trees(b) 种长法,
而且左右互不影响——你选左边的任何一种,都可以配右边的任何一种。
这是乘法原理。而不同的 a 给出的树一定不同(左子树的叶子数不一样),
所以各个 a 之间用加法。「独立选择用乘,互斥情况用加」。
a 的取值范围是什么?a 至少是 1(子树不能是空的,
满二叉树最小就是单个节点),b = n - a 也至少是 1,所以 a 最大是 n - 1。
写成 range(1, n)。
base case:n == 1 时只有一种(单节点),返回 1。
这个 base case 不能省——如果不写,num_trees(1) 会走到
for left_leaves in range(1, 1),循环零轮,返回 total = 0,
然后所有更大的 n 全部变成 0。这和 sum_tree 那种「循环零轮恰好正确」的情况相反:
这里循环零轮给出的是错误答案,所以必须显式拦截。
还有一个必须想清楚的问题:左右要不要区分?
比如 n = 3 时,a = 1, b = 2 和 a = 2, b = 1 算不算两种?
看 doctest 里 n = 3 的两张图——一张是 2 叶子子树在左,一张在右,
题目把它们算作两种不同的树。所以 range(1, n) 要跑完整个区间,
不能只跑一半再乘 2。
代码
def num_trees(n):
"""Returns the number of unique full binary trees with exactly n leaves. E.g.,
1 2 3 3 ...
* * * *
/ \ / \ / \
* * * * * *
/ \ / \
* * * *
>>> num_trees(1)
1
>>> num_trees(2)
1
>>> num_trees(3)
2
>>> num_trees(8)
429
"""
if n == 1:
return 1
# Split the n leaves between the left and right subtree of the root.
total = 0
for left_leaves in range(1, n):
total += num_trees(left_leaves) * num_trees(n - left_leaves)
return total
if n == 1: return 1—— 唯一的 base case,对应「单节点树」。total = 0—— 累加器初值。这里 0 是合法的初值,因为要做的是加法, 且n >= 2时循环至少跑一轮。(对比 Q3:那里之所以不能用固定初值, 是因为运算是未知的combiner。)range(1, n)—— 左子树的叶子数从 1 到n-1。 写成range(1, n+1)会让left_leaves = n, 于是num_trees(0)——那会走到range(1, 0)返回 0, 不报错但白算一轮;更糟的是若把 base 写成n <= 1返回 1,就会算错。num_trees(left_leaves) * num_trees(n - left_leaves)—— 乘法原理。 两个递归调用,参数都严格小于n(因为1 <= left_leaves <= n-1), 所以递归一定会终止。- 这是一个典型的树递归:一次调用派生出
2(n-1)次调用, 调用结构是一棵枝繁叶茂的树。num_trees(8)会重复计算num_trees(3)很多次—— 效率很差,但对 n = 8 完全够用。(想优化的话可以加记忆化,那是后面的内容。)
验证
先手算 num_trees(3):
num_trees(3):n != 1,total = 0,range(1, 3) = [1, 2]
left_leaves = 1
num_trees(1) = 1 (base case)
num_trees(3 - 1) = num_trees(2)
n = 2,range(1, 2) = [1]
left_leaves = 1: num_trees(1) * num_trees(1) = 1 * 1 = 1
total = 1 → 返回 1
贡献 1 * 1 = 1,total = 1
left_leaves = 2
num_trees(2) = 1
num_trees(1) = 1
贡献 1 * 1 = 1,total = 2
返回 2
得到 2,与 doctest 一致,也和上面画的两张图对应:
left_leaves = 1 那一项是「2 叶子子树在右」,
left_leaves = 2 那一项是「2 叶子子树在左」。
再往上推几层,验证 429 这个数确实会出现:
| n | 展开式 | 结果 |
|---|---|---|
| 1 | base case | 1 |
| 2 | 1·1 | 1 |
| 3 | 1·1 + 1·1 | 2 |
| 4 | 1·2 + 1·1 + 2·1 | 5 |
| 5 | 1·5 + 1·2 + 2·1 + 5·1 | 14 |
| 6 | 1·14 + 1·5 + 2·2 + 5·1 + 14·1 | 42 |
| 7 | 1·42 + 1·14 + 2·5 + 5·2 + 14·1 + 42·1 | 132 |
| 8 | 1·132 + 1·42 + 2·14 + 5·5 + 14·2 + 42·1 + 132·1 | 429 |
最后一行逐项相加:132 + 42 + 28 + 25 + 28 + 42 + 132 = 429。与 doctest 一致。 序列 1, 1, 2, 5, 14, 42, 132, 429 就是题面提到的「closed form solution」—— 卡塔兰数(Catalan numbers)。注意展开式左右对称,这正是「左右可以互换」的体现。
- 漏掉 base case:
num_trees(1)返回 0, 接着所有结果全是 0。而且不会报错,只会安静地给出全零。 - 用加法代替乘法:写成
num_trees(a) + num_trees(n-a), n = 3 得到(1+1) + (1+1) = 4。乘法原理没有想清楚。 - 把
total +=写成total =:只会保留最后一轮的值, n = 8 得到 132 而不是 429。 - 担心
range(1, n)重复计数而改成range(1, n // 2 + 1): 这会把左右镜像的两棵树当成一棵,n = 3 得到 1 而不是 2。
10. Q10(选做)only_paths
题目要什么
only_paths(t, n):返回一棵新树,它只保留 t 中那些
「位于某条根到叶的路径上、且该路径标签之和等于 n」的节点。
如果一条这样的路径都没有,返回 None。
>>> print_tree(only_paths(tree(5, [tree(2), tree(1, [tree(2)]), tree(1, [tree(1)])]), 7))
5
2
1
1
>>> t = tree(3, [tree(4), tree(1, [tree(3, [tree(2)]), tree(2, [tree(1)]), tree(5), tree(3)])])
>>> print_tree(only_paths(t, 7))
3
4
1
2
1
3
>>> print_tree(only_paths(t, 9))
3
1
3
2
5
>>> print(only_paths(t, 3))
None
三个关键点,缺一个就会写错:
- 路径必须从根一直走到叶。半路停下不算——第一个例子里
5 → 2之所以保留,是因为2是叶子且5 + 2 = 7; 而5 → 1这条(和为 6)不算,因为1不是叶子。 - 一个节点只要在至少一条合格路径上就保留。
第二个例子里节点
1同时通向2→1和3两条合格路径,所以它下面挂了两个分支。 - 没有任何合格路径时返回
None(第四个例子:t的根就是 3, 任何一条完整路径的和都大于 3)。print(None)显示None。
题面给了一个骨架,四个空要填:
if ____:
return t
new_branches = [____ for b in branches(t)]
if ____(new_branches):
return tree(label(t), [b for b in new_branches if ____])
怎么想到的
第一个念头通常是错的:想「先算出所有路径的和,再回头删节点」。 这需要两趟遍历,而且「删节点」在这个不可变的树抽象里做不到—— 你只能造新树,不能改旧树。
正确的思路是把问题重述成一个可以递归的形式。关键在于问一句:
当我从根走到某个节点 t 时,剩下的子问题长什么样?
假设根的标签是 5,目标和是 7。走进任何一个分支之后,
那个分支需要凑出的和就不再是 7,而是 7 - 5 = 2。
于是递归调用写成 only_paths(b, n - label(t))。
参数 n 从「总目标」变成了「剩余额度」,这样递归调用和原函数就是同一个形状了。
这个把「累积量」翻译成「剩余量」的手法,是递归题里最常用的一招——
因为递归函数只能看到自己的子树,看不到头顶上已经走过的路。
确定了这一点,四个空就基本自动了。
空 1:base case。 什么时候可以直接返回 t?
当 t 是叶子、并且它的标签恰好等于剩余额度 n 时——
说明这条路径正好走完且和对上了。写成 is_leaf(t) and label(t) == n。
注意这两个条件缺一不可。只写 label(t) == n:某个内部节点标签恰好等于剩余额度时
会被误当成终点,把它整个子树(包括不合格的部分)原样返回。
只写 is_leaf(t):所有叶子都保留,等于没过滤。
那不合格的叶子怎么办? 骨架里没有显式的「返回 None」。
看第三行到第五行:如果 t 是叶子但标签不等于 n,
branches(t) 是 [],new_branches 是 [],
any([]) 为 False,于是 if 不成立,
函数体走完却没有 return——Python 自动返回 None。
这正是我们要的。骨架的设计相当精巧:用「函数末尾隐式返回 None」来表达「这里没有合格路径」。
空 2:递归调用。 only_paths(b, n - label(t))。
它对每个分支返回「该分支中合格的部分」或者 None。
空 3:判断有没有分支活下来。 any。
只要有一个分支返回了非 None 的树,当前节点就在某条合格路径上,应当保留。
这里利用了真值性:一棵树是非空列表,为真;None 为假。
所以 any(new_branches) 恰好表达「至少有一棵子树活下来」。
空 4:过滤掉死掉的分支。 b is not None。
new_branches 里混着树和 None,不能直接交给 tree(...)——
构造器会断言 is_tree(None) 失败,抛
AssertionError: branches must be trees。
is not None 而不是直接 if b
两者在本题的测试下结果相同,但语义不同。if b 依赖真值性,
而理论上一棵树永远是非空列表,所以恒为真——除非它是 None。
写 b is not None 把意图说清楚了:我要排除的是「没有合格路径」这个信号,
而不是某种「空的树」(这个抽象里根本不存在空树)。
第三行的 any 用真值性是因为骨架限定了只能填一个函数名,
而 any 恰好可用。
代码
def only_paths(t, n):
"""Return a tree with only the nodes of t along paths from the root to a leaf of t
for which the node labels of the path sum to n. If no paths sum to n, return None.
>>> print_tree(only_paths(tree(5, [tree(2), tree(1, [tree(2)]), tree(1, [tree(1)])]), 7))
5
2
1
1
...
"""
if is_leaf(t) and label(t) == n:
return t
# Each branch only needs to make up the remaining sum n - label(t).
new_branches = [only_paths(b, n - label(t)) for b in branches(t)]
if any(new_branches):
return tree(label(t), [b for b in new_branches if b is not None])
is_leaf(t) and label(t) == n—— 唯一的成功终点。返回t本身 (叶子没有分支,直接复用不会有共享可变状态的问题)。n - label(t)—— 在推导式外面就固定了: 所有分支拿到的是同一个剩余额度。这是对的,因为兄弟分支互不影响,各自独立地去凑同一个数。any(new_branches)—— 注意不能写成len(new_branches) > 0:只要t有分支,new_branches就非空(里面可能全是None),那个判断永远为真。[b for b in new_branches if b is not None]—— 又一次my_filter。 这份 lab 前后呼应得很紧密。- 函数末尾没有
return—— 落到这里就返回None, 表示「这棵子树里没有任何合格路径」。这是有意为之,不是遗漏。
验证
追踪第一个 doctest:
only_paths(tree(5, [tree(2), tree(1, [tree(2)]), tree(1, [tree(1)])]), 7)。
原树 路径和
5 5+2 = 7 ✓
/ | \ 5+1+2 = 8 ✗
2 1 1 5+1+1 = 7 ✓
| |
2 1only_paths([5, [2], [1,[2]], [1,[1]]], 7)
is_leaf? 否
剩余额度 = 7 - 5 = 2
对三个分支各递归一次:
(A) only_paths([2], 2)
is_leaf([2]) 且 label = 2 == 2 → 返回 [2]
(B) only_paths([1,[2]], 2)
is_leaf? 否
剩余 = 2 - 1 = 1
only_paths([2], 1)
is_leaf 但 2 != 1,不进 base case
branches([2]) = [],new_branches = []
any([]) = False → 函数走完 → 返回 None
new_branches = [None]
any([None]) = False → 返回 None
(C) only_paths([1,[1]], 2)
is_leaf? 否
剩余 = 2 - 1 = 1
only_paths([1], 1)
is_leaf 且 1 == 1 → 返回 [1]
new_branches = [[1]]
any([[1]]) = True
过滤后仍是 [[1]]
返回 tree(1, [[1]]) = [1, [1]]
回到最外层:
new_branches = [ [2], None, [1,[1]] ]
any(...) = True (第一项 [2] 就是真值)
过滤:[b for b in new_branches if b is not None] = [ [2], [1,[1]] ]
返回 tree(5, [ [2], [1,[1]] ]) = [5, [2], [1,[1]]]
print_tree 打印这棵结果树:根 5 不缩进;
第一个分支 2 缩进两格;第二个分支 1 缩进两格,它的孩子 1 缩进四格。
输出正是:
5
2
1
1
与 doctest 一致。注意中间那条 5 → 1 → 2(和为 8)整条都被剪掉了——
连那个标签为 1 的中间节点也没有出现在结果里,因为它唯一的孩子返回了 None。
再看第四个 doctest only_paths(t, 3):t 的根标签是 3,
剩余额度 3 - 3 = 0。两个分支分别是 [4](叶子,4 != 0 → None)
和标签为 1 的子树(它的剩余额度变成 0 - 1 = -1,往下只会越来越负,
不可能有叶子标签等于负数,全线返回 None)。
于是 new_branches = [None, None],any 为 False,
最外层也返回 None,print 显示 None。
- 忘了过滤
None就丢给tree:AssertionError: branches must be trees。这个报错在树的题目里出现频率极高, 看到它就去检查「我传给构造器的列表里是不是混了非树的东西」。 - base case 写成
if label(t) == n: return t: 漏了is_leaf。第二个 doctest 里,节点3(在tree(3, [tree(2)])中) 在某一步剩余额度恰好也可能等于 3,就会被整棵原样返回,把不合格的子孙一起带进来。 - 递归时传
n而不是n - label(t): 每层都用原始目标,等于要求「每条从当前节点开始的路径和都等于 n」,语义完全变了。 - 用
all代替any: 那要求所有分支都合格,第一个 doctest 会因为中间分支返回None而整棵变成None。
整份作业回顾
Lab 4 表面上是十道零散的题,实际上只教了三个可迁移的动作。
动作一:把「对每个元素做点什么」写成一行
列表推导式在这份 lab 里出现了至少五次,形状完全一样:
| 出处 | 代码 | 在做什么 |
|---|---|---|
| Q1 | [fn(element) for element in seq] | 映射 |
| Q2 | [element for element in seq if pred(element)] | 筛选 |
Q8 sum_tree | [sum_tree(b) for b in branches(t)] | 对每个分支递归 |
Q8 balanced | [条件 for b in branches(t)] | 对每个分支求一个布尔值 |
| Q10 | [b for b in new_branches if b is not None] | 筛选 |
Q1 和 Q2 之所以放在这份 lab 的最前面,就是为了让你在写树递归之前,
把 [... for b in branches(t)] 这个句式练到不用想。
动作二:只透过接口看数据
Q4–Q6 用一个「偷换实现」的自动测试逼你养成习惯:
永远用选择器,不用下标。这个习惯在树的部分立刻收到回报——
树的抽象(label / branches / is_leaf)
让你可以写出「一棵树的和 = 根 + 各分支的和」这种和实现无关的句子;
如果你满脑子都是 t[0] 和 t[1:],递归就很难想清楚。
动作三:树递归的三个模板
| 模板 | 形状 | 本 lab 的例子 | 迁移到哪里 |
|---|---|---|---|
| 聚合 | f(t) = g(label(t), [f(b) for b in branches(t)]) | sum_tree、balanced | 求树的最大值、深度、节点数、是否含某值 |
| 造新树 | return tree(新标签, [处理过的分支]) | only_paths、copy_tree | 树的每个标签乘 2、剪枝、替换子树 |
| 计数分支 | total += f(小问题) * f(另一个小问题) | num_trees | 换零钱、划分数、路径计数 |
- 要不要写 base case? 看「循环零轮」时的返回值对不对。
sum_tree不需要(sum([]) = 0恰好对),balanced不需要(all([]) = True恰好对),num_trees需要(total = 0是错的),only_paths需要(叶子的成功条件和内部节点不同)。 - 递归参数怎么变? 如果题目给的是「从根开始累计」的量,
就把它翻译成「往下还差多少」——
only_paths的n - label(t)。 递归函数看不见头顶,只能靠参数捎带信息。 - 表达式写在推导式里面还是外面?
写在里面意味着「循环零轮时不会被求值」。
balanced里的branches(t)[0]必须在里面(否则叶子会IndexError),only_paths里的n - label(t)在里在外都行。
动手练习
做完 lab 之后,用同一套模板试试这几个(都能一两行写完):
1. max_label(t):返回树中最大的标签
def max_label(t):
return max([label(t)] + [max_label(b) for b in branches(t)])
把 label(t) 先放进列表里,这样叶子时 max([label(t)]) 也合法。
如果写成 max(label(t), max([...])),叶子会因为 max([]) 而报
ValueError: max() arg is an empty sequence。
这就是「循环零轮时会不会出事」这条判断的又一次应用。
2. height(t):返回根到最深叶子的边数(单节点树高度为 0)
def height(t):
if is_leaf(t):
return 0
return 1 + max([height(b) for b in branches(t)])
这里 base case 是必需的:叶子时 max([]) 会报错。
你也可以用 max([...], default=-1) 省掉 base case,但显式写更清楚。
3. double(t):返回一棵所有标签翻倍的新树
def double(t):
return tree(label(t) * 2, [double(b) for b in branches(t)])
「造新树」模板的最简形式。注意它和 copy_tree 只差一个 * 2。
不需要 base case:叶子时分支列表为空,tree(x, []) 就是一个叶子。
4. 为什么 my_reduce(lambda x, y: x - y, [10, 3, 2]) 是 5 而不是 9?
左折:先 10 - 3 = 7,再 7 - 2 = 5。
右折会得到 10 - (3 - 2) = 9。减法不满足结合律,
所以这是唯一能明确区分两种折叠方向的测试。
Q3 的第四个 doctest(x + 2 * y)起的就是同样的作用。