NumPy整向量索引:交叉熵梯度代码中grad[range(m),y] -= 1的作用
问题解答
1. grad[range(m),y] -= 1的语法逻辑
这行用的是NumPy的整数数组高级索引规则,针对二维数组grad的行为逻辑如下:
range(m)是长度为m的行索引序列,依次对应每个样本的行号(grad维度是样本数×类别数,每行对应一个样本的所有类别softmax输出)y是长度为m的列索引序列,每个元素对应该行样本的真实类别编号- 两个等长的一维索引数组配对后,会一次性选中所有
(行i, 列y[i])位置的元素,也就是每个样本真实类别对应的softmax输出值 - 最后的
-=1就是把所有选中的位置的数值各自减1,其他位置的数值保持不变
2. 是否符合「从grad中减去y」的预期
完全符合预期。
交叉熵损失对logits(也就是输入的X)的梯度公式为 softmax(X) - y的onehot编码矩阵。因为y的onehot编码只有真实类别位置为1,其他位置为0,这行操作等价于不用显式生成onehot矩阵,直接在softmax输出的对应位置减1,和完整做矩阵减法的结果完全一致,还能节省内存和计算开销,是工业界实现交叉熵梯度的标准优化写法。
内容的提问来源于stack exchange,提问作者Clebo Sevic
相关产品推荐
相关产品推荐

