跳转至

第 9 章 推导式

配套例题:PIO6 单组_一维数组、PIO8 单组_二维数组 来源:菜鸟教程 Python3 推导式、迭代器与生成器

推导式(comprehension)是 Python 最有辨识度的语法,也是竞赛代码里出现频率最高的结构之一。 它的价值不只是「短」——推导式把「追加到结果」这一步交给了专门的 LIST_APPEND 字节码, 省掉 res.append 的属性查找和函数调用,比等价的 for + append 快 30%–50% (循环体本身仍然是逐条执行的 Python 字节码,并没有下沉到 C 层)。

但它同样是最容易被滥用的语法。本章除了讲怎么写,还会认真讲什么时候不该写


9.1 四种推导式一览

形式 写法 结果类型 是否惰性
列表推导式 [f(x) for x in it] list 否,立刻全部算完
字典推导式 {k(x): v(x) for x in it} dict
集合推导式 {f(x) for x in it} set
生成器表达式 (f(x) for x in it) generator ,用一个算一个

没有「元组推导式」(x for x in a) 得到的是生成器,不是元组; 想要元组必须写 tuple(x for x in a)。菜鸟教程里把生成器表达式叫「元组推导式」, 这个叫法是错的,记住 () 出来的东西不能索引、不能 len、只能遍历一次。

>>> t = (x * x for x in range(5))
>>> t
<generator object <genexpr> at 0x...>
>>> len(t)
TypeError: object of type 'generator' has no len()

9.2 列表推导式

完整语法:

[ 表达式 for 变量 in 可迭代对象 if 条件 ]

书写顺序是「变换 for 遍历 if 过滤」,求值顺序正好倒过来: 先遍历,再过滤,最后才对留下来的元素做变换。

a = [1, -2, 3, -4, 5]

[x * x for x in a]                  # [1, 4, 9, 16, 25]        变换
[x for x in a if x > 0]             # [1, 3, 5]                过滤
[x * x for x in a if x > 0]         # [1, 9, 25]               先过滤再变换

竞赛里最常见的四种用法:

a = list(map(int, input().split()))          # 读一行数(这个更快,见 9.8)
a = [int(x) for x in input().split()]        # 等价写法

g = [[0] * m for _ in range(n)]              # 建 n×m 的二维数组(唯一正确写法)
idx = [i for i, x in enumerate(a) if x == target]   # 找出所有等于 target 的下标
s = "".join([str(x) for x in a])             # 拼接

_ 是一个普通变量名,约定表示「这个值我不用」。 [0] * m for _ in range(n) 里的 _ 就是循环计数器,用不上所以写 _

循环变量不会泄漏

Python 3 里推导式有自己的作用域,循环变量不会污染外层:

i = 100
a = [i for i in range(5)]
print(i)                # 100,没被改掉

这是 Python 3 相对 Python 2 的重要改动。但对普通 for 循环不成立—— for i in range(5): pass 之后 i 依然是 4。


9.3 条件:过滤 if 与三元 if-else 的位置不同

这是初学者最容易写错的一点:

[x for x in a if x > 0]                  # 过滤:if 在最后,没有 else
[x if x > 0 else 0 for x in a]           # 变换:三元表达式在最前,必须有 else
位置 含义 能否省 else
for 之后 过滤,决定这个元素要不要 必须省,写了 else 是语法错误
for 之前 三元表达式,决定这个元素变成什么 必须有 else

两者可以同时用(先过滤,再对留下来的做三元变换):

[x if x % 2 == 0 else -x for x in a if x > 0]

多个 if 表示「且」:

[x for x in a if x > 0 if x % 2 == 0]     # 等价于 if x > 0 and x % 2 == 0

9.4 多重循环与求值顺序

多个 for 子句从左到右书写,顺序和展开成嵌套循环时完全一致

[(i, j) for i in range(2) for j in range(3)]
# [(0,0), (0,1), (0,2), (1,0), (1,1), (1,2)]

等价于:

