从反向传播看 QKV:注意力只是计算图里的几个节点

从反向传播看 QKV:注意力只是计算图里的几个节点

一句话版 :继续沿用上一篇《反向传播通关指南》的框架------"每层只管局部导数,剩下的交给链式法则连乘"。Q、K、V 在结构上就是三个 Linear (也就是上一篇那块最熟悉的积木 y=wx+by=wx+by=wx+b,只是放大成了矩阵);attention 本身只是计算图里多出来的几个节点(矩阵乘 → 缩放 → softmax → 矩阵乘),每个节点照旧只有一个局部导数。真正神奇的地方在于:Q/K/V 的"语义角色"不是写死的,而是反向传播把梯度推了几万步之后"涌现"出来的

前置阅读:blog_backprop.md(反向传播 + autograd + minidnn)。

配套:transformer_architecture.md(Transformer 整体架构、KV Cache、硬件视角)。

适合读者 :已经看懂反向传播("三块积木 + 链式法则 + 计算图"),想从"训练/梯度"的角度而不是只看公式的角度理解 QKV 的人。

你将收获:一套"QKV 不神秘"的心智模型、一组算到底的具体数字(前向 + 反向)、以及一个能解释"为什么 Q/K/V 会变得有用"的关键洞察。


0. 前置:只靠反向传播文章的三个结论

blog_backprop.md 里带走三句话,后面全靠它们:

  1. 导数 = 谁影响谁:反向传播只关心相邻两层,每段当成单变量求导,然后连乘。
  2. 三块积木 :线性层 y=wx+by=wx+by=wx+b、激活层 a=tanh⁡(z)a=\tanh(z)a=tanh(z)、损失 L=(out−y)2L=(out-y)^2L=(out−y)2------任何网络都能用它们拼出来。
  3. autograd = 链式法则 × 计算图 :前向时每个运算记成一个节点(值 + 输入 + 反向函数),反向时按拓扑序倒着遍历、逐个调用反向函数。loss.backward() 和 175B 参数的大模型,机制完全一样。

本文要做的只有一件事:把第 2 条"积木"和第 3 条"计算图",原封不动地搬到 attention 上


1. 结构上:QKV 就是三个 Linear 层

先复习上一篇的第一块积木:

y=wx+b⟹∂y∂w=x,∂y∂b=1,∂y∂x=wy = wx + b \qquad\Longrightarrow\qquad \frac{\partial y}{\partial w}=x,\quad \frac{\partial y}{\partial b}=1,\quad \frac{\partial y}{\partial x}=wy=wx+b⟹∂w∂y=x,∂b∂y=1,∂x∂y=w

而 attention 的三个投影,数学形式一模一样

Q=XWQ,K=XWK,V=XWVQ = XW_Q,\qquad K = XW_K,\qquad V = XW_VQ=XWQ,K=XWK,V=XWV

其中 X∈RL×dmodelX \in \mathbb{R}^{L\times d_{model}}X∈RL×dmodel(LLL 个 token,每行一个 token 的向量),WQ/WK/WV∈Rdmodel×dheadW_Q/W_K/W_V \in \mathbb{R}^{d_{model}\times d_{head}}WQ/WK/WV∈Rdmodel×dhead 是三个独立的参数矩阵。

所以局部导数可以直接照抄上一篇的速查表 ,只是把标量乘法换成矩阵乘法(GEMM)------这正是 blog_backprop.md 第 4 节预告的"标量 → 矩阵,导数同一套":

上一篇(标量/向量) Q/K/V 版(矩阵)
dW=X⊤dZdW = X^\top dZdW=X⊤dZ dWQ=X⊤dQ,dWK=X⊤dK,dWV=X⊤dVdW_Q = X^\top dQ,\quad dW_K = X^\top dK,\quad dW_V = X^\top dVdWQ=X⊤dQ,dWK=X⊤dK,dWV=X⊤dV
dX=dZ W⊤dX = dZ\,W^\topdX=dZW⊤ dXQ=dQ WQ⊤dX_Q = dQ\,W_Q^\topdXQ=dQWQ⊤(Q 这一路)
db=∑ndZdb = \sum_n dZdb=∑ndZ dbQ=∑ndQdb_Q = \sum_n dQdbQ=∑ndQ(K、V 同理)

结论 1(本节核心)W_Q / W_K / W_V 和你手写的 W1 / W2 没有任何本质区别------它们都是"要更新的参数",更新规则都是

W←W−η ∂L∂WW \leftarrow W - \eta\,\frac{\partial L}{\partial W}W←W−η∂W∂L

