crossentropy=−log∑i=1neZi−max(Zi)eZtarget−max(Zi)
这里我们关注Zi变化对Loss的影响,可以看出,当i不等于target时,Zi只作用于分母
反之则同时作用于分子分母,导数为作用于分子和分母的导数之和
当 i=target
令 L=g1(x1)=−log(x1)
令 x1=g2(x2)=x2c1
c1=eZtarget−max(Zi)
令 x2=g3(x3)=x3+c2
c2为常量
令 x3=g4(x4)=ex4
令 x4=g5(Zi)=Zi−max(Zi)
令 sum=∑i=1neZi−max(Zi)
故 ∂Zi∂L=∂x1∂g1(x1)∂x2∂g2(x2)∂x3∂g3(x3)∂x4∂g4(x4)∂Zi∂g5(Zi)
∂x1∂g1(x1)=−x11
∂x2∂g2(x2)=−x22c1
∂x3∂g3(x3)=1
∂x4∂g4(x4)=ex4
∂Zi∂g5(Zi)=1
x1=sumc1
x2=sum
x4=Zi−max(Zi)
故
∂Zi∂L=x1x22c1ex4=sumeZi−max(Zi)
当 i=target
分母部分的导数同上
sumeZtarget−max(Zi)
下面计算分子部分p
令 g1(x1)=−log(x1)
令 x1=g2(x2)=sumx2
令 x2=g3(x3)=ex3
令 x3=g4(Ztarget)=Ztarget−max(Zi)
p=∂x1∂g1(x1)∂x2∂g2(x2)∂x3∂g3(x3)∂Ztarget∂g4(Ztarget)
∂x1∂g1(x1)=−x11
∂x2∂g2(x2)=sum1
∂x3∂g3(x3)=ex3
∂Ztarget∂g4(Ztarget)=1
故 p=−x1sumex3
x3=Ztarget−max(Zi)
x1=sumeZtarget−max(Zi)
故 p=−1
故整体的导数为 sumeZtarget−max(Zi)−1
∂Zi∂ce={sumeZi−max(Zi),sumeZtarget−max(Zi)−1,if i=targetif i=target
softmax(Zi)=∑j=1neZjeZi
令 softmax(Zi)=g1(x1,x2)=x2x1
令 sum=∑j=1neZj
∂Zi∂softmax(Zi)=∂x1∂g1(x1,x2)∂Zi∂x1+∂x2∂g1(x1,x2)∂Zi∂x2
同样考虑 i 是否等于 target的两种情况
当 i=target
∂x1∂g1(x1,x2)=x21
∂x2∂g1(x1,x2)=−x22x1
下面计算 ∂Zi∂x1
令 x1=g2(x3)=ex3
令 x3=g3(Zi)=Zi−max(Zi)
∂Zi∂x1=∂x3∂g2(x3)∂Zi∂g3(Zi)=ex3⋅1=eZi−max(Zi)
下面计算 ∂Zi∂x2
令 x2=g4(x4)=x4+c1
其中 c1 为常数
令 x4=g5(x5)=ex5
令 x5=g6(Zi)=Zi−max(Zi)
∂Zi∂x2=∂x4∂g4(x4)∂x5∂g5(x5)∂Zi∂g6(zt)
∂x4∂g4(x4)=1
∂x5∂g5(x5)=ex5=eZi−max(Zi)
∂Zi∂g6(zt)=1
∂Zi∂x2=eZi−max(Zi)
故 ∂Zi∂softmax(Zi)=x21⋅eZi−max(Zi)+(−x22x1)⋅eZi−max(Zi)
其中
x1=eZi−max(Zi)
x2=sum
故 ∂Zi∂softmax(Zi)=sumeZi−max(Zi)⋅(1−sumeZi−max(Zi))
又因为 softmax(Zi)=sumeZi
故 ∂Zi∂softmax(Zi)=softmax(Zi)⋅(1−softmax(Zi))
当 i=target
令 softmax(Ztarget)=g1(x1)=x1eZtarget−max(Zi)
令 x1=g2(x2)=x2+c1 其中 c1 为常数
令 x2=g3(x3)=ex3
令 x3=g4(Zi)=Zi−max(Zi)
∂Zi∂softmax(Ztarget)=x1∂g1(x1)∂x2∂g2(x2)∂x3∂g3(x3)∂Zi∂g4(Zi)
x1∂g1(x1)=−sum2eZtarget−max(Zi)
∂x2∂g2(x2)=1
∂x3∂g3(x3)=ex3=eZi−max(Zi)
∂Zi∂g4(Zi)=1
故
∂Zi∂softmax(Ztarget)=−sumeZtarget−max(Zi)⋅sumeZi−max(Zi)=−softmax(Ztarget)⋅softmax(Zi)
最终整理
∂Zi∂softmax(Ztarget)={−softmax(Ztarget)⋅softmax(Zi),softmax(Zi)⋅(1−softmax(Zi)),if i=targetif i=target
参考 https://zhuanlan.zhihu.com/p/634644501
Layernorm(xi)=γσxi−μ+β
σ=var+ϵ
var=n1∑i=1n(xi−μ)2
μ=n1∑i=1nxi
令 xi^=σxi−μ
我们最关注 ∂xj∂xi^ ,乘以 γ 和加上 β 的部分交给node.h中的矩阵乘法加法的自动求导即可
只用链式法则
令 xi=g1(x1,x2)
令 x1=xi−μ
令 x2=σ
故
∂xj∂xi^=∂x1∂g1(x1,x2)∂xj∂xi−μ+∂x2∂g1(x1,x2)∂xj∂σ (1)
∂x1∂g1(x1,x2)=x21=σ1 (2)
∂xj∂xi−μ=∂xj∂xi−∂xj∂μ
令 δij=∂xj∂xi
故
δij={0,1,if i=jif i=j
再计算 ∂xj∂μ=∂xj∂n1∑i=1nxi=n1
故 ∂xj∂xi−μ=∂xj∂xi−∂xj∂μ=δij−n1 (3)
∂x2∂g1(x1,x2)=−x22x1=−σ2xi−μ (4)
∂xj∂σ=∂var∂sigma∂xj∂var (5)
∂xj∂var=∂xj∂n1∑k=1n(xk−μ)2=n1∑k=1n2(xk−μ)∂xj∂xk−μ
深入分析 n1∑k=1n2(xk−μ)∂xj∂xk−μ
∂xj∂xk−μ={−n1,1−n1,if i=jif i=j
所以展开 ∑k=1n2(xk−μ)∂xj∂xk−μ=∑i=1i=kn−n1⋅2⋅(xk−μ)+2⋅(xj−μ)(1−n1)=2⋅(−n1∑i=1i=knxk+n1(n−1)μ+xj−n1⋅xj−μ+n1⋅μ)
观察第一项和第四项的和 −n1∑i=1i=knxk−n1⋅xj=−μ
观察第二项和第六项的和 n1(n−1)μ+n1⋅μ=μ
上面两个和相加消掉了,只剩下第三项和第五项
故 ∂xj∂var=n2(xj−μ)
又因为 ∂var∂σ=∂var∂(var+ϵ)=2(var+ϵ)1=2σ1
故 (5)= ∂xj∂σ=∂var∂sigma∂xj∂var=2σ1⋅n2(xj−μ)=nσ1(xj−μ)
将(2)(3)(4)(5) 带回 (1)
∂xj∂xi^=σ1(δij−n1)−σ2xi−μ⋅nσ1(xj−μ)
观察这里的第二项 σ2xi−μ⋅nσ1(xj−μ)=nσ1⋅σxi−μ⋅σxj−μ=nσ1⋅xi^⋅xj^
故
∂xj∂xi^=σ1(δij−n1)−nσ1⋅xi^⋅xj^=σδij−nσ1−nσ1xi^xj^
其中
δij={0,1,if i=jif i=j
具体实现参见代码, 考虑到精度问题,计算顺序有些微调整