在 Apache MXNet Gluon 中使用 KLDivLoss:原理、两种 `from_logits` 模式与实战陷阱

发布时间:2026/9/21 1:54:13
在 Apache MXNet Gluon 中使用 KLDivLoss:原理、两种 `from_logits` 模式与实战陷阱 在 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/mxnetKullback-LeiblerKL散度用于度量一个概率分布与另一个参考概率分布之间的差异在 MXNet Gluon 中通过KLDivLoss实现是变分自编码器VAE与 Trust Region Policy OptimizationTRPO等强化学习策略网络的常用训练损失。读完本文你将掌握KL 散度的数学定义与不对称性、KLDivLoss两种from_logits模式的正确用法与背后的log_softmax稳定性原理、以及分布支撑不一致common support与聚合方式aggregation两个高级陷阱。KL 散度的定义与不对称性Kullback-LeiblerKL散度衡量一个概率分布与第二个参考概率分布之间的差异程度。KL 散度值越小说明两个分布越相似由于该损失函数可微我们可以用梯度下降来最小化网络输出与目标分布之间的 KL 散度。典型应用场景包括变分自编码器VAE最小化隐变量后验分布与先验分布之间的差异强化学习策略网络如 Trust Region Policy OptimizationTRPO中约束新旧策略的差异。在 MXNet Gluon 中使用KLDivLoss即可比较类别分布categorical distributions。需要特别强调两点KL 散度是不对称的度量即KL(P,Q) ! KL(Q,P)。顺序很重要我们应该按照“预测分布 vs 目标分布”的顺序进行比较不能随意调换两个参数的位置。两种调用方式KLDivLoss的用法取决于from_logits参数的设置默认值为True。构造示例分布并可视化为直观理解我们先构造三个各含 4 个类别的类别分布dist_1、dist_2和dist_3from 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(top1) plt.bar(idx, dist_1, alpha0.5, colorblack) plt.bar(idx, dist_2, alpha0.5, coloraqua) plt.title(Distributions 1 2) plt.subplot(1,2,2) plt.ylim(top1) plt.bar(idx, dist_1, alpha0.5, colorblack) plt.bar(idx, dist_3, alpha0.5, coloraqua) plt.title(Distributions 1 3)从视觉上可以直观看出分布 1 与分布 2 比分布 1 与分布 3 更相似。接下来我们用KLDivLoss来量化验证这一结论。from_logitsTrue默认输入对数概率分布当使用默认的from_logitsTrue时需要注意输入约定预测值pred必须是取了对数的概率分布的参数logged probability distribution目标值label必须是概率分布的参数即未取对数。为什么推荐log_softmax而非softmax在实际网络中我们通常先对网络输出施加softmax得到概率分布但这样计算的梯度数值不稳定。更稳定的替代方案是使用log_softmax因此from_logitsTrue时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_logitsTrue) 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_logitsTrue时KLDivLoss的数学定义见 python/mxnet/gluon/loss.py 的 docstringL sum_i label_i * [log(label_i) - pred_i]其中pred为对数概率label为概率。对应的实现位于hybrid_forward中python/mxnet/gluon/loss.pydef hybrid_forward(self, F, pred, label, sample_weightNone): 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, axisself._batch_axis, excludeTrue)值得注意的是即便在from_logitsTrue模式下实现仍会对label施加F.log(label 1e-12)并加入1e-12的极小值平滑项——这是为了防止目标分布中出现 0 概率时log(0)产生-inf。这一点在后面“Common Support 陷阱”一节会再次呼应。from_logitsFalse让损失函数替你施加log_softmax另一种方式是不手动对网络输出施加log_softmax而是把这个操作交给损失函数内部完成。当KLDivLoss的from_logitsFalse时log_softmax会被施加在传入loss_fn的**第一个参数预测值**上。例如假设网络输出了如下未归一化的值特意选取使得对这些值施加softmax后恰好得到与dist_1相同的分布参数output mx.nd.array([0.39056206, 1.3068528, 0.39056206, -0.30258512])将其传给from_logitsFalse的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_logitsFalse) 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_logitsFalse时对应的数学定义见 python/mxnet/gluon/loss.pyprob softmax(pred) L sum_i label_i * [log(label_i) - log(prob_i)]axis参数默认-1仅在from_logitsFalse时生效用于指定施加softmax的维度。其余参数在两个模式下通用参数默认值说明from_logitsTrue预测值是否为对数概率通常来自log_softmaxaxis-1仅在from_logitsFalse时生效指定施加 softmax 的维度weightNone全局标量权重对整批损失统一缩放batch_axis0代表 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_logitsTrue时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, axisself._batch_axis, excludeTrue)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.05with_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_size10, label_namelabel) 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_epoch200, optimizer_params{learning_rate: 0.01}, eval_metricmx.metric.Loss(), optimizeradam) assert mod.score(data_iter, eval_metricmx.metric.Loss())[0][1] 0.05这个测试同时展示了两个值得留意的工程细节预测侧先做log_softmax再进KLDivLoss即采用默认的from_logitsTrue模式与文档推荐的数值稳定做法一致KLDivLoss是可混合符号hybrid的HybridBlock既可以直接以 NDArray 方式调用如本文前面的示例也可以嵌入mx.sym符号图并用mx.mod.Module训练说明该损失对命令式imperative与符号式symbolic两种编程范式均兼容。小结与使用建议综合文档示例与源码实现使用KLDivLoss时请遵循以下要点明确输入约定from_logitsTrue默认时预测值必须是log_softmax后的对数概率、目标值必须是概率from_logitsFalse时预测值为未归一化 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),仅供参考