Q/K/V 的名字(query/key/value)不是结构决定的,是用途决定的 。结构上,它们只是三个并排的 Linear ,从同一份输入 XXX 出发、各自线性投影到 dheadd_{head}dhead 维。
#mermaid-svg-mZNu6PvRFSfpOsnn{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-mZNu6PvRFSfpOsnn .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-mZNu6PvRFSfpOsnn .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-mZNu6PvRFSfpOsnn .error-icon{fill:#552222;}#mermaid-svg-mZNu6PvRFSfpOsnn .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-mZNu6PvRFSfpOsnn .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-mZNu6PvRFSfpOsnn .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-mZNu6PvRFSfpOsnn .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-mZNu6PvRFSfpOsnn .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-mZNu6PvRFSfpOsnn .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-mZNu6PvRFSfpOsnn .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-mZNu6PvRFSfpOsnn .marker{fill:#333333;stroke:#333333;}#mermaid-svg-mZNu6PvRFSfpOsnn .marker.cross{stroke:#333333;}#mermaid-svg-mZNu6PvRFSfpOsnn svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-mZNu6PvRFSfpOsnn p{margin:0;}#mermaid-svg-mZNu6PvRFSfpOsnn .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-mZNu6PvRFSfpOsnn .cluster-label text{fill:#333;}#mermaid-svg-mZNu6PvRFSfpOsnn .cluster-label span{color:#333;}#mermaid-svg-mZNu6PvRFSfpOsnn .cluster-label span p{background-color:transparent;}#mermaid-svg-mZNu6PvRFSfpOsnn .label text,#mermaid-svg-mZNu6PvRFSfpOsnn span{fill:#333;color:#333;}#mermaid-svg-mZNu6PvRFSfpOsnn .node rect,#mermaid-svg-mZNu6PvRFSfpOsnn .node circle,#mermaid-svg-mZNu6PvRFSfpOsnn .node ellipse,#mermaid-svg-mZNu6PvRFSfpOsnn .node polygon,#mermaid-svg-mZNu6PvRFSfpOsnn .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-mZNu6PvRFSfpOsnn .rough-node .label text,#mermaid-svg-mZNu6PvRFSfpOsnn .node .label text,#mermaid-svg-mZNu6PvRFSfpOsnn .image-shape .label,#mermaid-svg-mZNu6PvRFSfpOsnn .icon-shape .label{text-anchor:middle;}#mermaid-svg-mZNu6PvRFSfpOsnn .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-mZNu6PvRFSfpOsnn .rough-node .label,#mermaid-svg-mZNu6PvRFSfpOsnn .node .label,#mermaid-svg-mZNu6PvRFSfpOsnn .image-shape .label,#mermaid-svg-mZNu6PvRFSfpOsnn .icon-shape .label{text-align:center;}#mermaid-svg-mZNu6PvRFSfpOsnn .node.clickable{cursor:pointer;}#mermaid-svg-mZNu6PvRFSfpOsnn .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-mZNu6PvRFSfpOsnn .arrowheadPath{fill:#333333;}#mermaid-svg-mZNu6PvRFSfpOsnn .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-mZNu6PvRFSfpOsnn .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-mZNu6PvRFSfpOsnn .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-mZNu6PvRFSfpOsnn .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-mZNu6PvRFSfpOsnn .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-mZNu6PvRFSfpOsnn .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-mZNu6PvRFSfpOsnn .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-mZNu6PvRFSfpOsnn .cluster text{fill:#333;}#mermaid-svg-mZNu6PvRFSfpOsnn .cluster span{color:#333;}#mermaid-svg-mZNu6PvRFSfpOsnn div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-mZNu6PvRFSfpOsnn .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-mZNu6PvRFSfpOsnn rect.text{fill:none;stroke-width:0;}#mermaid-svg-mZNu6PvRFSfpOsnn .icon-shape,#mermaid-svg-mZNu6PvRFSfpOsnn .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-mZNu6PvRFSfpOsnn .icon-shape p,#mermaid-svg-mZNu6PvRFSfpOsnn .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-mZNu6PvRFSfpOsnn .icon-shape .label rect,#mermaid-svg-mZNu6PvRFSfpOsnn .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-mZNu6PvRFSfpOsnn .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-mZNu6PvRFSfpOsnn .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-mZNu6PvRFSfpOsnn :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 注意力引擎
输入 X (L×d)
Q = XW_Q

(Linear)
K = XW_K

(Linear)
V = XW_V

(Linear)
QKᵀ / √d
softmax
P·V
输出投影

上图中,带参数(要训练)的只有三个 Linear 和最后的输出投影 ;中间的 QKᵀ、√d、softmax、P·V 都是无参数节点------它们只是负责"怎么组合信息",本身没有权重可更新。


2. 一个具体数字例子:前向算到底

为了把"只是几个节点"落到实处,照抄上一篇"拿真实数字从头算"的风格。设 2 个 token、d=2d=2d=2:

X=1234,WQ=0.5−110.5,WK=1001,WV=0.5001X=\begin{bmatrix}1&2\\3&4\end{bmatrix},\qquad W_Q=\begin{bmatrix}0.5&-1\\1&0.5\end{bmatrix},\qquad W_K=\begin{bmatrix}1&0\\0&1\end{bmatrix},\qquad W_V=\begin{bmatrix}0.5&0\\0&1\end{bmatrix}X=1324,WQ=0.51−10.5,WK=1001,WV=0.5001

① 三次投影(就是三个 Linear 前向)

Q=XWQ=2.505.5−1,K=XWK=1234,V=XWV=0.521.54Q = XW_Q = \begin{bmatrix}2.5&0\\5.5&-1\end{bmatrix},\qquad K = XW_K = \begin{bmatrix}1&2\\3&4\end{bmatrix},\qquad V = XW_V = \begin{bmatrix}0.5&2\\1.5&4\end{bmatrix}Q=XWQ=2.55.50−1,K=XWK=1324,V=XWV=0.51.524

② 匹配打分:S=QK⊤/dS = QK^\top / \sqrt{d}S=QK⊤/d

这里的上标 ⊤^\top⊤ 是转置 (不是逆):把 KKK 沿主对角线翻转、行变列,好让 QQQ 的第 iii 行能和 KKK 的第 jjj 行对齐做内积。

K=1234  ⇒  K⊤=1324K = \begin{bmatrix}1&2\\3&4\end{bmatrix} \;\Rightarrow\; K^\top = \begin{bmatrix}1&3\\2&4\end{bmatrix}K=1324⇒K⊤=1234

矩阵乘法是"行 × 列,对应相乘再相加",逐个元素算:

QK⊤=2.505.5−11324=2.5×1+0×22.5×3+0×45.5×1+(−1)×25.5×3+(−1)×4=2.57.53.512.5QK^\top = \begin{bmatrix}2.5&0\\5.5&-1\end{bmatrix}\begin{bmatrix}1&3\\2&4\end{bmatrix} = \begin{bmatrix} 2.5\times1+0\times2 & 2.5\times3+0\times4\\ 5.5\times1+(-1)\times2 & 5.5\times3+(-1)\times4 \end{bmatrix} = \begin{bmatrix}2.5&7.5\\3.5&12.5\end{bmatrix}QK⊤=2.55.50−11234=2.5×1+0×25.5×1+(−1)×22.5×3+0×45.5×3+(−1)×4=2.53.57.512.5

