NumPy reshape() 实战:改变数组形状、增减维度与控制读取顺序

2026-07-20 23 预计阅读时间: 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.

预计阅读时间:7 分钟

numpy.reshape() 改变的是数组的维度组织方式,而不是元素数量。它常用于把一维数据整理成矩阵、为模型输入增加批次维度,或在图像与表格处理流程中合并多个轴。真正需要注意的并不是 API 本身,而是形状是否兼容、元素按什么顺序填入,以及结果是否与原数组共享内存。

reshape() 的核心约束:元素总数必须不变

假设数组包含 12 个元素,那么它可以变成 (3, 4)(2, 6)(2, 2, 3),但不能变成 (5, 3)。所有维度长度的乘积必须等于原数组的 size

下面的示例可以直接运行:

import numpy as np

values = np.arange(12)

matrix = values.reshape(3, 4)
cube = values.reshape(2, 2, 3)

print("原始形状:", values.shape)
print("矩阵形状:", matrix.shape)
print(matrix)
print("三维形状:", cube.shape)
print(cube)

输出中的元素没有增加或丢失,只是索引方式发生了变化。例如,matrix[1, 2] 对应原一维数组中的某个位置。

如果目标形状不匹配,NumPy 会抛出 ValueError

import numpy as np

values = np.arange(12)

try:
    values.reshape(5, 3)
except ValueError as exc:
    print(exc)

这类错误通常意味着上游数据量与代码假设不一致。生产代码中可以在变形前显式检查:

rows, columns = 3, 4
if values.size != rows * columns:
    raise ValueError(f"需要 {rows * columns} 个元素,实际得到 {values.size} 个")

matrix = values.reshape(rows, columns)

用 -1 自动推导维度

当某个轴的长度可以由元素总数推导时,可以用 -1 交给 NumPy 计算。这在批量处理数据时很实用,因为样本数量可能变化,而每条记录的字段数量固定。

import numpy as np

raw = np.arange(24)
records = raw.reshape(-1, 6)

print(records.shape)  # (4, 6)
print(records)

一个 reshape() 调用中最多只能出现一个 -1。写成 reshape(-1, -1) 会产生歧义,NumPy 无法同时推断两个轴。

增加维度也可以通过 reshape() 完成。例如,把长度为 4 的向量变成单行矩阵或单列矩阵:

import numpy as np

vector = np.array([10, 20, 30, 40])

row = vector.reshape(1, -1)
column = vector.reshape(-1, 1)

print("行向量:", row.shape)
print(row)
print("列向量:", column.shape)
print(column)

移除维度时同样要遵守元素总数不变的规则。比如 (1, 2, 3) 可以整理成 (2, 3)(6,)。如果目标只是删除长度为 1 的轴,np.squeeze() 往往比手写目标形状更清晰;如果要添加一个轴,np.expand_dims()None 索引也能更明确地表达意图。

C 顺序与 Fortran 顺序会改变填充方式

reshape()order 参数决定读取和放置元素时采用哪种索引顺序:

  • order="C":按行优先处理,默认行为。
  • order="F":按列优先处理。
  • order="A":根据输入数组的内存布局选择类似 C 或 Fortran 的顺序。

可以这样观察两种顺序的差异:

import numpy as np

values = np.arange(1, 7)

by_rows = values.reshape(2, 3, order="C")
by_columns = values.reshape(2, 3, order="F")

print("C 顺序:")
print(by_rows)

print("F 顺序:")
print(by_columns)

结果分别为:

C 顺序:
[[1 2 3]
 [4 5 6]]

F 顺序:
[[1 3 5]
 [2 4 6]]

这里的 order 描述的是元素索引与重排顺序,不应简单理解为“强制把结果转换成某种连续内存布局”。当数据来自 MATLAB、Fortran 或按列组织的外部格式时,明确指定顺序尤其重要。

reshape() 的结果可能共享原数组内存

reshape() 会尽量返回视图,从而避免复制数据,但这并不是所有输入布局下都能保证的契约。连续数组通常可以直接改变索引元数据;转置、切片等操作可能产生不连续数组,此时 NumPy 可能需要复制数据。

下面展示常见的共享内存情况:

import numpy as np

source = np.arange(6)
reshaped = source.reshape(2, 3)

reshaped[0, 0] = 99

print(source)             # 第一个元素通常也变为 99
print(np.shares_memory(source, reshaped))

如果业务逻辑要求结果与原数据完全隔离,应明确复制:

independent = source.reshape(2, 3).copy()
independent[0, 0] = -1

print(source)
print(independent)

不要仅通过修改后的现象猜测是否共享数据。需要依赖这一行为时,应使用 np.shares_memory() 检查,或者直接调用 .copy() 表达所有权边界。

在数据管道中如何稳妥使用

使用 reshape() 前,可以按下面的清单检查:

  1. 确认目标各维度乘积等于 array.size
  2. 动态维度只保留一个 -1,其余维度写清楚。
  3. 明确上游数据采用行优先还是列优先,必要时指定 order
  4. 不要默认结果一定是视图;涉及原地修改时检查共享内存。
  5. 只增减长度为 1 的轴时,考虑 expand_dims()squeeze(),让代码意图更直接。
  6. 面向机器学习模型时,同时验证轴的语义,例如 (batch, height, width, channels),而不只是验证元素数量。

reshape() 可以让同一段连续数据适配不同算法接口,但它不会理解“样本”“通道”或“时间步”的业务含义。形状在数学上合法,不代表语义一定正确。把轴定义写进变量名、断言和测试,才能避免代码正常运行却悄悄错位的数据问题。


相关推荐