Create your own
Lesson illustration

根据张量形状判断线性层与批次、序列维度的矩阵运算合法性

你好,欢迎开始这门面向大模型后训练的基础课程。后训练中的 SFT、偏好优化、PPO/GRPO、奖励模型训练,表面上是不同算法,但都依赖同一项基本能力:看到一个张量,能立刻判断每个轴代表什么、某个计算是否能执行,以及输出会是什么形状。

本模块先补齐这套“形状语言”。今天聚焦最常见的一类操作:线性层。你会学会判断批次维、序列维和特征维在 nn.Linear 与矩阵乘法中如何变化;这也是日后阅读 Transformer、训练框架和报错日志时最先需要用到的能力。


先把形状读成“有语义的坐标轴”

一个张量的 shape 是各轴长度组成的元组。例如:

在语言模型训练中,最常见的解释如下:

符号常见名称含义例子
batch size一次送入模型的样本数8 条对话
sequence length每条样本包含的 token 数512 个 token
hidden size每个 token 的隐藏表示长度4096
output features线性层输出特征数11008
vocabulary size词表大小151936

因此,形状为 [8, 512, 4096] 的隐藏状态张量,通常表示:

  • 条序列;
  • 每条序列有 个 token 位置;
  • 每个 token 位置有一个 维向量。

这里有一个必须建立的习惯:形状合法,不等于语义正确。

例如,[T, B, H][B, T, H] 都是三维张量,送入 nn.Linear(H, O) 都不会报错,因为最后一维都是 。但若某个模块约定第 0 轴是 batch,而你误把序列轴放在第 0 轴,模型虽然“能跑”,数据含义却已错位。后训练工程里,这类静默错误往往比直接报错更难排查。

为了建立对“矩阵形状”的直觉,可以先看一段简短的可视化说明。

Nonsquare matrices as transformations between dimensions | Chapter 8, Essence of linear algebra

观看 3Blue1Brown 的《Nonsquare matrices as transformations between dimensions》。它用几何语言解释了矩阵的行数对应输出维度、列数对应输入维度;这正是理解线性层权重形状的核心。

观看 三乘二矩阵,理解一个 3 \times 2 矩阵接收 2 维输入、产生 3 维输出;再观看 二乘三矩阵,对照理解输入和输出维度互换时矩阵形状为何也要互换。注意“列数匹配输入、行数决定输出”这一规律。


线性层:只检查最后一维,并保留前面的轴

PyTorch 的线性层写作:

nn.Linear(in_features, out_features)

若某个 token 的隐藏向量为

线性层参数为

则输出为

其中:

  • 输入维度是 ,所以权重的第二维必须是
  • 输出维度是 ,所以权重的第一维是
  • 输出向量的维度是

严格地说,含 bias 的 nn.Linear仿射变换,而非纯线性变换;但深度学习工程中通常仍称其为“线性层”。

PyTorch 以 (out_features, in_features) 存储权重。因此:

proj = nn.Linear(4096, 11008)

对应的参数形状是:

proj.weight.shape  # [11008, 4096]
proj.bias.shape    # [11008]

最关键的 API 规则是:

也就是说,nn.Linear最后一维视为输入特征维;此前的任意多个维度都会被保留。

Linear — PyTorch 2.9 documentation

阅读 PyTorch 官方的 Linear 文档,确认线性层的权重存储方式、输入输出形状约定,以及二维 batch 输入的基础示例。

在页面开头阅读 torch.nn.Linear 的定义、Parameters、Shape 与 Variables 部分。先用 in_features 与 out_features 的参数说明确认两者的角色;随后重点看 Shape 段落中输入 (*, H_in) 和输出 (*, H_out) 的约定:星号代表任意数量的前导维度,最后一维才必须匹配 in_features。最后阅读 Examples,观察 [128, 20] 如何得到 [128, 30]。

现在把这条规则用在语言模型隐藏状态上:

import torch
import torch.nn as nn

B, T, H, O = 4, 128, 768, 3072

hidden_states = torch.randn(B, T, H)
up_proj = nn.Linear(H, O)

output = up_proj(hidden_states)

print(hidden_states.shape)  # torch.Size([4, 128, 768])
print(output.shape)         # torch.Size([4, 128, 3072])

这个计算合法,因为:

  1. hidden_states 的最后一维是
  2. up_proj.in_features 也是
  3. batch 维 与序列维 原样保留;
  4. 输出特征维由 变成