再除以 d=2\sqrt{d}=\sqrt2d =2 防止数值爆炸:

S=QK⊤2=1.775.302.478.84S = \frac{QK^\top}{\sqrt2} = \begin{bmatrix}1.77&5.30\\2.47&8.84\end{bmatrix}S=2 QK⊤=1.772.475.308.84

SijS_{ij}Sij 的含义:第 iii 个 token 的 query 与第 jjj 个 token 的 key 的匹配度 (内积越大越"像")。比如 S12=7.5S_{12}=7.5S12=7.5:token 1 的 query 2.5, 02.5,\\,02.5,0 和 token 2 的 key 3, 43,\\,43,4 内积 =2.5×3+0×4=7.5=2.5\times3+0\times4=7.5=2.5×3+0×4=7.5,远大于它和自己 key 的内积 2.5×1+0×2=2.52.5\times1+0\times2=2.52.5×1+0×2=2.5。

③ softmax 归一化:P=softmax(S)P = \mathrm{softmax}(S)P=softmax(S)(按行)

P=0.0280.9720.0020.998P = \begin{bmatrix}0.028&0.972\\0.002&0.998\end{bmatrix}P=0.0280.0020.9720.998

softmax 就是"先指数、再按行归一化 ",拿第 1 行 1.77, 5.301.77,\\,5.301.77,5.30 算一遍:

e1.77=5.87,e5.30=200.3,分母=5.87+200.3=206.2e^{1.77}=5.87,\qquad e^{5.30}=200.3,\qquad \text{分母}=5.87+200.3=206.2e1.77=5.87,e5.30=200.3,分母=5.87+200.3=206.2

P11=5.87206.2≈0.028,P12=200.3206.2≈0.972P_{11}=\frac{5.87}{206.2}\approx0.028,\qquad P_{12}=\frac{200.3}{206.2}\approx0.972P11=206.25.87≈0.028,P12=206.2200.3≈0.972

读法:第 iii 行 = 第 iii 个 token 注意所有 token 的权重 ,每行加起来等于 1。上面这组数表明:token 2 几乎只注意自己(0.998),token 1 则主要注意 token 2(0.972)------这就是"注意力"这个名字的来源:不是所有 token 平均看待,而是按相关性加权

④ 加权求和:O=PVO = PVO=PV(输出是 V 的加权平均,权重就是 softmax 那一行)

O=PV=0.0280.9720.0020.9980.521.54=1.4723.9431.4983.997O = PV = \begin{bmatrix}0.028&0.972\\0.002&0.998\end{bmatrix}\begin{bmatrix}0.5&2\\1.5&4\end{bmatrix} = \begin{bmatrix}1.472&3.943\\1.498&3.997\end{bmatrix}O=PV=0.0280.0020.9720.9980.51.524=1.4721.4983.9433.997

还是"行 × 列":OOO 的每个元素 = PPP 的一行与 VVV 的一列点积。比如第 1 行第 1 列:

O00=0.028×0.5+0.972×1.5=1.472O_{00} = 0.028\times0.5 + 0.972\times1.5 = 1.472O00=0.028×0.5+0.972×1.5=1.472

------token 1 的输出 = 把自己对 token 2 的 V 拿大头(0.972),混上自己的一小部分(0.028)

一句话直觉:Q 说"我在找什么",K 说"我身上有什么标签",两者内积 = 匹配度;softmax 把匹配度变成权重;V 说"我能提供什么内容";最后按权重把 V 加权求和 = 按匹配度收集信息。


3. 反向传播穿过 attention:也只是几个局部导数

上一篇的灵魂是"每个运算 = 一个节点 + 一个存好的反向函数"。attention 拆开就是 5 个节点,每个节点都有现成的局部导数:

节点 前向 反向要算的局部导数
缩放内积 S=QK⊤/dS = QK^\top/\sqrt{d}S=QK⊤/d dQ=(dS/d) K,dK=(dS/d)⊤QdQ = (dS/\sqrt{d})\,K,\quad dK = (dS/\sqrt{d})^\top QdQ=(dS/d )K,dK=(dS/d )⊤Q
softmax P=softmax(S)P = \mathrm{softmax}(S)P=softmax(S) dS=P⊙(dP−1 (P⊤dP))dS = P\odot\big(dP - \mathbf{1}\,(P^\top dP)\big)dS=P⊙(dP−1(P⊤dP))
加权求和 O=PVO = PVO=PV dP=dO V⊤,dV=P⊤dOdP = dO\,V^\top,\quad dV = P^\top dOdP=dOV⊤,dV=P⊤dO

再往后,dQ,dK,dVdQ,dK,dVdQ,dK,dV 各自流进自己的 Linear,用第 1 节那张表(dWQ=X⊤dQdW_Q = X^\top dQdWQ=X⊤dQ 等)把梯度累积到 W_Q / W_K / W_V 上,W_K 还要额外累积到 XXX(dXK=dK WK⊤dX_K = dK\,W_K^\topdXK=dKWK⊤),因为 XXX 同时被三路共享。

注意 softmax 的反向长这样、而不是"每个元素除一下就完事"------因为 softmax 的每个输出 PijP_{ij}Pij 依赖同一行所有 输入 SikS_{ik}Sik(分母是整行求和)。但这也只是一个局部导数,公式查表即可,和 tanh⁡\tanhtanh 的 1−a21-a^21−a2 地位相同。

用第 2 节那组数字,把反向也走一遍 。设损失只关心 OOO 的 (0,0)(0,0)(0,0) 元素,L=(O00−2)2L = (O_{00} - 2)^2L=(O00−2)2,于是