res = []
for i in range(2):          # 左边的 for 在外层
    for j in range(3):      # 右边的 for 在内层
        res.append((i, j))

记忆法:把 for ... for ... 原样竖着抄下来,逐行加缩进,就是等价的嵌套循环。 这个规则对 if 也成立——if 写在哪个 for 后面,展开后就在那一层。

右边的 for 可以用左边的变量(反过来不行):

[(i, j) for i in range(4) for j in range(i)]      # j 依赖 i,合法
[(i, j) for j in range(i) for i in range(4)]      # ❌ NameError: i 未定义

这在枚举无序对时很常用:

pairs = [(i, j) for i in range(n) for j in range(i + 1, n)]     # 所有 i < j

嵌套推导式 ≠ 多重循环

「一个推导式里有多个 for」和「推导式的表达式部分又是一个推导式」是两回事

mat = [[1, 2, 3], [4, 5, 6]]

[x for row in mat for x in row]          # 多重 for:展平 → [1,2,3,4,5,6]
[[x * 2 for x in row] for row in mat]    # 嵌套推导式:保持形状 → [[2,4,6],[8,10,12]]

前者结果是一维的,后者结果还是二维的。判断方法:看最左边的表达式是不是方括号

转置矩阵

嵌套推导式的经典用途:

mat = [[1, 2, 3],
       [4, 5, 6]]

t = [[row[i] for row in mat] for i in range(len(mat[0]))]
# [[1, 4], [2, 5], [3, 6]]

理解顺序:外层先动。外层 for i in range(3) 每取一个 i, 内层就把所有行的第 i 列收集成一行。

竞赛里其实有更快的写法:

t = [list(col) for col in zip(*mat)]     # 用 zip 转置,C 层实现
t = list(zip(*mat))                      # 如果元组也能接受,连转换都省了

zip(*mat) 把每一行拆成参数传给 zipzip 天然按列打包。这是转置的标准惯用法。


9.5 字典推导式

{ 键表达式: 值表达式 for 变量 in 可迭代对象 if 条件 }
words = ["apple", "bob", "cat"]
{w: len(w) for w in words}                    # {'apple': 5, 'bob': 3, 'cat': 3}

# 值 → 下标的反查表(离散化的核心操作)
vals = sorted(set(a))
rank = {v: i for i, v in enumerate(vals)}

# 反转字典(键值互换,注意值必须可哈希且不重复)
inv = {v: k for k, v in d.items()}

# 过滤字典
big = {k: v for k, v in d.items() if v >= 3}

键重复时后写的覆盖先写的,不会报错。{x % 3: x for x in range(10)} 最终只有 3 个键,值是每组里最后一个。这个「静默覆盖」偶尔会掩盖 bug。

离散化是字典推导式在竞赛里的头号场景:

vals = sorted(set(a))                          # 去重排序
rank = {v: i for i, v in enumerate(vals)}      # 原值 → 排名
b = [rank[x] for x in a]                       # 映射后的数组,值域压到 [0, len(vals))

详见 41-桶计数与离散化


9.6 集合推导式

{x * x for x in range(-3, 4)}         # {0, 1, 4, 9}   自动去重
{x % 3 for x in a}                    # 所有余数

注意 {}空字典不是空集合,空集合只能写 set()

type({})            # <class 'dict'>
type(set())         # <class 'set'>

集合推导式和 set(生成器) 完全等价,写哪个都行:

{x for x in a if x > 0}          # 集合推导式
set(x for x in a if x > 0)       # 等价

9.7 生成器表达式:竞赛里最该用的那一个

把方括号换成圆括号,就从「一次性算完存起来」变成「用一个算一个」:

[x * x for x in range(10 ** 7)]      # 立刻分配 1000 万个元素,几百 MB 内存
(x * x for x in range(10 ** 7))      # 瞬间返回,几乎不占内存

当生成器表达式是函数的唯一参数时,外层括号可以省略

