损失函数在标量和矩阵上的求导对比

损失函数在标量和矩阵上的求导对比

从标量到矩阵时多了两件事 :① 同一份参数被多个样本共享 、② 同一个输入通向多个输出。这两件事都会带来"求和",而标量版因为只有一个样本、一个输出,根本看不到求和。下面把两张表放一起,再逐个拆给你看。

第一部分 标量 → 矩阵:核心概念

1. 先说重点:两张表,别混在一起

xxx 只是"局部导数",dW=X⊤dZdW=X^\top dZdW=X⊤dZ 是"损失梯度",两者不是一个层面。 之前把这两列并排放,容易让人以为 dWdWdW 直接等于 xxx,其实不是。要严格对应,得分两张表看:

1.1 表 1:局部导数(只问"output 对 input 的斜率",还没乘上游梯度 dzdzdz
标量(1 样本 × 1 输出) 矩阵(N 样本 × 多输出)
前向 y=wx+by = wx + by=wx+b Z=XW+bZ = XW + bZ=XW+b,X:(N,din), W:(din,dout)X{:}(N,d_{in}),\ W{:}(d_{in},d_{out})X:(N,din), W:(din,dout)
对权重 ∂y∂w=x\dfrac{\partial y}{\partial w}=x∂w∂y=x ∂Z∂W=X\dfrac{\partial Z}{\partial W}=X∂W∂Z=X
对偏置 ∂y∂b=1\dfrac{\partial y}{\partial b}=1∂b∂y=1 ∂Z∂b=1\dfrac{\partial Z}{\partial b}=1∂b∂Z=1
对输入 ∂y∂x=w\dfrac{\partial y}{\partial x}=w∂x∂y=w ∂Z∂X=W\dfrac{\partial Z}{\partial X}=W∂X∂Z=W

这层你之前的理解完全正确 :对权重的局部偏导"就是输入 xxx"。标量是单个数字 xxx,矩阵版就是整块输入 XXX------这才是真正互相对应的两列

1.2 表 2:损失梯度(链式法则把 LLL 一路连到参数,已经乘了上游梯度 dzdzdz
标量(N 样本共用一个 www) 矩阵(N 样本 × 多输出)
对权重 ∂L∂w=∑ndzn⋅xn\dfrac{\partial L}{\partial w}=\sum_n dz_n\cdot x_n∂w∂L=∑ndzn⋅xn dW=X⊤dZdW = X^\top dZdW=X⊤dZ
对偏置 ∂L∂b=∑ndzn\dfrac{\partial L}{\partial b}=\sum_n dz_n∂b∂L=∑ndzn db=∑n=1NdZn,:db = \sum\limits_{n=1}^{N} dZ_{n,:}db=n=1∑NdZn,:
对输入 ∂L∂xn=dzn⋅w\dfrac{\partial L}{\partial x_n}=dz_n\cdot w∂xn∂L=dzn⋅w dX=dZ W⊤dX = dZ\,W^\topdX=dZW⊤

dW=X⊤dZdW = X^\top dZdW=X⊤dZ 的真正对应物是 ∂L/∂w=∑ndznxn\partial L/\partial w=\sum_n dz_n x_n∂L/∂w=∑ndznxn,不是 ∂y/∂w=x\partial y/\partial w=x∂y/∂w=x。 后者只是"一段斜率",前者是"整串链式法则"。所以 dWdWdW 里那个 dzdzdz 不是凭空多出来的------它是把损失一路链到 WWW 时,前面那段已经传回来的梯度,带着它一起乘才叫"损失偏导"。
再次强调:下面两处为什么多出"求和",是因为表 2 才涉及"多个样本共享参数"。

2. 为什么多出"求和"?两个来源

2.1 来源①:参数被多个样本共享 → 对样本求和

标量版只有一个样本,www 只被用一次。但批量版里,同一个 www 被 NNN 个样本共用 (每个样本都是 zn=w xn+bz_n = w\,x_n + bzn=wxn+b)。

偏导 ∂L/∂w\partial L/\partial w∂L/∂w 问的是"www 动一点,LLL 动多少",而 www 影响每个 样本的 znz_nzn,所以要把所有样本的贡献都加起来:

∂L∂w=∑n∂L∂zn∂zn∂w=∑ndzn⋅xn\frac{\partial L}{\partial w}=\sum_n \frac{\partial L}{\partial z_n}\frac{\partial z_n}{\partial w} =\sum_n d z_n\cdot x_n∂w∂L=n∑∂zn∂L∂w∂zn=n∑dzn⋅xn

标量版 xxx 是单个数字,矩阵版 xnx_nxn 是 NNN 个样本拼成的列向量,∑ndznxn\sum_n d z_n x_n∑ndznxn 恰好就是两个向量的点积

X⊤⏟样本排成列  dZ⏟梯度排成列=∑nXn,k dZn,o\underbrace{X^\top}{\text{样本排成列}}\;\underbrace{dZ}{\text{梯度排成列}}=\sum_n X_{n,k}\,dZ_{n,o}样本排成列 X⊤梯度排成列 dZ=n∑Xn,kdZn,o

这就是 dW=X⊤dZdW=X^\top dZdW=X⊤dZ。转置的目的 :让"样本维 nnn"对齐、做内积消掉------因为 dwdwdw 里要对 nnn 求和。

偏置同理:bbb 也被所有样本共享,且 ∂zn/∂b=1\partial z_n/\partial b=1∂zn/∂b=1,所以 dbo=∑ndZn,o⋅1=∑ndZn,odb_o=\sum_n dZ_{n,o}\cdot 1=\sum_n dZ_{n,o}dbo=∑ndZn,o⋅1=∑ndZn,o。这就是那条"∑ndZn,:\sum_n dZ_{n,:}∑ndZn,:"。

2.2 深入看:为什么是 ∑ndzn⋅xn\sum_n dz_n\cdot x_n∑ndzn⋅xn

这条公式值得单独拆开。它其实在说一句话:www 同时影响 NNN 个样本的输出,一共 NNN 条"路径"通向 LLL,所以要把 NNN 条路径的贡献全加起来

① www 是怎么"分叉"的 :批量版每个样本用自己的输入 xnx_nxn,但共用同一个 www:

z1=w x1+b,z2=w x2+b,...,zN=w xN+bz_1 = w\,x_1 + b,\qquad z_2 = w\,x_2 + b,\qquad \dots,\qquad z_N = w\,x_N + bz1=wx1+b,z2=wx2+b,...,zN=wxN+b

所以数据流不是单线,而是从 www 分叉成 NNN 条:

text 复制代码
        ┌─> z1 ─> L
        ├─> z2 ─> L
   w ───┼─> z3 ─> L
        ├─> ...
        └─> zN ─> L

LLL 同时依赖所有 znz_nzn(比如 MSE:L=∑n(zn−yn)2L=\sum_n(z_n-y_n)^2L=∑n(zn−yn)2)。于是"www 动一点,LLL 动多少" = 每条路径都贡献一份,总共 NNN 份加起来。

② 这就是全导数的标准写法:当一个变量通过多条中间路径影响最终结果时,总导数 = 各路径偏导之和(链式法则的"全"版本):

∂L∂w=∑n=1N ∂L∂zn⋅∂zn∂w\frac{\partial L}{\partial w}=\sum_{n=1}^{N}\ \frac{\partial L}{\partial z_n}\cdot\frac{\partial z_n}{\partial w}∂w∂L=n=1∑N ∂zn∂L⋅∂w∂zn

每一项拆开看:

  • ∂L∂zn\dfrac{\partial L}{\partial z_n}∂zn∂L:第 nnn 条路径已流过的那段梯度,记为 dzndz_ndzn。它不是"从 LLL 一步算出",而是像剥洋葱一样从损失经过后面的层一层层传来的,但在这条公式里它就是一个已知的、由前面反向得到的数。
  • ∂zn∂w\dfrac{\partial z_n}{\partial w}∂w∂zn:第 nnn 条路径的局部斜率。因为 zn=wxn+bz_n = w x_n + bzn=wxn+b,对 www 求导就是 xnx_nxn(bbb 对 www 是常数)。

③ 拿 2 个样本走一遍 :设 x1=1, x2=3x_1=1,\ x_2=3x1=1, x2=3,w=2, b=0w=2,\ b=0w=2, b=0,目标 y1=3, y2=7y_1=3,\ y_2=7y1=3, y2=7,损失 L=(z1−y1)2+(z2−y2)2L=(z_1-y_1)^2+(z_2-y_2)^2L=(z1−y1)2+(z2−y2)2。

前向:

z1=2×1+0=2,z2=2×3+0=6,L=(2−3)2+(6−7)2=2z_1=2\times1+0=2,\qquad z_2=2\times3+0=6,\qquad L=(2-3)^2+(6-7)^2=2z1=2×1+0=2,z2=2×3+0=6,L=(2−3)2+(6−7)2=2

先算每个样本自己的梯度 dzndz_ndzn:

dz1=2(z1−y1)=2(2−3)=−2,dz2=2(z2−y2)=2(6−7)=−2dz_1=2(z_1-y_1)=2(2-3)=-2,\qquad dz_2=2(z_2-y_2)=2(6-7)=-2dz1=2(z1−y1)=2(2−3)=−2,dz2=2(z2−y2)=2(6−7)=−2

再把两条路径用链式法则合起来:

路径 dzndz_ndzn ∂zn/∂w=xn\partial z_n/\partial w=x_n∂zn/∂w=xn 该路径贡献
w→z1→Lw\to z_1\to Lw→z1→L −2-2−2 111 −2×1=−2-2\times1=-2−2×1=−2
w→z2→Lw\to z_2\to Lw→z2→L −2-2−2 333 −2×3=−6-2\times3=-6−2×3=−6

∂L∂w=(−2)+(−6)=−8\frac{\partial L}{\partial w}=(-2)+(-6)=-8∂w∂L=(−2)+(−6)=−8

④ 直接对 LLL 求导来验证 :L=(w−3)2+(3w−7)2L=(w-3)^2+(3w-7)^2L=(w−3)2+(3w−7)2(代入 x1=1,x2=3x_1=1,x_2=3x1=1,x2=3),

∂L∂w=2(w−3)+2(3w−7)⋅3\frac{\partial L}{\partial w}=2(w-3)+2(3w-7)\cdot3∂w∂L=2(w−3)+2(3w−7)⋅3

在 w=2w=2w=2 处:2(−1)+6(−1)=−82(-1)+6(-1)=-82(−1)+6(−1)=−8。两条路算出同一个 −8-8−8,说明"求和 NNN 条路径"是对的。

⑤ 一个容易踩的坑 :别把 dzndz_ndzn 当成"每个样本独立的损失"。dzndz_ndzn 是 ∂L/∂zn\partial L/\partial z_n∂L/∂zn,这里的 LLL 是整批 的损失。正因为 LLL 包含了所有样本,www 的梯度才会收集到每一份;如果只对单个样本求导(没有那个 ∑n\sum_n∑n),就丢掉了一半信息------那正是标量版和批量版的差别。

一句话:这条公式 = 链式法则的多路径求和版 。∑n\sum_n∑n 数的是"www 到 LLL 有几条路"------因为参数被 NNN 个样本共享,就有 NNN 条路,每条路的贡献 = 该样本的梯度 dzndz_ndzn × 该样本的局部斜率 xnx_nxn。

2.3 来源②:一个输入通向多个输出 → 对输出求和

标量版 xxx 只喂给一个输出 yyy。但矩阵版里,输入 Xn,kX_{n,k}Xn,k 同时参与该样本的每个输出 Zn,1,Zn,2,...Z_{n,1},Z_{n,2},...Zn,1,Zn,2,...(因为 Zn,o=∑kXn,kWk,o+boZ_{n,o}=\sum_k X_{n,k}W_{k,o}+b_oZn,o=∑kXn,kWk,o+bo)。

所以对输入求偏导时,要把每个输出方向的贡献都收回来:

∂L∂Xn,k=∑o∂L∂Zn,o∂Zn,o∂Xn,k=∑odZn,o Wk,o\frac{\partial L}{\partial X_{n,k}}=\sum_o \frac{\partial L}{\partial Z_{n,o}}\frac{\partial Z_{n,o}}{\partial X_{n,k}} =\sum_o dZ_{n,o}\,W_{k,o}∂Xn,k∂L=o∑∂Zn,o∂L∂Xn,k∂Zn,o=o∑dZn,oWk,o

右边就是 (dZ W⊤)n,k(dZ\,W^\top)_{n,k}(dZW⊤)n,k------对输出维 ooo 求和 。标量版的 www 是单个数字,矩阵版变成沿 WWW 的"一列"展开再求和,所以 WWW 要转置让 ooo 对齐做内积。

3. 常见疑问

3.1 疑问一:dbdbdb 在标量里不是 1 吗?向量里怎么变成求和了?

你记住的"=1=1=1"没错,但那是局部导数 ;dbdbdb 是损失梯度,两者不是一回事。

  • 局部导数 :∂Z/∂b=1\partial Z/\partial b = 1∂Z/∂b=1。矩阵版里每个输出对自己的偏置,斜率都是 1------确实"是一组 1"。
  • 损失梯度 :dbo=∑ndZn,o⋅1=∑ndZn,odb_o=\sum_n dZ_{n,o}\cdot 1=\sum_n dZ_{n,o}dbo=∑ndZn,o⋅1=∑ndZn,o。那个 1 还在,但它被上游梯度 dZdZdZ 乘了 ,还要对所有样本求和 (因为 bbb 被 NNN 个样本共享)。

用数字看:dZ=1,1⊤dZ=1,1^\topdZ=1,1⊤(2 个样本),

db=∑ndZn,:=1+1=2(不是 1)db=\sum_n dZ_{n,:}=1+1=2\quad(\text{不是 }1)db=n∑dZn,:=1+1=2(不是 1)

那为什么 blog_backprop.md 里 db2db_2db2 算出的是 1?因为那个例子里传到 out 的梯度正好是 1 (dout=1dout=1dout=1),于是 db2=dout×1=1×1=1db_2=dout\times1=1\times1=1db2=dout×1=1×1=1------它是"上游梯度恰好为 1"造成的,不是"bbb 的导数是 1" 。换个损失 L=(out−y)2L=(out-y)^2L=(out−y)2,dout=2(out−y)dout=2(out-y)dout=2(out−y),db2db_2db2 就不再是 1 了。

局部导数(一段) 损失梯度(整串)
标量 ∂y/∂b=1\partial y/\partial b = 1∂y/∂b=1 ∂L/∂b=dz⋅1=dz\partial L/\partial b = dz\cdot 1 = dz∂L/∂b=dz⋅1=dz
矩阵 ∂Z/∂b=1\partial Z/\partial b = 1∂Z/∂b=1 db=∑ndZn,:⋅1=∑ndZn,:db = \sum_n dZ_{n,:}\cdot 1 = \sum_n dZ_{n,:}db=∑ndZn,:⋅1=∑ndZn,:

一句话:"对 bbb 的偏导是 1"永远成立(局部),但 dbdbdb 是拿这个 1 去乘上游梯度 dZdZdZ、再对样本求和

3.2 疑问二:dbdbdb、dXdXdX 是干嘛的?不是主要去 dWdWdW 吗?

一个 Linear 层有两个参数 (WWW 和 bbb),不是一个!所以反向时这层要攒两个梯度:

  • dWdWdW:更新本层权重 WWW。
  • dbdbdb:更新本层偏置 bbb。
  • dXdXdX:不是参数 ,是一根"接力棒",专门把梯度传给上一层 (成为上一层的 dzdzdz)。

看这条链(blog_backprop.md 第 2 节那套):L→out→a(=tanh⁡z)→z→w1L\to out\to a(=\tanh z)\to z\to w_1L→out→a(=tanhz)→z→w1。反向时逐层往回走,关键在一层怎么"接上"上一层

这层反向算的 是给谁用的
dW, dbdW,\ dbdW, db 本层(更新参数)
dXdXdX 传给上一层 (成为上一层的 dzdzdz)

具体接法:本层的 dXdXdX 经过激活的导数(×(1−a2)\times(1-a^2)×(1−a2))就变成上一层的 dzdzdz。看 engine.py 里那两行反向:

python 复制代码
self.grad += out.grad @ other.data.T     # dX = dZ·Wᵀ  → 传给上一层
other.grad += self.data.T @ out.grad     # dW = Xᵀ·dZ  → 更新本层 W

为什么不能"只做 dWdWdW" :假设 2 层网络 X→Linear1→tanh⁡Linear2→outX\to\\text{Linear}_1\\to\\tanh\to\\text{Linear}_2\to outX→Linear1→tanhLinear2→out。在 Linear2\text{Linear}_2Linear2 算出 dX2dX_2dX2,它穿过 tanh⁡\tanhtanh 变成 Linear1\text{Linear}_1Linear1 的 dZ1dZ_1dZ1;Linear1\text{Linear}_1Linear1 再用这个 dZ1dZ_1dZ1 才能算出 dW1,db1dW_1,db_1dW1,db1。如果只算 dW2dW_2dW2 不算 dX2dX_2dX2,前面一层永远拿不到梯度、学不动

是不是参数 用途
dWdWdW 是(WWW) 更新本层权重
dbdbdb 是(bbb) 更新本层偏置
dXdXdX 不是 传给上一层 (作为它的 dzdzdz)

一句话:一层反向 = 给本层攒两个梯度(dW,dbdW,dbdW,db),再给上一层递一根接力棒(dXdXdX) 。dWdWdW 是"这层怎么改",dXdXdX 是"上一层的 dzdzdz 从哪来"------两者都要,链式法则才能一路传到底。

3.3 疑问三:反向传播是不是要"先全部前向,再统一反向"?

对,必须先整条前向跑完,再统一反向。 原因:反向的顺序是反过来依赖的------要算最前面一层的梯度,得先有它后面所有层传回来的梯度;而这些后面层的梯度,又要在前向真正跑到 loss 之后才知道。所以不能"边前向边反向",也不能"前向一层就反向一层"。

  • 前向一趟:从左往右,每层算一个值、存下来(顺便织成计算图)。
  • 反向一趟 :从 loss 出发,按拓扑序倒着走(输出 → 倒数第二层 → ... → 第一层),每层调用存好的反向函数。

为啥每层要"存值":反向算局部导数时要用到前向的值,而这些值只有前向时才便宜:

反向节点 需要的前向值 原因
tanh⁡\tanhtanh aaa 1−a21-a^21−a2 要用 aaa
softmax PPP dS=P⊙(dP−... )dS=P\odot(dP-\dots)dS=P⊙(dP−...) 要用 PPP
矩阵乘 Z=XWZ=XWZ=XW X, WX,\ WX, W dW=X⊤dZ, dX=dZ W⊤dW=X^\top dZ,\ dX=dZ\,W^\topdW=X⊤dZ, dX=dZW⊤ 要用 X,WX,WX,W

对照 blog_backprop.md 第 9 节的训练循环(顺序严格):

复制代码
pred = model(X)          # 1. 前向:一次跑到底,织图 + 存值
loss = mse_loss(pred,y)  # 2. 算损失
opt.zero_grad()          # 3. 清上次梯度
loss.backward()          # 4. 反向:从 loss 倒着遍历
opt.step()               # 5. 更新参数

看第 1 步和第 4 步:前向(1)先整个跑完,loss(2)也算出来,然后才 backward()backward() 不是"每层前向完就反向",而是等整条链铺好后一次性从后往前扫。以那个 out=w2tanh⁡(w1x+b1)+b2out = w_2\tanh(w_1x+b_1)+b_2out=w2tanh(w1x+b1)+b2 为例,反向顺序是 out→m→a→z→n1out\to m\to a\to z\to n_1out→m→a→z→n1。

一句话:反向传播 = 先把整条前向跑完(存好每层值、织好图),再统一从 loss 往反方向、按拓扑序倒着走一遍。 不能"边前向边反向",因为前面层的梯度要等后面层先算出来才拿得到。

3.4 疑问四:"传给上一层"是什么意思?上一层不是算过了吗?
  • 前向 :每层算出"值"(data),早就算好存起来了。所以上一层的值确实算过了,你说得对。
  • 但反向是另一趟(右→左),算的是"梯度",而上一层的梯度还没算

所以"传给上一层"的意思是:把梯度 dZdZdZ 递给它,好让它也能算出自己的 dW,db,dXdW,db,dXdW,db,dX------它缺的不是值,是梯度。

blog_backprop.md 的数字看一遍(前向值已存好:n1=-0.5, z=0.0, a=0.0, m=0.0, out=0.2):

反向到 算出 传给谁
out = m+b2 dm=1, db2=1 dm 传给 m
m = w2·a da=0.8, dw2=0 da=0.8 传给 a(上一层)
a = tanh(z) da=0.8dz=0.8 dz 传给 z
z = n1+b1 dz=0.8dn1=0.8, db1=0.8 dn1 传给 n1
n1 = w1·x dn1=0.8dw1=0.4 梯度落在参数上

看第 3 行:a 层的值 a=0.0 早就算好了,但它的反向梯度 dz 必须等 mda=0.8 传过来才能算 (要乘 tanh⁡\tanhtanh 的导数 1−a21-a^21−a2)。如果在 m 那里只算 dw2、不算 da,后面就断链,最前层的 dw1 永远是 0------前面学不动

一句话:前向一趟算"值"存好(左→右);反向一趟从后往前算"梯度",并把梯度递给前一层(右→左)。 上一层值早有了,但它缺梯度,而它恰好需要你此刻递过去的那份 dZdZdZ------所以"传给上一层"传的是梯度,不是值。

第二部分 线性层与矩阵乘:把求和写进矩阵乘

4. 把三条公式"展开"看个究竟

矩阵公式 展开成标量 算什么
dW=X⊤dZdW=X^\top dZdW=X⊤dZ (dW)k,o=∑nXn,kdZn,o(dW){k,o}=\sum_n X{n,k}dZ_{n,o}(dW)k,o=∑nXn,kdZn,o 每个权重 = 所有样本的"输入×梯度"求和(共享→求和)
db=∑ndZn,:db=\sum_n dZ_{n,:}db=∑ndZn,: dbo=∑ndZn,odb_o=\sum_n dZ_{n,o}dbo=∑ndZn,o 每个偏置 = 所有样本梯度求和(共享→求和)
dX=dZ W⊤dX=dZ\,W^\topdX=dZW⊤ (dX)n,k=∑odZn,oWk,o(dX){n,k}=\sum_o dZ{n,o}W_{k,o}(dX)n,k=∑odZn,oWk,o 每个输入 = 所有输出的"梯度×权重"求和(多输出→求和)

5. 推广到 O=PVO=PVO=PV(加权求和)的 dP,dVdP,dVdP,dV

前面 Z=XWZ=XWZ=XW 的 dW,dXdW,dXdW,dX 那套推导,原封不动搬到 O=PVO=PVO=PV 上即可------只是这里 PPP 和 VVV 都不是参数、都要往回传梯度 。前向 O=PVO=PVO=PV 逐元素是:

Oi,j=∑kPi,kVk,j,P:(L,L), V:(L,dhead), O:(L,dhead)O_{i,j}=\sum_k P_{i,k}V_{k,j},\qquad P{:}(L,L),\ V{:}(L,d_{head}),\ O{:}(L,d_{head})Oi,j=k∑Pi,kVk,j,P:(L,L), V:(L,dhead), O:(L,dhead)

5.1 推 dPdPdP:Pi,kP_{i,k}Pi,k 影响第 iii 行的所有输出

一个 Pi,kP_{i,k}Pi,k 出现在 OOO 的第 iii 行每个元素 里(每个 Oi,jO_{i,j}Oi,j 都含 Pi,kP_{i,k}Pi,k,系数是 Vk,jV_{k,j}Vk,j)。所以要对输出下标 jjj 求和

∂L∂Pi,k=∑j∂L∂Oi,j⋅∂Oi,j∂Pi,k=∑jdOi,j⋅Vk,j\frac{\partial L}{\partial P_{i,k}}=\sum_j \frac{\partial L}{\partial O_{i,j}}\cdot\frac{\partial O_{i,j}}{\partial P_{i,k}} =\sum_j dO_{i,j}\cdot V_{k,j}∂Pi,k∂L=j∑∂Oi,j∂L⋅∂Pi,k∂Oi,j=j∑dOi,j⋅Vk,j

右边正是矩阵乘 (dO V⊤)i,k(dO\,V^\top)_{i,k}(dOV⊤)i,k------输出维 jjj 被收缩掉

dP=dO V⊤dP = dO\,V^\topdP=dOV⊤

5.2 推 dVdVdV:Vk,jV_{k,j}Vk,j 影响第 jjj 列的所有输出

一个 Vk,jV_{k,j}Vk,j 出现在 OOO 的第 jjj 列每个元素 里(每个 Oi,jO_{i,j}Oi,j 都含 Vk,jV_{k,j}Vk,j,系数是 Pi,kP_{i,k}Pi,k)。所以要对行下标 iii 求和

∂L∂Vk,j=∑i∂L∂Oi,j⋅∂Oi,j∂Vk,j=∑idOi,j⋅Pi,k\frac{\partial L}{\partial V_{k,j}}=\sum_i \frac{\partial L}{\partial O_{i,j}}\cdot\frac{\partial O_{i,j}}{\partial V_{k,j}} =\sum_i dO_{i,j}\cdot P_{i,k}∂Vk,j∂L=i∑∂Oi,j∂L⋅∂Vk,j∂Oi,j=i∑dOi,j⋅Pi,k

右边正是矩阵乘 (P⊤dO)k,j(P^\top dO)_{k,j}(P⊤dO)k,j------行维 iii 被收缩掉

dV=P⊤dOdV = P^\top dOdV=P⊤dO

5.3 和 Z=XWZ=XWZ=XW 的 dW,dXdW,dXdW,dX 对照(一模一样)
前向 对左操作数求导 对右操作数求导
Z=XWZ = XWZ=XW dX=dZ W⊤dX = dZ\,W^\topdX=dZW⊤ dW=X⊤dZdW = X^\top dZdW=X⊤dZ
O=PVO = PVO=PV dP=dO V⊤dP = dO\,V^\topdP=dOV⊤ dV=P⊤dOdV = P^\top dOdV=P⊤dO

规律:对哪个操作数求导,就把另一个操作数放到外面做矩阵乘、转置让被收缩的维对齐 。dPdPdP 用 VVV(dP=dO V⊤dP=dO\,V^\topdP=dOV⊤)、dVdVdV 用 PPP(dV=P⊤dOdV=P^\top dOdV=P⊤dO)------左对左、右对右,各拿对方。

转置的用意和之前一样:把要"求和(收缩)"的那个维度对齐 。dPdPdP 想消掉输出维 jjj,所以 VVV 转置成 V⊤V^\topV⊤ 让 jjj 做内积;dVdVdV 想消掉行维 iii,所以 PPP 转置成 P⊤P^\topP⊤ 让 iii 做内积。

5.4 用数字验证

P=0.0280.9720.0020.998P=\begin{bmatrix}0.028&0.972\\0.002&0.998\end{bmatrix}P=0.0280.0020.9720.998,V=0.521.54V=\begin{bmatrix}0.5&2\\1.5&4\end{bmatrix}V=0.51.524,dO=−1.057000dO=\begin{bmatrix}-1.057&0\\0&0\end{bmatrix}dO=−1.057000

dP=dO V⊤=−1.0570000.51.524=−0.528−1.58500dP = dO\,V^\top=\begin{bmatrix}-1.057&0\\0&0\end{bmatrix}\begin{bmatrix}0.5&1.5\\2&4\end{bmatrix}=\begin{bmatrix}-0.528&-1.585\\0&0\end{bmatrix}dP=dOV⊤=−1.0570000.521.54=−0.5280−1.5850

dV=P⊤dO=0.0280.0020.9720.998−1.057000=−0.0300−1.0270dV = P^\top dO=\begin{bmatrix}0.028&0.002\\0.972&0.998\end{bmatrix}\begin{bmatrix}-1.057&0\\0&0\end{bmatrix}=\begin{bmatrix}-0.030&0\\-1.027&0\end{bmatrix}dV=P⊤dO=0.0280.9720.0020.998−1.057000=−0.030−1.02700

正好对上 blog_qkv.md 第 3 节表格里 dPdPdP、dVdVdV 的结果。

一句话:O=PVO=PVO=PV 的 dP,dVdP,dVdP,dV 和 Z=XWZ=XWZ=XW 的 dX,dWdX,dWdX,dW 是同一套链式法则------每个元素对哪一行/哪一列的输出都有贡献,就沿那个方向求和,求和写成矩阵乘 + 转置。

6. 一句话抓住本质

  • dWdWdW 和 dbdbdb:参数被大家用 → 把所有人的梯度收回来求和(横向:跨样本)。
  • dXdXdX:一个输入喂好几个输出 → 把这几个输出方向的梯度收回来求和(纵向:跨输出)。
  • 转置 不是魔法,它只是"把要消掉(求和)的那个维对齐",好让矩阵乘能顺带完成求和。标量版没有这些求和,是因为它既没有"多个样本"也没有"多个输出"。

用博客那组数字(X(2,3),W(3,1),dZ=1,1⊤X(2,3), W(3,1), dZ=1,1^\topX(2,3),W(3,1),dZ=1,1⊤)一验就对上:dW=X⊤dZ=5,7,9⊤dW=X^\top dZ=5,7,9^\topdW=X⊤dZ=5,7,9⊤、db=1+1=2db=1+1=2db=1+1=2、dX=dZ W⊤=0.5,−1,0.2dX=dZ\,W^\top=0.5,-1,0.2dX=dZW⊤=0.5,−1,0.2(两行相同)。标量手算和矩阵批量,就是同一套链式法则,只不过矩阵把"求和"写进了矩阵乘里。


第三部分 attention 各节点的求导

7. softmax 反向:dS=P⊙(dP−1(P⊤dP))dS = P\odot\big(dP-\mathbf{1}(P^\top dP)\big)dS=P⊙(dP−1(P⊤dP))

为什么不能逐元素 :softmax 是按行归一化 ,pj=esj∑keskp_j=\dfrac{e^{s_j}}{\sum_k e^{s_k}}pj=∑keskesj。分母是整行求和,所以每个输出 pjp_jpj 都依赖同一行所有输入 sks_ksk ------这就是它和 tanh⁡\tanhtanh(1−a21-a^21−a2 逐元素)的根本区别。

先拿一行做 ,丢掉行下标。对 sks_ksk 求导分两种:

∂pj∂sk=pj (δjk−pk)\frac{\partial p_j}{\partial s_k}=p_j\,(\delta_{jk}-p_k)∂sk∂pj=pj(δjk−pk)

这里的 δjk\delta_{jk}δjk 就是 Kronecker delta(克罗内克函数)------一个"两个下标相不相等"的开关:

δjk={1,j=k0,j≠k\delta_{jk}=\begin{cases}1, & j=k\\ 0, & j\neq k\end{cases}δjk={1,0,j=kj=k

把它按 jjj(行)、kkk(列)排开就是单位矩阵

δ=δ11δ12δ21δ22=1001\delta=\begin{bmatrix}\delta_{11}&\delta_{12}\\\delta_{21}&\delta_{22}\end{bmatrix}=\begin{bmatrix}1&0\\0&1\end{bmatrix}δ=δ11δ21δ12δ22=1001

主对角线(j=kj=kj=k)全是 1,其它位置(j≠kj\neq kj=k)全是 0。在求导里它负责"只保留 j=kj=kj=k 那一项,其余一律清零"。所以在 softmax 导数里:

情况 δjk\delta_{jk}δjk ∂pj/∂sk\partial p_j/\partial s_k∂pj/∂sk
j=kj=kj=k(动自己) 111 pj(1−pj)p_j(1-p_j)pj(1−pj)
j≠kj\neq kj=k(动别人) 000 −pjpk-p_jp_k−pjpk

一句话:δjk\delta_{jk}δjk 就是"等不等于"的记号,相等给 1、不等给 0,本质就是单位矩阵(III)------它负责在求和里挑出"自己匹配自己"那一项。

7.1 局部导数 ∂pj∂sk\dfrac{\partial p_j}{\partial s_k}∂sk∂pj 是怎么来的

这个局部导数本身也要用商法则 推,关键是分子里的指数是不是 sks_ksk 。写 pj=esjDp_j=\dfrac{e^{s_j}}{D}pj=Desj,D=∑mesmD=\sum_m e^{s_m}D=∑mesm(分母是整行和)。

对 sks_ksk 求导,分母 DDD 一定含 eske^{s_k}esk ,但分子 esje^{s_j}esj 只有当 j=kj=kj=k 时才含 sks_ksk。所以分两种情况:

情况一:j=kj=kj=k(对"自己"的指数求导) ------分子分母都依赖 sks_ksk,用完整商法则 (uv)′=u′v−uv′v2\left(\frac uv\right)'=\frac{u'v-uv'}{v^2}(vu)′=v2u′v−uv′,且 ∂D/∂sj=esj\partial D/\partial s_j=e^{s_j}∂D/∂sj=esj:

∂pj∂sj=esjD−esj⋅esjD2=esjD−esjesjD2=pj−pj2=pj(1−pj)\frac{\partial p_j}{\partial s_j}=\frac{e^{s_j}D-e^{s_j}\cdot e^{s_j}}{D^2}=\frac{e^{s_j}}{D}-\frac{e^{s_j}e^{s_j}}{D^2}=p_j-p_j^2=p_j(1-p_j)∂sj∂pj=D2esjD−esj⋅esj=Desj−D2esjesj=pj−pj2=pj(1−pj)

情况二:j≠kj\neq kj=k(对"别人"的指数求导) ------分子 esje^{s_j}esj 与 sks_ksk 无关(当常数,导数为 0),只有分母影响:

∂pj∂sk=0⋅D−esj⋅eskD2=−esjeskD2=−pjpk\frac{\partial p_j}{\partial s_k}=\frac{0\cdot D-e^{s_j}\cdot e^{s_k}}{D^2}=-\frac{e^{s_j}e^{s_k}}{D^2}=-p_jp_k∂sk∂pj=D20⋅D−esj⋅esk=−D2esjesk=−pjpk

合并成一个式子 :看 δjk\delta_{jk}δjk 的开关作用,两种情况恰好能统一:

pj(δjk−pk)={pj(1−pj),j=kpj(0−pk)=−pjpk,j≠kp_j(\delta_{jk}-p_k)=\begin{cases}p_j(1-p_j), & j=k\\ p_j(0-p_k)=-p_jp_k, & j\neq k\end{cases}pj(δjk−pk)={pj(1−pj),pj(0−pk)=−pjpk,j=kj=k

所以 ∂pj∂sk=pj(δjk−pk)\dfrac{\partial p_j}{\partial s_k}=p_j(\delta_{jk}-p_k)∂sk∂pj=pj(δjk−pk)。

用数字验证 (第 1 行,p1=0.028, p2=0.972p_1=0.028,\ p_2=0.972p1=0.028, p2=0.972):

j=k: ∂p1∂s1=p1(1−p1)=0.028×0.972=0.0272,j≠k: ∂p2∂s1=−p2p1=−0.972×0.028=−0.0272j=k:\ \frac{\partial p_1}{\partial s_1}=p_1(1-p_1)=0.028\times0.972=0.0272,\qquad j\neq k:\ \frac{\partial p_2}{\partial s_1}=-p_2p_1=-0.972\times0.028=-0.0272j=k: ∂s1∂p1=p1(1−p1)=0.028×0.972=0.0272,j=k: ∂s1∂p2=−p2p1=−0.972×0.028=−0.0272

直觉:sks_ksk 变大时分母 DDD 一定变大 (eske^{s_k}esk 变大),每个 pjp_jpj 都会被等比压低。若是自己(j=kj=kj=k)分子也涨得更凶,净效果是 pj(1−pj)p_j(1-p_j)pj(1−pj);若是别人,只有分母涨、无补偿,净效果是 −pjpk-p_jp_k−pjpk。δjk\delta_{jk}δjk 就是"这次动的是不是我自己"的开关。

反向:把这行所有输出对 sks_ksk 的贡献加起来

dsk=∑jdPj ∂pj∂sk=∑jdPj pj(δjk−pk)=pk dPk−pk∑jdPj pjds_k=\sum_j dP_j\,\frac{\partial p_j}{\partial s_k} =\sum_j dP_j\,p_j(\delta_{jk}-p_k) =p_k\,dP_k-p_k\sum_j dP_j\,p_jdsk=j∑dPj∂sk∂pj=j∑dPjpj(δjk−pk)=pkdPk−pkj∑dPjpj

7.2 逐步拆解:三个等号分别干嘛

上面这三个等号把三件事压在一行里,容易卡在第三个。逐个拆开:

第 ① 步:为什么开头要 ∑j\sum_j∑j

sks_ksk 不只影响 pkp_kpk 自己,它同时影响这一行的每一个 pjp_jpj (因为分母 ∑kesk\sum_k e^{s_k}∑kesk 是整行共享的)。所以"sks_ksk 动一点,损失动多少"要把每一路 pjp_jpj 的贡献都收回来------链式法则的"多路径求和":

dsk=∑jdPj⏟第 j 路的梯度×∂pj∂sk⏟第 j 路的局部斜率ds_k=\sum_j \underbrace{dP_j}{\text{第 }j\text{ 路的梯度}}\times\underbrace{\frac{\partial p_j}{\partial s_k}}{\text{第 }j\text{ 路的局部斜率}}dsk=j∑第 j 路的梯度 dPj×第 j 路的局部斜率 ∂sk∂pj

第 ② 步:把局部导数代入

∂pj∂sk=pj(δjk−pk)\dfrac{\partial p_j}{\partial s_k}=p_j(\delta_{jk}-p_k)∂sk∂pj=pj(δjk−pk),代入得 ∑jdPj pj(δjk−pk)\sum_j dP_j\,p_j(\delta_{jk}-p_k)∑jdPjpj(δjk−pk)。

第 ③ 步:把 ∑j\sum_j∑j 拆成两半(最容易卡在这)

因为 (δjk−pk)(\delta_{jk}-p_k)(δjk−pk) 是两项相减,可以把求和拆开:

∑jdPjpj(δjk−pk)=∑jdPjpj δjk⏟A−∑jdPjpj pk⏟B\sum_j dP_j p_j(\delta_{jk}-p_k)=\underbrace{\sum_j dP_j p_j\,\delta_{jk}}{\text{A}}-\underbrace{\sum_j dP_j p_j\,p_k}{\text{B}}j∑dPjpj(δjk−pk)=A j∑dPjpjδjk−B j∑dPjpjpk

  • 看 A :δjk\delta_{jk}δjk 是个"开关",只有 j=kj=kj=k 时等于 1,其余全是 0 。所以在 ∑j\sum_j∑j 里除了 j=kj=kj=k 那一项,其它全消失:∑jdPjpjδjk=dPk pk⋅1=pk dPk\sum_j dP_jp_j\delta_{jk}=dP_k\,p_k\cdot1=p_k\,dP_k∑jdPjpjδjk=dPkpk⋅1=pkdPk。
  • 看 B :pkp_kpk 不随 jjj 变,是常数,可以提出来:∑jdPjpj pk=pk∑jdPjpj\sum_j dP_jp_j\,p_k=p_k\sum_j dP_jp_j∑jdPjpjpk=pk∑jdPjpj。

合起来:pk dPk−pk∑jdPjpjp_k\,dP_k-p_k\sum_j dP_jp_jpkdPk−pk∑jdPjpj。

7.3 慢算一遍验证(第 1 行,k=1)

s=1.77,5.30s=1.77,5.30s=1.77,5.30,P=0.028,0.972P=0.028,0.972P=0.028,0.972,dP=−0.528,−1.585dP=-0.528,-1.585dP=−0.528,−1.585

先算两个局部导数 (j=1j=1j=1、j=2j=2j=2 分别对 s1s_1s1):

∂p1∂s1=p1(1−p1)=0.028×0.972=0.0272,∂p2∂s1=p2(0−p1)=0.972×(−0.028)=−0.0272\frac{\partial p_1}{\partial s_1}=p_1(1-p_1)=0.028\times0.972=0.0272,\qquad \frac{\partial p_2}{\partial s_1}=p_2(0-p_1)=0.972\times(-0.028)=-0.0272∂s1∂p1=p1(1−p1)=0.028×0.972=0.0272,∂s1∂p2=p2(0−p1)=0.972×(−0.028)=−0.0272

按 ① 求和

ds1=dP1∂p1∂s1+dP2∂p2∂s1=(−0.528)(0.0272)+(−1.585)(−0.0272)=−0.01436+0.04311≈0.029ds_1=dP_1\frac{\partial p_1}{\partial s_1}+dP_2\frac{\partial p_2}{\partial s_1}=(-0.528)(0.0272)+(-1.585)(-0.0272)=-0.01436+0.04311\approx0.029ds1=dP1∂s1∂p1+dP2∂s1∂p2=(−0.528)(0.0272)+(−1.585)(−0.0272)=−0.01436+0.04311≈0.029

按 ③ 的快捷式核对 (先算行总账 ∑jdPjpj=0.028(−0.528)+0.972(−1.585)≈−1.555\sum_j dP_jp_j=0.028(-0.528)+0.972(-1.585)\approx-1.555∑jdPjpj=0.028(−0.528)+0.972(−1.585)≈−1.555):

ds1=p1 dP1−p1∑jdPjpj=0.028(−0.528)−0.028(−1.555)≈0.029ds_1=p_1\,dP_1-p_1\sum_j dP_jp_j=0.028(-0.528)-0.028(-1.555)\approx0.029ds1=p1dP1−p1j∑dPjpj=0.028(−0.528)−0.028(−1.555)≈0.029

两条路都得到 ≈0.029\approx0.029≈0.029 ✓

记忆点:∑j\sum_j∑j 不是多余,是 sks_ksk 影响了整行 pjp_jpj ,得把整行梯度都收回来;δjk\delta_{jk}δjk 就是个开关 ,只在 j=kj=kj=k 时让 dPkpkdP_kp_kdPkpk 留下,其余全关掉;而 −pk-p_k−pk 那项跟 jjj 无关,直接提出来乘以整行总账。

写成整行(再恢复行下标)就是上面的公式:

dS=P⊙(dP−1(P⊤dP))⏟每行:dPk−∑jPjdPjdS = P\odot\underbrace{\big(dP-\mathbf{1}(P^\top dP)\big)}_{\text{每行:}dP_k-\sum_j P_j dP_j}dS=P⊙每行:dPk−∑jPjdPj (dP−1(P⊤dP))

两部分的含义:

  • P⊙dPP\odot dPP⊙dP :增大 sks_ksk 会直接抬高 pkp_kpk(自项);
  • −P⊙1(P⊤dP)-P\odot\mathbf{1}(P^\top dP)−P⊙1(P⊤dP) :增大 sks_ksk 会撑大分母、压扁同行的其它 pjp_jpj (竞争项)。P⊤dPP^\top dPP⊤dP 一次性算整行的"总梯度",1\mathbf{1}1 广播回这行每个位置,再按各自 PPP 摊回去。

一句直觉:softmax 的输出互相竞争(加起来=1),推高一个必然挤低别的;反向时既要算"自己涨跌",还要结清"挤了谁"。

验证 :P=0.028,0.972P=0.028,0.972P=0.028,0.972,dP=−0.528,−1.585dP=-0.528,-1.585dP=−0.528,−1.585,P⊤dP=0.028(−0.528)+0.972(−1.585)≈−1.555P^\top dP = 0.028(-0.528)+0.972(-1.585)\approx-1.555P⊤dP=0.028(−0.528)+0.972(−1.585)≈−1.555,于是

dP−1(P⊤dP)=−0.528+1.555, −1.585+1.555=1.027,−0.030dP-\mathbf{1}(P^\top dP)=-0.528+1.555,\\ -1.585+1.555=1.027,-0.030dP−1(P⊤dP)=−0.528+1.555, −1.585+1.555=1.027,−0.030

dS=P⊙1.027,−0.030=0.029,−0.029 ✓dS=P\odot1.027,-0.030=0.029,-0.029\ \checkmarkdS=P⊙1.027,−0.030=0.029,−0.029


8. 缩放内积 S=QK⊤/dS=QK^\top/\sqrt{d}S=QK⊤/d 的 dQ,dKdQ,dKdQ,dK

这其实就是 O=PVO=PVO=PV 那套矩阵乘 ,只是右边多了转置和缩放。分两步:先把缩放提出来,再算 QK⊤QK^\topQK⊤

原式 S=QK⊤dS=\dfrac{QK^\top}{\sqrt{d}}S=d QK⊤。设 M=QK⊤M=QK^\topM=QK⊤(不打分原始矩阵),则 S=M/dS=M/\sqrt{d}S=M/d 。由链式法则:

dM=∂L∂M=dSddM=\frac{\partial L}{\partial M}=\frac{dS}{\sqrt{d}}dM=∂M∂L=d dS

缩放是常数倍,只把梯度缩小 d\sqrt{d}d 倍 传进矩阵乘节点,不影响结构。后面就把 dMdMdM 当作矩阵乘 M=QK⊤M=QK^\topM=QK⊤ 的上游梯度。

8.1 M=QK⊤M=QK^\topM=QK⊤ 逐元素

Mij=∑kQik KjkQ:(L,d), K:(L,d), M:(L,L)M_{ij}=\sum_k Q_{ik}\,K_{jk}\qquad Q{:}(L,d),\ K{:}(L,d),\ M{:}(L,L)Mij=k∑QikKjkQ:(L,d), K:(L,d), M:(L,L)

8.2 推 dQ=dM KdQ = dM\,KdQ=dMK

QikQ_{ik}Qik 出现在 MMM 的第 iii 行所有元素 里(系数是 KjkK_{jk}Kjk),所以对 jjj 求和

∂L∂Qik=∑jdMij ∂Mij∂Qik=∑jdMij Kjk\frac{\partial L}{\partial Q_{ik}}=\sum_j dM_{ij}\,\frac{\partial M_{ij}}{\partial Q_{ik}}=\sum_j dM_{ij}\,K_{jk}∂Qik∂L=j∑dMij∂Qik∂Mij=j∑dMijKjk

右边就是 (dM K)ik(dM\,K)_{ik}(dMK)ik。代回 dM=dS/ddM=dS/\sqrt{d}dM=dS/d :

dQ=(dS/d) K\boxed{dQ = (dS/\sqrt{d})\,K}dQ=(dS/d )K

8.3 推 dK=(dS/d)⊤QdK = (dS/\sqrt{d})^\top QdK=(dS/d )⊤Q

KjkK_{jk}Kjk 出现在 MMM 的第 jjj 列所有元素 里(系数是 QikQ_{ik}Qik),所以对 iii 求和

∂L∂Kjk=∑idMij ∂Mij∂Kjk=∑idMij Qik\frac{\partial L}{\partial K_{jk}}=\sum_i dM_{ij}\,\frac{\partial M_{ij}}{\partial K_{jk}}=\sum_i dM_{ij}\,Q_{ik}∂Kjk∂L=i∑dMij∂Kjk∂Mij=i∑dMijQik

右边就是 (dM⊤Q)jk(dM^\top Q)_{jk}(dM⊤Q)jk。代回:

dK=(dS/d)⊤Q\boxed{dK = (dS/\sqrt{d})^\top Q}dK=(dS/d )⊤Q

规律和 O=PVO=PVO=PV 一致:对哪个操作数求导,就用另一个做矩阵乘、转置对齐要收缩的维 。dQdQdQ 用 KKK、dKdKdK 用 QQQ,各拿对方。

8.4 用第 2 节数字验证

dS=0.029−0.02900dS=\begin{bmatrix}0.029&-0.029\\0&0\end{bmatrix}dS=0.0290−0.0290,d=2\sqrt d=\sqrt2d =2 ,故 dM=dS/2=0.0205−0.020500dM=dS/\sqrt2=\begin{bmatrix}0.0205&-0.0205\\0&0\end{bmatrix}dM=dS/2 =0.02050−0.02050

Q=2.505.5−1Q=\begin{bmatrix}2.5&0\\5.5&-1\end{bmatrix}Q=2.55.50−1,K=1234K=\begin{bmatrix}1&2\\3&4\end{bmatrix}K=1324

dQ=dM K=0.0205−0.0205001234=−0.041−0.04100 ✓dQ=dM\,K=\begin{bmatrix}0.0205&-0.0205\\0&0\end{bmatrix}\begin{bmatrix}1&2\\3&4\end{bmatrix}=\begin{bmatrix}-0.041&-0.041\\0&0\end{bmatrix}\ \checkmarkdQ=dMK=0.02050−0.020501324=−0.0410−0.0410

dK=dM⊤Q=0.02050−0.020502.505.5−1=0.0510−0.0510 ✓dK=dM^\top Q=\begin{bmatrix}0.0205&0\\-0.0205&0\end{bmatrix}\begin{bmatrix}2.5&0\\5.5&-1\end{bmatrix}=\begin{bmatrix}0.051&0\\-0.051&0\end{bmatrix}\ \checkmarkdK=dM⊤Q=0.0205−0.0205002.55.50−1=0.051−0.05100

正好对上博客第 3 节表格。


9. 串起来看

attention 这一整段的反向,就是这几个"查表公式"按拓扑序倒着走:

text 复制代码
dO → [加权求和 O=PV] → dP, dV
                        ↓
    [softmax] → dS = P⊙(dP−1(PᵀdP))
                        ↓
    [QKᵀ/√d] → dQ = (dS/√d)K,dK = (dS/√d)ᵀQ

每个节点都只干"沿哪个维求和,就转置对齐哪个维"这一件事,没有魔法。


10. softmax 前向 vs 反向(整行耦合)

核心一句话:前向里每个 PjP_jPj 依赖整行的所有 sss;反向里每个 dSkdS_kdSk 也依赖整行的所有 dPdPdP。 两者都是"整行耦合",因为分母是整行求和。

10.1 一张表:前向 vs 反向
前向 S→PS\to PS→P 反向 dP→dSdP\to dSdP→dS
公式 Pj=esj∑keskP_j=\dfrac{e^{s_j}}{\sum_k e^{s_k}}Pj=∑keskesj dSk=Pk(dPk−∑jPjdPj)dS_k=P_k\big(dP_k-\sum_j P_j dP_j\big)dSk=Pk(dPk−∑jPjdPj)
每个输出依赖 整行所有 s1,s2,...s_1,s_2,\dotss1,s2,...(分母是整行和) 整行所有 dP1,dP2,...dP_1,dP_2,\dotsdP1,dP2,...(行总梯度包含它们)
是"逐元素"吗 不是 不是
为什么 esje^{s_j}esj 的分子 + ∑kesk\sum_k e^{s_k}∑kesk 的分母同时含多个 sss 每个 PjP_jPj 都含 sks_ksk,所以都往 dSkdS_kdSk 回传

前向"每个输出要看全行",反向就"每个梯度也要收全行"------这是同一件事(共享分母)的两副面孔。

10.2 用数字走一遍(第 1 行)

前向 :输入 s=1.77, 5.30s=1.77,\\ 5.30s=1.77, 5.30

步骤
指数 e1.77=5.87,e5.30=200.3e^{1.77}=5.87,\quad e^{5.30}=200.3e1.77=5.87,e5.30=200.3
分母(整行和 5.87+200.3=206.25.87+200.3=206.25.87+200.3=206.2
输出 P1=5.87206.2≈0.028,P2=200.3206.2≈0.972P_1=\dfrac{5.87}{206.2}\approx0.028,\quad P_2=\dfrac{200.3}{206.2}\approx0.972P1=206.25.87≈0.028,P2=206.2200.3≈0.972

注意:P1P_1P1 不仅取决于 s1s_1s1,也取决于 s2s_2s2 ------因为分母里有 es2e^{s_2}es2。这就是"整行耦合"。

反向 :上游梯度 dP=−0.528, −1.585dP=-0.528,\\ -1.585dP=−0.528, −1.585

步骤
行总梯度 ∑jPjdPj\sum_j P_jdP_j∑jPjdPj 0.028(−0.528)+0.972(−1.585)≈−1.5550.028(-0.528)+0.972(-1.585)\approx-1.5550.028(−0.528)+0.972(−1.585)≈−1.555
dS1dS_1dS1 P1(dP1−总梯度)=0.028(−0.528+1.555)≈0.029P_1(dP_1-\text{总梯度})=0.028(-0.528+1.555)\approx0.029P1(dP1−总梯度)=0.028(−0.528+1.555)≈0.029
dS2dS_2dS2 P2(dP2−总梯度)=0.972(−1.585+1.555)≈−0.029P_2(dP_2-\text{总梯度})=0.972(-1.585+1.555)\approx-0.029P2(dP2−总梯度)=0.972(−1.585+1.555)≈−0.029

结果 dS=+0.029, −0.029dS=+0.029,\\ -0.029dS=+0.029, −0.029

注意:dS1dS_1dS1 不仅来自 dP1dP_1dP1,也来自 dP2dP_2dP2 ------因为"行总梯度"把 dP2dP_2dP2 也收了进来。这正是反向里"整行耦合"的体现。

10.3 为什么反向是 Pk(dPk−∑jPjdPj)P_k(dP_k-\sum_jP_jdP_j)Pk(dPk−∑jPjdPj) 这个形状

拆成两项看:

dSk=Pk dPk⏟自项−Pk∑jPjdPj⏟竞争项dS_k=\underbrace{P_k\,dP_k}{\text{自项}}-\underbrace{P_k\sum_j P_j dP_j}{\text{竞争项}}dSk=自项 PkdPk−竞争项 Pkj∑PjdPj

  • 自项 PkdPkP_k dP_kPkdPk:sks_ksk 变大 → PkP_kPk 自己变大(来自 ∂Pk/∂sk=Pk(1−Pk)\partial P_k/\partial s_k=P_k(1-P_k)∂Pk/∂sk=Pk(1−Pk) 里那部分 Pk(1)P_k(1)Pk(1))。
  • 竞争项 −Pk∑jPjdPj-P_k\sum_j P_j dP_j−Pk∑jPjdPj:sks_ksk 变大 → 分母变大 → 同行的其它 PjP_jPj 都被压低 ,要把这些被挤掉的梯度收回来。∑jPjdPj\sum_j P_j dP_j∑jPjdPj 是"整行所有输出被挤压的总账",1\mathbf{1}1 广播回每个位置,再乘 PkP_kPk 按比例摊。
10.4 和前向对照的对称美感
前向 反向
动作 "先指数、再按行归一化"(一个 PjP_jPj 吃进全行 sss) "先算行总账、再按 PPP 摊回去"(一个 dSkdS_kdSk 收进全行 dPdPdP)
共同点 分母是整行和 → 整行耦合 行总梯度是整行和 → 整行耦合

一句话:softmax 前向把"一个输入"摊到整行输出;反向就把"整行梯度"收回成一个输入。 因为它按行归一化,所以前向、反向都是"整行一起算",永远不可能像 tanh⁡\tanhtanh 那样逐元素。

相关推荐
离凌寒35 分钟前
一、关于st上制作外部烧录算法时软件识别不到算法文件的问题总结
算法
青 春 记 忆2 小时前
LeetCode 283. 移动零|Python 解法详解
python·算法·leetcode
Hrain-AI2 小时前
企业 AI 治理运营怎么做:分级授权、Token 用量可观测与模型统一纳管
大数据·人工智能·算法
剑指offer.3 小时前
Linux多任务-线程篇
java·jvm·算法
-dzk-3 小时前
【矩阵】LC 54.螺旋矩阵
线性代数·矩阵
nike0good3 小时前
CF 126B(Password-z algorithm/exkmp)
开发语言·c++·算法
无定义_3 小时前
Bellman-Ford——贝尔曼福特算法
数据库·算法
HanhahnaH4 小时前
各数据结构操作的时间复杂度汇总
数据结构·算法
Yzzz-F4 小时前
CF2023D
c++·算法·dp