dO=∂L∂O=2(1.472−2)000=−1.057000dO = \frac{\partial L}{\partial O} = \begin{bmatrix}2(1.472-2)&0\\0&0\end{bmatrix} = \begin{bmatrix}-1.057&0\\0&0\end{bmatrix}dO=∂O∂L=2(1.472−2)000=−1.057000

从 dOdOdO 往回,每一步只算"局部斜率"(对照第 1 节的表逐项代入):

节点 局部导数代入 结果
O=PVO=PVO=PV → VVV dV=P⊤dOdV = P^\top dOdV=P⊤dO dV=−0.0300−1.0270dV = \begin{bmatrix}-0.030&0\\-1.027&0\end{bmatrix}dV=−0.030−1.02700
O=PVO=PVO=PV → PPP dP=dO V⊤dP = dO\,V^\topdP=dOV⊤ dP=−0.528−1.58500dP = \begin{bmatrix}-0.528&-1.585\\0&0\end{bmatrix}dP=−0.5280−1.5850
softmax dS=P⊙(dP−1(P⊤dP))dS = P\odot(dP - \mathbf{1}(P^\top dP))dS=P⊙(dP−1(P⊤dP)) dS=+0.029−0.02900dS = \begin{bmatrix}+0.029&-0.029\\0&0\end{bmatrix}dS=+0.0290−0.0290
QK⊤/dQK^\top/\sqrt{d}QK⊤/d → QQQ dQ=(dS/d) KdQ = (dS/\sqrt{d})\,KdQ=(dS/d )K dQ=−0.041−0.04100dQ = \begin{bmatrix}-0.041&-0.041\\0&0\end{bmatrix}dQ=−0.0410−0.0410
QK⊤/dQK^\top/\sqrt{d}QK⊤/d → KKK dK=(dS/d)⊤QdK = (dS/\sqrt{d})^\top QdK=(dS/d )⊤Q dK=0.0510−0.0510dK = \begin{bmatrix}0.051&0\\-0.051&0\end{bmatrix}dK=0.051−0.05100

再各走一步就到权重(三个 Linear 的局部导数):

dWV=X⊤dV=−3.1100−4.1670,dWQ=X⊤dQ=−0.041−0.041−0.082−0.082,dWK=X⊤dK=−0.1030−0.1030dW_V = X^\top dV = \begin{bmatrix}-3.110&0\\-4.167&0\end{bmatrix},\qquad dW_Q = X^\top dQ = \begin{bmatrix}-0.041&-0.041\\-0.082&-0.082\end{bmatrix},\qquad dW_K = X^\top dK = \begin{bmatrix}-0.103&0\\-0.103&0\end{bmatrix}dWV=X⊤dV=−3.110−4.16700,dWQ=X⊤dQ=−0.041−0.082−0.041−0.082,dWK=X⊤dK=−0.103−0.10300

然后和上一篇一模一样:W←W−η dWW \leftarrow W - \eta\,dWW←W−ηdW,梯度下降一步,损失下降,循环。

这个结果你也可以像上一篇第 4 节那样,用"数值梯度验证"(有限差分)复核------softmax 和注意力处处可导,两条独立路径应该对得上。整个 attention 路径上没有任何需要"特判"的魔法。


4. 关键洞察:Q/K/V 的角色是梯度"训练"出来的

这是从反向传播视角看 QKV 最大的收获。回到上一篇第 2 节:训练 = 反复问 ∂L/∂W\partial L/\partial W∂L/∂W,然后更新

  • 初始时刻W_Q / W_K / W_V 是随机数,此时注意力模式(PPP)毫无规律,softmax 近似均匀分布------所有 token 平均注意,等于没注意。
  • 每一轮 :前向算出 LLL 和 PPP;反向把梯度流回三个投影;更新 WQ,WK,WVW_Q,W_K,W_VWQ,WK,WV。
  • 梯度在"教"它们什么 :损失 LLL 隐含着"这个位置该参考哪些位置的信息"(比如"it"后面应该参考"the cat")。梯度 ∂L/∂WQ\partial L/\partial W_Q∂L/∂WQ 就朝能产生这种注意力模式的方向去推投影权重。
  • 几万步之后 :WQW_QWQ 学会了"提出对的查询",WKW_KWK 学会了"给出对的标签",WVW_VWV 学会了"给出有用的内容"。

