跳到正文

Python 条件判断与循环:if、for、while 与推导式

Python 的分支和循环语法一眼就能看懂,缩进代替花括号这一点适应几分钟就够。真正需要花时间的是另外三件事:非布尔值在 if 里怎么判定、for 后面接的是什么、以及什么时候该把循环写成推导式。

if 后面的条件不要求是 bool,任何对象都有真值。规则是:空的、零的、None 算假,其余算真。

for v in [0, 0.0, "", [], {}, set(), None, 0j, float("nan")]:
print(repr(v), bool(v))
0 False
0.0 False
'' False
[] False
{} False
set() False
None False
0j False
nan True

最后一行是这份清单里唯一的意外:float("nan") 的真值是 True。NaN 在数值上不等于任何东西(nan == nanFalse),但它是个非零浮点数,真值判定只看向量本身,不看向量代表什么。用 if value: 来判断「缺失值」在遇到 NaN 时会静默走错分支,必须写成 if pd.isna(value):

NumPy 数组的真值判定是另一套规则,多元素数组直接报错:

import numpy as np
a = np.array([1, 2, 3])
if a:
pass
ValueError: The truth value of an array with more than one element is ambiguous. Use a.any() or a.all()

报错比给一个错误的答案好。数组上要判断「是否有任意元素为真」用 a.any(),判断「是否全部为真」用 a.all()。这两个函数在 NumPy 里的用法见 /python/basics/numpy/

x = 7
if x > 10:
r = "big"
elif x > 5:
r = "medium"
else:
r = "small"
print(r)
medium

Python 没有 switch,多分支用 elif 串起来。条件不需要括号,但冒号和缩进不能少。同一层级的缩进量必须一致——混用 Tab 和空格是初学者最常见的报错来源,把编辑器的「Tab 自动转空格」打开可以避免。

if 块是语句而不是表达式,不能直接赋值。要在一行里做选择,用条件表达式:

print("high" if x > 5 else "low")
high

Python 还支持链式比较,3 < y < 10 等价于 3 < y and y < 10,但中间那个 y 只求值一次:

y = 6
print(3 < y < 10)
True

这一点在写区间判断时省事,也避免了 y 是个有副作用的函数调用时被跑两遍。

for 后面跟的是可迭代对象,不是下标序列。字符串、列表、元组、字典、集合、文件对象、生成器都可以直接遍历:

for ch in "abc":
print(ch)
a
b
c

遍历字典时,默认拿到的是键。要同时拿键和值,用 .items()

d = {"a": 1, "b": 2}
for k in d:
print(k, d[k])
for k, v in d.items():
print(k, v)
a 1
b 2
a 1
b 2

两种写法输出一样,但 for k, v in d.items() 少一次哈希查找,也更清楚。字典在遍历过程中不能增删键,需要边遍历边改的话先 list(d.items()) 把结果固定下来。

循环变量在循环结束后依然存在,这一点和 R 一样:

for i in range(5):
pass
print(i)
4

如果函数体里已经有一个叫 i 的变量,循环会覆盖它。变量名尽量用有意义的名字,别用 ij 当全局变量。

range 生成整数序列,range(5) 是 0 到 4,右端点取不到:

print(list(range(5)))
print(list(range(2, 10, 3)))
[0, 1, 2, 3, 4]
[2, 5, 8]

需要下标时不必写 for i in range(len(xs))enumerate 直接同时给出下标和值。start=1 让编号从 1 开始,跟论文里的表格编号对齐:

nm = ["mpg", "wt", "hp"]
for i, name in enumerate(nm, start=1):
print(i, name)
1 mpg
2 wt
3 hp

要并行遍历多个序列,用 zip:以最短的那个为准,长度不一致时不会报错,多出来的元素被静默丢掉。

vals = [21.0, 3.2, 147.0]
for name, v in zip(nm, vals):
print(name, v)
mpg 21.0
wt 3.2
hp 147.0

静默截断是 zip 最容易出问题的地方。两个列表长度本该相等却不等时,你拿到的是少了几行结果的输出,没有任何提示。Python 3.10 起可以写 zip(nm, vals, strict=True),长度不一致时直接抛 ValueError。做数据核对时建议默认加上 strict=True

