使用交叉熵损失函数时,如何处理log(0)的计算问题
刚好最近也在梳理二分类场景下的损失函数逻辑,我把这段内容整理成清晰的结构,方便你理解:
二分类场景下的交叉熵损失函数详解
1. 核心概念说明
- 真实标签Y:二分类场景下仅取0或1两个值,代表样本的真实类别
- 预测概率predY:由神经网络输出的
logits(分类前的原始输出)经过sigmoid激活转换得到,取值范围(0,1),表示样本被预测为类别1的概率 - 样本数量m:当前训练batch中包含的样本总数
2. 关键函数与损失计算实现
首先是sigmoid激活函数的实现,用来把logits映射到概率区间:
import numpy as np def sigmoid(X): return 1/(1 + np.exp(-X))
通过sigmoid得到预测概率:
# logits是神经网络最后一层的原始输出 predY = sigmoid(logits)
接下来是交叉熵损失的计算逻辑:
# 计算每个样本的交叉熵损失项 loss = np.multiply(np.log(predY), Y) + np.multiply((1 - Y), np.log(1 - predY)) # 计算整个batch的平均损失(cost) cost = -np.sum(loss)/m
3. 简单逻辑解释
交叉熵的本质是衡量真实标签和预测概率的匹配程度:
- 当真实标签Y=1时,损失项简化为
-np.log(predY):predY越接近1,损失值越小,模型预测越准确 - 当真实标签Y=0时,损失项简化为
-np.log(1-predY):predY越接近0,损失值越小,模型预测越准确
最后对所有样本的损失求和取平均,得到的cost就是我们训练时需要最小化的目标值。
内容的提问来源于stack exchange,提问作者GabrielChu
相关产品推荐
相关产品推荐