所以 Q/K/V 不是被谁"设计"成 query/key/value 的------是反向传播把这三个普通参数矩阵 训练成了有这个分工 。它们只是三个并排的 Linear 层,只是恰好处在"匹配"(QKᵀ)和"取值"(×V)这两个运算的两侧,于是梯度就把语义装进了它们。
#mermaid-svg-KpwdwqPNugbxkwF3{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-KpwdwqPNugbxkwF3 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-KpwdwqPNugbxkwF3 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-KpwdwqPNugbxkwF3 .error-icon{fill:#552222;}#mermaid-svg-KpwdwqPNugbxkwF3 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-KpwdwqPNugbxkwF3 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-KpwdwqPNugbxkwF3 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-KpwdwqPNugbxkwF3 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-KpwdwqPNugbxkwF3 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-KpwdwqPNugbxkwF3 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-KpwdwqPNugbxkwF3 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-KpwdwqPNugbxkwF3 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-KpwdwqPNugbxkwF3 .marker.cross{stroke:#333333;}#mermaid-svg-KpwdwqPNugbxkwF3 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-KpwdwqPNugbxkwF3 p{margin:0;}#mermaid-svg-KpwdwqPNugbxkwF3 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-KpwdwqPNugbxkwF3 .cluster-label text{fill:#333;}#mermaid-svg-KpwdwqPNugbxkwF3 .cluster-label span{color:#333;}#mermaid-svg-KpwdwqPNugbxkwF3 .cluster-label span p{background-color:transparent;}#mermaid-svg-KpwdwqPNugbxkwF3 .label text,#mermaid-svg-KpwdwqPNugbxkwF3 span{fill:#333;color:#333;}#mermaid-svg-KpwdwqPNugbxkwF3 .node rect,#mermaid-svg-KpwdwqPNugbxkwF3 .node circle,#mermaid-svg-KpwdwqPNugbxkwF3 .node ellipse,#mermaid-svg-KpwdwqPNugbxkwF3 .node polygon,#mermaid-svg-KpwdwqPNugbxkwF3 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-KpwdwqPNugbxkwF3 .rough-node .label text,#mermaid-svg-KpwdwqPNugbxkwF3 .node .label text,#mermaid-svg-KpwdwqPNugbxkwF3 .image-shape .label,#mermaid-svg-KpwdwqPNugbxkwF3 .icon-shape .label{text-anchor:middle;}#mermaid-svg-KpwdwqPNugbxkwF3 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-KpwdwqPNugbxkwF3 .rough-node .label,#mermaid-svg-KpwdwqPNugbxkwF3 .node .label,#mermaid-svg-KpwdwqPNugbxkwF3 .image-shape .label,#mermaid-svg-KpwdwqPNugbxkwF3 .icon-shape .label{text-align:center;}#mermaid-svg-KpwdwqPNugbxkwF3 .node.clickable{cursor:pointer;}#mermaid-svg-KpwdwqPNugbxkwF3 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-KpwdwqPNugbxkwF3 .arrowheadPath{fill:#333333;}#mermaid-svg-KpwdwqPNugbxkwF3 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-KpwdwqPNugbxkwF3 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-KpwdwqPNugbxkwF3 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-KpwdwqPNugbxkwF3 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-KpwdwqPNugbxkwF3 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-KpwdwqPNugbxkwF3 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-KpwdwqPNugbxkwF3 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-KpwdwqPNugbxkwF3 .cluster text{fill:#333;}#mermaid-svg-KpwdwqPNugbxkwF3 .cluster span{color:#333;}#mermaid-svg-KpwdwqPNugbxkwF3 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-KpwdwqPNugbxkwF3 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-KpwdwqPNugbxkwF3 rect.text{fill:none;stroke-width:0;}#mermaid-svg-KpwdwqPNugbxkwF3 .icon-shape,#mermaid-svg-KpwdwqPNugbxkwF3 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-KpwdwqPNugbxkwF3 .icon-shape p,#mermaid-svg-KpwdwqPNugbxkwF3 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-KpwdwqPNugbxkwF3 .icon-shape .label rect,#mermaid-svg-KpwdwqPNugbxkwF3 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-KpwdwqPNugbxkwF3 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-KpwdwqPNugbxkwF3 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-KpwdwqPNugbxkwF3 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 注意力模式逐渐从'均匀'变成'有选择'
随机初始化

W_Q W_K W_V
前向:算出注意力 P 和损失 L
反向:梯度流回

dW_Q dW_K dW_V
更新:W ← W − η·dW
最终:Q 会找、K 会给标签、V 会给内容

(角色涌现)

这和你上一篇"梯度是负的 → 说明 w1w_1w1 该变大"是同一件事:梯度不解释语义,但它会找到让损失下降的方向。attention 只是让这条方向里多了一类"选择性组合"的算子。


5. 为什么是 QKV 这种结构?(还有别的吗)

一句话版 :QKV 不是唯一可能的结构,而是"选择性收集信息"这个需求下,表达能力 × 可训练性 × 硬件速度三者权衡的产物。想知道它好在哪,最好的办法是想象"去掉 Q"或"去掉 V"会退化到什么程度------那正是其它结构存在的理由。

5.1 先看需求:attention 要解决两个子问题

每个 token 输出时要做两件不同的事:

  1. 决定"该参考谁、参考多少" ------ 匹配/打分(对应 Q⋅KQ\cdot KQ⋅K);
  2. 把参考对象的"内容"拿过来 ------ 取值/收集(对应 ×V\times V×V)。

所以结构天生需要两条不同的信息通道:一条管"谁重要"(Q/K),一条管"给什么内容"(V)。

5.2 Q 和 K 为什么要分开:匹配必须是非对称的

看打分矩阵 Sij=Qi⋅KjS_{ij} = Q_i \cdot K_jSij=Qi⋅Kj:SijS_{ij}Sij(i 问 j)和 SjiS_{ji}Sji(j 问 i)相互独立、可以不一样------语言正是这样(词序、方向、主谓宾角色都带着方向性)。

如果不用 Q/K 分开,直接用输入自己打分 S=XX⊤S = XX^\topS=XX⊤:

  • 对称 :Sij=SjiS_{ij} = S_{ji}Sij=Sji,注意力方向无法区分;
  • 对角占优:自己和自己的内积最大,softmax 后每行几乎只注意自己 → 退化成"没注意别人";
  • 不可学习 :XX⊤XX^\topXX⊤ 没有可训练参数,学不出"该找谁"。

而 Q=XWQ, K=XWKQ = XW_Q,\ K = XW_KQ=XWQ, K=XWK 是两个独立可训练 的投影,等于给每个 token 两顶帽子:"我想找什么"(Q)"我身上有什么标签"(K)。一个 token 可以"自己想要猫",但"身上贴着狗",别人就能按"狗"这个标签找到它------这正是第 4 节"角色由梯度训练出来"的落点。

5.3 V 为什么要单独:匹配和内容是两种职责

V=XWVV = XW_VV=XWV 让"哪些位置值得被看"(由 SSS 决定)和"这些位置实际提供什么"(由 VVV 决定)解耦:

  • 一个 token 可以"特别显眼、被大家注意"(K 有强标签),但实际内容(V)是另一套向量;
  • 没有 V、直接用输入加权(O=PXO = PXO=PX)也能做,但"被注意的资格"和"携带的内容"被迫共用同一个向量,互相牵制、表达能力受限。

