在 Apache MXNet Gluon 中使用 KLDivLoss:原理、两种from_logits模式与实战陷阱
【免费下载链接】mxnetLightweight, Portable, Flexible Distributed/Mobile Deep Learning with Dynamic, Mutation-aware Dataflow Dep Scheduler; for Python, R, Julia, Scala, Go, Javascript and more项目地址: https://gitcode.com/gh_mirrors/mxnet1/mxnet
Kullback-Leibler(KL)散度用于度量一个概率分布与另一个参考概率分布之间的差异,在 MXNet Gluon 中通过KLDivLoss实现,是变分自编码器(VAE)与 Trust Region Policy Optimization(TRPO)等强化学习策略网络的常用训练损失。读完本文你将掌握:KL 散度的数学定义与不对称性、KLDivLoss两种from_logits模式的正确用法与背后的log_softmax稳定性原理、以及分布支撑不一致(common support)与聚合方式(aggregation)两个高级陷阱。
KL 散度的定义与不对称性
Kullback-Leibler(KL)散度衡量一个概率分布与第二个参考概率分布之间的差异程度。KL 散度值越小,说明两个分布越相似;由于该损失函数可微,我们可以用梯度下降来最小化网络输出与目标分布之间的 KL 散度。典型应用场景包括:
- 变分自编码器(VAE):最小化隐变量后验分布与先验分布之间的差异;
- 强化学习策略网络:如 Trust Region Policy Optimization(TRPO)中约束新旧策略的差异。
在 MXNet Gluon 中,使用KLDivLoss即可比较类别分布(categorical distributions)。需要特别强调两点:
- KL 散度是不对称的度量:即
KL(P,Q) != KL(Q,P)。顺序很重要,我们应该按照“预测分布 vs 目标分布”的顺序进行比较,不能随意调换两个参数的位置。 - 两种调用方式:
KLDivLoss的用法取决于from_logits参数的设置(默认值为True)。
构造示例分布并可视化
为直观理解,我们先构造三个各含 4 个类别的类别分布dist_1、dist_2和dist_3:
from matplotlib import pyplot as plt import mxnet as mx import numpy as np idx = np.array([1, 2, 3, 4]) dist_1 = np.array([0.2, 0.5, 0.2, 0.1]) dist_2 = np.array([0.3, 0.4, 0.1, 0.2]) dist_3 = np.array([0.1, 0.1, 0.1, 0.7]) plt.figure(figsize=(10,5)) plt.subplot(1,2,1) plt.ylim(top=1) plt.bar(idx, dist_1, alpha=0.5, color='black') plt.bar(idx, dist_2, alpha=0.5, color='aqua') plt.title('Distributions 1 & 2') plt.subplot(1,2,2) plt.ylim(top=1) plt.bar(idx, dist_1, alpha=0.5, color='black') plt.bar(idx, dist_3, alpha=0.5, color='aqua') plt.title('Distributions 1 & 3')从视觉上可以直观看出,分布 1 与分布 2 比分布 1 与分布 3 更相似。接下来我们用KLDivLoss来量化验证这一结论。
from_logits=True(默认):输入对数概率分布
当使用默认的from_logits=True时,需要注意输入约定:
- 预测值(pred)必须是取了对数的概率分布的参数(logged probability distribution);
- 目标值(label)必须是概率分布的参数(即未取对数)。
为什么推荐log_softmax而非softmax
在实际网络中,我们通常先对网络输出施加softmax得到概率分布,但这样计算的梯度数值不稳定。更稳定的替代方案是使用log_softmax,因此from_logits=True时KLDivLoss期望的正是这类取了对数的预测值。此外,训练时通常处理的是批次数据,因此预测与目标都需要有批次维度(默认是第一维)。
由于本例中我们处理的本身就是分布,无需再施加 softmax,只需对分布取log。同时即使处理单个分布,也要为它构造出批次维度:
def kl_divergence(dist_a, dist_b): # 添加批次维度 pred_batch = mx.nd.array(dist_a).expand_dims(0) target_batch = mx.nd.array(dist_b).expand_dims(0) # 对分布取对数 pred_batch = pred_batch.log() # 创建损失(假定预测分布已取对数) loss_fn = mx.gluon.loss.KLDivLoss(from_logits=True) divergence = loss_fn(pred_batch, target_batch) return divergence.asscalar()分别计算分布 1 与分布 2、分布 3、以及自身之间的散度:
print("Distribution 1 compared with Distribution 2: {}".format( kl_divergence(dist_1, dist_2))) print("Distribution 1 compared with Distribution 3: {}".format( kl_divergence(dist_1, dist_3))) print("Distribution 1 compared with Distribution 1: {}".format( kl_divergence(dist_1, dist_1)))结果与预期一致:分布 1 与 2 的 KL 散度小于分布 1 与 3;且分布与自身的 KL 散度为 0。
源码视角:公式与实现
from_logits=True时,KLDivLoss的数学定义(见 python/mxnet/gluon/loss.py 的 docstring):
L = sum_i label_i * [log(label_i) - pred_i]其中pred为对数概率,label为概率。对应的实现位于hybrid_forward中(python/mxnet/gluon/loss.py):
def hybrid_forward(self, F, pred, label, sample_weight=None): if not self._from_logits: pred = F.log_softmax(pred, self._axis) loss = label * (F.log(label + 1e-12) - pred) loss = _apply_weighting(F, loss, self._weight, sample_weight) return F.mean(loss, axis=self._batch_axis, exclude=True)值得注意的是,即便在from_logits=True模式下,实现仍会对label施加F.log(label + 1e-12)并加入1e-12的极小值平滑项——这是为了防止目标分布中出现 0 概率时log(0)产生-inf。这一点在后面“Common Support 陷阱”一节会再次呼应。
from_logits=False:让损失函数替你施加log_softmax
另一种方式是:不手动对网络输出施加log_softmax,而是把这个操作交给损失函数内部完成。当KLDivLoss的from_logits=False时,log_softmax会被施加在传入loss_fn的**第一个参数(预测值)**上。
例如,假设网络输出了如下未归一化的值(特意选取,使得对这些值施加softmax后恰好得到与dist_1相同的分布参数):
output = mx.nd.array([0.39056206, 1.3068528, 0.39056206, -0.30258512])将其传给from_logits=False的KLDivLoss,由于损失函数内部会施加log_softmax,得到的dist_1与dist_2之间的 KL 散度应当与前面完全相同:
def kl_divergence_not_from_logits(dist_a, dist_b): # 添加批次维度 pred_batch = mx.nd.array(dist_a).expand_dims(0) target_batch = mx.nd.array(dist_b).expand_dims(0) # 创建损失(由损失函数内部施加 log_softmax) loss_fn = mx.gluon.loss.KLDivLoss(from_logits=False) divergence = loss_fn(pred_batch, target_batch) return divergence.asscalar()print("Distribution 1 compared with Distribution 2: {}".format( kl_divergence_not_from_logits(output, dist_2)))两种模式的实现对照
从 KLDivLoss.hybrid_forward 的源码可以看出两种模式的唯一区别就在于是否执行pred = F.log_softmax(pred, self._axis)。from_logits=False时对应的数学定义(见 python/mxnet/gluon/loss.py):
prob = softmax(pred) L = sum_i label_i * [log(label_i) - log(prob_i)]axis参数(默认-1)仅在from_logits=False时生效,用于指定施加softmax的维度。其余参数在两个模式下通用:
| 参数 | 默认值 | 说明 |
|---|---|---|
from_logits | True | 预测值是否为对数概率(通常来自log_softmax) |
axis | -1 | 仅在from_logits=False时生效,指定施加 softmax 的维度 |
weight | None | 全局标量权重,对整批损失统一缩放 |
batch_axis | 0 | 代表 mini-batch 的维度 |
pred与label的形状可以任意,只要元素总数相同即可;输出 loss 的形状为(batch_size,),除batch_axis之外的维度会被平均掉(见 KLDivLoss 的 docstring)。sample_weight支持逐元素加权,需可广播到与pred相同的形状,例如pred形状为(64, 10)时,sample_weight可传(64, 1)来按样本加权(加权逻辑见_apply_weighting)。
高级陷阱一:Common Support(分布支撑不一致)
偶尔你会遇到KLDivLoss给出异常结果的情况,最常见的问题之一是:所比较的两个分布的支撑(support)不一致。这里的“支撑”指分布中概率非零的那些取值。前面所有示例恰好具有相同的支撑,但现实中很可能出现某些类别概率为 0 的情况:
dist_4 = np.array([0, 0.9, 0, 0.1])print("Distribution 4 compared with Distribution 1: {}".format( kl_divergence(dist_4, dist_1)))可以看到结果是nan——这显然会在计算梯度时引发问题。原因在于,当预测分布(或目标分布)中的某个类别概率为 0 时,公式中的log(0)会产生-inf,进而导致整个损失为nan。
一种常见的应对方案是:给所有概率都加上一个极小值epsilon。事实上,KLDivLoss的内部实现已经对目标分布做了这一步——在 hybrid_forward 中使用F.log(label + 1e-12),即对 label 施加了1e-12的平滑项。因此若nan来自预测侧的 0 概率(from_logits=True时pred为log(0) = -inf),就需要在送入损失函数前自行对预测分布做类似的处理。
高级陷阱二:Aggregation(聚合方式与定义差异)
KLDivLoss的结果与 KL 散度的“教科书定义”之间还有一个细微差异:聚合类别贡献的方式。尽管真正的定义是对各类别贡献求和,但 MXNet Gluon 的默认行为是沿批次维度取平均。因此KLDivLoss的输出会比真实定义小,缩小的倍数为类别数。
验证如下,先按定义手工计算真实散度:
true_divergence = (dist_2*(np.log(dist_2)-np.log(dist_1))).sum() print('true_divergence: {}'.format(true_divergence))再对比KLDivLoss的结果:
num_categories = dist_1.shape[0] divergence = kl_divergence(dist_1, dist_2) print('divergence: {}'.format(divergence)) print('divergence * num_categories: {}'.format(divergence * num_categories))可以看到divergence * num_categories与true_divergence一致。这一点可以从实现确认:hybrid_forward末尾使用F.mean(loss, axis=self._batch_axis, exclude=True)(python/mxnet/gluon/loss.py),即对除批次维度外的所有维度取平均,而非求和。如果你在论文复现或跨框架对齐时需要严格的求和语义,请记得自行乘以类别数。
端到端验证:仓库中的 KL 散度训练测试
仓库的单元测试 tests/python/unittest/test_loss.py 提供了一个端到端训练验证:构造 20 个样本、每个样本 10 维特征的随机输入,目标为 2 类 softmax 概率分布,将网络输出log_softmax后的符号与KLDivLoss()组合成make_loss,再用mx.mod.Module配合 Adam 优化器训练 200 轮,最终断言训练损失小于0.05:
@with_seed() def test_kl_loss(): N = 20 data = mx.random.uniform(-1, 1, shape=(N, 10)) label = mx.nd.softmax(mx.random.uniform(0, 1, shape=(N, 2))) data_iter = mx.io.NDArrayIter(data, label, batch_size=10, label_name='label') output = mx.sym.log_softmax(get_net(2)) l = mx.symbol.Variable('label') Loss = gluon.loss.KLDivLoss() loss = Loss(output, l) loss = mx.sym.make_loss(loss) mod = mx.mod.Module(loss, data_names=('data',), label_names=('label',)) mod.fit(data_iter, num_epoch=200, optimizer_params={'learning_rate': 0.01}, eval_metric=mx.metric.Loss(), optimizer='adam') assert mod.score(data_iter, eval_metric=mx.metric.Loss())[0][1] < 0.05这个测试同时展示了两个值得留意的工程细节:
- 预测侧先做
log_softmax再进KLDivLoss,即采用默认的from_logits=True模式,与文档推荐的数值稳定做法一致; KLDivLoss是可混合符号(hybrid)的HybridBlock,既可以直接以 NDArray 方式调用(如本文前面的示例),也可以嵌入mx.sym符号图并用mx.mod.Module训练,说明该损失对命令式(imperative)与符号式(symbolic)两种编程范式均兼容。
小结与使用建议
综合文档示例与源码实现,使用KLDivLoss时请遵循以下要点:
- 明确输入约定:
from_logits=True(默认)时预测值必须是log_softmax后的对数概率、目标值必须是概率;from_logits=False时预测值为未归一化 logits,由损失内部施加log_softmax(此时可用axis指定 softmax 维度)。 - 注意方向性:KL 散度不对称,
KLDivLoss(pred, label)中预测与目标不可随意互换。 - 警惕 0 概率:目标侧已内置
1e-12平滑,但预测侧若出现 0 概率仍可能产生nan,需要自行平滑。 - 聚合语义:默认沿批次维度取平均,输出比求和定义小“类别数”倍,跨实现对比时需换算。
- 加权能力:可通过
weight(全局标量)与sample_weight(逐元素、可广播)灵活调整损失贡献。
相关代码入口:损失实现 python/mxnet/gluon/loss.py#L408-L480,训练验证 tests/python/unittest/test_loss.py#L132-L146。
【免费下载链接】mxnetLightweight, Portable, Flexible Distributed/Mobile Deep Learning with Dynamic, Mutation-aware Dataflow Dep Scheduler; for Python, R, Julia, Scala, Go, Javascript and more项目地址: https://gitcode.com/gh_mirrors/mxnet1/mxnet
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考