\[ \begin{align}\begin{aligned}\newcommand{\ba}{\boldsymbol{a}} \newcommand{\bb}{\boldsymbol{b}} \newcommand{\be}{\boldsymbol{e}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bw}{\boldsymbol{w}} \newcommand{\bx}{\boldsymbol{x}} \newcommand{\by}{\boldsymbol{y}} \newcommand{\bz}{\boldsymbol{z}} \newcommand{\bd}{\boldsymbol{d}} \newcommand{\bv}{\boldsymbol{v}} \newcommand{\bs}{\boldsymbol{s}}\\\newcommand{\btheta}{\boldsymbol{\theta}} \newcommand{\bbeta}{\boldsymbol{\beta}} \newcommand{\bgamma}{\boldsymbol{\gamma}} \newcommand{\bsigma}{\boldsymbol{\sigma}} \newcommand{\md}{\mbox{d}} \newcommand{\bmu}{\boldsymbol{\mu}} \newcommand{\bone}{\boldsymbol{1}} \newcommand{\bzero}{\boldsymbol{0}} \newcommand{\bepsilon}{\boldsymbol{\epsilon}} \newcommand{\bphi}{\boldsymbol{\phi}} \newcommand{\bh}{\boldsymbol{h}} \newcommand{\bc}{\boldsymbol{c}} \newcommand{\br}{\boldsymbol{r}} \newcommand{\bQ}{\boldsymbol{Q}} \newcommand{\bK}{\boldsymbol{K}} \newcommand{\bV}{\boldsymbol{V}} \newcommand{\bSigma}{\boldsymbol{\Sigma}} \newcommand{\bg}{\boldsymbol{g}} \newcommand{\bxi}{\boldsymbol{\xi}} \newcommand{\bvarepsilon}{\boldsymbol{\varepsilon}} \newcommand{\bdelta}{\boldsymbol{\delta}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bJ}{\boldsymbol{J}} \newcommand{\bp}{\boldsymbol{p}} \newcommand{\bi}{\boldsymbol{i}} \newcommand{\bo}{\boldsymbol{o}} \newcommand{\bE}{\boldsymbol{E}} \newcommand{\bH}{\boldsymbol{H}} \newcommand{\bL}{\boldsymbol{L}} \newcommand{\bu}{\boldsymbol{u}} \newcommand{\bLambda}{\boldsymbol{\Lambda}} \newcommand{\trans}{^{\rm\scriptsize T}} \newcommand{\var}{\mathrm{var}}\\\newcommand{\bA}{\boldsymbol{A}} \newcommand{\bB}{\boldsymbol{B}} \newcommand{\bC}{\boldsymbol{C}} \newcommand{\bD}{\boldsymbol{D}} \newcommand{\bG}{\boldsymbol{G}} \newcommand{\bI}{\boldsymbol{I}} \newcommand{\bM}{\boldsymbol{M}} \newcommand{\bP}{\boldsymbol{P}} \newcommand{\bS}{\boldsymbol{S}} \newcommand{\bU}{\boldsymbol{U}} \newcommand{\bW}{\boldsymbol{W}} \newcommand{\bX}{\boldsymbol{X}} \newcommand{\bY}{\boldsymbol{Y}} \newcommand{\bZ}{\boldsymbol{Z}} \newcommand{\cotp}{\textcolor[RGB]{48,209,88}{TP}} \newcommand{\cotn}{\textcolor[RGB]{100,210,255}{TN}} \newcommand{\cofp}{\textcolor[RGB]{94,92,230}{FP}} \newcommand{\cofn}{\textcolor[RGB]{191,90,242}{FN}}\\\newcommand{\numcotp}{\textcolor[RGB]{48,209,88}{50}} \newcommand{\numcotn}{\textcolor[RGB]{100,210,255}{30}} \newcommand{\numcofp}{\textcolor[RGB]{94,92,230}{10}} \newcommand{\numcofn}{\textcolor[RGB]{191,90,242}{10}} \DeclareMathOperator*{\argmin}{arg\,min}\end{aligned}\end{align} \]

Python 基本命令#

学习目标与记号#

  1. 准确解释 Python 中列表、元组、字典和集合的基本用法,并区分内容可以修改和不能修改的对象。

  2. 根据公式和张量维度分析广播规则,并识别常见实现错误;

  3. Python 实现或验证向量化,通过受控实验解释结果。

  本节沿用统一记号:普通小写字母表示标量,粗体小写字母表示向量,粗体大写字母表示矩阵或高阶张量;样本或时间编号写作下标;转置写作 \(\trans\)⁠。除非另有说明,批量样本按行存放。正文与练习中的程序都应同时检查数值结果和数组维度。本节练习的参考答案见 Python 基础命令与 NumPy 广播答案⁠。

  在本节中,我们将对本课程中常用的一些 Python 数据结构及命令进行简单的回顾。关于 Python 基础命令的详细介绍,可参见 Python Tutorial 等。

数据类型#

  数据在编程中无处不在,无论是进行简单的数学运算,还是处理复杂的数据分析,数据类型都是必须要掌握的基本内容。Python 作为一门功能强大的编程语言,提供了多种不同的数据类型,主要包括整数(integer)、浮点数(floating point)、字符串(string)和布尔值(boolean)。

整型变量(integer)

  在 Python 中,内置整数类型 int 用于表示正整数、负整数和零。它采用任意精度表示,整数的数值范围不受 32 位或 64 位等固定字长限制。随着整数位数增加,Python 会为其分配更多内存。因此,实际能够处理的整数大小仍受可用内存和计算时间限制。整数运算不会因为超出固定字长而发生通常意义上的整数溢出,但所需的存储空间和计算时间会随整数规模增加。下面把两个整数变量相加,输入是 age=25height_in_cm=175预期打印整数和 200。

1 age = 25  # 定义整数变量
2 height_in_cm = 175  # 定义整数变量
3
4 total = age + height_in_cm  #整数运算
5 print("Sum:", total)  #输出
Sum: 200

  + 在两个 int 之间执行整数加法,print 同时输出说明文字和结果。这个例子只演示数据类型和运算,不表示年龄与身高在实际问题中具有可相加的意义。