加权求和 O=PVO = PVO=PV 本质上是一个可学习的"软查表":先定位(Q/K),再取值(V)。

5.4 其它结构真的存在

结构 打分方式 优点 缺点
QKV(本文) QK⊤/d→softmaxQK^\top/\sqrt d \to \mathrm{softmax}QK⊤/d →softmax 非对称、可学习、一个 GEMM 搞定、GPU 极快 O(L²) 复杂度;三次投影的参数开销
加性注意力(Bahdanau) v⊤tanh⁡(W1hi+W2sj)v^\top\tanh(W_1 h_i + W_2 s_j)v⊤tanh(W1hi+W2sj) 早年主流、灵活 不是纯矩阵乘,硬件不友好、慢
无投影自相似 XX⊤XX^\topXX⊤ 零参数、最省 对称、对角占优、不可学习(退化)
线性注意力 ϕ(Q)ϕ(K)⊤\phi(Q)\phi(K)^\topϕ(Q)ϕ(K)⊤(去掉 softmax) 可结合律 → O(L) 复杂度 表达不了尖锐的"赢者通吃"
MQA / GQA 多头之间共享 K/V KV Cache 省内存、推理快 每头表达能力略降

QKV 选的是点积(乘性)注意力 这一支:矩阵乘 = GEMM = 硬件最爱,加 d\sqrt dd 缩放是为了数值稳定。

5.5 优劣势清单

优势

  • 非对称匹配 → 能表达方向性、词序、角色关系;
  • Q/K/V 三个独立投影 → 梯度能分别"教"它们扮演不同角色(呼应第 4 节);
  • 结构统一 → 自注意力、跨注意力(Q 来自序列 A,K/V 来自序列 B)、多头,都只是改输入;
  • 计算全是矩阵乘,GPU 友好。

代价

  • 三次投影:输入 ddd 维 → 3 个 d×dd\times dd×d 投影,参数和算力不便宜;
  • O(L²) 的注意力矩阵:序列一长就爆炸,所以才有 FlashAttention、稀疏/线性注意力、GQA 等后续改造;
  • 结构假定了"匹配 → 收集"这个配方是对的;位置信息要另加位置编码(RoPE 等);
  • 超参数多(dheadd_{head}dhead、nheadsn_{heads}nheads...),调起来麻烦。

一句话直觉:Q 负责问、K 负责被找到、V 负责给内容。想验证这个设计值钱在哪,就试着"去掉 Q"或"去掉 V",看它会退化到什么程度------那是替代结构存在的理由。


6. 主流开源模型在这个基础上升级了什么?

一句话版 :从 2017 到现在的旗舰开源模型(Llama、Mistral、Qwen、DeepSeek...),没有推翻 QKV ,而是围着它的几个痛点做手术------O(L²) 复杂度、位置编码、KV Cache 内存、数值稳定。你学会的这套前向 + 反向,换上它们的"外壳"后依然成立。

6.1 先给一张总览:痛点 → 主流解法

博客里的痛点(第 5.5 节) 主流开源模型的解法
O(L²) 的注意力矩阵 FlashAttention、滑动窗口、稀疏注意力、线性注意力
位置信息要另加编码 RoPE(旋转位置编码)+ YaRN/NTK 长上下文缩放
三次投影 / KV 开销大 KV Cache、MQA → GQA、DeepSeek 的 MLA
d\sqrt dd 只能粗略压内积 QK-Norm(对 Q、K 先归一化)
单头表达能力有限 多头 + GQA 分组(本质还是那套 QKV)

6.2 推理内存:GQA 是默认项,MLA 更进一步

  • KV Cache:生成时把历史 K/V 存下来、不重算------推理标配。
  • MQA / GQA (Llama 2/3、Mistral、Qwen):多头之间共享 K/V (GQA 是每 ggg 个头一组共享),KV Cache 大幅缩小。打分逻辑 QK⊤QK^\topQK⊤ 完全不变,只是 K/V 的头数变少了。
  • MLA(DeepSeek-V2/V3):把 K/V 先压进一个低维"潜向量",用的时候再展开,KV Cache 缩到约 1/10。等于"KV 也加了个瓶颈层"。

6.3 位置编码:RoPE 取代正弦/绝对位置

  • RoPE(旋转位置编码) :把位置信息旋转进 Q/K 的内积 ------QQQ、KKK 各乘一个随位置旋转的矩阵,内积天然带"相对位置"。现在开源旗舰几乎全套用它(呼应第 5.5 节"位置要另加编码",它就是那个答案)。
  • 长上下文 :RoPE 不能无限外推,于是有 YaRN、NTK 缩放、dynamic RoPE------训练 4K 的模型推理能撑到 128K+。

6.4 复杂度 O(L²):软硬兼施

  • FlashAttention (工程层,不改模型):不把 L×LL\times LL×L 矩阵物化到显存,分块算 + 反向重算,速度显存双赢。PyTorch 的 sdpa 默认就是它。
  • 滑动窗口(Mistral):每 token 只看附近窗口,外加少量全局层。
  • 稀疏注意力(DeepSeek-V3.2 / NSA):对打分做 top-k 选择,只算重要的那部分------就是第 5.4 节表格里"稀疏/线性注意力"那一行,现在真落地了。

6.5 数值稳定与结构级创新

  • QK-Norm (Qwen3、Gemma 2):算 QK⊤QK^\topQK⊤ 之前 先对 Q、K 做归一化------比 d\sqrt dd 更彻底地控制内积尺度(第 2 节那个 d\sqrt dd 的进阶版)。
  • 模型块层面:Pre-norm + RMSNorm + SwiGLU(Llama 系)已成新基线,但不动 attention 本身。
  • MoE (DeepSeek-V3、Mixtral):换的是 FFN(专家路由),attention 还是普通 GQA/MLA------QKV 这块地基至今没人推翻
  • Hybrid(Jamba):attention 与状态空间模型(Mamba)交替堆叠,长序列便宜、短序列精确。
  • MTP(DeepSeek-V3):一次预测多个下一个 token,纯训练技巧。

