Softmax+CrossEntropy反向传播异常:梯度恒为0求解决方案
问题根源
你写的softmax_backward函数逻辑完全错误——它计算的是Softmax输出对输入logits的偏导数之和,而这个求和结果必然是0:
对每个输出维度i,计算过程为:
sum_j [∂p_i/∂x_j] = sum_j [p_i(δ_ij - p_j)] = p_i*(1 - sum_j p_j) = p_i*(1-1) = 0
这直接导致derivative向量全为0,后续和gradient相乘得到的delta自然也全是0,最终所有权重、偏置的梯度都为0,损失根本没法更新。
另外你搞反了反向传播的计算逻辑:Softmax层的反向传播需要把**上层传来的梯度(损失对Softmax输出的梯度)**和Softmax的Jacobian矩阵做矩阵乘法,而不是单独算Softmax的导数再点乘。
正确解决方案
Softmax的反向传播不需要计算完整的Jacobian矩阵,尤其是和CrossEntropy损失配合时,有极简的公式可用:
如果用的是CrossEntropy损失,损失对Softmax输入logits的梯度直接等于Softmax输出概率减去one-hot格式的目标标签。
就算不是CrossEntropy损失,正确的Softmax反向传播公式也应该是:
对每个输入维度i,损失对x_i的梯度为:
∇L/∂x_i = p_i * (gradient[i] - sum_j (gradient[j] * p_j))
修正后的代码示例
首先重写softmax_backward,让它接收上层梯度和Softmax的输出概率:
fn softmax_backward(&self, gradient: &Vec<f32>, probability: &Vec<f32>) -> Vec<f32> { // 计算梯度与概率的加权和 let grad_p_sum: f32 = gradient.iter() .zip(probability.iter()) .map(|(g, p)| g * p) .sum(); // 逐维度计算损失对logits的梯度 probability.iter() .enumerate() .map(|(i, &p)| p * (gradient[i] - grad_p_sum)) .collect() }
然后修正layer_backward的逻辑——注意原来的代码把output传给softmax_backward是错误的(原函数参数是logits),现在直接传入Softmax的输出概率:
fn layer_backward( &self, gradient: &Vec<f32>, input: &Vec<f32>, output: &Vec<f32> ) -> (Vec<Vec<f32>>, Option<Vec<f32>>, Vec<f32>) { // 用上层梯度和Softmax输出计算delta(损失对logits的梯度) let delta: Vec<f32> = self.softmax_backward(gradient, output); // 计算权重梯度 let weight_gradient: Vec<Vec<f32>> = delta .iter() .map(|d| input.iter().map(|i| i * d).collect()) .collect(); // 计算偏置梯度 let bias_gradient: Option<Vec<f32>> = self.bias.as_ref().map(|_| delta.clone()); // 计算输入梯度 let input_gradient: Vec<f32> = (0..input.len()) .map(|i| delta .iter() .zip(self.weights.iter()) .map(|(d, w)| d * w[i]) .sum()) .collect(); (weight_gradient, bias_gradient, input_gradient) }
CrossEntropy场景的额外优化
如果你的损失确实是CrossEntropy,上层传来的gradient其实是-target / output(target是one-hot向量),代入上面的公式会直接简化为output - target,可以跳过softmax_backward直接计算:
// 仅当使用CrossEntropy损失时可用,target是one-hot格式的标签向量 let delta: Vec<f32> = output.iter() .zip(target.iter()) .map(|(&p, &t)| p - t) .collect();
这样能省掉求和步骤,效率更高。
验证步骤
- 确认
gradient参数是正确的:它必须是损失对当前层输出(Softmax概率)的梯度,不能传错值。 - 检查修正后的
delta:训练初期delta应该是非零的,确保梯度能正常传递到权重和偏置。
内容的提问来源于stack exchange,提问作者Hallvard