浮点型变量(float)

  在 Python 中,浮点数 用有限位二进制近似表示实数。许多十进制小数,例如 0.1无法被二进制浮点数精确表示,因此计算结果可能包含微小的舍入误差。科学计算中的浮点结果通常只能近似相等。比较两个结果时,应根据数值的大小、计算过程和精度要求设置允许的误差范围:对于接近零的数值,主要检查绝对误差;对于较大的数值,还应检查相对误差。因此,不宜直接用 == 判断两个计算得到的浮点数是否相等,可使用 np.isclose()np.allclose() 进行近似比较。下面把两个浮点数相乘,预期输出 price * weight 的近似值。

1 price = 19.99  # 定义浮点数变量
2 weight = 72.5  # 定义浮点数变量
3
4 total_cost = price * weight  #浮点数运算
5 print("Total Cost:", total_cost)  # 输出(输出结果可能因浮点数精度问题略有不同)
Total Cost: 1449.2749999999999

  * 执行浮点乘法,结果仍为 float输出的末尾小数可能体现二进制表示造成的舍入误差;实际金额计算通常还需要明确币种、单位和舍入规则,不能直接把这个教学运算当作财务结算程序。

  关于整型和浮点型数据的详细介绍,请参见 Numeric Types⁠。

字符串(string)

  在 Python 中,字符串 是一种用于表示文本数据的数据类型。字符串由一系列字符组成,可以包含字母、数字、符号以及空格等,用于存储和操作文本信息。Python 中的字符串是不可变的,这意味着一旦字符串被创建,它的内容就不能被改变。字符串在 Python 编程中非常常用,无论是处理用户输入、文件操作还是网络数据传输,都离不开字符串的处理。Python 提供了丰富的字符串操作方法和函数,使得字符串的处理变得非常方便和高效。下面输入两个字符串并用 + 拼接,预期打印 Hello, World!

1a = "Hello, "  # 定义字符串变量
2b = "World!"  # 定义字符串变量
3
4c = a+b  #字符串拼接
5print(c)
Hello, World!

  对字符串使用 + 会按顺序连接文本并创建新字符串,不会修改原来的 ab大量片段的拼接通常更适合使用 str.join以避免反复创建中间字符串。

  关于字符串的详细介绍,请参见 Text Sequence Type⁠。

