在 Apache MXNet Gluon 中使用 KLDivLoss:原理、两种 `from_logits` 模式与实战陷阱
2026/9/21 1:54:10 网站建设 项目流程

在 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)。需要特别强调两点:

  1. KL 散度是不对称的度量:即KL(P,Q) != KL(Q,P)。顺序很重要,我们应该按照“预测分布 vs 目标分布”的顺序进行比较,不能随意调换两个参数的位置。
  2. 两种调用方式KLDivLoss的用法取决于from_logits参数的设置(默认值为True)。

构造示例分布并可视化

为直观理解,我们先构造三个各含 4 个类别的类别分布dist_1dist_2dist_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时,需要注意输入约定:

  1. 预测值(pred)必须是取了对数的概率分布的参数(logged probability distribution);
  2. 目标值(label)必须是概率分布的参数(即未取对数)。

为什么推荐log_softmax而非softmax

在实际网络中,我们通常先对网络输出施加softmax得到概率分布,但这样计算的梯度数值不稳定。更稳定的替代方案是使用log_softmax,因此from_logits=TrueKLDivLoss期望的正是这类取了对数的预测值。此外,训练时通常处理的是批次数据,因此预测与目标都需要有批次维度(默认是第一维)。

由于本例中我们处理的本身就是分布,无需再施加 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,而是把这个操作交给损失函数内部完成。当KLDivLossfrom_logits=False时,log_softmax会被施加在传入loss_fn的**第一个参数(预测值)**上。

例如,假设网络输出了如下未归一化的值(特意选取,使得对这些值施加softmax后恰好得到与dist_1相同的分布参数):

output = mx.nd.array([0.39056206, 1.3068528, 0.39056206, -0.30258512])

将其传给from_logits=FalseKLDivLoss,由于损失函数内部会施加log_softmax,得到的dist_1dist_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_logitsTrue预测值是否为对数概率(通常来自log_softmax
axis-1仅在from_logits=False时生效,指定施加 softmax 的维度
weightNone全局标量权重,对整批损失统一缩放
batch_axis0代表 mini-batch 的维度

predlabel的形状可以任意,只要元素总数相同即可;输出 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=Truepredlog(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_categoriestrue_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

这个测试同时展示了两个值得留意的工程细节:

  1. 预测侧先做log_softmax再进KLDivLoss,即采用默认的from_logits=True模式,与文档推荐的数值稳定做法一致;
  2. KLDivLoss是可混合符号(hybrid)的HybridBlock,既可以直接以 NDArray 方式调用(如本文前面的示例),也可以嵌入mx.sym符号图并用mx.mod.Module训练,说明该损失对命令式(imperative)与符号式(symbolic)两种编程范式均兼容。

小结与使用建议

综合文档示例与源码实现,使用KLDivLoss时请遵循以下要点:

  1. 明确输入约定from_logits=True(默认)时预测值必须是log_softmax后的对数概率、目标值必须是概率;from_logits=False时预测值为未归一化 logits,由损失内部施加log_softmax(此时可用axis指定 softmax 维度)。
  2. 注意方向性:KL 散度不对称,KLDivLoss(pred, label)中预测与目标不可随意互换。
  3. 警惕 0 概率:目标侧已内置1e-12平滑,但预测侧若出现 0 概率仍可能产生nan,需要自行平滑。
  4. 聚合语义:默认沿批次维度取平均,输出比求和定义小“类别数”倍,跨实现对比时需换算。
  5. 加权能力:可通过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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询