You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

TensorFlow Keras Sequential模型:损失函数形式不同训练结果差异问题

问题原因与解决方案

这种差异其实是Keras处理字符串形式和函数形式损失函数的逻辑不同导致的,结合你的场景,具体原因可以拆解为以下几点:

1. 损失函数的from_logits参数自动适配差异

当你传入字符串形式的"categorical_crossentropy"时,Keras会在模型编译阶段自动检查你的模型输出层是否包含激活函数(比如softmax):

  • 如果你的模型最后一层没有加softmax(输出是原始logits),Keras会自动为损失函数设置from_logits=True,此时计算的是logits与真实标签之间的交叉熵,损失值会处于正常范围(比如你看到的初始损失约41)。

而当你直接传入函数tf.keras.losses.categorical_crossentropy时,它使用的是默认参数from_logits=False:

  • 这意味着损失函数会假设你的模型输出已经经过了softmax激活(概率分布),但如果你的模型输出是原始logits,此时计算交叉熵会得到异常大的数值(因为logits的范围没有被压缩到[0,1],错误分类的情况下交叉熵会被放大),也就是你看到的训练5个epoch后损失约192的情况。

2. 多进程模式对损失函数初始化的影响

当use_multiprocessing=True时,Keras在多进程数据加载的场景下,内部会重新初始化损失函数的上下文,这时候会触发和字符串形式类似的自动适配逻辑,从而让函数形式的损失函数也能正确识别logits输入,最终消除了两种方式的差异。反之,单进程模式下没有这个触发逻辑,差异就会显现。

3. 指标的类似差异原因

你提到指标用字符串形式也有类似差异,原理和损失函数完全一致:

  • 字符串形式的"categorical_accuracy"会自动适配logits/概率输出;
  • 直接传入tf.keras.metrics.categorical_accuracy函数时,默认是基于概率分布计算准确率,如果输入是logits,计算结果会不准确或异常。

解决方案

要消除两种方式的差异,你可以显式为函数形式的损失函数和指标设置from_logits=True:

  • 对于损失函数,使用类实例并指定参数:
    loss_param1 = tf.keras.losses.CategoricalCrossentropy(from_logits=True)
    
    或者用lambda包装函数:
    loss_param1 = lambda y_true, y_pred: tf.keras.losses.categorical_crossentropy(y_true, y_pred, from_logits=True)
    
  • 对于指标,同样使用类实例:
    metrics=[tf.keras.metrics.CategoricalAccuracy(from_logits=True)]
    

这样无论你用字符串形式还是显式的函数/类实例形式,训练结果都会保持一致,和use_multiprocessing的设置无关。

内容的提问来源于stack exchange,提问作者delyeet

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.14 08:21:10