布尔值(boolean)

  在 Python 中,布尔值 是一种用于表示真(True)或假(False)两种状态的数据类型。布尔值是逻辑运算的基础,广泛用于条件判断、循环控制以及逻辑表达式中。在 Python 里,布尔值不仅可以直接使用 TrueFalse 来表示,还可以通过逻辑运算符(如 andornot对布尔值或其他可转化为布尔值的表达式进行组合和运算。下面输入“是学生”和“有票”两个状态,预期 or 得到 True并进入 if 的学生分支。

 1is_student = True  # 定义布尔变量
 2has_ticket = False  # 定义布尔变量
 3
 4
 5can_enter = is_student or has_ticket  #布尔值或运算
 6print("Can Enter:", can_enter)  # 输出运算结果
 7
 8if is_student:  #在 if 语句中使用布尔值
 9    print("The student can enter.")
10else:
11    print("The person needs a ticket to enter.")
Can Enter: True
The student can enter.

  or 只要有一个条件为真就返回真;ifis_student 为真时执行第一条分支。真实准入规则通常不能仅靠这两个布尔量描述,还需处理身份验证、缺失信息等情况。

  关于布尔值的详细介绍,请参见 Boolean Type⁠。

  变量名应清晰、简洁并能表达用途。按照 Python 的常用风格,变量和函数使用蛇形命名法(如 user_name,类名使用大驼峰命名法(PascalCase,如 UserProfile,常量通常使用全大写形式(如 MAX_SIZE。循环下标可使用 ij 等短名称,但普通变量应避免含义不明的缩写。还应避免覆盖内置名称,例如不要把变量命名为 sumlistdict

基本数据结构#

  Python 提供列表(list)、元组(tuple)、集合(set)和字典(dictionary)等内置数据结构。列表有序且可变,元组有序且不可变;集合中的元素不重复,适合成员测试和集合运算;字典以键值对存储数据,键必须唯一且可哈希。自 Python 3.7 起,字典保证保留插入顺序,但代码通常仍应通过键而不是位置访问字典元素。本课程会使用列表组织多层网络的中间结果,使用元组返回多个值,并使用字典保存模型参数或缓存。

列表(list)

  Python 中的 列表 是一种非常基础且强大的数据结构,它是一种有序、可变的数据集合。列表可以存储任意类型的对象,包括数字、字符串、甚至是其他列表,这使得它非常灵活。列表中的元素按照插入的顺序进行排列,并且每个元素都有一个对应的索引,索引从 0 开始⁠。这使得我们可以通过索引来快速访问、修改或删除列表中的元素。此外,列表还支持 切片 操作,我们可以通过指定起始和结束索引来获取列表的一个子序列。下面分别输入整数、字符串和混合类型列表,预期输出一个索引元素和两个切片;数值批量计算则通常应使用 NumPy 数组或 torch.Tensor

1numbers = [1, 2, 3, 4, 5]  #创建一个包含整数的列表
2print(numbers[2])  #列表索引
3
4fruits = ['apple', 'banana', 'cherry']  #创建一个包含字符串的列表
5print(fruits[:2])  #列表索引
6
7mixed_list = [1, 'hello', 3.14, True]  #创建一个包含混合类型的列表
8print(mixed_list[1:3])  #列表索引
3
['apple', 'banana']
['hello', 3.14]

  索引 2 取得第三个元素 3;切片 [:2] 取得前两个元素,[1:3] 取得下标 1 和 2,但不含终点 3。切片会创建一个新列表;这里没有演示列表的修改、删除或嵌套引用。

元组(tuple)

  元组Python 中的基本数据结构,其是由小括号括起来的若干元素组成。Python 规定,元组的元素不能更改。在这里,我们不对元组进行详细的解释,只介绍其 打包拆包 的功能。下面把三个整数打包为元组,再拆包到三个变量,预期依次打印元组及其中的 1、2、3。

1a = 1,2,3 #打包
2print(a)
3
4x1, x2, x3 = a #拆包
5print(x1)
6print(x2)
7print(x3)
(1, 2, 3)
1
2
3

  在上面的例子中,我们用命令 a = 1,2,3 表示一个打包过程,即我们将等式右端用逗号隔开的三个数字打包成一个元组 (1,2,3)再将此元组赋值给变量 a我们用命令 x1, x2, x3 = a 表示一个拆包过程,即我们将元组 a 拆成三个元素,分别赋值给变量 x1, x2, x3需要指出的是,有时候我们并不需要对拆包之后的元素都赋值给相应变量。例如,我们只关心元组 a 中的第三个值,而对前两个数值并不关心。此时,我们可以用符号 _ 代替前两个变量;请参见下例。

1 a = 1,2,3 #打包
2 print(a)
3
4 _, _, x3 = a #拆包
5 print(x3)
(1, 2, 3)
3

  两个下划线接收但不再使用前两个元素,x3 得到第三个元素并打印 3。下划线只是“此值有意忽略”的命名惯例,并不会阻止该值被赋给变量;拆包两侧的元素数量仍须匹配。

  在本课程中,我们将利用元组的打包作为函数的输出,而利用拆包过程对函数的结果进行赋值。在下例中,函数输入数值数组 x输出由总体均值和总体方差组成的二元组;示例输入为 [0,1,2,3,4]随后把两个结果拆到独立变量中。

 1 import numpy as np
 2 def statistics(x):
 3     # x: 一个包含样本数据的数组
 4     mean_x = np.mean(x) #均值
 5     var_x = np.var(x) #方差
 6     return mean_x, var_x #将均值和方差结果打包输出
 7
 8 x = np.arange(5)
 9 print(x)
10
11 mean_x, var_x = statistics(x) #将函数结果拆包,并进行相应的赋值
12 print(mean_x)
13 print(var_x)
[0 1 2 3 4]
2.0
2.0

  return mean_x, var_x 会自动打包两个标量,调用处再按相同数量拆包。这里 np.var 默认除以样本数,计算的是 ddof=0 的总体方差;若要估计总体方差的无偏样本估计,通常应显式设置 ddof=1

  关于列表和元组的详细介绍,请参见 Sequence Types⁠。

集合(set)

  Python 中的集合是一种无序、不重复的数据结构,它主要用于快速成员关系测试和消除重复元素。集合中的元素没有固定顺序,且每个元素都是唯一的,这保证了数据的互异性。集合支持多种操作,如添加元素、删除元素以及进行集合运算(如交集、并集、差集等)。下面输入含重复值的集合及 set1set2预期展示自动去重、增删元素,以及并集、交集和差集。

 1 my_set = {1, 2, 3, 2, 3, 4}  #重复的元素会被自动去除
 2 print(my_set)
 3
 4 my_set.add(5)   #添加元素到集合中
 5 print("添加元素", my_set)  # 输出可能是: {1, 2, 3, 4, 5}(顺序可能不同)
 6
 7 my_set.remove(3)  #删除元素
 8 print("删除元素", my_set)  # 输出可能是: {1, 2, 4, 5}
 9
10 set1 = {1, 2, 3}
11 set2 = {3, 4, 5}
12 print("并集",set1.union(set2))  # 并集: {1, 2, 3, 4, 5}
13 print("交集",set1.intersection(set2))  # 交集: {3}
14 print("差集",set1.difference(set2))  # 差集: {1, 2}
{1, 2, 3, 4}
添加元素 {1, 2, 3, 4, 5}
删除元素 {1, 2, 4, 5}
并集 {1, 2, 3, 4, 5}
交集 {3}
差集 {1, 2}

  初始重复元素只保留一份;addremove 会原地修改 my_set后三个方法则返回新的集合结果。集合的打印顺序没有保证,remove 一个不存在的元素还会抛出 KeyError若希望不存在时也不报错可使用 discard

  关于集合的详细介绍,请参见 Set Types⁠。

字典(dictionary)

  字典是 Python 中另一个重要的数据结构。一个字典由 键:值 对组成。下面输入以参数名为键的字典,预期展示按键读取、添加键值对,以及由层编号动态拼出键名。

1 dic = {"b1": 1, "W1": 2, "b2": 3, "W2": 4}  # 定义字典
2 print(dic)
3 print(dic["b1"]) #根据键"b1",调取其对应的值
4
5 dic["b3"] = 5 #对字典新增加一个键:值对
6 print(dic)
7
8 i = 1
9 print(dic["W"+str(i)]) #通过字符串处理生成键值,并调取相应的值。
{'b1': 1, 'W1': 2, 'b2': 3, 'W2': 4}
1
{'b1': 1, 'W1': 2, 'b2': 3, 'W2': 4, 'b3': 5}
2

  在上例中,我们首先生成了一个包含四个 键:值 对的字典,并利用相应的键值调取了键值为 b1 对应的具体值。命令 dic["b3"] = 5 会原地增加一个键值对;dic["W"+str(i)] 则把 i=1 拼成键 W1 并读出数值 2。若所查键不存在,方括号访问会抛出 KeyError如果要遍历字典中的所有键值对,我们可以使用如下命令。

 1 # 遍历键
 2 for key in dic.keys():
 3     print(key)
 4
 5 # 遍历值
 6 for value in dic.values():
 7     print(value)
 8
 9 # 遍历键值对
10 for key, value in dic.items():
11     print(key, value)
b1
W1
b2
W2
b3
1
2
3
4
5
b1 1
W1 2
b2 3
W2 4
b3 5

  三个循环分别输出全部键、全部值和全部键值对;每次循环的输入都是前例已添加 b3dic现代 Python 按插入顺序遍历字典,但程序表达逻辑时仍应依赖键名,而不应把当前显示顺序当作排序结果。

  在本课程中,我们将着重利用字典记录模型参数值,并对其进行迭代更新。关于字典的详细介绍,请参见 Mapping Types⁠。

NumPy 简介#

  NumPyPython 科学计算生态的基础库,核心对象 numpy.ndarray 用于表示同质的多维数组。数组通常在连续或规则分块的内存中存储数据,并把循环下沉到编译后的底层实现,因此向量化运算往往比逐元素的 Python 循环更简洁、更高效。本节介绍数组的维度、轴、逐元素运算、矩阵运算和广播机制。

  在 Python 中,加载库(也称为模块)是一个简单而常见的操作,它允许你访问库提供的函数、类和变量等。要加载一个库,我们需要使用 import 语句。 通常,我们将用 np 简化 NumPy下面的代码没有数值输入或可见输出,只把模块加载到当前环境,供后面的数组示例调用。

1 import numpy as np #加载numpy库

  执行后,np 指向 NumPy 模块,例如可用 np.array 创建数组。导入成功本身不会生成数据;若环境尚未安装 NumPy这条语句会抛出 ModuleNotFoundError

NumPy 数组#

  NumPy 数组(numpy.ndarrayPython 列表(list都支持索引和切片,但用途不同。一个数组通常具有统一的数据类型,并通过 ndarray.shape 记录各轴长度;列表可以混合保存任意对象,并适合频繁增删元素。数值计算时,应优先使用数组和向量化运算。下例展示几种常用的数组创建方式。

备注

  数组的轴数、维度与数据类型这三个概念容易混淆:

  • ndarray.ndim 是轴的个数。例如,矩阵有两个轴,因此 ndim == 2

  • ndarray.shape 是由各轴长度组成的元组。例如,(1000, 6) 表示 1 000 行、6 列。

  • ndarray.dtype 是元素的数据类型。例如,float64 表示每个元素通常用 64 位浮点数存储。

  本章用“轴数”表示 ndim用“维度”表示 shape用“维度”表示向量空间或特征空间的维度。计算前先写清这些信息,可以避免大量广播和矩阵乘法错误。

例2.1

  下面不接收外部输入,而是分别创建一维整数序列、零矩阵、全 1 矩阵和随机矩阵;预期打印各数组的内容,并用分隔线区分结果。

 1 # 以下命令我们会在后续课程中经常用到
 2
 3 x1 = np.arange(4)  # 生成数组 [0, 1, 2, 3]
 4 print(x1)
 5 print("*"*6)
 6
 7 x2 = np.zeros((4, 2))  # 生成 4×2 的零矩阵
 8 print(x2)
 9 print("*"*6)
10
11 x3 = np.ones((4, 2))  # 生成 4×2 的全 1 矩阵
12 print(x3)
13 print("*"*6)
14
15 x4 = np.ones_like(x2)  # 生成与 x2 同形的全 1 矩阵
16 print(x4)
17 print("*"*6)
18
19 x7 = np.random.normal(size=(4, 2))  # 生成元素服从标准正态分布的 4×2 矩阵
20 print(x7)
21 print("*"*6)
[0 1 2 3]
******
[[0. 0.]
 [0. 0.]
 [0. 0.]
 [0. 0.]]
******
[[1. 1.]
 [1. 1.]
 [1. 1.]
 [1. 1.]]
******
[[1. 1.]
 [1. 1.]
 [1. 1.]
 [1. 1.]]
******
[[ 1.31352423  1.74422157]
 [ 1.51887538 -0.42156299]
 [-0.54097554  1.35351482]
 [ 0.64925537  1.25253604]]
******

  arange 按顺序生成 0 至 3,zerosones 按指定维度填充值,ones_like 复用 x2 的维度和数据类型,normal 产生随机数。最后一组数每次运行通常不同;这些创建方式适合演示和初始化,但大型数组仍会按元素占用内存。

  当给定一个数组时,我们通常需要对其特定位置的元素进行操作。下面先输入一个随机的 \(4\times2\) 数组,再分别取第一行、第一列及单个元素;预期两个“第一行”写法得到相同的一维结果。

 1 x1 = np.random.normal(size=(4, 2))  # 生成元素服从标准正态分布的 4×2 矩阵
 2 print(x1)
 3 print("*"*6)
 4
 5 print(x1[0,:]) #取出x1的第一行
 6 print("*"*6)
 7
 8 print(x1[0]) #同样取出x1的第一行
 9 print("*"*6)
10
11 print(x1[:,0]) #取出x1的第一列
12 print("*"*6)
13
14 print(x1[0,1]) #取出x1中第一行第二列元素
[[ 0.51257249  0.38971561]
 [-0.38564445 -0.69905686]
 [-1.06047377 -0.70143226]
 [ 0.22137465 -0.06028292]]
******
[0.51257249 0.38971561]
******
[0.51257249 0.38971561]
******
[ 0.51257249 -0.38564445 -1.06047377  0.22137465]
******
0.3897156115095367

  x1[0, :]x1[0] 都保留第 0 行的两个元素,x1[:, 0] 取得第 0 列,x1[0, 1] 返回一个标量。基本切片通常是原数组的视图,后续若修改切片可能同时改变原数组;随机输入也使具体打印值不可预先固定。

广播机制#

  在 Python 中,广播机制(Broadcasting)是一种强大的功能,特别是在处理 NumPy 数组时,它能够极大地简化代码并提高计算效率。广播机制允许 NumPy 在进行数组运算时,对维度不同的数组进行操作。通过广播机制,NumPy 会自动将较小的数组“扩展”到较大数组的维度上,以便进行元素间的运算。这种机制不需要显式地复制数据,而是通过调整数组的维度来实现。本节内容主要参考 Broadcasting⁠。

  当两个数组维度相同时,+-*/ 等运算符默认执行逐元素运算。广播机制把这一规则推广到某些维度不同但相互兼容的数组。

  下面输入两个形状均为 (3,) 的数组,先打印维度,再逐元素相加;预期输出 [3., 4., 5.]此时尚不需要扩展任何轴。

1a = np.array([1.0, 2.0, 3.0])
2b = np.array([2.0, 2.0, 2.0])
3
4print(a.shape)
5print(b.shape)
6print(a+b)
(3,)
(3,)
[3. 4. 5.]

  两个数组在每个位置一一对应,因此结果仍为 (3,)+ 不是向量点积;数组长度不相等且不能广播时,运算会报维度不兼容错误。

  然而,要求两个数组具有相同维度是具有明显缺陷的。例如,有时候我们只希望对以上数组 a 中的每个元素加上相同的因子 c下面仍输入 a=[1,2,3]但只用标量 c=2预期输出与上例相同的 [3.,4.,5.]用于展示标量广播。

1 a = np.array([1.0, 2.0, 3.0])
2 c = 2.0
3
4 print(a+c)
[3. 4. 5.]

  结合以上两个例子,我们可以直观地认为在第二个例子中,标量 c 的值被 拉伸 成了一个与数组 a 具有相同维度的数组;换句话说,我们通过 复制 c 的值若干次,形成了一个与数组 a 相同维度的新数组,然后进行代数运算。当然,这种基于 拉伸 或者 复制 的想法只是理解广播机制的直观,NumPy 利用了更加高效的方法避免了由于复制同一数值而带来的内存和计算效率的损失。尽管我们用了两种方法实现将数组 a 的每个元素加上相同的元素,但第二种方法的计算效率更高、对内存要求更低。下图 展示了第二种计算的直观理解。

a.shape = (3,)       c.shape = ()
[1, 2, 3]      +          2
                     ↓ 广播到 (3,)
[1, 2, 3]      +    [2, 2, 2]    →    [3, 4, 5]

  对数组 ab 进行逐元素运算时,NumPy 从两个 shape 元组的最右侧开始逐轴比较。某一对轴长度满足下列任一条件时,称其兼容:

  1. 两个维度相等,

  2. 其中一个数组的维度等于 1。

  例如,(3,1)(1,3) 在两个轴上都兼容,广播结果的维度为 (3,3)(3,4)(4,3) 的最右侧轴长度分别为 4 和 3,既不相等也不为 1,因此不兼容。只要存在一对不兼容的轴,运算就会引发 ValueError

  若两个数组的轴数不同,可在较短的 shape 左侧补 1 后再比较。例如,维度 (256,256,3)(3,) 可视为 (256,256,3)(1,1,3)结果维度为 (256,256,3)维度 (8,1,6,1)(7,1,5) 则按 (8,1,6,1)(1,7,1,5) 比较,结果维度为 (8,7,6,5)这里“补 1”只是理解规则的方式,并不会实际改变原数组。

  当我们已经知道结果的维度时,我们还需要知道 NumPy 是如何通过广播机制得到最终结果的。这个过程与数组与标量间的二元运算 例子 中展示的过程基本一致。下面第一个例子输入 \(3\times3\) 矩阵和长度为 3 的向量,预期把向量加到矩阵的每一行,并输出 \(3\times3\) 结果。

1 a = np.array([[2.0, 3.0, 4.0],
2               [3.0, 4.0, 5.0],
3               [4.0, 5.0, 6.0]])
4 b = np.array([3.0, 4.0, 5.0])
5
6 print(a.shape)
7 print(b.shape)
8 print(a+b)
(3, 3)
(3,)
[[ 5.  7.  9.]
 [ 6.  8. 10.]
 [ 7.  9. 11.]]

  在以上这个例子中,数组 a 的维度为 \((3,3)\) 但数组 b 的维度为 \((3,)\)⁠。我们首先将数组 b 的维度从左段补 1 扩充为 \((1,3)\)⁠。通过比较扩充后的两个数组可知,他们的维度在各个位置上均是兼容的。通过观察可知,在左边第一个维度上,数组 a 的维度为 3,而数组 b 的维度为 1。因此,直观上理解,NumPy 将数组 b 沿着第一个维度,将该数组复制两次,得到一个与数组 a 相同的数组。此时,我们便可在每个位置上对两个元素进行相加,得到最终结果了。需要指出的是,NumPy 在运算时,不会复制数组以节省内容和提高计算效率。下图 展示了本例子计算过程的直观理解。

a.shape = (3, 3)     b.shape = (3,)
                     视为 (1, 3)
                     沿第 0 轴广播
结果维度:(3, 3)

  第二个广播例子输入形状为 (3,1) 的列数组 a 和形状为 (3,) 的一维数组 b预期两者分别沿不同轴扩展,打印形状 (3,3) 的两两求和结果。

1 a = np.array([[1.0],[2.0], [3.0]])
2 b = np.array([3.0, 4.0, 5.0])
3
4 print(a.shape)
5 print(b.shape)
6 print(a+b)
(3, 1)
(3,)
[[4. 5. 6.]
 [5. 6. 7.]
 [6. 7. 8.]]

  在上例中,数组 a 的维度为 \((3,1)\) 但数组 b 的维度为 \((3,)\)⁠。首先,我们将数组 b 的维度从左段补 1 扩充为 \((1,3)\)⁠。通过比较,我们可知两个数组在各个位置上的维度是兼容的。但通过比较我们可知,我们需要同时复制两个数组,只不过数组 a 需要沿着第二个维度复制,而数组 b 需要沿着第一个维度复制。最终结果的维度为 \((3,3)\)⁠。下图 展示了本例子计算过程的直观理解。

a.shape = (3, 1)     b.shape = (3,)
                     视为 (1, 3)
广播后维度:
    (3, 1) → (3, 3)
    (1, 3) → (3, 3)

  广播本身并不是错误,但不符合预期的维度可能产生“能够运行却含义错误”的结果。编写代码时应明确记录各数组的维度,并在关键位置使用 assert 检查。下例中 b.shape == (3,)c.shape == (1,3)两者结果相同,但 c 更明确地表达了“行向量”的含义。

 1 a = np.array([[1.0],[2.0], [3.0]])
 2 b = np.array([3.0, 4.0, 5.0])
 3 c = np.array([3.0, 4.0, 5.0]).reshape((1,3))
 4
 5 print("数组a的维度为:"+str(a.shape))
 6 print("数组b的维度为:"+str(b.shape))
 7 print("数组c的维度为:"+str(c.shape))
 8
 9 print("a+b的结果为:")
10 print(a+b)
11 print("a+c的结果为:")
12 print(a+c)
数组a的维度为:(3, 1)
数组b的维度为:(3,)
数组c的维度为:(1, 3)
a+b的结果为:
[[4. 5. 6.]
 [5. 6. 7.]
 [6. 7. 8.]]
a+c的结果为:
[[4. 5. 6.]
 [5. 6. 7.]
 [6. 7. 8.]]

  b 与显式行向量 c 都从右侧匹配 a所以 a+ba+c 均得到 (3,3) 且数值相同。这个对照只能说明当前运算的广播语义一致;在矩阵乘法、拼接等其他操作中,(3,)(1,3) 并不总能互换。

常用命令#

  NumPy 除了以上介绍的基于二元运算的广播机制外,还有很多高效的基于元素的运算。下面输入 a=[1,2,3]分别计算每个元素的指数、平方、平方根和正弦;四项输出都保持形状 (3,)

1 a = np.array([1.0, 2.0, 3.0])
2
3 print(np.exp(a)) #求指数
4 print(a**2) #求幂次
5 print(np.sqrt(a)) #开根号
6 print(np.sin(a)) #求正弦
[ 2.71828183  7.3890561  20.08553692]
[1. 4. 9.]
[1.         1.41421356 1.73205081]
[0.84147098 0.90929743 0.14112001]

  这些函数逐元素作用,不会在元素之间求和。sqrt 要求实数输入非负,否则在实数数据类型下会产生 nan 警告;exp 对很大的输入还可能溢出,因此实际模型中要关注数值范围。

  我们介绍两个命令,包括 np.mean() 以及 np.var()下例固定种子后输入 10 个标准正态随机数,预期先打印数组,再打印一个均值和一个方差标量。

 1 np.random.seed(1)
 2 a = np.random.normal(size=10)
 3
 4 print("数组a为:")
 5 print(a)
 6
 7 print("数组a的均值为:")
 8 print(np.mean(a)) #求均值
 9
10 print("数组a的方差为:")
11 print(np.var(a)) #求方差
数组a为:
[ 1.62434536 -0.61175641 -0.52817175 -1.07296862  0.86540763 -2.3015387
  1.74481176 -0.7612069   0.3190391  -0.24937038]
数组a的均值为:
-0.09714089080609986
数组a的方差为:
1.4182393613078983

  固定种子使本环境下的数组便于复现;不指定 axis 时,meanvar 聚合全部元素。np.var 默认使用 ddof=0所以这里输出的是按 10 作分母的总体方差,而不是按 9 作分母的无偏样本方差。

  在实际应用中,我们通常将训练集表示成矩阵的形式。沿用 按行组成矩阵的记号⁠,对于训练集 \(\{\bx_i:i=1,\ldots,n\}\) 而言,本书统一记

\[\bX=[\bx_1\trans;\ldots;\bx_n\trans] \in\mathbb{R}^{n\times d},\]

其中,\(n\) 表示训练集的规模,\(d\) 表示特征的维度。该式把 \(n\) 个特征列向量分别转置后按行排列,而不是把它们按列拼接。

  在实际操作中,我们往往需要对训练集中每个特征求均值或者方差。下面输入一个随机的 \(1000\times6\) 数据矩阵,使用 axis=0 分别聚合每一列;预期输出两个长度为 6 的数组,对应六个特征的均值与方差。

1 np.random.seed(1)
2 X = np.random.normal(size=(1000, 6))
3
4 print("数组X的列均值为:")
5 print(np.mean(X, axis=0)) #求均值
6
7 print("数组X的列方差为:")
8 print(np.var(X, axis=0)) #求方差
数组X的列均值为:
[ 0.01679263 -0.01111949  0.01457063  0.02767384  0.03295546 -0.00981674]
数组X的列方差为:
[1.05076426 1.07445554 0.96374936 0.94109343 1.02405612 0.95722715]

  第 0 轴是样本所在的行轴,消去它后每列只留下一个统计量。随机数据来自标准正态分布,因此结果通常在 0 和 1 附近,但有限样本不会恰好等于理论值;方差仍采用 ddof=0

  参数 axis 指定被聚合并从结果中消去的轴。下面把整数 0 至 3 整形成 \(2\times2\) 矩阵,并分别沿第 0、1 轴求均值;预期两项输出都为长度 2,但分别表示列均值和行均值。

1 X = np.arange(4).reshape(2,2)
2
3 print(X)
4
5 print(np.mean(X, axis=0))
6 print(np.mean(X, axis=1))
[[0 1]
 [2 3]]
[1. 2.]
[0.5 2.5]

  axis=0 把两行压缩为每列一个数,axis=1 把两列压缩为每行一个数。仅凭输出维度相同不能判断语义,实际代码还应结合“样本在哪一轴”的约定解释结果。

  将以上讨论拓展到三维数组:下面把 0 至 7 整形成 (2,2,2)分别沿三条轴求均值;每次都会消去指定轴,预期三个结果的形状均为 (2,2)

1 X = np.arange(8).reshape(2,2,2)
2
3 print(X)
4
5 print("下面是沿着三个维度的均值求解结果:")
6 print(np.mean(X, axis=0))
7 print(np.mean(X, axis=1))
8 print(np.mean(X, axis=2))
[[[0 1]
  [2 3]]

 [[4 5]
  [6 7]]]
下面是沿着三个维度的均值求解结果:
[[2. 3.]
 [4. 5.]]
[[1. 2.]
 [5. 6.]]
[[0.5 2.5]
 [4.5 6.5]]

  在上例中,np.mean(X,axis =0) 结果中,对应与第一行第一列的元素为 X[0,0,0]X[1,0,0] 的均值。而 np.mean(X,axis =1) 结果中,对应与第一行第一列的元素为 X[0,0,0]X[0,1,0] 的均值。

  默认情况下,聚合运算会删除指定轴。例如,若 X.shape == (1000,6)np.mean(X, axis=0).shape == (6,)设置 keepdims=True 会保留该轴并把长度设为 1,此时结果维度为 (1,6)保留轴有助于后续广播,也能让“按哪一轴计算”的意图更清楚。对三维数组 (2,2,2) 而言,分别沿三个轴求均值且保留轴,结果维度依次为 (1,2,2)(2,1,2)(2,2,1)

 1 np.random.seed(1)
 2 X = np.random.normal(size=(1000, 6))
 3
 4 mean1 = np.mean(X,axis=0)
 5 mean2 = np.mean(X,axis=0,keepdims =True)
 6
 7 print(mean1)
 8 print(mean2)
 9
10 print(mean1.shape)
11 print(mean2.shape)
[ 0.01679263 -0.01111949  0.01457063  0.02767384  0.03295546 -0.00981674]
[[ 0.01679263 -0.01111949  0.01457063  0.02767384  0.03295546 -0.00981674]]
(6,)
(1, 6)

  mean1 删除样本轴,形状为 (6,)mean2 保留该轴为长度 1,形状为 (1,6)两者包含相同的六个列均值,只是维度不同。keepdims=True 特别适合把统计量广播回原数组;若后续接口只接受一维结果,则应使用不保留轴的形式。

  下面把 0 至 7 整形成输入 X.shape == (2,2,2)分别沿第 0、1、2 轴求均值并保留轴;预期输出形状依次为 (1,2,2)(2,1,2)(2,2,1)

 1 X = np.arange(8).reshape(2,2,2)
 2
 3 print(X)
 4
 5 print("下面是沿着三个维度的均值求解结果:")
 6 mean0 = np.mean(X,axis=0, keepdims=True)
 7 mean1 = np.mean(X,axis=1, keepdims=True)
 8 mean2 = np.mean(X,axis=2, keepdims=True)
 9
10 print(mean0)
11 print(mean1)
12 print(mean2)
13
14 print("下面是沿着三个维度的均值求解结果的维度:")
15 print(mean0.shape)
16 print(mean1.shape)
17 print(mean2.shape)
[[[0 1]
  [2 3]]

 [[4 5]
  [6 7]]]
下面是沿着三个维度的均值求解结果:
[[[2. 3.]
  [4. 5.]]]
[[[1. 2.]]

 [[5. 6.]]]
[[[0.5]
  [2.5]]

 [[4.5]
  [6.5]]]
下面是沿着三个维度的均值求解结果的维度:
(1, 2, 2)
(2, 1, 2)
(2, 2, 1)

  每次运算只把被聚合的那条轴长度改为 1,其余两条轴保持不变,打印的数值则是对应两个元素的均值。保留轴便于结果与原三维数组继续广播,但这里所有轴长度恰好都为 2,实际任务中仍应先根据输入的 shape 推导结果,不能记忆本例的固定数值。

  Hadamard 乘积也称逐元素乘积。若 \(\bA=(a_{ij})\)\(\bB=(b_{ij})\) 均属于 \(\mathbb{R}^{m\times n}\)⁠,则 \(\bC=\bA\odot\bB\) 的元素为 \(c_{ij}=a_{ij}b_{ij}\)⁠。在 NumPy 中使用 A * B 计算逐元素乘积;若维度不同但兼容,则先应用广播规则。不要把 * 与矩阵乘法运算符 @ 混淆。

  Python 中的矩阵往往是一个二维 NumPy 数组,其运算是通过 NumPy 库来实现的。NumPy 提供了丰富的矩阵操作函数和方法,使得矩阵运算变得简单而高效。矩阵的基本运算包括加法、减法、乘法、转置、求逆、行列式计算等。例如,使用 +- 运算符可以进行矩阵的加法和减法运算;使用运算符 @ 等进行矩阵乘法运算;使用 .T 属性可以获取矩阵的转置;使用 np.linalg.inv() 函数可以计算矩阵的逆(如果矩阵可逆);使用 np.linalg.det() 函数可以计算矩阵的行列式。我们简单介绍本课程中常用的三种基本的矩阵运算,包括 Hadamard 乘法、标准矩阵乘法以及矩阵转置。

  标准矩阵乘法是指按照矩阵乘法的规则,第一个矩阵的行与第二个矩阵的列对应元素相乘后求和。下面输入两个 \(2\times2\) 矩阵,用三种接口计算同一乘积;预期 CDE 都是相同的 \(2\times2\) 矩阵。

 1 A = np.array([[1, 2], [3, 4]])
 2 B = np.array([[5, 6], [7, 8]])
 3
 4 C = A @ B  #使用 @ 操作符进行标准矩阵乘法
 5 D = np.matmul(A, B)  #或者使用 np.matmul() 函数
 6 E = np.dot(A,B) #或者使用 np.dot() 函数
 7
 8 print(C)
 9 print(D)
10 print(E)
[[19 22]
 [43 50]]
[[19 22]
 [43 50]]
[[19 22]
 [43 50]]

  三种写法在这个二维例子中等价,可用于交叉核对结果;乘法要求左矩阵的列数等于右矩阵的行数。对更高维数组,np.dotnp.matmul 的轴规则并不完全相同,因此深度学习代码通常更明确地使用 @matmul

  矩阵转置是指将矩阵的行和列互换。下面输入一个 \(2\times3\) 矩阵,用 .Tnp.transpose 得到两个 \(3\times2\) 输出,预期二者完全相同。

1 A = np.array([[1, 2, 3], [4, 5, 6]])
2
3 B = A.T
4 C = np.transpose(A)
5
6 print(B)
7 print(C)
[[1 4]
 [2 5]
 [3 6]]
[[1 4]
 [2 5]
 [3 6]]

  对二维数组,两种写法都交换行轴和列轴,通常返回共享原数据的视图而不是独立副本。对一维数组,A.T 的形状仍为 (n,)若需要显式行、列向量应使用 reshape高维 transpose 还应明确给出轴顺序。

Shiny 交互演示:NumPy 广播与线性回归损失

  交互页面的第一个标签页会从最右侧开始逐轴比较两个数组的维度,并显示广播后的实际计算结果;第二个标签页用同一组随机生成的训练数据展示线性回归的拟合直线与平均平方损失曲面。广播页面用于核对数组维度与逐元素运算,损失曲面页面则用于理解模型参数与损失函数之间的关系。

点击打开“NumPy 广播与线性回归损失”交互演示

核心推导与实现核验#

核心关系

\[(a_m,\ldots,a_1)\ \text{与}\ (b_n,\ldots,b_1)\ \text{从末轴向前逐轴相容}.\]

  推导路径。 把两个维度在左侧补 1,再从最右轴逐轴比较;每一轴只有“相等”或“其中之一为 1”时才可广播,输出轴长取两者较大者。

关键条件

  广播只复制逻辑视图:在每个相容轴上,长度为 1 的数组可重复使用同一元素;若两轴都大于 1 且不相等,就无法建立逐元素对应,因此规则既充分又必要。

数据规模

  若数据包含 \(n\) 个样本和 \(d\) 个特征,数据矩阵 \(\bX\) 的大小为 \(n\times d\)⁠,每个特征的均值 \(\bmu\) 共有 \(d\) 个数。从每一行减去这 \(d\) 个均值后,结果仍为 \(n\times d\)⁠;长度为 \(n\) 的向量对应样本而不是特征,不能直接完成这一操作。

常见误区

  最常见错误是把 axis 理解成“保留的轴”,或误以为所有维度不同的数组都能自动扩展;修改切片视图还可能连带修改原数组。

动手检查

  枚举长度为 1、相等和不相容的轴组合;可广播情形与显式 tile 结果一致,不相容情形必须抛出 ValueError

本节小结#

  1. 广播从末轴向前比较,而不是按元素总数判断。

  2. axis、keepdims 和数组维度必须一起理解。

  3. 向量化通常更快,但仍需关注视图、复制和内存峰值。

综合练习#

  以下练习均以给定数组为起点。除非题目另有说明,请先执行 import numpy as np并在程序中输出关键结果及其维度。遇到正文中尚未介绍的命令,请主动查阅官方文档或其他可靠资料进行学习。全部参考答案见 Python 基础命令与 NumPy 广播答案⁠。

  1. 数组创建。 给定列表 values = [3, 1, 4, 1, 5, 9]创建数据类型为 int64 的一维数组 a再将其重塑为维度为 (2, 3) 的数组 A输出 AA.ndimA.shapeA.dtype并使用断言检查结果。

  2. 数据类型。 给定 a = np.array([1, 2, 3], dtype=np.int64)b = np.array([0.5, 1.5, 2.5], dtype=np.float64)计算 c = a + b 并检查 c.dtype随后将二者都转换为 float32比较转换前后的 nbytes

  3. 维度与轴。 给定 X = np.arange(24).reshape(2, 3, 4)输出它的轴数、维度和元素总数;分别沿第 0、1、2 轴求和,并检查三个结果的维度。

  4. 数组索引。 给定 A = np.arange(1, 21).reshape(4, 5)分别取出第 3 行、第 2 列以及第 4 行第 5 列的元素。再使用负索引取得右下角元素,并验证两种取法得到相同结果。

  5. 数组切片。 对数组 A = np.arange(1, 21).reshape(4, 5)取出第 2 至第 4 行与第 2 至第 4 列构成的子数组;再分别取出所有偶数编号的列和按相反顺序排列的各行。

  6. 视图与复制。 给定 A = np.arange(12).reshape(3, 4)B = A[:, 1:3]C = B.copy()分别修改 BC 的一个元素,观察 A 是否变化,并使用 np.shares_memory() 检查共享内存关系。

  7. 布尔索引。 给定 A = np.arange(1, 21).reshape(4, 5)取出所有能被 3 整除的元素;再复制得到 BB 中大于 15 的元素设为 0,确认原数组 A 未被修改。

  8. 整数数组索引。 给定 X = np.arange(16).reshape(4, 4)rows = np.array([0, 3, 1])cols = np.array([2, 0, 3])分别计算 X[rows, cols]X[np.ix_(rows, cols)]比较二者的结果和维度。

  9. 重塑与转置。 给定 x = np.arange(24)将其重塑为 X.shape == (2, 3, 4)再通过 transpose 得到 Y.shape == (3, 2, 4)用一个具体元素验证转置前后的轴对应关系。

  10. 数组拼接与拆分。 给定 A = np.ones((2, 3), dtype=int)B = np.full((2, 3), 2)分别按行和按列拼接这两个数组。再把按列拼接的结果拆回 AB并检查恢复结果。

  11. 逐元素运算。 给定 x = np.array([1.0, 2.0, 4.0])y = np.array([2.0, 4.0, 8.0])分别计算逐元素加法、乘法、除法以及 x 的平方,并检查每个结果的维度。

  12. 广播机制:行向量。 给定 X = np.arange(12).reshape(3, 4)b = np.array([10, 20, 30, 40])计算 Y = X + b再使用 np.tile() 显式构造与 X 同形的数组,验证两种计算结果相同。

  13. 广播机制:列向量。 给定 X = np.arange(12).reshape(3, 4)c = np.array([[100], [200], [300]])计算 X + c再从一维数组 np.array([100, 200, 300]) 出发,通过重塑得到同样结果。

  14. 广播兼容性。 分别尝试对维度为 (2, 1, 4)(1, 3, 1)(5, 1)(4,)(2, 3)(3, 2) 的数组执行加法。编写函数输出可广播时的结果维度,并在不兼容时捕获 ValueError

  15. 按轴聚合。 给定 X = np.arange(1, 13).reshape(3, 4)分别计算每一行的和、每一列的平均值以及整个数组的最大值,并用断言检查结果维度。

  16. 保留轴与标准化。 给定 X = np.array([[1., 2., 5., 7.], [3., 4., 9., 11.], [5., 8., 13., 15.]])使用 axis=0keepdims=True 计算列均值与列标准差,再对每一列进行标准化。检查均值、标准差和标准化结果的维度,并验证标准化后各列均值接近 0、标准差接近 1。

  17. 矩阵乘法。 给定 A = np.array([[1., 2., 3.], [4., 5., 6.]])B = np.array([[1., 2.], [3., 4.], [5., 6.]])使用 @np.matmul()np.einsum() 三种方法计算矩阵乘积,验证结果一致,并说明它与逐元素乘法的区别。

  18. 随机数组与可复现性。 分别使用两个由 np.random.default_rng(2026) 创建的随机数生成器产生维度为 (3, 4) 的标准正态随机数组,验证两个数组完全相同;再更换种子,确认结果发生变化。

  19. 广播机制:成对距离。 给定 X = np.array([[0., 0.], [1., 0.], [0., 2.]])Y = np.array([[0., 1.], [2., 0.]])利用增加长度为 1 的轴和广播机制,计算每个 X 中样本到每个 Y 中样本的平方欧氏距离,得到维度为 (3, 2) 的距离矩阵,并找出每行距离最小值的列索引。

  20. 综合应用。 给定特征矩阵 X = np.array([[1., 2., 3.], [2., 4., 5.], [4., 5., 7.], [5., 8., 9.]])权重 w = np.array([0.5, -1., 2.]) 和偏置 b = 0.25先按列标准化 X再计算 scores = Z @ w + b最后以 0 为阈值得到布尔预测数组。输出所有关键中间量的维度,并用断言检查计算过程。