一句话直觉:你刚学会的 QKV 是所有主流模型 attention 的地基;升级都发生在"地基之上"------位置怎么进(RoPE)、KV 怎么省(GQA/MLA)、L×L 怎么躲(FlashAttention/稀疏)、数值怎么稳(QK-Norm)。


7. 一张图看懂:Attention + FFN 组成一个 Transformer Block

把前 6 节拼起来看------一个 Transformer Block = 注意力子层(QKV)+ 前馈子层(FFN),两件事交替、残差相连:
#mermaid-svg-FUutpfr3nQdViRO3{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-FUutpfr3nQdViRO3 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-FUutpfr3nQdViRO3 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-FUutpfr3nQdViRO3 .error-icon{fill:#552222;}#mermaid-svg-FUutpfr3nQdViRO3 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-FUutpfr3nQdViRO3 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-FUutpfr3nQdViRO3 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-FUutpfr3nQdViRO3 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-FUutpfr3nQdViRO3 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-FUutpfr3nQdViRO3 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-FUutpfr3nQdViRO3 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-FUutpfr3nQdViRO3 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-FUutpfr3nQdViRO3 .marker.cross{stroke:#333333;}#mermaid-svg-FUutpfr3nQdViRO3 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-FUutpfr3nQdViRO3 p{margin:0;}#mermaid-svg-FUutpfr3nQdViRO3 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-FUutpfr3nQdViRO3 .cluster-label text{fill:#333;}#mermaid-svg-FUutpfr3nQdViRO3 .cluster-label span{color:#333;}#mermaid-svg-FUutpfr3nQdViRO3 .cluster-label span p{background-color:transparent;}#mermaid-svg-FUutpfr3nQdViRO3 .label text,#mermaid-svg-FUutpfr3nQdViRO3 span{fill:#333;color:#333;}#mermaid-svg-FUutpfr3nQdViRO3 .node rect,#mermaid-svg-FUutpfr3nQdViRO3 .node circle,#mermaid-svg-FUutpfr3nQdViRO3 .node ellipse,#mermaid-svg-FUutpfr3nQdViRO3 .node polygon,#mermaid-svg-FUutpfr3nQdViRO3 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-FUutpfr3nQdViRO3 .rough-node .label text,#mermaid-svg-FUutpfr3nQdViRO3 .node .label text,#mermaid-svg-FUutpfr3nQdViRO3 .image-shape .label,#mermaid-svg-FUutpfr3nQdViRO3 .icon-shape .label{text-anchor:middle;}#mermaid-svg-FUutpfr3nQdViRO3 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-FUutpfr3nQdViRO3 .rough-node .label,#mermaid-svg-FUutpfr3nQdViRO3 .node .label,#mermaid-svg-FUutpfr3nQdViRO3 .image-shape .label,#mermaid-svg-FUutpfr3nQdViRO3 .icon-shape .label{text-align:center;}#mermaid-svg-FUutpfr3nQdViRO3 .node.clickable{cursor:pointer;}#mermaid-svg-FUutpfr3nQdViRO3 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-FUutpfr3nQdViRO3 .arrowheadPath{fill:#333333;}#mermaid-svg-FUutpfr3nQdViRO3 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-FUutpfr3nQdViRO3 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-FUutpfr3nQdViRO3 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-FUutpfr3nQdViRO3 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-FUutpfr3nQdViRO3 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-FUutpfr3nQdViRO3 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-FUutpfr3nQdViRO3 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-FUutpfr3nQdViRO3 .cluster text{fill:#333;}#mermaid-svg-FUutpfr3nQdViRO3 .cluster span{color:#333;}#mermaid-svg-FUutpfr3nQdViRO3 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-FUutpfr3nQdViRO3 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-FUutpfr3nQdViRO3 rect.text{fill:none;stroke-width:0;}#mermaid-svg-FUutpfr3nQdViRO3 .icon-shape,#mermaid-svg-FUutpfr3nQdViRO3 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-FUutpfr3nQdViRO3 .icon-shape p,#mermaid-svg-FUutpfr3nQdViRO3 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-FUutpfr3nQdViRO3 .icon-shape .label rect,#mermaid-svg-FUutpfr3nQdViRO3 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-FUutpfr3nQdViRO3 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-FUutpfr3nQdViRO3 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-FUutpfr3nQdViRO3 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} ① 横向:token 之间通信
② 纵向:单个 token 内部加工
FFN 前馈子层 ------ 想清楚 / 存知识
W1(d → 4d)
激活

GELU / SiLU / SwiGLU
W2(4d → d)
注意力子层(QKV)------ 谁该参考谁
Q = XW_Q
S = QKᵀ/√d

打分
K = XW_K
P = softmax(S)

权重
O = P·V

加权求和
V = XW_V
输入 x(L 个 token,每行一个向量)

  • 位置编码(RoPE)
  • 残差(x + O)
    归一化
  • 残差
    归一化
    输出 → 下一个 Block(重复 N 次)

这张图里的两个关键澄清:

  1. token 是什么 :token 不是权重矩阵,而是输入数据 ------序列里的一个单元(一个词/子词/图像块)。XXX 的每一行就是一个 token 的向量;WQ/WK/WVW_Q/W_K/W_VWQ/WK/WV(以及 FFN 的 W1,W2W_1,W_2W1,W2)才是权重矩阵(参数) 。本文例子里 2 个 token,就是 XXX 有 2 行。
  2. 左半是本文 (QKV,你已会算),右半是上一篇的 MLP(FFN):注意力负责"token 之间通信",FFN 负责"每个 token 独自加工"。整个 Transformer = 这两件事交替 N 次。

