CS 61A  /  作业解析
LAB 04

Lab 4:序列、树递归与树ok 10 项通过

从「把列表变成另一个列表」,到「把一棵树变成另一棵树」——这份 lab 把递归从数字搬到了数据结构上。

对应讲次:Lecture 6 树递归、Lecture 7 序列与容器、Lecture 9 树 官方题面:cs61a.org/lab/lab04 代码:labs/lab04/lab04.py

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']):

逐步推演
1 调用发生,形参绑定:fn → 那个 lambda 函数对象,seq → ['cs61a', 'summer', '2023']。
2 开始求值推导式,建一个空的新列表 []。
3 element = 'cs61a',求值 fn('cs61a')。进入 lambda 体,执行 print('cs61a'):屏幕出现 cs61a,print 返回 None;lambda 把这个 None 作为自己的返回值。新列表变成 [None]。
4 element = 'summer',同理,屏幕出现 summer,列表变成 [None, None]。
5 element = '2023',屏幕出现 2023,列表变成 [None, None, None]。
6 遍历结束,推导式的值是 [None, None, None],被 return 出去。
7 交互式解释器收到一个非 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_mapfn(element) ← 变换for element in seq无 —— 全都要
my_filterelement ← 原样for element in seqif 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]):

elementelement + 5% 3pred 结果新列表
160True[1]
271False[1]
382False[1]
490True[1, 4]
5101False[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):

角色函数作用
构造器 constructormake_city(name, lat, lon)把名字、纬度、经度打包成一个 city
选择器 selectorget_name(city)取出名字
选择器 selectorget_lat(city)取出纬度
选择器 selectorget_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. 把坐标包装成一个 citymake_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)):

逐步推演(change_abstraction.changed 为 True)
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)

三条必须记住的事实:

核心结论
  1. 树被表示成一个列表:第 0 个元素是标签,从第 1 个起全是分支。 所以 branches(t) 返回的是一个列表,里面每个元素本身又是一棵树。
  2. branches 参数必须是一个列表。tree(1, tree(2)) 是错的, tree(1, [tree(2)]) 才对。因为构造器会 for branch in branches 逐个断言。
  3. 叶子(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。逐条讲为什么。

1 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 存在的全部意义——它把「忘了套方括号」这个错误 在构造的瞬间就抓住,而不是等到你几十行之后取分支时才莫名其妙地崩。
2 t = tree(1, [tree(2)]) → Nothing。 这次分支列表是 [[2]],唯一的元素 [2] 通过了 is_tree, 构造成功,t 是 [1, [2]]。 但整行是一条赋值语句(assignment statement),不是表达式。 语句不产生值,交互式解释器什么也不显示。 这就是为什么答案是 Nothing 而不是 [1, [2]]。 顺带说明第 1 行为什么会 Error 而不是 Nothing:赋值语句要先求值右边, 求值过程中就炸了,压根轮不到赋值。
3 label(t) → 1。 t 是 [1, [2]],label 取 t[0],即 1。
4 label(branches(t)[0]) → 2。 从里往外算:branches(t) 是 t[1:],即 [[2]]—— 一个装着一棵树的列表。 [0] 取出第一棵(也是唯一一棵)分支,得到 [2]。 label([2]) 取 [2][0],即 2。 这里最容易混的是「分支列表」和「一个分支」的区别: branches(t) 是列表,branches(t)[0] 才是树。
5 len(x) → 1。 x = branches(t) = [[2]],长度 1,因为 t 只有一个分支。 注意别把它和 len(t) 混了:len([1, [2]]) 也是 2, 因为标签也占一格。len(branches(t)) 才是「有几个分支」。
6 is_leaf(x[0]) → True。 x[0] 是 [2]。branches([2]) 是 [2][1:],即 []。 not [] 为 True(空列表是假值)。所以它是叶子。
7 label(t) + label(branch) → 3。 branch = x[0] = [2]。label(t) = 1,label(branch) = 2, 1 + 2 = 3。这行想强调的是:标签是普通的值,可以直接参与算术; 而树本身不行——t + branch 会得到 [1, [2], 2](列表拼接), 是个毫无意义的东西,而且跨越了抽象屏障。
8 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]]]]
1 打印每个分支的标签 → 5 换行 8。 branches(t) 是 [b1, b2]。循环两轮: b = b1 时 label(b1) = 5,打印 5; b = b2 时 label(b2) = 8,打印 8。 for 语句本身不产生值,所以除了这两行打印之外没有别的输出。
2 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)。
3 [label(b) + 100 for b in branches(t)] → [105, 108]。 这就是 Q1 的 my_map,只不过映射的是树的分支: 5 + 100 = 105,8 + 100 = 108。 这次是表达式不是语句,所以解释器把返回的列表显示出来。 「对每个分支做点什么,收成一个列表」——记住这个形状,Q8 马上就要用。
4 [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 = 15

8.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:题目要什么

判断一棵树是否「平衡」。定义有两条,缺一不可:

  1. t 的每个分支的总和(sum_tree)都相等;
  2. 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,第一条就不满足。

注意 doctest 里的名字复用

第二个例子写的是 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展开式结果
1base case1
21·11
31·1 + 1·12
41·2 + 1·1 + 2·15
51·5 + 1·2 + 2·1 + 5·114
61·14 + 1·5 + 2·2 + 5·1 + 14·142
71·42 + 1·14 + 2·5 + 5·2 + 14·1 + 42·1132
81·132 + 1·42 + 2·14 + 5·5 + 14·2 + 42·1 + 132·1429

最后一行逐项相加: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

三个关键点,缺一个就会写错:

  1. 路径必须从根一直走到叶。半路停下不算——第一个例子里 5 → 2 之所以保留,是因为 2 是叶子且 5 + 2 = 7; 而 5 → 1 这条(和为 6)不算,因为 1 不是叶子。
  2. 一个节点只要在至少一条合格路径上就保留。 第二个例子里节点 1 同时通向 2→1 和 3 两条合格路径,所以它下面挂了两个分支。
  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 时,剩下的子问题长什么样?

关键一步:把 n 变成「还差多少」

假设根的标签是 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  1
逐步推演(递归展开)
only_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换零钱、划分数、路径计数
三条最值得带走的判断
  1. 要不要写 base case? 看「循环零轮」时的返回值对不对。 sum_tree 不需要(sum([]) = 0 恰好对), balanced 不需要(all([]) = True 恰好对), num_trees 需要(total = 0 是错的), only_paths 需要(叶子的成功条件和内部节点不同)。
  2. 递归参数怎么变? 如果题目给的是「从根开始累计」的量, 就把它翻译成「往下还差多少」——only_paths 的 n - label(t)。 递归函数看不见头顶,只能靠参数捎带信息。
  3. 表达式写在推导式里面还是外面? 写在里面意味着「循环零轮时不会被求值」。 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)起的就是同样的作用。