sum(x * x for x in a)             # ✅ 推荐,不用写 sum((x*x for x in a))
max(len(s) for s in words)
any(x < 0 for x in a)
"\n".join(str(x) for x in a)

有其它参数时不能省:

sum((x * x for x in a), 0)        # 有第二个参数,括号必须留

三个必须用生成器的场景

1. sum / max / min / any / all 的参数。 中间列表纯属浪费,sum([...]) 会先建一个完整列表再求和:

total = sum([x * x for x in a])       # ❌ 多分配一个长度 n 的列表
total = sum(x * x for x in a)         # ✅

2. any / all 需要短路。 这是正确性以外的性能关键点

if any(x < 0 for x in a):        # ✅ 找到第一个负数就停
if any([x < 0 for x in a]):      # ❌ 必须把整个列表算完才开始判断

数组 \(10^6\) 长、第 1 个元素就是负数时,两者差 \(10^6\) 倍。

3. 大数据量输出。

sys.stdout.write("\n".join(str(x) for x in ans))

反过来的陷阱:生成器只能遍历一次。

g = (x for x in range(5))
print(sum(g))       # 10
print(sum(g))       # 0   ← 已经耗尽了!

需要多次遍历就老老实实用列表。这个坑在「先求 max 再求这些元素的和」这类代码里高频出现。

生成器表达式的第一个可迭代对象是立即求值的

一个细节但会咬人:

a = [1, 2, 3]
g = (x for x in a)
a = [9, 9, 9]           # 重新绑定 a
list(g)                 # [1, 2, 3]   ← 用的是创建时的那个列表

for 后面最外层的可迭代对象在创建生成器时就被求值并保存了,其余部分才是惰性的。


9.8 性能:推导式 vs 循环 vs map

\(n = 10^6\) 的量级下,几种写法的相对耗时(同一台机器上的典型比值):

写法 相对耗时 说明
list(map(int, data)) 1.0× 最快,全程 C 层,无 Python 字节码循环
[int(x) for x in data] 约 1.5× 每次迭代要走一遍字节码
res = [] + for + res.append(...) 约 2.2× 多了属性查找和方法调用
res = [] + for + res += [x] 约 2.5× 还要建临时列表

结论:

  • 函数已经现成时(intstrabs),用 maplist(map(int, ...))
  • 需要表达式(x * 2 + 1)或带过滤时,用推导式maplambda 反而更慢, 因为每次调用 lambda 都是一次 Python 函数调用。
list(map(int, data))                  # ✅ 内建函数用 map
list(map(lambda x: x * 2, a))         # ❌ lambda + map,慢
[x * 2 for x in a]                    # ✅ 表达式用推导式
  • 循环体带 append 的写法,能改推导式就改,白捡 40% 提速。

为什么推导式快?CPython 为列表推导式生成了专门的 LIST_APPEND 字节码, 直接调底层的列表追加,跳过了「查 res.append 属性 → 构造调用帧 → 调用」这三步。


9.9 什么时候不该用推导式

推导式不是越多越好。以下五种情况请老实写循环。

1. 需要 break 提前退出。 推导式没有 break

# ❌ 想在找到第一个满足条件的元素后停下,推导式做不到
first = [x for x in a if check(x)][0]     # 会把整个 a 扫完,还可能 IndexError

# ✅ 用生成器 + next,带默认值
first = next((x for x in a if check(x)), None)

# ✅ 或者直接写循环
for x in a:
    if check(x):
        first = x
        break

2. 循环体有副作用。 推导式的返回值被丢弃时,它就是在滥用:

[print(x) for x in a]            # ❌ 建了一个全是 None 的列表
for x in a: print(x)             # ✅

3. 逻辑超过一行能读懂的程度。 三层 for 加两个 if 的推导式, 调试时无法在中间打印,也无法下断点。竞赛中读不懂自己的代码就是灾难。

4. 需要复用中间结果。 3.8+ 可以用海象运算符,但可读性差:

[y for x in a if (y := f(x)) > 0]        # 能写,但不好读