while 在条件为真时反复执行,适合「不知道要跑几轮」的场景:

n = 1
while n * 2 < 100:
n = n * 2
print(n)
64

continue 跳过本轮剩下的语句,break 跳出整个循环:

out = []
for i in range(1, 11):
if i % 2 == 0:
continue
if i > 7:
break
out.append(i)
print(out)
[1, 3, 5, 7]

这两个关键字只作用于最内层循环,Python 没有 break 2 这种写法。要一次跳出两层,可以把内层循环包成函数用 return 返回,比设标志变量干净。

while 要自己兜住死循环的风险:条件永远为真时循环不会停。写的时候顺手加一个轮数上限,比如 while cond and iter_count < 10000,比半夜发现脚本卡住划算。

forwhile 可以带一个 else 子句,它在循环正常结束(没有被 break 中断)时执行。该语法在别的语言里没有,但用来表达「找遍了都没找到」很自然:

target = 8
for i in range(2, 10):
if i == target:
print("found", i)
break
else:
print("not found")
found 8

target 换成 99,else 分支就会执行:

target = 99
for i in range(2, 10):
if i == target:
print("found", i)
break
else:
print("not found")
not found

读代码时注意:else 属于 for 而不属于 if,缩进层级能看出来。该写法比在循环外设一个 found = False 标志变量少两行,也不用担心忘记更新标志。

列表推导式把「对每个元素做变换」写成一个表达式:

print([i ** 2 for i in range(5)])
print([i for i in range(10) if i % 2 == 0])
[0, 1, 4, 9, 16]
[0, 2, 4, 6, 8]

把方括号换成花括号,得到字典和集合推导式:

print({i: i ** 2 for i in range(4)})
print({abs(i) for i in range(-3, 1)})
{0: 0, 1: 1, 2: 4, 3: 9}
{0, 1, 2, 3}

把方括号换成圆括号,得到生成器表达式。它不一次性构造列表,而是按需产出元素,适合直接喂给 summaxany 这类聚合函数:

print(sum(i ** 2 for i in range(100)))
total = sum(i ** 2 for i in range(1000))
print(total)
328350
332833500

嵌套循环也能压进一个推导式,for 的先后顺序与写成嵌套循环时一致:

m = [[1, 2], [3, 4]]
print([x for row in m for x in row])
[1, 2, 3, 4]

推导式的可读性上限大约是两个 for 加一个 if。再复杂就该退回普通循环——压成一行省下的那点行数,不足以补偿读代码时的停顿。

关于推导式比 for 快,流传的数字经常被夸大。让四种写法做同一件事(累加 0 到 999 的平方和),交错轮流测量,每种跑 9 轮后取中位数:

import statistics
import timeit
variants = [
("for 循环累加", "total = 0\nfor i in range(1000):\n total += i * i\n"),
("生成器表达式", "total = sum(i * i for i in range(1000))"),
("列表推导式", "total = sum([i * i for i in range(1000)])"),
("NumPy 向量化", "total = int((np.arange(1000) ** 2).sum())"),
]
results = {name: [] for name, _ in variants}
for _ in range(9):
for name, stmt in variants:
t = min(timeit.repeat(stmt, setup="import numpy as np", number=200, repeat=3))
results[name].append(t / 200)
for name, _ in variants:
xs = [v * 1e6 for v in results[name]]
print(f"{name:<16}{statistics.median(xs):8.2f} us")

一次代表性输出:

for 循环累加 191.65 us
生成器表达式 208.44 us
列表推导式 151.83 us
NumPy 向量化 6.02 us

把这段代码在同一台机器上反复跑,得到的倍数关系是这样的:

写法 相对 for 循环
生成器表达式 0.92x ~ 1.09x
列表推导式 1.06x ~ 1.31x
NumPy 向量化 12x ~ 49x

结论有两条。第一,列表推导式确实比显式 for 快,但幅度在 10% 到 30% 之间,不是网上常见的「快好几倍」;生成器表达式和 for 循环基本持平,谁快谁慢落在噪声范围里。把推导式当成性能优化手段会失望,它的价值在可读性:同样的意图写得更短,读的人不用在脑子里跟踪一个累加变量。第二,真正的数量级差距来自 NumPy,稳定在十倍以上,因为整条运算落在编译好的 C 代码里,完全不进 Python 解释器。

