博主介绍:程序喵大人

在这里插入图片描述
第 01 章:模型到底是什么?把模型拆成了参数和网络结构,数据在传送带上流动。但那些流动的"数据箱"到底是什么格式?装了多少东西?怎么描述它的规格?这一章专门回答这个问题。

你做 AI Infra,打交道最多的不是模型的数学原理,而是数据的形状。报错了先看 shape,优化了先算 shape,做显存估算要从 shape 算起。

这一章从最基础的"一个数"开始,一路讲到 forward pass 里 shape 怎么一层层变化。读完之后,你看到任何一行模型代码里的 .shape 输出都不会再陌生。

一、从标量到 Tensor:维度阶梯

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

台阶从低到高,每上一级加一个维度。(很多概念其实我们在学校的时候已经学过了,我们再来复习一遍!)

第 0 级是标量(scalar),shape 是 (),就是一个裸数字,比如 3.14。没有任何方向,没有任何结构,最简单的形态。

第 1 级是向量(vector),shape 是 (N,),一排数字。比如一个 5 维向量 shape 是 (5,),里面有 5 个数。词向量、隐藏层的一个时间步输出,都是向量。

第 2 级是矩阵(matrix),shape 是 (M, N),行列二维排列。第 01 章:模型到底是什么?里的参数矩阵 W 就是这一级,shape (3, 2) 表示 3 行 2 列共 6 个数。

第 3 级及以上是 tensor,shape 有三个或更多数字。(3, 4, 3) 就是 3 层、每层 4 行、每行 3 列。深度学习里几乎所有数据都活在这一级:batch of sequences、batch of images、batch of embeddings,全是三维或四维的 tensor。

shape 里有几个数,就是几维,这是最直接的读法。每加一个维度,就多一层嵌套。

二、Shape 怎么读在这里插入图片描述

:(batch, seq, hidden)

拿一个具体的 tensor 来看吧:shape (4, 128, 768)

这是一个三维货柜箱:正面 4 层(batch),侧面每层 128 行(seq),顶面每行 768 列(hidden)。

读法:从左到右 = 从最外层维度到最内层维度。

  • 第 1 个数 4 是 batch,最外层——一次送进来 4 个样本
  • 第 2 个数 128 是 seq,中间层——每个样本有 128 个位置(token)
  • 第 3 个数 768 是 hidden,最内层——每个位置用 768 个数字表示特征

总元素数 = 4 × 128 × 768 = 393,216 个浮点数。FP16 存储的话,就是 393,216 × 2 字节 ≈ 786 KB。

这个 (batch, seq, hidden) 的结构在 LLM 的 forward pass 里到处出现。你第一次在代码里看到 tensor.shape = torch.Size([4, 128, 768]),对照这张图,每个数字对应什么就一目了然。

注意:名字只是习惯,shape 本身不带名字。 (4, 128, 768)里的 4 不会告诉你"我是 batch",这个语义是你自己加的,要靠上下文理解。

三、Batch 维度:GPU 同时处理多少个样本

在这里插入图片描述

为什么 AI Infra 工程师看到 shape 第一眼先看 batch?因为 batch 直接决定了 GPU 在做什么。

画面是 8 条并行传送带同时把样本送进 MatMul 工位,8 条传送带同时出结果。batch=8 意味着 GPU 并行处理 8 个样本,而不是排队一个一个来。

GPU 有成千上万个并行计算单元,batch=1 时大多数单元闲着,算力浪费;batch=32 或更大时,这些单元才能充分跑起来。batch 越大,GPU 利用率越高,单位时间内处理的样本越多(吞吐量越大)。

但 batch 也直接乘进了显存。在前向传播(Forward Pass)过程中,GPU 需要保存每一层计算出来的中间结果(即激活值,Activation),以便在反向传播(Backward Pass)中计算梯度。这些激活值 Tensor 的 shape 每一维都包含 batch。

其显存计算公式为:单层激活值显存 = batch × seq × hidden × 字节数。这意味着激活值占用的显存大小与 batch size 呈完全线性的正比关系。

我们可以算一笔账:

假设一个大语言模型的隐藏层维度 hidden = 4096,输入序列长度 seq = 2048,使用 FP16(半精度浮点数,每个元素占 2 字节)进行训练。

  • 当 batch = 1 时,单层该特征的激活值显存为:

1 × 2048 × 4096 × 2  字节 ≈ 16  MB 1 \times 2048 \times 4096 \times 2 \text{ 字节} \approx 16 \text{ MB} 1×2048×4096×2 字节16 MB

  • 当 batch = 32 时,该激活值显存直接暴涨了 32 倍:

32 × 2048 × 4096 × 2  字节 ≈ 512  MB 32 \times 2048 \times 4096 \times 2 \text{ 字节} \approx 512 \text{ MB} 32×2048×4096×2 字节512 MB

这仅仅是某一层中某一个 Tensor 的大小。实际上,一个模型包含几十层,每层又有注意力机制的 QKV、MLP 的中间放大层(通常是 4 倍 hidden)等多个激活值 Tensor,它们累加起来的显存总量极其惊人。batch 从 1 涨到 32,激活值显存就会从几个 GB 猛增至上百 GB,极易导致显存溢出(OOM)。

因此,训练和推理服务在 batch 策略以及显存瓶颈上有着本质的不同:

训练阶段:因为需要反向传播,GPU 必须在显存中保留所有层的中间激活值,显存压力极大。为了提升吞吐量并稳定梯度,通常需要配合激活值重计算(Activation Recomputation)等技术,在显存允许的上限内尽量把 batch 推大。

推理阶段:推理不需要反向传播,因此普通激活值用完即释放,不随层数累加,显存开销极小。在生产环境中,为了提高 GPU 利用率,服务框架(如 vLLM)通常会使用连续批处理(Continuous Batching)技术,把不同用户的多个并发请求动态拼成一个 batch 来处理。此时限制最大 batch 数量的不再是激活值,而是用于存放上下文历史特征的 KV Cache,它同样会随着 batch size 的增大而线性暴涨。

这就是为什么 AI Infra 工程师第一眼看 shape 看 batch——batch 是吞吐和显存之间最大的调节旋钮。

四、Reshape 与 transpose:同一份数据,换个排法

在这里插入图片描述

24 个编号小球,两种完全不同的"换个排法"方式。

Reshape 是换形状,不改顺序。原始 shape (6, 4) 的托盘,reshape 成 (24,) 就是拍平成一排,reshape 成 (2, 3, 4) 就是叠成两层。总数不变,顺序不变,只是改了盒子的分格方式。 1 号球始终在 2 号球前面,无论换成什么形状。

Transpose 是交换维度,会真的改变数据的排列顺序。(6, 4) transpose 成 (4, 6),原来按行排的现在按列排,1 号球的邻居变了。维度顺序变了,数据的物理访问路径也变了。

这个区别有严重的实际后果。虽然这两种操作底层的一维物理内存都没有移动,但逻辑读取顺序和物理存放顺序的关系变了。

我们可以用一个简单的例子来推演:假设物理内存里挨个存了 6 个数 [1, 2, 3, 4, 5, 6],原始逻辑 shape 是 (2, 3)(即 2 行 3 列)。

  • Reshape 成 (3, 2):按物理内存顺序重新画格子,变成 3 行 2 列。此时你按逻辑一行行读,读出来的依然是 1,2,3,4,5,6。物理内存没动,逻辑顺序也没违背物理顺序,所以数据依然是连续的(contiguous)。
  • Transpose 翻转成 (3, 2):直接把原来的行变成列。此时你再按新的逻辑一行行读,读出来的变成了 1,4,2,5,3,6。逻辑上相邻的两个数(比如 1 和 4),在底层的物理内存里中间隔了别的数字,变成了跳跃的。这打破了物理顺序,也就是非连续的(non-contiguous)。

在代码里,tensor.reshape(new_shape)tensor.transpose(dim0, dim1)tensor.permute(...) 是完全不同的操作,不要混用。

五、Contiguous 与 stride:内存里到底怎么排

在这里插入图片描述

上一张图说 transpose 后内存不连续。这张图我们把内存的实际布局展开来看看。

地板上的一排地砖就是物理内存,格子 0 到 15,每格存一个数。

stride 的本质就是在回答一个问题:“当我在逻辑矩阵里往下走一格,或者往右走一格时,底层的物理内存指针需要跨过多少个格子?”

我们把 shape = (2, 8) 的 16 个数拆开看,它们在物理内存上是老老实实排成一条直线的:0, 1, 2... 15

Contiguous(连续)的情况:

  • 往右走一步(沿列方向): 比如从数字 0 走到数字 1,它们在物理内存里也是紧挨着的。所以逻辑上往右走一步,物理内存只需跳 1 格。
  • 往下走一步(沿行方向): 比如从第 1 行的数字 0 直直往下走到第 2 行的数字 8。在 1D 物理内存里,0 和 8 中间隔了第一行的其他数字。物理指针必须大步流星跨过整整 8 格才能到达。
  • 连起来,这个矩阵的 stride 就是 (8, 1)。只要你顺着往右读(不停跳 1 格),移动轨迹就完全顺应物理内存的天然顺序,GPU 读起来极其高效。

非 Contiguous(转置后)的情况:

同样这 16 个数,做完 transpose 后,逻辑 shape 变成了 (8, 2),但物理内存没动。

  • 此时 stride 变成了 (1, 8)。这说明原本挨着的数在逻辑上被翻转了。现在你往右走一步去读第二列时,物理指针要大幅跳跃 8 格。
  • 因为读取时要在内存里来回倒腾,访问路径变成了锯齿状,GPU 的内存访问效率就会随之下降。

只要数据是 Contiguous 的,它的 stride 从左到右一定是单调递减的,这代表它的逻辑访问顺序和物理存放顺序完美吻合。

遇到 RuntimeError: non-contiguous tensor 报错,通常调用 .contiguous() 把数据重新整理成连续布局,或者用 .reshape() 隐式触发同样的效果。

六、Broadcasting:shape 不一样也能算

两个 shape 不同的 tensor 直接相加,而且不报错,并且自动出结果——这就是 broadcasting。在这里插入图片描述

画面里:shape (4, 3) 的大箱子和 shape (3,) 的小托盘,要做逐元素相加。小托盘只有 1 行 3 列,大箱子有 4 行 3 列。broadcasting 的做法是把小托盘逻辑上复制 4 份,凑成 (4, 3) 的尺寸,再做加法。注意"逻辑复制"——不真实分配内存,只是在读取时重复读那 3 个数。

Broadcasting 的核心底线是:系统可以帮你自动补齐数据,前提是补齐方式必须“毫无歧义”。它绝不负责猜你的意图。

具体规则是:两个 shape 从右往左对齐,逐维度比较。

要么维度相等,要么其中一边必须是 1(或者缺失)。因为只有当数据只有 1 份时,“直接无脑复制多份”才是唯一合乎逻辑的操作。

几个推演例子(假设要和 (4, 3) 的大矩阵相加):

  • 加上 (3,)(相当于 1 行 3 列):系统毫无歧义地把这 1 行原样复制 4 份,凑成 4 行去相加。✓
  • 加上 (4, 1)(4 行 1 列):系统毫无歧义地把这 1 列原样复制 3 份,凑成 3 列去相加。✓
  • 加上 (2, 3)(2 行 3 列):系统彻底懵了。为了凑齐 4 行,它是该把这两行循环重复一次(Row0, Row1, Row0, Row1)?还是各自复制两遍(Row0, Row0, Row1, Row1)?因为无法毫无歧义地推断你的意图,系统拒绝瞎猜,直接抛出报错。✗

正因为 Broadcasting 会静默地在后台扩充数据,一旦你的 shape 阴差阳错碰巧符合了规则(比如 1 变成了大数字),它就不会报错,而是默默给出一个巨大且完全不符合你预期的结果。所以在深度学习里,遇到奇怪的输出或显存暴涨,第一反应永远是去查验一下参与运算的 tensor shape。

七、一次 forward,shape 怎么一路变下去

在这里插入图片描述

现在,我们把前面零碎的 shape 知识串起来,放进大模型的一次完整 Forward Pass 里,看看这块三维数据是怎么流动的。

假设初始输入是 (batch=4, seq=128, hidden=768)

  • 经过工位 A(Embedding 查表):shape 不变,依然是 (4, 128, 768)。这一步只是把干巴巴的 token ID 替换成了 768 维的稠密特征向量。位置数和批量都没变。

  • 经过工位 B(Linear 矩阵乘 / FFN 层):hidden 维度发生巨变,(4, 128, 768)(4, 128, 3072)

    • 机制上:当前数据的最后一维 768 会和权重矩阵 (768, 3072) 做乘法对齐,输出的最后一维变成 3072。前两维(batch 和 seq)属于外层循环,直接穿透不受影响。
    • 直觉上:模型通过放大维度,临时向系统申请了更宽广的“操作空间”和更大的“脑容量”。它把原本高度揉杂的 768 维特征在这里拆解、铺开,以便更精细地解耦特征、激活预训练时存下的庞大知识库。
  • 经过工位 C(LayerNorm 或激活函数):shape 保持 (4, 128, 3072) 不变。这类算子属于逐元素操作(Element-wise),只负责调整具体的数值大小来做归一化或非线性过滤,绝对不会去动数据容器的形状。

这就是大模型内部 shape 流动的基本铁律: 外层的批量(batch)和序列(seq)通常是一路穿透到底的;真正的形变全发生在最内层(hidden),且基本上只有矩阵乘法(MatMul)有资格去改变它。

在日常手撕模型或调试 bug 时,逐层打印 tensor.shape 永远是最高效的排查手段。形状对不上,比数值不对要容易发现得多,它能帮你在一秒钟内锁定到底是哪一层算子出了问题。

码字不易,欢迎大家点赞,关注,评论,谢谢!

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