一句话:attention 决定"该听谁的",FFN 决定"听完怎么想";token 是数据(XXX 的行),QKV/FFN 里的 WWW 才是要训练的权重。

还有一个关键:图里的"+ 残差"是给梯度修的"直达高速路"。

每个子层都套着 out=x+F(x)out = x + F(x)out=x+F(x)(FFF 是 attention 或 FFN)。它存在的意义:

  • 反向时梯度至少是 1 :对"加法节点"求导得 ∂out/∂x=1+∂F/∂x\partial out/\partial x = 1 + \partial F/\partial x∂out/∂x=1+∂F/∂x,那个 "111" 来自直达路------梯度能绕过子层直接传回前一层,不会随着深度连乘而消失;
  • 每层只需学"增量" :最坏情况 F=0F=0F=0,输出仍是 xxx(等价恒等映射),先不掉分、再慢慢精修;
  • 所以 Transformer 能堆几十上百层:没有残差,深网梯度会消失、前面学不动(这是 ResNet 2015 的关键贡献,Transformer 全盘继承)。

一句话:attention 决定"该听谁的",FFN 决定"听完怎么想",残差决定"梯度能不能一路穿回去";token 是数据(XXX 的行),QKV/FFN 里的 WWW 才是要训练的权重。


8. 收尾:和上一篇的"下一步"接上

readme.md 第 4 节说过:把 MLP 的 W 换成 Q/K/V 投影、把 h@W 换成 softmax(QKᵀ/√d)V,attention 也只是计算图里的几个节点,反向传播完全一样自动。本文把这句话兑现了:

上一篇的 MLP Transformer 的 attention
W1, W2(2 个 Linear) W_Q, W_K, W_V + 输出投影(4 个 Linear)
z=wx+b→tanh⁡→outz=wx+b \to \tanh \to outz=wx+b→tanh→out QK⊤/d→softmax→×VQK^\top/\sqrt{d} \to \mathrm{softmax} \to \times VQK⊤/d →softmax→×V
局部导数:dW=X⊤dZdW=X^\top dZdW=X⊤dZ 同一套,加上 softmax / 矩阵乘的局部导数
训练循环:前向 → 损失 → 反向 → 更新 完全一样

完整的心智模型一句话

QKV = 你早已会算的三个 Linear 层(上一篇的第一块积木,变大了);attention = 计算图里多出的几个无参数节点(每个节点只需一个局部导数);它们的语义角色 = 梯度下降几万步之后"涌现"的结果。 你手写的 60 行 autograd,放到 attention 上一样能自动算出 ∂L/∂WQ\partial L/\partial W_Q∂L/∂WQ。

下一步建议(呼应 readme.md 的路线图):动手写一个 attention_backprop.py 小实验------用 numpy 手写 softmax(QKᵀ/√d)V 的前向 + 反向,再用有限差分验证梯度,然后用它替换 backprop_lab.py 里的线性层,你就能亲眼看到"注意力的梯度到底从哪来"。


附:本文用到的公式速查

积木 局部导数
线性层 Y=XW+bY = XW + bY=XW+b dW=X⊤dY, db=∑ndY, dX=dY W⊤dW = X^\top dY,\ \ db=\sum_n dY,\ \ dX = dY\,W^\topdW=X⊤dY, db=∑ndY, dX=dYW⊤
缩放内积 S=QK⊤/dS = QK^\top/\sqrt{d}S=QK⊤/d dQ=(dS/d)K, dK=(dS/d)⊤QdQ = (dS/\sqrt{d})K,\ \ dK = (dS/\sqrt{d})^\top QdQ=(dS/d )K, dK=(dS/d )⊤Q
softmax P=softmax(S)P = \mathrm{softmax}(S)P=softmax(S) dS=P⊙(dP−1(P⊤dP))dS = P\odot(dP - \mathbf{1}(P^\top dP))dS=P⊙(dP−1(P⊤dP))
加权求和 O=PVO = PVO=PV dP=dO V⊤, dV=P⊤dOdP = dO\,V^\top,\ \ dV = P^\top dOdP=dOV⊤, dV=P⊤dO

规则没变:每一层只管它对"直接输入"的导数,剩下的交给链式法则连乘 ;要更新的参数永远在分母里(∂L/∂WQ\partial L/\partial W_Q∂L/∂WQ),分子永远是损失 LLL。

相关推荐
磁场转动100万匹40 分钟前
深度学习入门:从神经网络到反向传播的完整解析
人工智能·深度学习·神经网络
啥都想学点的研究生1 小时前
一篇文章讲清楚:超参数的选择方法——交叉验证和网格搜索
人工智能·深度学习·机器学习
点PY1 小时前
《一种基于深度学习的超分辨率重建方法及系统》专利精读
人工智能·深度学习·超分辨率重建
阿龙AI日记2 小时前
快速入门深度学习01:神经元和神经网络
人工智能·深度学习·神经网络
I Promise342 小时前
Fast‑BEV 完整实战教程(智驾工程视角)
深度学习·神经网络·目标检测·目标跟踪·自动驾驶
神经星星3 小时前
「TVM教程」理解 Relax 抽象层
人工智能·深度学习
薛定e的猫咪3 小时前
【大模型量化】使用 llama.cpp 完成量化、本地推理与服务化部署
人工智能·深度学习·算法·llama
薛定e的猫咪3 小时前
【源码解读版】ContraBAR:用 CPC 对比学习替代变分推断做贝叶斯元 RL
人工智能·深度学习·学习·算法·机器学习
现代野蛮人3 小时前
【深度学习实验】—— 基于 LSTM 的年度医疗费用回归预测
深度学习·回归·lstm