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

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

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

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

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

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

1.1 表 1:局部导数(只问"output 对 input 的斜率",还没乘上游梯度 dzdzdz
标量(1 样本 × 1 输出)矩阵(N 样本 × 多输出)
前向y=wx+by = wx + by=wx+bZ=XW+bZ = XW + bZ=XW+bX:(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}=xwy=x∂Z∂W=X\dfrac{\partial Z}{\partial W}=XWZ=X
对偏置∂y∂b=1\dfrac{\partial y}{\partial b}=1by=1∂Z∂b=1\dfrac{\partial Z}{\partial b}=1bZ=1
对输入∂y∂x=w\dfrac{\partial y}{\partial x}=wxy=w∂Z∂X=W\dfrac{\partial Z}{\partial X}=WXZ=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_nwL=ndznxndW=X⊤dZdW = X^\top dZdW=XdZ
对偏置∂L∂b=∑ndzn\dfrac{\partial L}{\partial b}=\sum_n dz_nbL=ndzndb=∑n=1NdZn,:db = \sum\limits_{n=1}^{N} dZ_{n,:}db=n=1NdZn,:
对输入∂L∂xn=dzn⋅w\dfrac{\partial L}{\partial x_n}=dz_n\cdot wxnL=dznwdX=dZ W⊤dX = dZ\,W^\topdX=dZW

dW=X⊤dZdW = X^\top dZdW=XdZ 的真正对应物是 ∂L/∂w=∑ndznxn\partial L/\partial w=\sum_n dz_n x_nL/w=ndznxn,不是 ∂y/∂w=x\partial y/\partial w=xy/w=x 后者只是"一段斜率",前者是"整串链式法则"。所以 dWdWdW 里那个 dzdzdz 不是凭空多出来的——它是把损失一路链到 WWW 时,前面那段已经传回来的梯度,带着它一起乘才叫"损失偏导"。

再次强调:下面两处为什么多出"求和",是因为表 2 才涉及"多个样本共享参数"。

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

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

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

偏导 ∂L/∂w\partial L/\partial wL/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_nwL=nznLwzn=ndznxn

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

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

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

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

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

这条公式值得单独拆开。它其实在说一句话: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 条:

        ┌─> 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(znyn)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}wL=n=1N znLwzn

每一项拆开看:

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

③ 拿 2 个样本走一遍:设 x1=1, x2=3x_1=1,\ x_2=3x1=1, x2=3w=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=(z1y1)2+(z2y2)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=(23)2+(67)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(z1y1)=2(23)=2,dz2=2(z2y2)=2(67)=2

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

路径dzndz_ndzn∂zn/∂w=xn\partial z_n/\partial w=x_nzn/w=xn该路径贡献
w→z1→Lw\to z_1\to Lwz1L−2-22111−2×1=−2-2\times1=-22×1=2
w→z2→Lw\to z_2\to Lwz2L−2-22333−2×3=−6-2\times3=-62×3=6

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

④ 直接对 LLL 求导来验证L=(w−3)2+(3w−7)2L=(w-3)^2+(3w-7)^2L=(w3)2+(3w7)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)\cdot3wL=2(w3)+2(3w7)3

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

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

一句话:这条公式 = 链式法则的多路径求和版∑n\sum_nn 数的是"wwwLLL 有几条路"——因为参数被 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,kL=oZn,oLXn,kZn,o=odZn,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 = 1Z/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,o1=ndZn,o。那个 1 还在,但它被上游梯度 dZdZdZ 乘了,还要对所有样本求和(因为 bbbNNN 个样本共享)。

用数字看: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=ndZn,:=1+1=2(不是 1)

那为什么 blog_backprop.mddb2db_2db2 算出的是 1?因为那个例子里传到 out 的梯度正好是 1dout=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=(outy)2dout=2(out−y)dout=2(out-y)dout=2(outy)db2db_2db2 就不再是 1 了。

