在入门篇中,我们学会了如何创建张量、进行基本的运算。但你是否想过:

  • 为什么 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。它只做了两件事:

  1. 修改 Shape:把形状从 (2, 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 的处理逻辑是:

  1. 广播批次维度:就像广播加法一样,对齐前面的维度。
  2. 锁定最后两维:对最后两维执行标准的矩阵乘法。
# 模拟一个 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 学习基础。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