NumPy 是 Python 科学计算的基石。你可以把它理解为 "带高级操作的、多维的、同类型数据的数组"。
PyTorch 的张量(Tensor)设计,几乎就是照搬了 NumPy 的这套逻辑,然后把运算搬到了 GPU 上。所以搞懂 NumPy 的重点,PyTorch 的核心也就通了。
它的核心重点,可以浓缩成 "一个对象、两大杀招、三个关键":
1. 一个核心对象:ndarray(多维数组)
所有操作都围绕它展开。你只需要死死记住它的 3 个关键属性:
-
shape(形状) :几行几列,比如(3, 4)。注意: 一维数组的形状是(3,)(还记得那个逗号吗?)。 -
dtype(数据类型) :是int还是float。打印带不带点.,就是由它决定的。 -
ndim(维度数):数组有几根轴。
2. 两大杀招(NumPy 灵魂所在)
杀招一:向量化运算(彻底扔掉 for 循环)
NumPy 最牛的地方是,它对数组的操作默认都是**"逐元素"**的,而且底层用 C 语言实现,速度极快。
-
你可以直接写
arr * 2,它会把每个元素都乘以 2。 -
你可以直接写
arr1 + arr2,它会把对应位置的元素相加。 -
支持
np.exp(arr)、np.log(arr)、np.sin(arr)等数学函数,直接作用在每一个元素上。
杀招二:广播机制(Broadcasting)
这和你在 PyTorch 里学到的完全一样 。形状不同的数组做运算时,NumPy 会自动把小的数组"逻辑上"拉伸,规则依然是:从右往左对齐,维度要么相等,要么一方为 1 。
(比如形状 (3,1) 和 (1,4) 运算,会自动变成 (3,4))。
3. 三个关键操作(日常高频使用)
① 索引与切片(Indexing / Slicing)
和 Python 列表一样灵活,且支持"花式索引"(用列表取数)。
-
arr[0]:取第一行。 -
arr[:, 0]:取所有行的第一列(就是你刚才发的那张图里的操作)。 -
arr[arr > 5]:布尔索引,直接取出所有大于 5 的数(极其好用!)。
② 形状变换(Reshape 与 Transpose)
这俩的区别,和我们之前聊的张量规则完全相同:
-
reshape:改变形状,不改数据顺序(按行先拉平再填充)。 -
.T或transpose():转置/交换轴,改变读取顺序(行列互换)。 -
注意 :
reshape通常返回"视图"(View,共享内存),改了一个另一个会变;如果想独立拷贝,用.copy()。
③ 聚合与轴(Axis)
NumPy 中有一个极其重要的参数 axis(轴),它决定了运算的方向:
-
axis=0:沿着行方向(垂直)压缩,意思是"跨行操作",结果会减少行数。 -
axis=1:沿着列方向(水平)压缩,意思是"跨列操作",结果会减少列数。
经典例子:
python
arr = np.array([[1, 2, 3], [4, 5, 6]])
arr.sum(axis=0) # 输出 [5, 7, 9] (把两行对应列相加,行没了)
arr.sum(axis=1) # 输出 [6, 15] (把两行各自内部的三列相加,列没了)
🔗 与 PyTorch 的"桥梁"(重点)
NumPy 和 PyTorch 可以无缝互转,但有个极大陷阱:
-
CPU 上的张量 转 NumPy:
tensor.numpy(),两者共享内存(改一个,另一个会变)。 -
NumPy 转张量:
torch.from_numpy(arr),同样共享内存。 -
GPU 上的张量 :必须先
.cpu()再转 NumPy,而且这时通常会复制一份数据,不再共享。
⚠️ 新手最易犯的错误
NumPy 默认是**"按行主序(C-order)"** 存储数据。当你使用了 arr.T(转置)后,数据在内存里变得不连续 。此时如果贸然调用 .reshape(),极大概率会报错或得到乱码数据。解决办法 :先调用 np.ascontiguousarray(arr) 再 reshape。
一句话总结 :NumPy 的重点就是 "把数组当数字一样做数学运算" 。只要记住 shape、dtype、axis 三个关键词,剩下的操作(求均值、标准差、矩阵乘法 @)都能顺着这些逻辑推导出来。你对 PyTorch 的理解已经在这里了,NumPy 只是把 .cuda() 去掉后的"陆地版本"。😊