局部导数(一段)损失梯度(整串)
标量∂y/∂b=1\partial y/\partial b = 1y/b=1∂L/∂b=dz⋅1=dz\partial L/\partial b = dz\cdot 1 = dzL/b=dz1=dz
矩阵∂Z/∂b=1\partial Z/\partial b = 1Z/b=1db=∑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 疑问二:dbdbdbdXdXdX 是干嘛的?不是主要去 dWdWdW 吗?

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

  • 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_1Louta(=tanhz)zw1。反向时逐层往回走,关键在一层怎么"接上"上一层

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

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

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[Linear1tanh][Linear2]out。在 Linear2\text{Linear}_2Linear2 算出 dX2dX_2dX2,它穿过 tanh⁡\tanhtanh 变成 Linear1\text{Linear}_1Linear1dZ1dZ_1dZ1Linear1\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),再给上一层递一根接力棒(dXdXdXdWdWdW 是"这层怎么改",dXdXdX 是"上一层的 dzdzdz 从哪来"——两者都要,链式法则才能一路传到底。

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

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

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

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

反向节点需要的前向值原因
tanh⁡\tanhtanhaaa1−a21-a^21a2 要用 aaa
softmaxPPPdS=P⊙(dP−… )dS=P\odot(dP-\dots)dS=P(dP) 要用 PPP
矩阵乘 Z=XWZ=XWZ=XWX, WX,\ WX, WdW=X⊤dZ, dX=dZ W⊤dW=X^\top dZ,\ dX=dZ\,W^\topdW=XdZ, 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_1outmazn1

一句话:反向传播 = 先把整条前向跑完(存好每层值、织好图),再统一从 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+b2dm=1, db2=1dm 传给 m
m = w2·ada=0.8, dw2=0da=0.8 传给 a(上一层)
a = tanh(z)da=0.8dz=0.8dz 传给 z
z = n1+b1dz=0.8dn1=0.8, db1=0.8dn1 传给 n1
n1 = w1·xdn1=0.8dw1=0.4梯度落在参数上

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

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

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

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

矩阵公式展开成标量算什么
dW=X⊤dZdW=X^\top dZdW=XdZ(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=XWdW,dXdW,dXdW,dX 那套推导,原封不动搬到 O=PVO=PVO=PV 上即可——只是这里 PPPVVV 都不是参数、都要往回传梯度。前向 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=kPi,kVk,j,P:(L,L), V:(L,dhead), O:(L,dhead)

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

一个 Pi,kP_{i,k}Pi,k 出现在 OOOiii 行每个元素里(每个 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,kL=jOi,jLPi,kOi,j=jdOi,jVk,j

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

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

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

一个 Vk,jV_{k,j}Vk,j 出现在 OOOjjj 列每个元素里(每个 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,jL=iOi,jLVk,jOi,j=idOi,jPi,k

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

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

5.3 和 Z=XWZ=XWZ=XWdW,dXdW,dXdW,dX 对照(一模一样)
前向对左操作数求导对右操作数求导
Z=XWZ = XWZ=XWdX=dZ W⊤dX = dZ\,W^\topdX=dZWdW=X⊤dZdW = X^\top dZdW=XdZ
O=PVO = PVO=PVdP=dO V⊤dP = dO\,V^\topdP=dOVdV=P⊤dOdV = P^\top dOdV=PdO

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

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

5.4 用数字验证

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

dP=dO V⊤=[−1.057000][0.51.524]=[−0.528−1.58500]dP = 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.057000][0.521.54]=[0.52801.5850]

dV=P⊤dO=[0.0280.0020.9720.998][−1.057000]=[−0.0300−1.0270]dV = 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=PdO=[0.0280.9720.0020.998][1.057000]=[0.0301.02700]

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

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

6. 一句话抓住本质

  • dWdWdWdbdbdb参数被大家用 → 把所有人的梯度收回来求和(横向:跨样本)。
  • 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=XdZ=[5,7,9]db=1+1=2db=1+1=2db=1+1=2dX=dZ W⊤=[0.5,−1,0.2]dX=dZ\,W^\top=[0.5,-1,0.2]dX=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(dP1(PdP))

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

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

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

这里的 δ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_kpj/sk
j=kj=kj=k(动自己)111pj(1−pj)p_j(1-p_j)pj(1pj)
j≠kj\neq kj=k(动别人)000−pjpk-p_jp_kpjpk

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

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

这个局部导数本身也要用商法则推,关键是分子里的指数是不是 sks_ksk。写 pj=esjDp_j=\dfrac{e^{s_j}}{D}pj=DesjD=∑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)=v2uvuv,且 ∂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)sjpj=D2esjDesjesj=DesjD2esjesj=pjpj2=pj(1pj)

情况二:j≠kj\neq kj=k(对"别人"的指数求导)——分子 esje^{s_j}esjsks_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_kskpj=D20Desjesk=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(δjkpk)={pj(1pj),pj(0pk)=pjpk,j=kj=k

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

用数字验证(第 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: s1p1=p1(1p1)=0.028×0.972=0.0272,j=k: s1p2=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(1pj);若是别人,只有分母涨、无补偿,净效果是 −pjpk-p_jp_kpjpkδ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=jdPjskpj=jdPjpj(δjkpk)=pkdPkpkjdPjpj

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

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

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

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 路的局部斜率skpj

第 ② 步:把局部导数代入

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

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

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

∑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}}jdPjpj(δjkpk)=AjdPjpjδjkBjdPjpjpk

  • 看 Aδjk\delta_{jk}δjk 是个"开关",只有 j=kj=kj=k 时等于 1,其余全是 0。所以在 ∑j\sum_jj 里除了 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_kjdPjpjδjk=dPkpk1=pkdPk
  • 看 Bpkp_kpk 不随 jjj 变,是常数,可以提出来:∑jdPjpj pk=pk∑jdPjpj\sum_j dP_jp_j\,p_k=p_k\sum_j dP_jp_jjdPjpjpk=pkjdPjpj

合起来:pk dPk−pk∑jdPjpjp_k\,dP_k-p_k\sum_j dP_jp_jpkdPkpkjdPjpj

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

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

先算两个局部导数j=1j=1j=1j=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.0272s1p1=p1(1p1)=0.028×0.972=0.0272,s1p2=p2(0p1)=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=dP1s1p1+dP2s1p2=(0.528)(0.0272)+(1.585)(0.0272)=0.01436+0.043110.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.555jdPjpj=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=p1dP1p1jdPjpj=0.028(0.528)0.028(1.555)0.029

两条路都得到 ≈0.029\approx0.0290.029

记忆点:∑j\sum_jj 不是多余,是 sks_ksk 影响了整行 pjp_jpj,得把整行梯度都收回来;δjk\delta_{jk}δjk 就是个开关,只在 j=kj=kj=k 时让 dPkpkdP_kp_kdPkpk 留下,其余全关掉;而 −pk-p_kpk 那项跟 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每行:dPkjPjdPj(dP1(PdP))

两部分的含义:

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

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

验证P=[0.028,0.972]P=[0.028,0.972]P=[0.028,0.972]dP=[−0.528,−1.585]dP=[-0.528,-1.585]dP=[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.555PdP=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.030]dP-\mathbf{1}(P^\top dP)=[-0.528+1.555,\ -1.585+1.555]=[1.027,-0.030]dP1(PdP)=[0.528+1.555, 1.585+1.555]=[1.027,0.030]

dS=P⊙[1.027,−0.030]=[0.029,−0.029] ✓dS=P\odot[1.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/ddQ,dKdQ,dKdQ,dK

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

原式 S=QK⊤dS=\dfrac{QK^\top}{\sqrt{d}}S=dQK。设 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=ML=ddS

缩放是常数倍,只把梯度缩小 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=kQikKjkQ:(L,d), K:(L,d), M:(L,L)

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

QikQ_{ik}Qik 出现在 MMMiii 行所有元素里(系数是 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}QikL=jdMijQikMij=jdMijKjk

右边就是 (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 出现在 MMMjjj 列所有元素里(系数是 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}KjkL=idMijKjkMij=idMijQik

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

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

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

8.4 用第 2 节数字验证

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

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

dQ=dM K=[0.0205−0.020500][1234]=[−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.020500.02050][1324]=[0.04100.0410] 

dK=dM⊤Q=[0.02050−0.02050][2.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=dMQ=[0.02050.020500][2.55.501]=[0.0510.05100] 

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


9. 串起来看

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

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 PSP反向 dP→dSdP\to dSdPdS
公式Pj=esj∑keskP_j=\dfrac{e^{s_j}}{\sum_k e^{s_k}}Pj=keskesjdSk=Pk(dPk−∑jPjdPj)dS_k=P_k\big(dP_k-\sum_j P_j dP_j\big)dSk=Pk(dPkjPjdPj)
每个输出依赖整行所有 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.30]s=[1.77,\ 5.30]s=[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.870.028,P2=206.2200.30.972

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

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

步骤
行总梯度 ∑jPjdPj\sum_j P_jdP_jjPjdPj0.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_1dS1P1(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_2dS2P2(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.029]dS=[+0.029,\ -0.029]dS=[+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(dPkjPjdPj) 这个形状

拆成两项看:

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竞争项PkjPjdPj

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

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

评论
成就一亿技术人!
拼手气红包6.0元
还能输入1000个字符
 
 条评论被折叠 查看
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值