从概念上说,线性层对每一个位置的 维向量应用同一组参数。也可以把 [B, T, H] 暂时想成 [B*T, H],经过线性层后得到 [B*T, O],再恢复成 [B, T, O]。但实现时不需要手动展平,nn.Linear 已经支持任意数量的前导维度。

相反,下面的计算必然非法:

bad_hidden_states = torch.randn(4, 128, 512)
output = up_proj(bad_hidden_states)

因为输入最后一维是 ,而线性层要求 。无论 batch size 和 sequence length 是多少,都无法弥补这个不匹配。


Transformer 中的批次维、序列维与特征维

下面这张 Transformer 架构图值得作为今后读代码时的形状参考。

该图展示了 Transformer 编码器和解码器中的典型张量形状:输入 token 经 embedding 与位置编码后成为 `(bs, N, d_model)`,解码器侧为 `(bs, M, d_model)`,最终线性层把最后的 `d_model` 维映射为词表维度。

以 decoder-only 大语言模型为例,一条常见的形状链是:

模块或张量形状发生了什么
input_ids[B, T]每个位置是一个整数 token ID
token embedding 输出[B, T, H]每个 token ID 被查表为 维向量
Transformer block 输出[B, T, H]注意力和 MLP 改变表示内容,但通常保持形状
LM head 输出 logits[B, T, V]每个 token 位置得到对整个词表的打分

注意:embedding 层和线性层都可使张量最后一维变化,但机制不同。

  • Embedding 将离散 token ID 映射成连续向量,常见输入是 [B, T],输出是 [B, T, H]
  • 线性层接收已经是浮点表示的张量,例如 [B, T, H],输出 [B, T, O]

在后训练中,最常见的最终线性投影是语言模型头:

它为每个序列位置、每个词表 token 产生一个 logit。之后的 softmax、标签移位与交叉熵损失会在后续课程中展开。

还可以用 MLP 层观察特征维如何变化。一个典型 Transformer MLP 中可能有:

经过上投影层得到:

经过激活函数后形状不变,随后经下投影层回到:

这解释了为何残差连接通常合法:主分支输出和残差分支输入都具有 [B, T, H],能够逐元素相加。


矩阵乘法:中间维必须相同

线性层的核心是矩阵乘法。对于二维矩阵:

乘积合法:

判断规则只有一句:

左侧矩阵的列数,必须等于右侧矩阵的行数。

其中 是被消去的“中间维”。输出保留左矩阵的行数 与右矩阵的列数

例如:

而:

不合法,因为相邻的维度 不相等。

这恰好解释了一个常见困惑。若权重以 PyTorch 的格式存储:

而输入以“每行一个样本”的格式表示:

那么应使用:

因为:

所以,下面手写实现与 nn.Linear(H, O) 的权重方向一致:

x = torch.randn(4, 128, 768)
layer = nn.Linear(768, 3072)

y_manual = x @ layer.weight.T + layer.bias
y_module = layer(x)

print(y_manual.shape)  # torch.Size([4, 128, 3072])

实际工程中优先使用 layer(x)torch.nn.functional.linear;手动写矩阵乘法主要用于理解、调试和验证自定义实现。

torch.matmul@ 对高阶张量采用“最后两维是矩阵维”的规则。因而:

对于左侧张量,最后两维 [T, H] 可视为一个 矩阵,前面的 轴则表示有 组这样的矩阵。

这与模型语义中“序列轴”的叫法并不冲突:

  • nn.Linear 而言,BT 都只是被保留的前导轴;
  • matmul 而言,T 在这次操作中扮演左矩阵的行数;
  • 对注意力而言,T 还会表示 query 或 key 的位置数。

同一个轴的角色取决于当前操作,不能只凭轴的位置机械记忆。

例如,注意力分数的形状常见为:

因此:

这里两个 分别对应 query 位置与 key 位置。我们之后会手工计算注意力;现在只需能检查:相乘的中间维都是 ,因此这一矩阵乘法合法。


广播:它能帮助逐元素运算,不能修复矩阵乘法

线性层公式中的 bias:

看起来像把 [B, T, O][O] 相加。它之所以合法,依赖的是 broadcasting(广播)

广播用于逐元素加、减、乘、除等操作。比较两个形状时,从最右侧向左对齐;每一对对应维度必须:

  1. 数值相等;或
  2. 其中一方为

例如:

会把 [O] 看作 [1,1,O],然后在 batch 和 sequence 方向复制,因此结果是 [B,T,O]

