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

Google Colab中ResNet50在TPU训练时损失与精度为NaN,CPU正常

TPU训练ResNet50出现NaN值的问题排查与解决

问题场景

在Google Colab中使用v2-8 TPU加速器训练ResNet50模型,输入为5000张形状(224, 224, 3)的归一化图片,数据已验证无NaN、无穷值及类别不平衡问题。CPU环境下训练完全正常,但切换到TPU后,训练过程中损失和精度变为NaN。训练代码如下:

INPUT_SHAPE = (224, 224, 3)

with strategy.scope():
    base_model = ResNet50(weights='imagenet', include_top=False, input_shape=INPUT_SHAPE)
    base_model.trainable = False
    model = tf.keras.models.Sequential([
        base_model,
        tf.keras.layers.GlobalAveragePooling2D(),
        tf.keras.layers.Dense(1024, activation='relu'),
        tf.keras.layers.Dense(6, activation='sigmoid') 
    ])
    model.compile(optimizer='adam', 
                  loss='binary_crossentropy',
                  metrics=['accuracy'])
    
model.fit(
    X_train, 
    y_train, 
    epochs=10, 
    validation_data=(X_val, y_val),
    batch_size=32
)

核心原因

  1. 有效批量过大引发梯度爆炸:TPU采用分布式训练策略,v2-8包含8个核心,代码中设置的batch_size=32会被放大为32*8=256的有效批量。Adam优化器默认学习率(1e-4)在大批量下会导致梯度更新幅度过大,引发梯度爆炸,最终出现NaN。
  2. 数据类型不匹配:TPU对浮点数据类型兼容性更严格,若输入数据为float64类型,会导致计算过程中出现数值异常。
  3. 损失函数数值稳定性问题:sigmoid激活配合binary_crossentropy计算时,当输出接近0或1时,log运算会出现数值下溢,TPU的浮点计算特性会放大这一问题,产生NaN。

解决方法

1. 调整学习率适配大有效批量

将Adam优化器的学习率缩小至原来的1/8(对应TPU核心数),例如设置为1e-5:

optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5)

2. 统一数据类型为float32

确保输入数据和标签都转换为TPU更友好的float32类型:

X_train = X_train.astype('float32')
X_val = X_val.astype('float32')
y_train = y_train.astype('float32')
y_val = y_val.astype('float32')

3. 使用数值稳定的损失函数变体

改用from_logits=True的BinaryCrossentropy损失,同时移除最后一层的sigmoid激活(改用线性激活),避免sigmoid与交叉熵计算时的数值不稳定:

# 模型最后一层修改
tf.keras.layers.Dense(6, activation='linear')
# 编译时的损失设置
loss=tf.keras.losses.BinaryCrossentropy(from_logits=True)

4. 添加梯度裁剪

在优化器中启用梯度裁剪,限制梯度的最大范数,防止梯度爆炸:

optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5, clipnorm=1.0)

5. 确保TPU策略正确初始化

在代码开头添加完整的TPU初始化流程,避免策略配置不完整导致的异常:

import tensorflow as tf
resolver = tf.distribute.cluster_resolver.TPUClusterResolver()
tf.config.experimental_connect_to_cluster(resolver)
tf.tpu.experimental.initialize_tpu_system(resolver)
strategy = tf.distribute.TPUStrategy(resolver)

修改后的完整代码示例

import tensorflow as tf
from tensorflow.keras.applications import ResNet50

# 初始化TPU策略
resolver = tf.distribute.cluster_resolver.TPUClusterResolver()
tf.config.experimental_connect_to_cluster(resolver)
tf.tpu.experimental.initialize_tpu_system(resolver)
strategy = tf.distribute.TPUStrategy(resolver)

INPUT_SHAPE = (224, 224, 3)
# 转换数据类型为float32
X_train = X_train.astype('float32')
X_val = X_val.astype('float32')
y_train = y_train.astype('float32')
y_val = y_val.astype('float32')

with strategy.scope():
    base_model = ResNet50(weights='imagenet', include_top=False, input_shape=INPUT_SHAPE)
    base_model.trainable = False
    model = tf.keras.models.Sequential([
        base_model,
        tf.keras.layers.GlobalAveragePooling2D(),
        tf.keras.layers.Dense(1024, activation='relu'),
        tf.keras.layers.Dense(6, activation='linear') 
    ])
    model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5, clipnorm=1.0), 
                  loss=tf.keras.losses.BinaryCrossentropy(from_logits=True),
                  metrics=['accuracy'])
    
model.fit(
    X_train, 
    y_train, 
    epochs=10, 
    validation_data=(X_val, y_val),
    batch_size=32
)

内容的提问来源于stack exchange,提问作者احمد القيسي

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 01:13:25