从数组到广播:用一组可运行实验打牢 NumPy 基础

2026-07-30 15 预计阅读时间: 1 分钟
来源: realpython.com AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:8 分钟

NumPy 的难点通常不是记住函数名,而是准确判断数组的形状、轴、索引结果和数据类型。数组看似只是一个数字表格,但一次 axis 选择、一次广播或一个布尔掩码,就可能改变计算的范围与结果。

下面围绕数组、轴、广播、索引、掩码和数据类型,建立一套可以直接运行验证的思考方法。

先看形状,再看数值

阅读 NumPy 代码时,先写出每个数组的 shape,再推导结果。对于二维数组,shape 通常写成 (行数, 列数)

import numpy as np

scores = np.array([
    [82, 91, 76],
    [88, 79, 95],
])

print(scores.shape)   # (2, 3)
print(scores.ndim)    # 2
print(scores.size)    # 6

这里有 2 行、3 列,共 6 个元素。ndim 表示维度数量,不是元素数量。

索引会不会保留维度,是一个常见分界点:

print(scores[0].shape)      # (3,)
print(scores[0:1].shape)    # (1, 3)
print(scores[:, 1].shape)   # (2,)
print(scores[:, 1:2].shape) # (2, 1)

整数索引通常会移除被选中的轴,切片则通常保留该轴。这个区别会直接影响后续广播和矩阵运算。

axis 到底消掉了哪一维

理解聚合操作时,可以把 axis 读成“沿着这条轴计算,并在结果中消掉它”。

import numpy as np

sales = np.array([
    [10, 20, 30],
    [40, 50, 60],
])

print(sales.sum())         # 210
print(sales.sum(axis=0))   # [50 70 90]
print(sales.sum(axis=1))   # [ 60 150]

原数组形状是 (2, 3)

  • axis=0 消掉第 0 维,结果形状为 (3,),得到每一列的总和。
  • axis=1 消掉第 1 维,结果形状为 (2,),得到每一行的总和。

需要保留维度以便继续广播时,可以使用 keepdims=True

row_mean = sales.mean(axis=1, keepdims=True)
print(row_mean)
# [[20.]
#  [50.]]

centered = sales - row_mean
print(centered)
# [[-10.   0.  10.]
#  [-10.   0.  10.]]

row_mean 的形状是 (2, 1),因此它可以逐行扩展到 (2, 3)。如果省略 keepdims=True,均值的形状会变成 (2,),无法按预期与 (2, 3) 对齐。

广播不是复制,而是形状兼容

广播允许不同形状的数组参与逐元素运算。判断时从形状的最右侧开始比较,每一维必须满足以下条件之一:

  • 两个维度相等;
  • 其中一个维度是 1
  • 某一侧缺少该维度,可视为 1

下面的例子为每一列增加不同的偏移量:

import numpy as np

measurements = np.array([
    [1.0, 2.0, 3.0],
    [4.0, 5.0, 6.0],
])
offset = np.array([0.1, 0.2, 0.3])

adjusted = measurements + offset
print(adjusted)
# [[1.1 2.2 3.3]
#  [4.1 5.2 6.3]]

measurements 的形状是 (2, 3)offset 的形状是 (3,)。从右向左比较时,两个数组最后一维都是 3,因此 offset 可以用于每一行。

如果想给每一行增加不同的值,需要显式构造 (2, 1)

row_offset = np.array([100, 200])[:, np.newaxis]
print((measurements + row_offset).shape)  # (2, 3)
print(measurements + row_offset)

遇到广播错误时,不要反复尝试转置。先打印参与运算的形状,再决定应该用切片、np.newaxis 还是 reshape 表达真实意图。

索引与掩码:筛选结果会变成什么形状

布尔掩码适合表达数据条件。比较操作会生成一个与原数组形状相同的布尔数组:

import numpy as np

temperatures = np.array([18.5, 21.0, 27.3, 16.8, 30.1])
mask = temperatures >= 25

print(mask)
# [False False  True False  True]

hot = temperatures[mask]
print(hot)  # [27.3 30.1]

多个条件需要使用逐元素运算符 &|~,并给每个比较表达式加括号:

comfortable = temperatures[
    (temperatures >= 18) & (temperatures <= 27)
]
print(comfortable)  # [18.5 21. ]

不能在这里使用 Python 的 and,因为 and 要求把整个数组解释为一个布尔值,而 NumPy 条件通常包含多个布尔元素。

二维数组使用布尔掩码时还要注意:对整个数组应用同形状掩码,筛选结果通常会被收集成一维数组。

matrix = np.array([[1, 8, 3], [9, 2, 7]])
selected = matrix[matrix > 5]

print(selected)        # [8 9 7]
print(selected.shape)  # (3,)

如果目标是保留原形状,可以使用 np.where

kept_shape = np.where(matrix > 5, matrix, 0)
print(kept_shape)
# [[0 8 0]
#  [9 0 7]]

数据类型会影响结果,而不只是内存

NumPy 数组通常具有统一的 dtype。混合整数和浮点数时,NumPy 会选择能够容纳这些值的公共类型:

import numpy as np

values = np.array([1, 2, 3.5])
print(values.dtype)  # 通常为 float64

类型转换需要留意截断和溢出:

prices = np.array([19.9, 20.1, 255.8])
print(prices.astype(np.int32))
# [ 19  20 255]

浮点数转整数会截掉小数部分,并不是四舍五入。如果业务含义需要四舍五入,应先明确执行:

rounded = np.rint(prices).astype(np.int32)
print(rounded)  # [ 20  20 256]

无符号小整数尤其需要谨慎。它们节省内存,但可表示范围有限,不适合未经检查的算术计算。

一份可直接运行的综合练习

将下面内容保存为 numpy_practice.py,先预测每个结果,再运行核对。环境中若尚未安装 NumPy,可执行 python -m pip install numpy

import numpy as np

x = np.array([
    [2, 4, 6],
    [1, 3, 5],
], dtype=np.float64)

column_mean = x.mean(axis=0, keepdims=True)
normalized = x - column_mean
mask = normalized > 0

print("shape:", x.shape)
print("column mean:\n", column_mean)
print("normalized:\n", normalized)
print("positive values:", normalized[mask])
print("result dtype:", normalized.dtype)

assert column_mean.shape == (1, 3)
assert normalized.shape == x.shape
assert mask.dtype == np.bool_
assert np.allclose(normalized.sum(axis=0), 0)

运行命令:

python numpy_practice.py

这段练习同时覆盖了二维数组、列方向聚合、维度保留、广播、布尔掩码、浮点类型和数值断言。

建立稳定的 NumPy 排错顺序

面对不熟悉的 NumPy 表达式,可以依次检查:

  1. 写出每个输入的 shapendimdtype
  2. 判断整数索引是否移除了某个轴。
  3. 对聚合操作确认 axis 消掉的是哪一维。
  4. 从最右侧比较广播维度。
  5. 区分“筛选元素”和“保持原形状替换元素”。
  6. 在类型转换前确认截断、范围和精度是否符合业务含义。
  7. assertnp.allclose 验证形状与数值不变量。

NumPy 基础是否扎实,不取决于记住多少 API,而取决于能否在运行代码前推导形状、轴和类型。把这些推导写成小实验和断言,数据处理代码会更容易阅读,也更不容易产生隐蔽错误。


相关推荐