系列 《算法工程师的数学:读公式不卡壳的最小集》 第 1 / 9 篇
下一篇:内积、范数与余弦相似度
打开任何一个大模型的结构图,看到的全是矩阵乘法:一个 token(模型处理文字的最小单位,大致是一个词或半个词——”我爱北京”约 3 个 token——不是一字一个,常用词”北京”通常整个是一个 token,”我”“爱”各一个;反过来英文长词 “tokenizer” 会被切成两三个。切法由分词器(tokenizer)的词表决定,不同模型不同;模型读到的不是字,而是 token 在词表里的编号)变成一个向量,向量乘一个矩阵变成另一个向量,再乘一个矩阵……几十层之后输出下一个 token 的概率。所以线性代数在 AI 里的第一个用处不是什么高深的定理,而是两条极其朴素的规则:形状规则(什么形状乘什么形状得到什么形状)与成本规则(这一乘要算多少次)。会了这两条,读模型结构图、算训练成本、判断”这个改动为什么慢”都有了起点。
本篇从零建立这两条规则。全篇的核心问题是:
一、总览
1. 本文的对象
Llama-3-8B 的一层里最常见的一步:一个 token 的表示(一个长 4096 的向量)乘一个 \(4096 \times 4096\) 的权重矩阵。本文要把这一步的每个字都解释清楚——”长 4096 的向量”是什么、”\(4096 \times 4096\) 的矩阵”是什么、”乘”是怎么定义的、算了多少次——然后把它推广到一层、一个模型。
2. 本文的章节安排
| 章 | 主题 | 内容 |
|---|---|---|
| 二 | 从标量到张量 | 标量、向量、矩阵、张量各是什么、怎么写、代码里是什么形状 |
| 三 | 矩阵乘法与形状规则 | 定义、内维必须相同、三种看法、单位矩阵与转置 |
| 四 | 成本规则 | FLOPs 是什么、\(2mnk\) 从哪来、代 Llama-3-8B 的数字 |
| 五 | 训练代码里的形状操作 | 批维度、逐元素运算、广播、reshape |
| 六 | 一层 Transformer 有哪些矩阵 | 七个权重矩阵、一层的参数量与 FLOPs、为 L4 的算账铺路 |
| 七 | 本文小结 | |
| 八 | 自测 | 五道题 |
二、从标量到张量
1. 四个称呼,一个东西
标量、向量、矩阵,看起来是不同的东西,其实都可以统一到张量(Tensor)这个概念下。在深度学习中,可以先把张量理解为多维数值数组:没有轴的是标量,一个轴的是向量,两个轴的是矩阵,更多轴的则称为高阶张量。
在 PyTorch 中,它们都可以用同一种类型 torch.Tensor 表示,下面这些称呼的区别主要体现在形状(shape)和轴数上:
| 名字 | 例子 | 轴的数量(张量维数) | 代码里的形状(shape) | 底层本质 |
|---|---|---|---|---|
| 标量 | 3.7 |
0 | () |
零维张量 |
| 向量 | (0.2, -1.1, 0.5) |
1 | (3,) |
一维张量 |
| 矩阵 | [[1, 2, 3], [4, 5, 6]] |
2 | (2, 3),即 2 行 3 列 |
二维张量 |
| 高阶张量 | 一批 32 张 28×28 的灰度图 | 3 | (32, 28, 28) |
本例为三维张量 |
标量(scalar)是一个数:学习率 0.001、loss 1.83、温度 0.7。
向量(vector)是一组有序的数,写作 a = (a₁, a₂, …, a_d),其中 d 是向量的长度,也常叫向量的维度。注意这里“维度”的两种用法:长 4096 的向量在数学上叫“4096 维向量”,但在代码里是形状为 (4096,) 的一维张量——前者数分量,后者数轴。
在 AI 里,一个 token 在模型内部通常用一个向量表示。例如,Llama-3-8B 的每个 token 的主干隐藏表示是一个长 4096 的向量,这个 4096 叫隐藏维度(hidden size),常写作 d 或 d_model。“隐藏”沿用神经网络的术语,指输入与输出之间的内部表示,并不是说这些值无法查看。为什么 token 的向量也算”内部”:模型从外面接到的只是 token 的编号(一个整数,比如”北京”→ 20302),第一步就把编号查表换成一个 4096 维的向量——从这一步起,直到最后一层输出下一个 token 的概率之前,所有向量都只在模型内部流动,外面既不输入它们也不读取它们,所以整段都叫隐藏表示,它的宽度就叫隐藏维度。
向量作为抽象对象不必区分横竖,代码中的一维张量也没有行、列之分。为了使用矩阵乘法,可以把它写成形状为 (d, 1) 的列向量,或 (1, d) 的行向量。教材常用列向量,写 Wx;xᵀ 中的 T 表示转置,把列向量变成行向量,反之亦然。
本文按 token 沿行排列的习惯写 xW。两种约定可以表示同一个线性变换,但切换约定时,权重矩阵也要转置:(Wx)ᵀ = xᵀWᵀ,不能只把乘法顺序调换。PyTorch 两种写法都支持;一维张量在 matmul 中的处理规则将在后文说明,第三章第 4 节会对照这两种约定。
矩阵(matrix)是一张二维数表。矩阵 W 有 m 行、n 列,形状就是 (m, n);第 i 行、第 j 列的元素写作 Wᵢⱼ。在大语言模型里,大部分参数存放在矩阵中,例如嵌入矩阵和线性层的权重;也有归一化权重、偏置等向量参数。一个 4096 × 4096 的矩阵包含 16,777,216 个数,约 16.8M 参数。
张量(tensor)在本文的深度学习语境中,可以理解为按若干个轴组织起来的数值数组。三维及以上的张量,可以直观地看作“把矩阵再堆一维或多维”。例如,一批 32 个句子,每句统一为 512 个 token,每个 token 用长 4096 的向量表示,合起来就是形状为 (32, 512, 4096) 的三维张量,三个轴分别表示批次、序列、特征。
各个轴怎样参与计算,取决于具体运算。“最后两维做矩阵乘法,前面的维度组织批次”是 torch.matmul 在两个输入都至少二维时的规则,不是所有张量运算的规则。这里用 @ 表示矩阵乘法;* 则表示逐元素乘法。
例如,输入 X 的形状为 (32, 512, 4096),权重 W 的形状为 (4096, 4096):
(32, 512, 4096) @ (4096, 4096) → (32, 512, 4096)
这在数学上相当于做 32 份 (512, 4096) @ (4096, 4096):每句话的 512 个 token 都乘同一个 W。前面的 32 用来索引批次,本次运算不会沿批次轴求和,也不会混合不同句子的数据。这里的“32 份”描述的是计算含义,不代表底层一定调用 32 次独立运算。
两个输入都带批次维度时也是如此:
(32, 512, 4096) @ (32, 4096, 64) → (32, 512, 64)
这次是 32 对矩阵分别相乘,右边每一批可以是不同的矩阵。一般规则是:
- 矩阵维度:左边最后一维必须等于右边倒数第二维,也就是
(m, k) @ (k, n) → (m, n)。这里的 k 必须相等,不能靠广播匹配。 - 批次维度:去掉最后两维后,剩余形状从右向左对齐,缺失的维度按
1处理;每对大小必须相等,或者其中一个为1。 - 结果形状:广播后的批次形状,再接上
(m, n)。
所以,输入的维数和批次大小不必完全相同:
| 左边形状 | 右边形状 | 结果 |
|---|---|---|
(5, 32, 512, 4096) |
(32, 4096, 64) |
(5, 32, 512, 64) |
(5, 1, 512, 4096) |
(1, 32, 4096, 64) |
(5, 32, 512, 64) |
(2, 32, 512, 4096) |
(5, 32, 4096, 64) |
报错:批次大小 2 与 5 不兼容 |
简单说:对这类 matmul,最后两维负责矩阵乘法,前面的维度负责批次对应与广播。如果要指定其他轴之间的对应关系或收缩求和,可以使用 einsum、tensordot 等运算(L1 第二篇第四章)。第五章第 1 节再展开批维度的账。
对后端工程师,一个有用的理解是:常见的稠密张量由底层存储和描述访问方式的元数据组成。除了数据类型、设备等信息,关键元数据包括:
shape:每个轴有多长。stride:沿每个轴走一步,要在底层存储中跨过多少个元素。storage_offset:张量的起点相对底层存储偏移多少个元素。
对于按通常行优先方式连续存储、起始偏移为 0 的 (32, 512, 4096) 张量,stride 为 (512*4096, 4096, 1),元素 [b, t, i] 位于底层存储的第 b*512*4096 + t*4096 + i 个位置,按零开始计数。可以把它类比为一个附带形状、步长和偏移信息的 ByteBuffer。
但张量不一定连续,也不一定独占一份存储:转置、切片后的张量可以共享原有数据,只改变访问方式。因此,“几维”主要描述数据的组织方式,不代表内存中真的存在多维格子;而 shape 与 stride 会影响是否能直接改变视图,以及运算效率。
这也解释了第五章的几个行为:view 在布局兼容时只改元数据,不复制数据;reshape 能共享存储时共享,否则可能复制;广播本身通常无需实际复制数据,但运算结果仍可能需要新内存。这些概念也是 Infra 地图 03 系列介绍 TensorImpl 与 storage 的起点。
2. 怎么读形状
读形状的习惯是本篇最想建立的东西。看到任何一个量,先问”它是几维的、每一维是多少、每一维代表什么”。
| 量 | 几维 | 每一维代表什么 |
|---|---|---|
| \(x \in \mathbb{R}^{4096}\) | 1 | 一个 token 的表示:4096 个数 |
| \(X \in \mathbb{R}^{512 \times 4096}\) | 2 | 一句话 512 个 token:512 行,每行是一个 token 的表示 |
| \(W \in \mathbb{R}^{4096 \times 4096}\) | 2 | 一个权重矩阵:把 4096 维映到 4096 维 |
X[32, 512, 4096] |
3 | 32 句话,每句 512 个 token,每个 token 4096 维 |
约定:本系列把”一批 token”写成每行一个 token 的矩阵 \(X \in \mathbb{R}^{T \times d}\)——\(T\) 是一句话里的 token 数(序列长度,seq),\(d\) 是每个 token 的维度(hidden)。这是 PyTorch 布局 [batch, seq, hidden] 去掉第一维的样子:三维的 [B, T, d] 里 \(B\) 是第几句话、\(T\) 是句内第几个 token、\(d\) 是特征——最后一维永远是特征。公式里通常省掉 \(B\),因为矩阵乘法对每句话做的事完全一样(上一节说的”有很多份”),写出来只是多一层括号;代码里 \(B\) 是真实存在的,第五章第 1 节讨论它怎么进成本账。
三、矩阵乘法与形状规则
1. 定义
矩阵 \(A \in \mathbb{R}^{m \times k}\) 与 \(B \in \mathbb{R}^{k \times n}\) 的乘积 \(C = AB \in \mathbb{R}^{m \times n}\),每个元素是
\[C_{ij} = \sum_{l=1}^{k} A_{il}\, B_{lj}\]——\(A\) 的第 \(i\) 行与 \(B\) 的第 \(j\) 列对应位置相乘再加起来。画出来:
注意三个矩形的边长:\(A\) 的宽(\(k\) 列)等于 \(B\) 的高(\(k\) 行)——这就是下一节”内维相同”的几何形象;\(C\) 的高取自 \(A\)、宽取自 \(B\)。
一个能手算的例子(\(m = 2, k = 3, n = 2\)):
\[\begin{pmatrix} 1 & 2 & 3 \\ 4 & 5 & 6 \end{pmatrix} \begin{pmatrix} 1 & 0 \\ 0 & 1 \\ 1 & 1 \end{pmatrix} = \begin{pmatrix} 1\cdot1 + 2\cdot0 + 3\cdot1 & 1\cdot0 + 2\cdot1 + 3\cdot1 \\ 4\cdot1 + 5\cdot0 + 6\cdot1 & 4\cdot0 + 5\cdot1 + 6\cdot1 \end{pmatrix} = \begin{pmatrix} 4 & 5 \\ 10 & 11 \end{pmatrix}\]左边 \([2 \times 3]\),右边 \([3 \times 2]\),结果 \([2 \times 2]\)。
2. 形状规则
从定义直接得到本篇第一条规则:
\([m, k] \times [k, n] \to [m, n]\):左边的列数必须等于右边的行数(内维 \(k\) 相同),结果取左边的行数与右边的列数。
内维不同就不能乘——这是读模型代码时第一个要检查的东西,也是 RuntimeError: mat1 and mat2 shapes cannot be multiplied (512x4096 and 1024x4096) 这类报错的全部含义:512×4096 的东西不能乘 1024×4096 的东西,因为 4096 ≠ 1024。
两个推论:
- 矩阵乘法不交换:\(AB \ne BA\),一般连形状都对不上(\([m,k] \times [k,n]\) 能乘,\([k,n] \times [m,k]\) 除非 \(n = m\) 否则不能)。
- 结合律成立:\((AB)C = A(BC)\)。这在算成本时有用:同样的结果,先乘哪两个成本可能差很多(第四章)。
3. 三种看法
同一个矩阵乘法可以从三个角度看,各在不同场合有用。三种看法说的是同一个计算,只是把注意力放在结果的不同部分上:一次看整个输出向量、一次看每一行、一次看单个元素。下面都用上一节那个能手算的例子
\[A = \begin{pmatrix} 1 & 2 & 3 \\ 4 & 5 & 6 \end{pmatrix},\quad B = \begin{pmatrix} 1 & 0 \\ 0 & 1 \\ 1 & 1 \end{pmatrix},\quad AB = \begin{pmatrix} 4 & 5 \\ 10 & 11 \end{pmatrix}\]来说。
看法一:一个向量过一个矩阵是一次”变换”——看整个输出向量。 只取 \(A\) 的第一行 \(x = (1, 2, 3)\),它乘 \(B\) 得到 \((4, 5)\):一个 3 维向量进去,一个 2 维向量出来,\(B\) 把它从 3 维”变换”到了 2 维。神经网络的每一个”线性层”(nn.Linear)就是这个:输入 in_features 维,输出 out_features 维,权重是一个 [out, in] 的矩阵,一个 token 的向量过一层就是被变换一次。
看法二:一批向量过同一个矩阵是”把每一行分别变换”——看每一行。 \(A\) 有两行 \((1, 2, 3)\) 与 \((4, 5, 6)\),\(AB\) 的两行 \((4, 5)\) 与 \((10, 11)\) 恰好是这两行各自乘 \(B\) 的结果——把第二行改成任何数,第一行的结果 \((4, 5)\) 不变。所以 \(Y = XW\) 里 \(X\) 的每一行是一个 token,\(Y\) 的每一行是对应 token 变换后的结果,各行互不影响。这是为什么一句话里 512 个 token 可以一起算、一批 32 句话也可以一起算:它们只是同一个矩阵乘法的不同行,摞在一起算是为了快,不是因为它们之间有什么关系。
看法三:\(C_{ij}\) 是 \(A\) 的第 \(i\) 行与 \(B\) 的第 \(j\) 列的”相似度”——看单个元素。 \(C_{11} = 4\) 是 \(A\) 的第 1 行 \((1, 2, 3)\) 与 \(B\) 的第 1 列 \((1, 0, 1)\) 对应位置相乘再加:\(1 \cdot 1 + 2 \cdot 0 + 3 \cdot 1 = 4\)。这个”对应位置相乘再加”叫两个向量的内积(第二篇专门讲它)。为什么说它是”相似度”?看一个更小的例子:
\[Q = \begin{pmatrix} 1 & 0 \\ 0 & 1 \end{pmatrix},\quad K = \begin{pmatrix} 1 & 0 \\ 1 & 1 \end{pmatrix},\quad QK^T = \begin{pmatrix} 1 & 1 \\ 0 & 1 \end{pmatrix}\]\(Q\) 的两行是两个 query 向量 \(q_1 = (1, 0)\)、\(q_2 = (0, 1)\),\(K\) 的两行是两个 key 向量 \(k_1 = (1, 0)\)、\(k_2 = (1, 1)\);\(K^T\) 把 key 竖起来当列,于是 \((QK^T)_{ij}\) 就是 \(q_i\) 与 \(k_j\) 的内积。读这张 \(2 \times 2\) 表:\(q_1\) 与 \(k_1\) 方向相同,内积 1;\(q_2\) 与 \(k_1\) 互相垂直,内积 0;\(q_2\) 与 \(k_2\) 有一半重合,内积 1。方向越接近内积越大,垂直为 0,相反为负——这就是”相似度”的含义。attention 正是这样用它:每一行是一个 token 的 query,每一列是一个 token 的 key,\((QK^T)_{ij}\) 越大,”第 \(i\) 个 token 该看第 \(j\) 个 token”就越多。一次矩阵乘法算出了所有 token 对之间的相似度,这是看法三的用处。
4. 转置与单位矩阵
转置 \(A^T\) 把行列互换:\(A \in \mathbb{R}^{m \times n}\) 则 \(A^T \in \mathbb{R}^{n \times m}\),\((A^T)_{ij} = A_{ji}\)。一条常用规则:\((AB)^T = B^T A^T\)——顺序反过来(用形状检查:\(AB\) 是 \([m, n]\),转置是 \([n, m]\),\(B^T A^T\) 是 \([n, k] \times [k, m] = [n, m]\),对上了)。
转置在代码里到处出现,有一个值得记住的具体例子:PyTorch 的 nn.Linear(in, out) 把权重存成 [out, in],前向算的是 \(y = x W^T + b\)。所以文档里写 \(xW^T\)、论文里写 \(Wx\)、本文写 \(XW\),指的是同一件事,只是把向量摆成行还是列、把矩阵存成 [in, out] 还是 [out, in] 的差别。读代码时用形状规则对一下就不会乱。
单位矩阵 \(I\) 是对角线全 1、其余全 0 的方阵,\(IA = AI = A\)——乘它什么都不变。它在第三篇(正交矩阵 \(R^T R = I\))和 L3(残差连接的 Jacobian 是 \(I + J\))里出现。
四、成本规则:这一乘要算多少
1. FLOPs 是什么
FLOP 是 floating-point operation,一次浮点运算——一次加法或一次乘法各算一个。FLOPs(复数)是运算的总次数,是衡量”这一步要算多少”的标准单位。注意区分 FLOPs(次数,衡量工作量)与 FLOPS(每秒次数,衡量硬件速度,如”一张 H100 约 989 TFLOPS”);后面各篇说”FLOPs”都指前者。
2. \(2mnk\) 从哪来
回到定义:\(C = AB\),\(C\) 有 \(m \times n\) 个元素,每个元素 \(C_{ij} = \sum_{l=1}^{k} A_{il} B_{lj}\) 要做 \(k\) 次乘法和 \(k - 1\) 次加法,约 \(2k\) 次运算。所以
\[\text{FLOPs}(A B) \approx m \cdot n \cdot 2k = 2mnk\]这是本篇第二条规则:
\([m, k] \times [k, n]\) 的矩阵乘法约需 \(2mnk\) FLOPs——三个维度的乘积再乘 2。
“约”是因为把 \(k - 1\) 次加法算成了 \(k\) 次;硬件上一次”乘加”(FMA)本来就是一条指令,所以工业界统一按 2 计。
写成代码,\(\sum\) 就是最内层循环,三个维度就是三层循环,\(2mnk\) 就是最内层那行执行的次数乘 2:
for i in range(m): # C 的每一行
for j in range(n): # C 的每一列
for l in range(k): # 内维:A 的第 i 行 · B 的第 j 列
C[i][j] += A[i][l] * B[l][j] # 一次乘 + 一次加 = 2 FLOPs
后面所有关于矩阵乘法的公式——形状规则(A[i][l] 与 B[l][j] 共用同一个 l,所以 \(A\) 的列数必须等于 \(B\) 的行数)、\(C_{ij}\) 是行 \(i\) 与列 \(j\) 的内积、转置为什么出现——都能在这三行里指出来。GPU 上的实现当然不是这样写的(Infra 地图的 05 系列讲它怎么分块、怎么并行),但算的东西一样,FLOPs 也一样。
3. 代一个数字
Llama-3-8B 的隐藏维度 \(d = 4096\)(config.json 里的 hidden_size)。attention 里把输入投影成 query 的权重 \(W_Q\) 是一个 \(4096 \times 4096\) 的矩阵。一个 token 是一个长 4096 的行向量,经过 \(W_Q\) 就是一次 \([1, 4096] \times [4096, 4096]\) 的矩阵乘;prefill 4096 个 token 时,输入摞成 \([4096, 4096]\),同一个 \(W_Q\) 不变。把 \(m, k, n\) 对上去,两条规则各走一遍:
一个 token: x [1 × 4096] × W_Q [4096 × 4096] → q [1 × 4096]
m = 1 k = 4096, n = 4096
FLOPs = 2 m n k = 2 × 1 × 4096 × 4096 ≈ 33.5 M
4096 个 token: X [4096 × 4096] × W_Q [4096 × 4096] → Q [4096 × 4096]
m = 4096 k = 4096, n = 4096
FLOPs = 2 × 4096 × 4096 × 4096 ≈ 137 G ← m 大 4096 倍,FLOPs 也大 4096 倍
两个观察,后面每一层都会用到:
- 一个 token 过一个 \(d \times d\) 的矩阵,FLOPs 是 \(2d^2\),恰好是这个矩阵参数量(\(d^2\))的两倍。 推而广之:一个 token 过整个模型的 FLOPs 约等于参数量的两倍,\(2N\)——前提是每个参数都参与了这个 token 的计算。稠密模型(Llama 这类)几乎如此,只有词嵌入例外:它是按编号查一行,不是矩阵乘(第六章的表里单列了它);MoE 模型(如 DeepSeek-V3)每个 token 只经过被选中的少数专家,这时要用”激活参数量”代替 \(N\),L4 第八篇讲。这是 L4 里”训练 FLOPs \(\approx 6ND\)“(\(D\) 个 token,前向 \(2N\)、反向 \(4N\))的起点。
- token 数 \(m\) 是线性的。 4096 个 token 就是 4096 倍——所以算力账里”处理了多少 token”是最重要的变量。
4. 结合律与成本
\((AB)C\) 与 \(A(BC)\) 结果相同,成本可以差很多。设 \(B \in \mathbb{R}^{4096 \times 16}\)、\(C \in \mathbb{R}^{16 \times 4096}\),\(x \in \mathbb{R}^{1 \times 4096}\),要算 \(x(BC)\):
| 顺序 | 第一步 | FLOPs | 第二步 | FLOPs | 合计 |
|---|---|---|---|---|---|
| 先乘两个矩阵 | \(BC\):\([4096 \times 16] \times [16 \times 4096] \to [4096 \times 4096]\) | \(2 \times 4096 \times 4096 \times 16 \approx 537\) M | \(x(BC)\):\([1 \times 4096] \times [4096 \times 4096]\) | \(2 \times 4096 \times 4096 \approx 33.5\) M | ≈ 570 M |
| 先让向量过瘦矩阵 | \(xB\):\([1 \times 4096] \times [4096 \times 16] \to [1 \times 16]\) | \(2 \times 4096 \times 16 \approx 131\) K | \((xB)C\):\([1 \times 16] \times [16 \times 4096] \to [1 \times 4096]\) | \(2 \times 16 \times 4096 \approx 131\) K | ≈ 262 K(少两千倍) |
这正是第三篇 LoRA 的形状:\(\Delta W = BA\) 两个瘦矩阵。每一步都在线算时(训练、多租户共用底座切换 adapter),先让输入过瘦的那个,不要每次都乘成大矩阵;训完只服务一个 adapter 时则反过来——一次性算出 \(W' = W + BA\) 合并进权重(第三篇 §六),之后推理零额外开销。两种做法不矛盾:前者省的是每步重复算 \(BA\) 的 570 MFLOPs,后者是把它算一次存下来。
五、训练代码里的形状操作
1. 批维度
真实代码里的输入不是 \([T, d]\) 而是 \([B, T, d]\):\(B\) 句话、每句 \(T\) 个 token、每个 \(d\) 维。矩阵乘法只作用于最后两维,前面的维度是”批”——同一件事做 \(B\) 遍。所以 torch.matmul(X, W) 在 \(X\) 是 \([B, T, d]\)、\(W\) 是 \([d, d']\) 时输出 \([B, T, d']\),FLOPs 是 \(B \times 2Td d'\)。读形状时把前面的批维度”括起来”,只看最后两维是否满足 \([\cdot, k] \times [k, \cdot]\)。
用小数字走一遍:\(B = 2\) 句话、每句 \(T = 3\) 个 token、每个 token \(d = 4\) 维,\(X\) 的形状是 \([2, 3, 4]\),一共 24 个数;\(W\) 是 \([4, 5]\)。X @ W 做的是:取第 1 句话的 \([3, 4]\) 乘 \(W\) 得 \([3, 5]\),再取第 2 句话的 \([3, 4]\) 乘同一个 \(W\) 得 \([3, 5]\),叠起来就是 \([2, 3, 5]\)。FLOPs 是 \(2 \times (2 \times 3 \times 4 \times 5) = 240\)——括号里是 \(B \times T \times d \times d'\),前面的 2 是”一乘一加”。\(B\) 和 \(T\) 在这里的地位一样:都只是”这件事做多少遍”,所以成本里它们只是线性因子,真正决定一次矩阵乘法形状的只有 \(d\) 和 \(d'\)。
2. 逐元素运算
两个形状相同的矩阵可以逐元素相加、相乘(记作 \(A + B\)、\(A \odot B\),后者叫 Hadamard 积,代码里就是 A * B)。激活函数(ReLU、GELU、SiLU)、残差连接(\(x + f(x)\))、门控(Llama 的 MLP 里 silu(gate) * up)都是逐元素运算。用 \(2 \times 2\) 的例子看清”逐元素”三个字:
每个位置只和自己对应位置的那个数打交道,没有任何”行乘列再求和”,所以形状不变、也不需要 \(k\) 这个中间维。写成循环就是 for i: for j: C[i][j] = A[i][j] * B[i][j]。它和第三章的内积是亲戚:两个向量的内积 = 先逐元素相乘,再把结果全部加起来——\(\langle a, b \rangle = \sum_i a_i b_i\),(a * b).sum();矩阵乘法的每一格又是一个内积。所以三者的成本是一级一级上去的:逐元素 \(mn\) 次乘法;内积多一步求和;矩阵乘法要做 \(mn\) 个长度为 \(k\) 的内积,\(2mnk\)。逐元素运算比矩阵乘法少一个 \(k\) 的因子——这是为什么算模型成本时通常只数矩阵乘法。
3. 广播
形状不完全相同的两个张量做逐元素运算(相加、相乘、比较……)时,框架按广播(broadcasting)规则把小的那个”复制”到大的形状:从最后一维往前对齐,每一维要么相等、要么其中一个是 1(或不存在)。它解决的是一个天天遇到的需求——给一句话里每个 token 的向量都加上同一个偏置 b,不想真的把 b 抄 512 份再相加。
| 表达式 | 结果 | 发生了什么 |
|---|---|---|
X[512, 4096] + b[4096] |
[512, 4096] |
b 被复制 512 行,加到每个 token 上(Linear 的偏置) |
X[32, 512, 4096] * s[32, 512, 1] |
[32, 512, 4096] |
每个 token 乘自己的一个标量(比如归一化的缩放) |
X[512, 4096] + Y[4096, 512] |
报错 | 4096 ≠ 512 且都不是 1 |
广播不是数学概念,是 NumPy / PyTorch 的约定,但它决定了代码里大量”看起来形状不对却能跑”的表达式在做什么;L1 工具箱第二篇会专门练它。”复制”打了引号:框架并不真的把 b 复制 512 份,而是把它在那一维的 stride 设为 0——同一段内存被读 512 次,写成循环就是 for t: for i: Y[t][i] = X[t][i] + b[i],b[i] 的下标里没有 t。广播是零内存开销的,代价只在计算。顺带一提,Spark / Flink 里也有”广播变量”(把一份小数据发到集群的每台机器上),和这里只是同名:张量广播发生在同一进程的一次运算里,只是按形状规则对齐下标,没有任何数据搬运。
4. reshape 与 view
同样 \(B \times T \times d\) 个数,可以看成 \([B, T, d]\)、\([B \cdot T, d]\) 或 \([B, T, h, d/h]\)(多头 attention 把 \(d\) 拆成 \(h\) 个头)。这些操作不改变任何数值,只改变”怎么读”。PyTorch 的 view 只改 shape 与 stride 两组元数据,一个字节都不搬;reshape 在内存连续时等价于 view,不连续时才复制。转置也只是交换两维的 stride——所以转置之后的张量”不连续”,再 view 会报错,要先 .contiguous()。这是训练代码里最常见的一类报错,知道存法就知道原因。多头 attention 的全部形状变化就是一串 reshape 与转置,用 \(T = 3\)、\(d = 4\)、\(h = 2\) 的玩具尺寸把每一步画出来:
- \(Q = X W_Q\),形状 \([B, T, d]\):每个 token 一行 \(d\) 个数;
- reshape 成 \([B, T, h, d_h]\)(\(d_h = d / h\)):不动任何数,只是把每行的 \(d\) 个数”按顺序切成 \(h\) 段”,第 0 段归头 0、第 1 段归头 1……;
- 转置 第 1、2 维得到 \([B, h, T, d_h]\):把”每个 token 的各个头”重排成”每个头的所有 token”,于是头 \(i\) 拿到一个自己的 \([T, d_h]\) 小矩阵;
- 每个头独立做 \([T, d_h] \times [d_h, T] \to [T, T]\) 的 \(QK^T\),得到一张”第 \(i\) 个 token 看第 \(j\) 个 token 多少分”的表;\(B\) 句话 × \(h\) 个头,共 \(B \times h\) 个这样的小乘法。
拆头的意义在第 4 步:\(h\) 个头各算一张分数表,可以各自关注不同的关系(一个头看语法、一个头看指代……);成本上,\(h\) 个 \([T, d_h]\) 的小乘法加起来与一个 \([T, d]\) 的大乘法相同(\(h \times 2T^2 d_h = 2T^2 d\))。Llama-3-8B:\(h = 32\) 个头,每头 \(d_h = 4096 / 32 = 128\)。
六、一层 Transformer 有哪些矩阵
1. 七个权重矩阵
先说一层在做什么,再数它有哪些矩阵。一个 Transformer 层只有两个部件,每个 token 的向量依次经过它们:attention(注意力)让每个 token 看一眼句子里的其他 token、把有用的信息加到自己身上——它是唯一一处 token 之间发生交流的地方;MLP(多层感知机,两三个矩阵乘法夹一个非线性函数)对每个 token 各自做一次变换,token 之间互不影响。两个部件的输出都加回输入(残差连接,\(x + f(x)\),L3 会讲它为什么必要)。把这条路径上的每个矩阵标出来:
%%{init: {"flowchart": {"wrappingWidth": 180}}}%%
%% 图:一个 Transformer 层的七个权重矩阵:attention 部件四个、MLP 部件三个,以 Llama-3-8B 的形状标注
flowchart TB
X["x [1, 4096]"]
subgraph ATTN["attention 部件:token 之间交流"]
direction TB
P["× W_Q → [1, 4096]<br/>× W_K → [1, 1024]<br/>× W_V → [1, 1024]"]
A["attention:QKᵀ、softmax(·)V<br/>用自己的 Q 看其他 token 的 K、V<br/>无参数,FLOPs 随 T 增长"]
O["× W_O → [1, 4096]"]
P --> A --> O
end
subgraph MLP["MLP 部件:每个 token 各自变换"]
direction TB
G["× W_gate → [1, 14336]<br/>× W_up → [1, 14336]"]
S["silu(gate) ⊙ up 逐元素"]
D["× W_down → [1, 4096]"]
G --> S --> D
end
X --> P
O --> ADD1(("+")) --> H["h [1, 4096]"]
X -. 残差 .-> ADD1
H --> G
D --> ADD2(("+")) --> Y["下一层的 x [1, 4096]"]
H -. 残差 .-> ADD2
先用一组小数字把这条路径走一遍——取 \(d = 4\)、MLP 中间维 \(d_{ff} = 8\)、只有 1 个头,一个 token 的向量 \(x = (1, 0, 2, 1)\):
| 步骤 | 运算 | 输入形状 | 权重形状 | 输出形状 | 这一步在做什么 |
|---|---|---|---|---|---|
| ① | \(q = x W_Q\),\(k = x W_K\),\(v = x W_V\) | [1, 4] |
三个 [4, 4] |
三个 [1, 4] |
把同一个 token 投影成”我要找什么 / 我是什么 / 我能提供什么”三个向量 |
| ② | 与句中 \(T\) 个 token 的 \(k\) 算内积、softmax、加权求和它们的 \(v\) | [1, 4] 与 [T, 4] |
无 | [1, 4] |
唯一一处看别的 token;没有参数 |
| ③ | \(\times W_O\),加回 \(x\) | [1, 4] |
[4, 4] |
[1, 4] |
把”看到的东西”变回自己的坐标系,残差相加得 \(h\) |
| ④ | \(g = h W_{\text{gate}}\),\(u = h W_{\text{up}}\) | [1, 4] |
两个 [4, 8] |
两个 [1, 8] |
放大到 2 倍宽(Llama 是 3.5 倍) |
| ⑤ | \(\text{silu}(g) \odot u\) | 两个 [1, 8] |
无 | [1, 8] |
逐元素:门控,决定每个通道留多少 |
| ⑥ | \(\times W_{\text{down}}\),加回 \(h\) | [1, 8] |
[8, 4] |
[1, 4] |
压回原宽度,残差相加,交给下一层 |
参数只在 ①③④⑥ 四步里出现:\(3 \times 16 + 16 + 2 \times 32 + 32 = 160\) 个数,FLOPs 恰好是它的两倍 320(再加 ② 里随 \(T\) 变的那部分)。把 4 换成 4096、8 换成 14336、1 个头换成 32 个 query 头 + 8 个 KV 头,就是下面 Llama-3-8B 的表。七个带 \(W\) 的框就是这一层的全部参数;其余框(attention 内部、逐元素运算、加法)没有参数。用形状规则逐个读 Llama-3-8B 的这七个矩阵(\(d = 4096\);MLP 中间维度 \(d_{ff} = 14336\),先把向量放大到 3.5 倍再压回来;attention 用 GQA——32 个 query 头,但 key / value 只有 8 个头、由每 4 个 query 头共用,所以 \(W_K, W_V\) 的输出维是 \(8 \times 128 = 1024\) 而不是 4096。为什么要共用:推理时每个 token 的 K、V 要存下来反复用,头少四倍存的就少四倍,L4 第七篇算这笔账):
| 矩阵 | 形状 [in, out] |
参数量 | 一个 token 的 FLOPs(2 × 参数量) |
|---|---|---|---|
| \(W_Q\) | 4096 × 4096 | 16.8 M | 33.5 M |
| \(W_K\) | 4096 × 1024 | 4.2 M | 8.4 M |
| \(W_V\) | 4096 × 1024 | 4.2 M | 8.4 M |
| \(W_O\) | 4096 × 4096 | 16.8 M | 33.5 M |
| \(W_{\text{gate}}\) | 4096 × 14336 | 58.7 M | 117.4 M |
| \(W_{\text{up}}\) | 4096 × 14336 | 58.7 M | 117.4 M |
| \(W_{\text{down}}\) | 14336 × 4096 | 58.7 M | 117.4 M |
| 一层合计 | 218 M | 436 M | |
| × 32 层 | 6.98 B | 14.0 G | |
| + 词嵌入 | 128256 × 4096 | 0.53 B | (查表,不算矩阵乘) |
| + 输出层 | 4096 × 128256 | 0.53 B | 1.05 G |
| 总计 | ≈ 8.03 B | 一个 token 前向 ≈ 15 G FLOPs ≈ 2N | |
| 表外:attention 的 \(QK^T\) 与 \(\text{softmax}(\cdot)V\) | 每层 \([T, d_h] \times [d_h, T]\) 与 \([T, T] \times [T, d_h]\),32 个头 | 0(无参数) | 每层 \(4Td\):\(T = 512\) 时 8.4 M(≈ 一层权重的 2%),\(T = 8192\) 时 134 M(≈ 31%),\(T = 128\)K 时 2.1 G(≈ 5 倍) |
三件事从这张表直接读出来:
- 参数量的大头在 MLP(三个 \(4096 \times 14336\) 占一层的 81%),不在 attention;
- 一个 token 过整个模型约 \(2N\) FLOPs——每个参数用一次乘一次加,与第四章的观察一致;
- 上图 attention 框里的 \(QK^T\) 与 \(\text{softmax}(\cdot) V\) 那两个矩阵乘不在七个矩阵里——它们乘的是数据和数据(Q 乘 K、分数乘 V),没有参数,所以表的主体没有它们,单列在表外一行:一个 token 在每层要与全部 \(T\) 个 token 各算一次内积再各加权一次,FLOPs 是每层 \(4Td\)、与参数量无关,随 \(T\) 线性增长(整句话一起算就是 \(T^2\))。\(T = 512\) 时只占一层权重 FLOPs 的 2%,可以忽略;\(T = 128\)K 时是权重的 5 倍,成了大头——这就是长上下文贵在哪里。
这七个矩阵只是一层 block 里带参数的部分。整台 decoder-only Transformer 还有几样不在这七个里的部件——两张 embedding 表、LayerNorm 的缩放与偏置、lm_head——它们和这七个矩阵怎么对应,以及同一套矩阵在 GPT-2 small 里叫什么、形状多大,见《Transformer 与 LLM》第一篇《从一句话到下一个 token》第一章「六种部件」后面的对应表。
2. 为 L4 铺路
到这里,形状规则与成本规则已经足够读懂 L4《Transformer 与 LLM》第十篇对整个模型的算账——那一篇的每一行都是本篇两条规则的重复应用:数出每个矩阵的形状、乘 2、乘 token 数、加上 attention 的平方项、乘 3 变成训练。这里只要求会两条规则;到了 L4,它们会被用来回答”训一个 8B 模型要多少 GPU 小时”。
七、本文小结
- 标量 / 向量 / 矩阵 / 张量是 0 / 1 / 2 / 多维的数表;一个 token 在模型里是一个向量(Llama-3-8B:4096 维),模型的参数是一堆矩阵;张量的矩阵乘法只作用于最后两维。
- 形状规则:\([m, k] \times [k, n] \to [m, n]\),内维必须相同。矩阵乘法不交换、满足结合律;\((AB)^T = B^T A^T\)。
- 成本规则:\(2mnk\) FLOPs。一个 token 过一个 \(d \times d\) 矩阵是 \(2d^2\),过整个模型约 \(2N\)(\(N\) 是参数量);token 数是线性因子。
- 结合律让同一个结果的成本差几千倍——LoRA 在线用时先过瘦矩阵;只服务一个 adapter 时一次性合并 \(W + BA\)。
- 批维度”括起来”只看最后两维;逐元素运算与广播的成本比矩阵乘法少一个 \(k\) 因子;reshape 不改数值只改读法。
- 一层 Llama-3-8B 有七个权重矩阵、约 218M 参数,81% 在 MLP;全模型 8.03B,一个 token 前向约 15 GFLOPs。
八、自测
-
\(A \in \mathbb{R}^{3 \times 5}\),\(B \in \mathbb{R}^{5 \times 2}\):\(AB\) 是什么形状?\(BA\) 能不能算?\((AB)^T\) 等于什么?
答案
\([3, 2]\);不能,\(2 \ne 3\);\(B^T A^T\),形状 \([2, 3]\)。
-
nn.Linear(4096, 14336)的权重存成什么形状?一个 token 过它是多少 FLOPs?一句 2048 个 token 呢?答案
[14336, 4096];\(2 \times 4096 \times 14336 \approx 117\) M;乘 2048 约 240 G。 -
\(X \in \mathbb{R}^{512 \times 4096}\) 与 \(b \in \mathbb{R}^{4096}\) 相加,结果是什么形状?与 \(c \in \mathbb{R}^{512}\) 相加呢?
答案
\([512, 4096]\)(\(b\) 广播到每一行);报错,4096 ≠ 512。
-
把 \([B, T, 4096]\) 的张量拆成 32 个头,每个头多少维?\(QK^T\) 对一个头、一句 \(T\) 个 token 是多少 FLOPs(用 \(T\) 表示)?
答案
128 维;\([T, 128] \times [128, T]\) 是 \(2 \times 128 \times T^2 = 256\,T^2\)。
-
一个 70B 参数的模型,一个 token 前向大约多少 FLOPs?处理 1 万亿(\(10^{12}\))个 token 呢?
答案
\(2N = 1.4 \times 10^{11}\);乘 \(10^{12}\) 得 \(1.4 \times 10^{23}\)——这就是训练 FLOPs 里”前向”的那一份,训练还要乘 3。
下一篇讲向量之间怎么比较:内积、范数与余弦相似度——attention 的 score、检索的相似度、正则化项与量化误差用的是同一套语言。
-
能,两条规则就够:形状规则 \([m, k] \times [k, n] \to [m, n]\)(内维必须相同;批维度括起来只看最后两维)与成本规则 \(2mnk\) FLOPs(每个输出元素一次乘加)。逐元素运算比矩阵乘少一个 \(k\) 因子,所以算账时只数矩阵乘。详见第三章、第四章。 ↩
-
能。Llama-3-8B:一个 token 过一个 \(4096 \times 4096\) 的 \(W_Q\) 是 \(2 \times 4096^2 \approx 33.5\) M FLOPs;一层七个矩阵 218 M 参数、436 M FLOPs;全模型 8.03 B 参数,一个 token 前向约 15 G FLOPs \(\approx 2N\)。详见第六章。 ↩
系列 《算法工程师的数学:读公式不卡壳的最小集》 第 1 / 9 篇
下一篇:内积、范数与余弦相似度
本文由 arganzheng 创作,采用 CC BY 4.0 许可协议。在保留原文作者、署名以及完整原文链接(https://arganzheng.life/vectors-matrices-shapes-and-flops.html)的前提下,欢迎各种形式的转载、翻译或商业引用。
COMMENTS
评论存放在 GitHub Discussions, 用 GitHub 账号登录即可发表,支持 Markdown。 想针对正文某句话说?选中那段文字,点浮出的「评论」即可划线评论;觉得哪里写错了,发表时勾上「同时提交 Issue」。 有人回复你时 GitHub 会按你的通知设置发邮件,不用守在这里。