PyTorch进阶:透视张量内存与矩阵乘法
在入门篇中,我们学会了如何创建张量、进行基本的运算。但你是否想过:
- 为什么
view()改变形状非常快,而reshape()有时会变慢? - 为什么
transpose()转置后的张量,在内存中其实是“乱序”的? torch.matmul到底是如何处理高维张量的?
今天,我们将深入 PyTorch 的底层,通过存储和步幅的视角,彻底看懂张量的本质。
实际上,一个张量背后包含三个关键概念:
Storage(存储) + Stride(步长) + Shape(形状)
理解这三者,再结合矩阵乘法,你会对神经网络的计算过程有更深刻的认识。
一、什么是Storage(存储)
Storage 是张量数据在内存中的一维连续存储空间。
比如一个二维张量:
[1234]
\begin{bmatrix}
1 & 2 \\
3 & 4
\end{bmatrix}
[1324]
在内存中其实就是:
[1,2,3,4]
[1,2,3,4]
[1,2,3,4]
也就是说:
所有张量本质都是一段线性内存 + 解释方式
二、什么是Stride(步长)
Stride 决定了:
“当你在某一维移动一步,内存地址跳多少”
例如:
示例:2×2矩阵
[1234]
\begin{bmatrix}
1 & 2 \\
3 & 4
\end{bmatrix}
[1324]
shape = (2,2) # 形状
stride = (2,1) # 步长
解释:
- 行移动:跳 2 个元素
- 列移动:跳 1 个元素
所以访问规则是:
index=i×stride0+j×stride1
\mathrm{i n d e x}=i \times s t r i d e_{0}+j \times s t r i d e_{1}
index=i×stride0+j×stride1
为了理解这个公式,我们需要把张量想象成一栋公寓楼。
- Storage (内存):整栋楼的所有房间被拉成一条长走廊,门牌号从 0 到 N 依次排列
- Tensor (张量):这栋楼有 3 层(行),每层有 4 个房间(列)
- 目标 :寻找住在第 i 层,第 j 号房间的人(数据)
公式解读:
物理门牌号=(层数i×每层房间数)+房间号j
\mathrm{物理门牌号}=(层数i \times 每层房间数)+房间号j \
物理门牌号=(层数i×每层房间数)+房间号j
- stride0s t r i d e_{0}stride0 (行的步幅):就是“每层有多少个房间”。如果你想下一层楼,你得在长走廊里跨过这一整层的房间数。
- stride0s t r i d e_{0}stride0(列的步幅):就是“房间之间的间距”,通常是 1,因为同一层的房间是挨着的。
为什么 Stride 很重要?
因为它允许我们:
- 不复制数据实现 reshape
- 实现 transpose(转置)
- 高效切片(slice)
例如转置:
[1234]→[1324]
\begin{bmatrix}
1 & 2 \\
3 & 4
\end{bmatrix} →
\begin{bmatrix}
1 & 3 \\
2 & 4
\end{bmatrix}
[1324]→[1234]
- Storage(物理内存):
[1, 2, 3, 4](数据是连续存放的) - Shape(形状): 2 行 2 列
- Stride (步幅): (2,1)
stride[0] = 2:往下走一行(比如从1到3),需要跨过2个元素。stride[1] = 1:往右走一列(比如从1到2),需要跨过1个元素。
进行转置操作时,PyTorch 没有把数据拿出来重新排成 1, 3, 2, 4。它只做了两件事:
- 修改 Shape:把形状从
(2, 2)变成(2, 2)(虽然数字没变,但逻辑变了)。 - 交换 Stride:把步幅从
(2, 1)变成 (1, 2)。
三、矩阵乘法的本质
理解了 Stride,我们再来看矩阵乘法。在 PyTorch 中,矩阵乘法不仅仅是数学公式,它还涉及到维度的广播和内存布局。
1. 核心规则:内积
矩阵乘法 C=A×BC=A×B 的核心是:A 的行 与 B 的列 做点积。
前提条件:A 的最后一维长度必须等于 B 的倒数第二维长度。
- A:(N,K)A:(N,K)A:(N,K)
- B:(K,M)B:(K,M)B:(K,M)
- Result:(N,M)Result:(N,M)Result:(N,M)
2. torch.mm vs torch.matmul
我们在基础篇提到过它们的区别,现在用底层视角再看一遍:
torch.mm:它是严格的。它只接受 2D 张量,效率极高,但不支持广播。torch.matmul:它是智能的。- 如果输入是 1D,它会自动补维度(向量点积)。
- 如果输入是 3D+,它会广播。
3. 高维矩阵乘法(批量矩阵乘法)
以 3 维张量相乘作为例子 (Batch, N, K) @ (Batch, K, M)
PyTorch 的处理逻辑是:
- 广播批次维度:就像广播加法一样,对齐前面的维度。
- 锁定最后两维:对最后两维执行标准的矩阵乘法。
# 模拟一个 Batch Size 为 2 的数据
A = torch.randn(2, 3, 4) # 2个矩阵,3行4列
B = torch.randn(2, 4, 5) # 2个矩阵,4行5列
C = torch.matmul(A, B)
print(C.shape) # torch.Size([2, 3, 5])
四、进阶练习题
为了巩固今天的内容,请尝试解答以下问题(可以先思考,再看答案)。
练习题 1:计算 Stride
有一个张量 x 形状为 (2, 3, 4),它是连续存储的。请问 x.stride() 的输出是什么?
答案: (12, 4, 1)
解析:
- 这是一个 3维 张量。
- 第2维(最内层):移动 1 步,内存跳过 1 个元素 → 1。
- 第1维(中间层):移动 1 步,跳过一整个“行”(长度为4) → 4。
- 第0维(最外层):移动 1 步,跳过一整个“块”(3行4列)→ 3 * 4 = 12。
练习题 2:转置后的 Stride
接上题,如果对 x 执行 x.transpose(0, 2),新的 Stride 是多少?
答案: (1, 4, 12)
解析:
transpose(0, 2)交换了第 0 维和第 2 维的索引。- 原来的 Stride 是
(12, 4, 1)。 - 交换位置后,Stride 变为
(1, 4, 12)。 - 注意:此时张量在内存中变得不连续。
练习题 3:矩阵乘法维度
代码 torch.randn(10, 5) @ torch.randn(5,) 的输出形状是什么?
答案: torch.Size([10])
解析:
- 左边是矩阵
(10, 5)。 - 右边是向量
(5,)。 - 根据
matmul规则,向量被视为列向量(5, 1)。 - 运算:
(10, 5) @ (5, 1) -> (10, 1)。 - 特殊规则:如果最后结果是
(N, 1)这种单列矩阵,matmul会自动压缩掉最后的 1,变成(10,)的向量。
练习题 4:连续性问题
x = torch.randn(3, 4)
y = x.transpose(0, 1)
z = y.view(12) # 这里会发生什么?
答案: 报错 RuntimeError: view size is not compatible with input tensor's size and stride...
解析:
y是转置后的张量,其 Stride 为(1, 3)(假设原 Stride 为(4, 1)),它是非连续的。view()要求底层数据必须是连续的,以便重新解释形状。- 修正:使用
z = y.contiguous().view(12)。
本文章仅用于 个人深度学习技术学习、知识点梳理与交流探讨,聚焦 PyTorch 张量的存储(Storage)、步幅(Stride)及矩阵乘法底层逻辑,所有内容均为个人学习过程中的理解与总结,不用于任何商业用途、盈利活动及正式教学场景。
因本人对 PyTorch 底层原理的认知仍在逐步提升中,文章内容(包括知识点解读、代码示例、练习题解析等)可能存在疏漏、错误或表述不严谨之处,不代表 PyTorch 官方技术标准。
在此恳请各位学习者理性看待,若发现文中错误、有不同理解或补充建议,欢迎留言交流、批评指正,共同探讨、共同进步,一起夯实 PyTorch 学习基础。
更多推荐



所有评论(0)