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

Keras(TensorFlow)训练出现loss: nan问题排查求助

问题描述

依赖配置(requirements.txt):

absl-py==2.0.0
astunparse==1.6.3
cachetools==5.3.1
certifi==2023.7.22
charset-normalizer==3.2.0
flatbuffers==23.5.26
gast==0.4.0
google-auth==2.23.0
google-auth-oauthlib==1.0.0
google-pasta==0.2.0
grpcio==1.58.0
h5py==3.9.0
idna==3.4
keras==2.13.1
libclang==16.0.6
Markdown==3.4.4
MarkupSafe==2.1.3
numpy==1.24.3
oauthlib==3.2.2
opt-einsum==3.3.0
packaging==23.1
pandas==2.1.1
protobuf==4.24.3
pyasn1==0.5.0
pyasn1-modules==0.3.0
python-dateutil==2.8.2
pytz==2023.3.post1
requests==2.31.0
requests-oauthlib==1.3.1
rsa==4.9
six==1.16.0
tensorboard==2.13.0
tensorboard-data-server==0.7.1
tensorflow==2.13.0
tensorflow-estimator==2.13.0
tensorflow-intel==2.13.0
tensorflow-io-gcs-filesystem==0.31.0
termcolor==2.3.0
typing_extensions==4.5.0
tzdata==2023.3
urllib3==1.26.16
Werkzeug==2.3.7
wrapt==1.15.0 

运行的Keras回归代码:

import pandas as pd
import numpy as np

from keras.models import Sequential
from keras.layers import Dense
from keras.optimizers import SGD

# csv load
csv_file = 'point.csv'

# df change
df = pd.read_csv(csv_file)

# df numpy change
sampledata = df.to_numpy()

# a, b columns
a = sampledata[:, [0, 1]]
b = sampledata[:, 2]

# print(np.array(a))
# print(np.array(b))

a_data = np.array(a)
b_data = np.array(b)

print("===============")

print('x', a.shape, 'y', b.shape)

model = Sequential()

model.add(Dense(1, input_shape=(2, ), activation='linear'))

model.compile(optimizer=SGD(learning_rate=1e-2), loss='mse')

model.summary()

hist = model.fit(a, b, epochs=1000)

# result print
 
csv_file = 'point_copy.csv'
df = pd.read_csv(csv_file)
sampledata = df.to_numpy()

a = sampledata[:, [0, 1]]

a_data = np.array(a)

result = model.predict(a_data)

print(result)

训练输出(loss始终为nan):

Model: "sequential"
_________________________________________________________________
 Layer (type)                Output Shape              Param #
=================================================================
 dense (Dense)               (None, 1)                 3

=================================================================
Total params: 3 (12.00 Byte)
Trainable params: 3 (12.00 Byte)
Non-trainable params: 0 (0.00 Byte)
_________________________________________________________________
Epoch 1/1000
27/27 [==============================] - 0s 1ms/step - loss: nan
Epoch 2/1000
27/27 [==============================] - 0s 1ms/step - loss: nan
Epoch 3/1000
27/27 [==============================] - 0s 1ms/step - loss: nan

训练数据point.csv约842行,特征为date、time,标签point为大数值;预测数据为point_copy.csv,怀疑数据量级过大导致loss为nan,需排查原因并给出解决方法。

原因分析
  • 大数值引发数值溢出:标签point为大数值时,MSE损失是预测值与真实值差的平方,数值量级会急剧放大,超出浮点数表示范围直接变成nan;同时SGD计算梯度时也会因大数值溢出,导致参数变为nan,后续迭代全部失效。
  • 特征未做预处理:date、time若以原始大数值(如时间戳、日期整数)输入,会和标签的大数值共同放大计算范围,加剧溢出问题。
  • 学习率不匹配:当前1e-2的学习率对大数值数据过高,参数更新步幅太大,一步就跳出有效数值范围,直接出现nan。
  • 数据存在无效值:若point.csv中存在nan、inf等无效值,会直接导致计算异常。
解决方法

1. 数据标准化/归一化

将特征和标签缩放至小范围(如[-1,1]或[0,1]),消除量级差异:

from sklearn.preprocessing import StandardScaler

# 特征标准化
scaler_x = StandardScaler()
a_scaled = scaler_x.fit_transform(a)

# 标签标准化
scaler_y = StandardScaler()
b_scaled = scaler_y.fit_transform(b.reshape(-1, 1)).flatten()

# 用缩放后的数据训练
hist = model.fit(a_scaled, b_scaled, epochs=1000)

# 预测时先缩放输入,再反缩放结果得到真实值
a_test_scaled = scaler_x.transform(a_data)
result_scaled = model.predict(a_test_scaled)
result = scaler_y.inverse_transform(result_scaled)
print(result)

2. 调整学习率

配合数据预处理降低学习率,避免参数更新步幅过大:

model.compile(optimizer=SGD(learning_rate=1e-5), loss='mse')

3. 检查并清理无效数据

确认数据中是否存在nan、inf等无效值,及时处理:

# 检查缺失值和无穷值
print(df.isnull().sum())
print(np.isinf(sampledata).any())

# 删除缺失值行
df = df.dropna()
# 或用均值填充缺失值
df = df.fillna(df.mean())

4. 替换为自适应优化器

使用Adam等自适应学习率优化器,自动调整参数更新步幅:

from keras.optimizers import Adam

model.compile(optimizer=Adam(learning_rate=1e-3), loss='mse')

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 23:38:11