5. 会一次性吃掉大量内存。 \(n \times m = 10^7\) 的二维列表在 Python 里 每个 int 对象至少 28 字节,列表本身每个槽 8 字节,很容易撞 MLE。 只需要遍历时改用生成器;需要数值数组时考虑 array 模块或按行处理。

判断口诀:推导式回答的是「用旧序列造一个新序列」。 要做的不是这件事,就别用推导式。


9.10 例题

PIO6 单组_一维数组(入门)

第一行 \(n\ (1 \le n \le 10^5)\),第二行 \(n\) 个整数 \(a_i\ (1 \le a_i \le 10^9)\), 求数组元素之和。 题面见 PIO6 原题(牛客)

三种写法,都能过,但速度差一截:

input()                                        # n 读掉但用不上
print(sum(map(int, input().split())))          # ✅ 最快
input()
a = [int(x) for x in input().split()]          # 推导式,慢约 50%
print(sum(a))
input()
s = 0
for x in input().split():                      # 显式循环,最慢
    s += int(x)
print(s)

两个要点:

  • n 是多余的,但那一行必须读掉。C++ 需要 n 控制读入次数,Python 直接 split 整行。
  • 结果最大是 \(10^5 \times 10^9 = 10^{14}\),超过 32 位。C++ 要开 long longPython 的 int 无上限,什么都不用做。这是 Python 在这类题上的天然优势。

PIO8 单组_二维数组(入门)

第一行 \(n, m\ (1 \le n, m \le 10^3)\),接下来 \(n\) 行每行 \(m\) 个整数,求所有元素之和。 题面见 PIO8 原题(牛客)

既然只求总和,不要真的建二维表——把所有 token 一口气加起来:

import sys

data = sys.stdin.buffer.read().split()
n, m = int(data[0]), int(data[1])
# 前 2 个 token 是 n 和 m,矩阵元素紧随其后共 n*m 个;
# 切片左闭右开,[2, 2 + n*m) 恰好把它们全部框住,行边界在求和时无关紧要
print(sum(map(int, data[2:2 + n * m])))

如果后续确实需要按行列访问(比如求前缀和、做 DP),用嵌套推导式建表:

g = [list(map(int, data[2 + i * m: 2 + (i + 1) * m])) for i in range(n)]

这里外层是推导式、内层是 map——内层用 map 而不是推导式,因为 int 是现成的内建函数。

本章最重要的一个坑:二维数组的错误初始化。

g = [[0] * m] * n         # ❌ n 行指向同一个列表对象!
g[0][0] = 1
print(g)                  # 每一行的第 0 个元素全变成了 1

g = [[0] * m for _ in range(n)]     # ✅ 推导式每次求值都新建一个列表

* 复制的是引用,不是内容。[0] * m 之所以安全,是因为 int 不可变; [列表] * n 就会出事。这是 Python 算法题最经典的 WA 来源之一, 详见 05-列表

三维同理:

dp = [[[0] * c for _ in range(b)] for _ in range(a)]

完整题解:solutions/PIO6.pysolutions/PIO8.py


9.11 本章速查

需求 写法
读一行整数 list(map(int, input().split()))(比推导式快)
变换 [f(x) for x in a]
过滤 [x for x in a if cond]if,不能带 else
条件变换 [x if cond else y for x in a],三元在,必须带 else
多重循环 for 从左到右 = 从外到内
展平二维 [x for row in mat for x in row]
建二维数组 [[0] * m for _ in range(n)]绝不用 [[0] * m] * n
转置 list(zip(*mat))
离散化 rank = {v: i for i, v in enumerate(sorted(set(a)))}
求和/判存在 sum(...)any(...) 里用生成器,不要加方括号
找第一个满足的 next((x for x in a if cond), None)
空集合 set(),不是 {}
元组 没有元组推导式,要写 tuple(...)
生成器 只能遍历一次;第一个可迭代对象立即求值
不该用推导式 需要 break、有副作用、逻辑过长、内存过大