文章
Tranforms讲义
Loading collection...
带3d模型的参考学习:https://bbycroft.net/llm
19 世纪中后期麦克斯韦和玻尔兹曼共同提出玻尔兹曼分布,引入概率统计的方法来解释热力学现象,奠定了统计力学基础。1985 年 Hinton 引入统计力学的玻尔兹曼分布来构建神经网络概率,发明了玻尔兹曼机。1989 年,John Bridle 正式定义 Softmax 函数并将其应用于深度学习解决概率分布。Transformer 架构将 Softmax 作为注意力机制核心。
一、 基石:从独热编码到几何空间 (Token, Embedding, Geometry)
1. 数学本质
文档指出,计算机无法理解文本。初始输入是独热(One-hot)向量 ei∈RV(V为词表大小)。但独热向量正交,无法度量词语相似性。
数学推理:我们引入嵌入矩阵 E∈RV×d。给定独热向量 ei,其对应的词向量为:
Embedding=ETei∈Rd
这本质上是一个查表操作。通过这个映射,每个词变成了 d 维空间中的一个点。衡量词义相似度转化为向量点积(余弦相似度):
similarity(v,w)=v⋅w=∥v∥∥w∥cosθ
至此,语言变成了几何。
2. 句子矩阵
输入句子有 n 个 token,嵌入维度 d,则输入表示为 X∈Rn×d(行是token,列是特征)。这是我们进入Attention的起点。
二、 核心灵魂:自注意力与缩放点积 (Scaled Dot-Product Attention)
这是全文档最重要的数学推导部分,必须完全吃透。
1. 为什么需要 Q、K、V?
朴素想法:用 XXT 计算相似度,再用 AX 加权求和。
问题:同一个 X 同时扮演“被查询者(Query)”、“被匹配者(Key)”和“被提取者(Value)”三个角色,角色冲突。
解决方案:学习三个投影矩阵 WQ,WK,WV∈Rdmodel×dk。
Q=XWQ,K=XWK,V=XWV
此时,qiTkj 表示“Token i 的查询”与“Token j 的键”的匹配度。
2. 关键数学推理:为什么要除以 dk?(极其重要!)
文档给出了严谨的概率论推导,这是Transformer的点睛之笔。我们梳理一下:
- 问题:计算 sij=qiTkj=∑m=1dkqmkm。
- 假设:训练初期,qm 和 km 相互独立,均值为0,方差为1。
- 求方差:
单个乘积的方差:Var(qmkm)=E[qm2]E[km2]−(E[qm]E[km])2=1×1−0=1。
由于各项独立,点积和的方差为:
Var(qTk)=m=1∑dkVar(qmkm)=dk
这意味着,点积的标准差为 dk。
- 后果:当 dk=64 时,得分通常在 ±8 左右;当 dk=4096 时,得分可达 ±64。Softmax 对极大的输入极其敏感,会导致梯度进入饱和区(即 e64 远大于其他项,输出接近 One-hot,梯度接近 0)。
- 修正:除以标准差 dk,使方差重新变为 1:
Var(dkqTk)=dk1×dk=1
这样,得分被稳定在合理的数值范围,Softmax 的梯度保持健康。
3. 完整公式(注意文档第8页的括号笔误)
文档写的是 Softmax(dkQKTV),这在数学上会误解为 V 也被乘在指数里。正确的数学写法是:
Attention(Q,K,V)=Softmax(dkQKT)V
先算注意力权重矩阵 A∈Rn×n,再用 A 对 V 进行加权行混合。
三、 多角度洞察:多头与 MLP 的分工
1. 多头注意力 (Multi-Head)
单头注意力只能给出一种分布。文档举了 “she” 的例子,它需要同时关注主语、动词、从句。
数学上,设 h=8 个头,每个头有独立的 WQ(r),WK(r),WV(r)。如果 dmodel=512,则每个头 dk=512/8=64。
计算完各头后,拼接(Concat)得到 H∈Rn×512,最后经过输出投影 WO∈R512×512 混合各头信息:
MultiHead(X)=Concat(head1,…,head8)WO
2. MLP 层的非线性数学本质
文档点明了一个极其重要的线性代数事实:
- Attention:计算 Y=AX。这是行混合(Token之间混合),但每列(特征)独立变化,本质是线性操作(加权平均)。
- MLP:对每个 Token 独立应用 f(x)=W2σ(W1x+b1)+b2。这是列混合(特征之间混合),且引入了非线性激活函数(如 GELU)。
没有 MLP 的非线性,多层线性层会坍缩成一层(因为 W2(W1x)=(W2W1)x)。正是激活函数的存在,赋予了模型复杂的特征变换能力。
四、 训练动力学:残差连接与反向传播的数学流
1. 残差连接(Residual)的梯度救赎
深层网络面临梯度消失。文档给出了关键的雅可比矩阵推导:
不加残差:∂X∂Y=∏∂Fi−1∂Fi(连乘,极易趋于0)。
加残差(设 Y=X+F(X)):
∂X∂Y=I+∂X∂F
即使 ∂X∂F 很小,恒等映射 I 保证了梯度至少为1,能够无损地传回浅层。这是训练数十层甚至上百层网络的数学基石。
2. 反向传播的核心矩阵求导(必须掌握)
文档推导了神经网络中最核心的线性层梯度公式。设 Y=XW,损失为 L。
- 对权重 W 的梯度(用于更新参数):
∂W∂L=XT∂Y∂L
维度验证:XT∈Rd×n,∂Y∂L∈Rn×dk,相乘得 Rd×dk,与 W 同维。
- 对输入 X 的梯度(用于传回上一层):
∂X∂L=∂Y∂LWT
因此,在反向传播通过 Q,K,V 时,梯度会分别回传至 WQ,WK,WV,并累加回传给 X:
∂X∂L=∂Q∂LWQT+∂K∂LWKT+∂V∂LWVT
3. 损失函数(Cross-Entropy)的梯度推导
文档完美地推导了 Logits(z)的梯度。设真实标签为 y(One-hot),模型输出概率为 p=Softmax(z)。
损失函数:L=−∑yilogpi。因为 y 是独热的,假设正确类别为 k,则 L=−logpk。
将 pk=∑jezjezk 代入:
L=−zk+log(j∑ezj)
对 z 求偏导(分 i=k 和 i=k 两种情况),结果极其优雅:
∂z∂L=p−y
这个结果告诉我们在反向传播时,预测概率减去真实标签就是传递给最后一层的梯度信号,简单、直接、高效。
五、 总结:宏观视角下的 Transformer
通过这份文档的数学解读,我们可以清晰地给 Transformer 画个像:
- Attention(行操作):解决“向谁看”的问题。通过 QKT/dk 计算关系矩阵,用 Softmax 归一化,再对 V 进行线性加权平均。
- MLP(列操作):解决“怎么想”的问题。通过升维(d→4d)和非线性激活,挖掘特征间的复杂逻辑。
- 残差连接:解决“怎么深”的问题。利用 I+∂X∂F 保证梯度畅通,让几百层网络得以训练。
- 训练目标:基于信息论,最小化交叉熵等价于最大化正确 token 的似然概率,而反向传播通过 ∂z∂L=p−y 将误差转化为具体的参数更新方向。
第一个数学困惑:点积(Dot Product)到底凭什么衡量“注意力”?
你的直觉:两个数相乘就能代表“关注”?太抽象了。
手算破局(几何视角):
想象你在黑板上画两个箭头(向量)。点积 a⋅b a⋅b 其实在计算:“箭头的长度” × “它们方向一致的程度”。
- 假设 a=(1,0) a=(1,0)(指向正右方),b=(1,0) b=(1,0)(也指向正右方)。点积 =1×1+0×0=1=1×1+0×0=1。方向完全一致,非常关注!
- 假设 a=(1,0) a=(1,0)(指右),b=(0,1) b=(0,1)(指上)。点积 =1×0+0×1=0=1×0+0×1=0。方向垂直,几乎无关,不关注。
- 假设 a=(1,0) a=(1,0)(指右),b=(−1,0) b=(−1,0)(指左)。点积 =1×(−1)=−1=1×(−1)=−1。方向相反,反向排斥。
所以,点积就是一个粗糙但有效的“关系测量仪”。在 Transformer 里,qiTkj qiTkj 其实是在问:“我的查询方向”和“你的键方向”对齐了多少?”数值越大,我分给你的注意力权重就越大。
第二个数学困惑:为什么要除以 dkdk?(最经典的晕点)
你的直觉:为什么要除以一个开根号?这个数从哪来的,像是硬凑的。
手算破局(方差游戏):
先假设只有 dk=2 dk=2 维(方便手算)。每个维度里都是随机的小数,比如 q=(0.5,0.5) q=(0.5,0.5),k=(0.5,0.5) k=(0.5,0.5)。
点积结果:0.5×0.5+0.5×0.5=0.25+0.25=0.5。这个数不大,Softmax 算起来很顺。
现在把场景换成大模型:dk=4096 dk=4096 维(很大),每个维度仍是 0.5 左右的小数。
点积结果:相当于把 4096 个 0.25 加起来!0.25×4096=1024。
问题来了:
把 1024 喂给 Softmax(算 e1024 e1024)时,数值会大到溢出;Softmax 的输出几乎只剩 0 或 1,梯度直接消失,模型就学不动了。
怎么破?
既然问题来自维度 dk dk 太大导致点积规模膨胀,那就在点积后除以 √dk*dk*,把数值拉回到可训练的范围。
10244096=102464=16
4096
1024=641024=16
你会看到,巨大的 1024 被拉回到 16。16 放进 Softmax(算 e16 e16)依然偏大,但至少计算可控、梯度仍然存在。
记住这句人话:除以 dkdk 不是为了装深沉,而是为了防止“4096 个小数加起来变成天文数字”,把 Softmax 撑死。
第三个数学困惑:反向传播的梯度 p−yp−y 是什么意思?
你的直觉:导数公式太抽象,看不出它到底在“干什么”。
手算破局(猜词游戏):
假设词表只有 3 个词:
[猫, 狗, 鸟]。正确答案是
猫(所以真实标签 y=[1,0,0] y=[1,0,0])。模型刚开始乱猜,算出来的概率 p=[0.2,0.5,0.3] p=[0.2,0.5,0.3](模型更倾向于“狗”)。
套用终极梯度公式(文档第 21 页):∂L∂z=p−y∂z∂L=p−*y*。
我们来算三个位置的梯度:
- 对
猫(正确答案):p−y=0.2−1=−0.8 p−y=0.2−1=−0.8。 - 负数意味着什么?意味着要提高
猫的分数(梯度下降沿负梯度方向走,负负得正)。模型收到的信号是:“给猫加分!”
- 对
狗(错误答案):p−y=0.5−0=+0.5 p−y=0.5−0=+0.5。 - 正数意味着什么?意味着要降低
狗的分数。模型收到的信号是:“别瞎猜狗了,减分!”
- 对
鸟(错误答案):p−y=0.3−0=+0.3 p−y=0.3−0=+0.3。 - 同样要降低
鸟的分数,但力度比“狗”小一些。
你看,这个公式 p−y p−y 多直白。
它就是一个“纠错信号”:
- 猜高了(p=0.5 p=0.5 但正确答案不是它)→→ 梯度为正 →→ 往下调。
- 猜低了(p=0.2 p=0.2 但正确答案是它)→→ 梯度为负 →→ 往上调。
不需要死记硬背:它就是个“高了减,低了加”的自动调节器。
第四个数学困惑:残差连接 X+F(X)X+F(X) 为什么能治梯度消失?
你的直觉:不就是多加了个 X X 吗?凭什么几百层网络就能跑起来?
手算破局(传话游戏):
假设没有残差(只有 Y=F(X) Y=F(X)),梯度像传话游戏:每传一层,信息就模糊一点(相当于乘以 0.1)。
- 第 1 层梯度:0.1
- 第 10 层传回第 1 层时:0.1^10=0.0000000001(几乎消失,底层收不到信号,学不动)。
有了残差(Y=X+F(X) Y=X+F(X)),梯度就变成了“直达电梯 + 模糊楼梯”:1+0.1。
- 第 1 层梯度:1.1
- 第 10 层传回第 1 层时:1.1^10≈2.59(不仅没消失,甚至还有点放大)。
记住这句人话:残差连接就像在 50 层大楼里装了一部直达电梯。哪怕其他楼层(F(X) F(X))的电梯坏了(梯度接近 0),你仍然能坐这部直达电梯(X X)把信号从顶楼无损送回一楼。
总复习:用“人话”把数学串一遍
- 点积 (QKTQKT):就是用尺子量两个向量的方向;越对齐,越“关注”。
- 除以 dk*dk*:维度太高时,点积会膨胀成“大数炸弹”;除以 64 让它回到可训练的范围。
- 梯度 p−yp−*y*:就是“高了减,低了加”的傻瓜纠错器。
- 残差 (+X+X):训练大楼里的直达电梯,保证底层永远能收到顶层的信号。