绝对耗时随机器和负载波动很大,上面同一段代码在不同时刻测出的 for 循环耗时从 53 微秒到 191 微秒都有过,所以有价值的是倍数关系而不是具体数值。要判断自己代码里的瓶颈,用 timeit 在你自己的数据上测一次,别套用别人机器上的数字。

选法就清楚了:结果能用一条 NumPy 表达式算出来,用 NumPy;只是对序列做变换或过滤,用推导式,理由是读起来更紧凑;循环里有副作用(读写文件、调用外部程序、上一轮结果决定下一轮输入)才写 for

顺带一个容易踩的细节:for i in np.arange(100000)for i in range(100000) 慢约一倍(多次测量落在 2.0x 上下)。迭代 np.arange 拿到的是 NumPy 标量对象而不是 Python 整数,每轮都要额外做一次对象构造与比较;range 则按需生成 Python 整数。需要循环时用 range,需要向量化时用 np.arange,两者别混。

变异系数(coefficient of variation, CV)是标准差除以均值,用来比较量纲不同的几组数据的离散程度。对 iris 的四个数值列各算一个,先用循环:

import pandas as pd
from sklearn.datasets import load_iris
iris = load_iris(as_frame=True).frame
num_cols = ["sepal length (cm)", "sepal width (cm)", "petal length (cm)", "petal width (cm)"]
cv = {}
for col in num_cols:
v = iris[col]
cv[col] = v.std() / v.mean()
print(pd.Series(cv).round(3))
sepal length (cm) 0.142
sepal width (cm) 0.143
petal length (cm) 0.470
petal width (cm) 0.636
dtype: float64

同一件事用字典推导式一行写完:

cv2 = {col: iris[col].std() / iris[col].mean() for col in num_cols}
print(pd.Series(cv2).round(6))
sepal length (cm) 0.141711
sepal width (cm) 0.142564
petal length (cm) 0.469744
petal width (cm) 0.635551
dtype: float64

Series.std() 默认 ddof=1,算的是样本标准差,和 R 的 sd() 一致,所以这里的结果可以直接和 R 版对照——R 语言算出来的四个值是 0.1417113、0.1425642、0.4697441、0.6355511,完全吻合。如果换成 NumPy 的 np.std(),默认 ddof=0 算总体标准差,数值会略小,做统计推断时要注意该差别。

结果读起来很清楚:花萼长度和宽度的变异系数都在 0.14 左右,花瓣长度 0.47、花瓣宽度 0.64。花瓣的离散程度是花萼的四倍多,所以用花瓣尺寸做品种判别比用花萼尺寸有效得多。R 语言的 /r/basics/control-flow/sapply 做了同一件事,两边对照能看出向量化和推导式处理的是同一类问题。

边遍历边删除,结果不对 在列表上做 removepop 会让后面的元素前移,而迭代器按位置推进,于是紧跟被删元素的下一个元素被跳过:

xs = [1, 2, 2, 3]
for v in xs:
if v == 2:
xs.remove(v)
print(xs, len(xs))
[1, 2, 3] 3

里面还剩一个 2。正确做法是用推导式生成新列表,[v for v in xs if v != 2] 得到 [1, 3];确实要在原列表上改,就倒着遍历或用 while 手动控制下标。

for i in range(len(xs)) 又长又容易越界 直接 for v in xs 遍历值,需要下标就用 enumerate(xs)。反过来,如果你在循环里只用到 xs[i],那 i 本身就是多余的。

zip 少给了几行 两个序列长度不等时 zip 按短的截断,不报错。核对数据时加 strict=True,让它在长度不匹配时直接失败。

推导式里做了副作用 推导式适合纯变换。在推导式里写 print、写文件、改外部变量,读代码的人会以为它只是个构造表达式的语句。这类逻辑放回 for 循环。

控制流写顺之后,下一步是把重复的逻辑收进函数:/python/basics/functions/