下面的视频用二维张量加一维向量展示了同一规则。

PyTorch Course (2022), Part 1: Tensors

观看 Mr. P Solver 的《PyTorch Course, Part 1: Tensors》中的广播示例。它解释了为什么 [6, 5] 和 [5] 可以逐元素相加,这与线性层中 [B, T, O] 加 [O] 的 bias 操作完全同构。

观看 二维广播示例。重点看讲解者如何从最右侧比较维度,以及长度为 5 的向量如何沿着行方向被复用。将其中的 [6, 5] 替换为 [B, T, O],将 [5] 替换为 [O],即可得到线性层 bias 的广播过程。

一个很常见的后训练代码问题是:你想为 batch 中每条样本乘一个系数,但写出了错误形状。

rewards = torch.randn(B, T)
sample_weights = torch.randn(B)

直接计算 rewards * sample_weights 通常不符合预期,因为 [B] 会优先对齐 rewards 的最后一维 ,而不是第 0 维的

应显式把权重写成:

weighted_rewards = rewards * sample_weights[:, None]

这时形状是:

[B, T] * [B, 1]

广播发生在长度为 的序列维上,每条样本使用自己的权重。

对三维隐藏状态 [B, T, H] 做按样本缩放,则可使用:

scales = torch.randn(B, 1, 1)
scaled_hidden_states = hidden_states * scales

广播很强大,但务必记住它的边界:

  • 广播可以使逐元素运算合法;
  • 广播不改变矩阵乘法的中间维匹配要求;
  • HO 不相等时,不能指望广播使 X @ W 自动正确;
  • 显式使用 unsqueezeNonereshape,通常比依赖“碰巧对齐”更安全。

一套可用于读代码和排错的形状检查法

以后看到任意训练代码中的一行张量操作,可以按下面五步检查。

  1. 写出每个张量的完整 shape。
    不满足于“这是 logits”或“这是 reward”;直接标记为 [B, T, V][B][B, T] 等。

  2. 为每个轴补上语义。
    确认它是 batch、sequence、head、hidden、vocab,还是一个临时矩阵维度。

  3. 识别操作类别。
    逐元素操作检查广播;矩阵乘法检查相邻中间维;线性层检查输入最后一维;拼接则检查除拼接轴外的所有维度。

  4. 在执行前预测输出 shape。
    例如看到 nn.Linear(768, 3072) 作用于 [2, 1024, 768],应先得出 [2, 1024, 3072]

  5. 用断言把推理写进代码。

assert hidden_states.ndim == 3
assert hidden_states.shape[-1] == up_proj.in_features

projected = up_proj(hidden_states)

assert projected.shape[:2] == hidden_states.shape[:2]
assert projected.shape[-1] == up_proj.out_features

在需要排查框架调用链或自定义 loss 时,形状日志也非常有效:

print(f"input_ids: {input_ids.shape}")
print(f"hidden_states: {hidden_states.shape}")
print(f"logits: {logits.shape}")
print(f"labels: {labels.shape}")

尤其要警惕两类错误:

  • 显式报错:如 mat1 and mat2 shapes cannot be multiplied,通常是线性层输入特征维或矩阵乘法中间维不匹配。
  • 静默错误:如 batch 与 sequence 轴被交换、广播对齐到了错误轴、mask 的 shape 虽可广播但语义错位。它们不会停止训练,却会让结果异常。

小结

今天建立了几条后训练工程中高频使用的形状规则:

  • [B, T, H] 通常表示一批 token 的隐藏状态, 分别承担不同语义。
  • nn.Linear(H, O) 只要求输入最后一维等于 ,并将最后一维替换为
  • PyTorch 线性层权重的存储形状是 [O, H];以行向量形式计算时使用
  • 二维矩阵乘法要求左矩阵列数等于右矩阵行数;高阶 matmul 则在最后两维应用这一规则。
  • 广播服务于逐元素运算,例如线性层 bias 的 [O] 加到 [B, T, O];它不能修复错误的矩阵乘法。
  • 合法的 shape 不保证语义正确,因此要同时检查轴长度与轴含义。

下一课会把这些形状规则落实到具体数字上:手工计算向量点积和小型矩阵乘法。那时你会更清楚地看到,线性层输出的每一个元素究竟是怎样由输入特征与权重共同计算出来的。

Can't find a good explanation? Sign up and we'll make it for you

Sign up