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

TensorFlow图像字幕模型训练出现NaN损失的原因及解决问询

图像字幕训练中NaN损失问题的原因与解决方法

问题背景

基于TensorFlow训练图像字幕模型,采用Flickr 8K数据集:

  • 图像特征:通过ResNet50提取,形状为(m,49,2048),预存储后用于训练
  • 文本处理:用GloVe 6B 300d向量构建词嵌入矩阵,通过StringLookup层处理字幕,训练集字幕形状(m,37),验证集(m,32)

小样本训练(1/5张图像)时模型可正常过拟合,但全量训练(约6000张图像)时,训练至epoch中途(处理约1000张图像后)会出现NaN损失,更换数据集起始位置训练仍会复现,已确认数据无异常。添加调试回调后发现,logits数值持续增大,loss先变为inf后转为NaN。

模型代码

def model_build():
    strategy = tf.distribute.MirroredStrategy()
    with strategy.scope():
        image = tf.keras.Input((49, 2048))
        input_caption = tf.keras.Input((None,))

        x_image = Dense(1024, activation='relu')(image)
        x_image = Dense(512, activation='relu')(x_image)

        embedding_layer = Embedding(400004, 300, trainable=False, mask_zero=False)
        embedding_layer.build((None,))
        embedding_layer.set_weights([emb_matrix])

        x_caption = embedding_layer(input_caption)
        x_caption = LSTM(512, return_sequences=True)(x_caption)

        attention = MultiHeadAttention(num_heads=1, key_dim=64)(query=x_caption, value=x_image)

        x = tf.keras.layers.Add()([x_caption, attention])
        x = LayerNormalization(epsilon=1e-6)(x)
        x = tf.keras.layers.Dropout(0.3)(x)

        x = LSTM(256, return_sequences=True)(x)
        x = tf.keras.layers.Dropout(0.3)(x)

        logits = Dense(400004, activation='linear',name="logits_layer")(x)
        logits = tf.keras.layers.Lambda(lambda t: tf.clip_by_value(t, -10.0, 10.0))(logits)

        model = tf.keras.Model(inputs=[image, input_caption], outputs=logits)
        model.compile(optimizer=Adam(learning_rate=1e-4, clipnorm=1.0),
                      loss=SparseCategoricalCrossentropy(from_logits=False, ignore_class=0),
                      metrics=[masked_accuracy])
    return model

训练日志片段

history=model.fit(
        x=[train_images,train_input_captions],y=train_label_captions,
        epochs=50,
        batch_size=8,
        validation_data=([dev_images,dev_input_captions],dev_label_captions),
        callbacks=[NaNLossCallback(),debug_callback]
    )

Epoch 1/50
I0000 00:00:1749020366.186489    1026 cuda_dnn.cc:529] Loaded cuDNN version 90300
I0000 00:00:1749020366.445219    1028 cuda_dnn.cc:529] Loaded cuDNN version 90300
Batch 0: Logits max = 0.0634, min = -0.0696
1/708 ━━━━━━━━━━━━━━━━━━━━ 2:16:45 12s/step - loss: 12.8995 - masked_accuracy:0.0000e+00Batch 1: Logits max = 0.0622, min = -0.0707
...
120/708 ━━━━━━━━━━━━━━━━━━━━ 3:41 376ms/step - loss: 12.8935 - masked_accuracy: 0.0118Batch 120: Logits max = 3.4171, min = -2.2954
121/708 ━━━━━━━━━━━━━━━━━━━━ 3:40 376ms/step - loss: 12.8935 - masked_accuracy: 0.0118Batch 121: Logits max = 3.4450, min = -2.3163
122/708 ━━━━━━━━━━━━━━━━━━━━ 3:40 376ms/step - loss: inf - masked_accuracy: 0.0118    Batch 122: Logits max = 3.4731, min = -2.3371
123/708 ━━━━━━━━━━━━━━━━━━━━ 3:40 376ms/step - loss: inf - masked_accuracy: 0.0118Batch 123: Logits max = 3.5013, min = -2.3580
124/708 ━━━━━━━━━━━━━━━━━━━━ 3:39 376ms/step - loss: inf - masked_accuracy: 0.0118NaN loss at batch 124
Batch 124: Logits max = 3.5296, min = -2.3789
708/708 ━━━━━━━━━━━━━━━━━━━━ 78s 94ms/step - loss: nan - masked_accuracy: 0.0121 - val_loss: nan - val_masked_accuracy: nan

原因分析

  • 损失函数参数不匹配:模型输出的是未经过softmax的logits,但SparseCategoricalCrossentropy设置了from_logits=False,会强制对logits做softmax计算。当logits数值增大时,softmax的结果会出现极端值(接近0或1),计算交叉熵时会触发log(0)的情况,直接产生inf进而转为NaN。
  • padding未被正确屏蔽:字幕序列存在0填充,但Embedding层设置mask_zero=False,导致padding位置的向量参与了LSTM和注意力计算,引入无效梯度,加剧训练过程中的数值不稳定。
  • 梯度累积风险:尽管设置了clipnorm=1.0,但全量数据下样本多样性带来的梯度波动更大,LSTM与注意力层的组合可能导致梯度在传播过程中被放大,进而推动logits持续增大。

解决方法

  1. 修正损失函数参数:将from_logits改为True,让TensorFlow内部处理logits到概率的转换,避免手动softmax带来的数值不稳定:
    loss=SparseCategoricalCrossentropy(from_logits=True, ignore_class=0)
    
  2. 启用mask_zero屏蔽padding:修改Embedding层参数,让0填充的位置不参与后续计算,减少无效梯度:
    embedding_layer = Embedding(400004, 300, trainable=False, mask_zero=True)
    
  3. 降低学习率:将当前1e-4的学习率调整为5e-5或1e-5,减小梯度更新幅度,避免参数波动过大。
  4. 增强梯度裁剪:把clipnorm从1.0降低到0.5,进一步限制梯度范数,防止梯度爆炸。
  5. 检查自定义指标:确保masked_accuracy正确忽略class 0的padding部分,避免指标计算中出现NaN干扰训练流程。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 23:30:53