tidymodels中MLP模型损失函数出现NaN的原因排查求助
MLP训练损失与MAE全为NaN的问题排查与解决
任务与数据背景
使用电子游戏销售数据集,构建MLP模型预测日本销量(JP_Sales),排除特征包括Rank、Name及与目标高度相关的Global_Sales。
预处理代码
vgames <- read_csv('data/vgsales.csv', show_col_types = FALSE, col_types = list( Year = col_date("%Y") )) %>% mutate( Platform = factor(Platform), Genre = factor(Genre), Publisher = factor(Publisher) ) vgames_model <- vgames %>% select(-c(Rank, Name, Global_Sales)) # 训练测试集划分 vgames_split <- vgames_model %>% initial_split() vgames_training <- vgames_split %>% training() vgames_testing <- vgames_split %>% testing() # 交叉验证折数 vgames_folds <- vgames_training %>% vfold_cv(v = 10) # 预处理流程 vgames_recipe <- vgames_training %>% recipe(formula = JP_Sales ~ .) %>% step_normalize(all_numeric_predictors()) %>% step_date(Year, features = c("year"), keep_original_cols = FALSE) %>% step_dummy(all_nominal()) %>% step_zv(all_numeric_predictors())
预处理后数据情况
预处理后生成570个特征,示例如下:
# A tibble: 12,448 × 570 NA_Sales EU_Sales Other_…¹ JP_Sa…² Year_…³ Platf…⁴ Platf…⁵ Platf…⁶ Platf…⁷ Platf…⁸ Platf…⁹ Platf…˟ Platf…˟ Platf…˟ <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> <dbl> 1 -0.272 -0.279 -0.240 0 2006 0 0 0 0 0 0 0 0 0 2 0.145 0.258 0.0629 0 2012 0 0 0 0 0 0 0 0 0 3 -0.198 -0.241 -0.189 0.07 2008 0 0 0 1 0 0 0 0 0 4 -0.149 -0.260 -0.189 0 2010 0 0 0 1 0 0 0 0 0 5 -0.149 -0.0679 -0.0380 0 2006 0 0 0 0 0 0 0 0 0 6 -0.296 -0.183 -0.189 0 2015 0 1 0 0 0 0 0 0 0 7 3.32 1.05 0.315 1.81 1988 0 0 0 0 0 0 0 0 0 8 -0.308 -0.260 -0.240 0 2016 0 0 0 0 0 0 0 0 0 9 -0.321 -0.202 -0.240 0 2015 0 0 0 0 0 0 0 0 0 10 -0.112 -0.145 -0.139 0 2010 0 0 0 0 0 0 0 0 0 # … with 12,438 more rows, 556 more variables: Platform_N64 <dbl>, Platform_NES <dbl>, Platform_NG <dbl>, # Platform_PC <dbl>, Platform_PCFX <dbl>, Platform_PS <dbl>, Platform_PS2 <dbl>, Platform_PS3 <dbl>, # Platform_PS4 <dbl>, Platform_PSP <dbl>, Platform_PSV <dbl>, Platform_SAT <dbl>, Platform_SCD <dbl>, # Platform_SNES <dbl>, Platform_TG16 <dbl>, Platform_Wii <dbl>, Platform_WiiU <dbl>, Platform_WS <dbl>, # Platform_X360 <dbl>, Platform_XB <dbl>, Platform_XOne <dbl>, Genre_Adventure <dbl>, Genre_Fighting <dbl>, # Genre_Misc <dbl>, Genre_Platform <dbl>, Genre_Puzzle <dbl>, Genre_Racing <dbl>, Genre_Role.Playing <dbl>, # Genre_Shooter <dbl>, Genre_Simulation <dbl>, Genre_Sports <dbl>, Genre_Strategy <dbl>, … # ℹ Use `print(n = ...)` to see more rows, and `colnames()` to see all variable names
模型训练问题
定义并拟合MLP模型时,所有epoch的损失与MAE均返回NaN:
nn <- mlp(epochs = 20) %>% set_engine('keras', verbose = 1, metrics = c("mae"), optimizer = 'adam', loss = 'mean_absolute_error') %>% set_mode('regression') nnwf <- workflow() %>% add_model(nn) %>% add_recipe(vgames_recipe) nnwf %>% fit(vgames_training)
训练输出:
... Epoch 16/20 389/389 [==============================] - 1s 1ms/step - loss: nan - mae: nan Epoch 17/20 389/389 [==============================] - 1s 1ms/step - loss: nan - mae: nan Epoch 18/20 389/389 [==============================] - 1s 2ms/step - loss: nan - mae: nan Epoch 19/20 389/389 [==============================] - 1s 2ms/step - loss: nan - mae: nan Epoch 20/20 389/389 [==============================] - 1s 1ms/step - loss: nan - mae: nan
已尝试调整归一化时机、降低学习率、移除日期列,均未解决问题。
问题原因分析
- 缺失值未处理:原数据集Year列存在大量缺失值,转成日期类型后变为NA,后续预处理未对缺失值进行填充或移除,导致特征中存在NA,训练时触发NaN计算。
- 特征维度爆炸:Publisher类别过多(原数据集有数百个发行商),哑变量编码后生成大量稀疏特征,加上提取后的Year为原始年份数值(如2006)未被归一化,导致输入特征量级差异极大,模型梯度更新时出现数值溢出。
- 目标变量分布极端:JP_Sales存在大量0值和少量极端大值,右偏分布可能导致损失计算过程中出现异常。
解决方案
1. 修复预处理流程,处理缺失值与特征归一化
修改预处理代码,先处理缺失值,合并低频次类别减少特征维度,确保所有数值特征被归一化:
vgames_recipe <- vgames_training %>% recipe(formula = JP_Sales ~ .) %>% # 用中位数填充Year缺失值 step_impute_median(Year) %>% # 先提取年份特征,再归一化所有数值变量 step_date(Year, features = c("year"), keep_original_cols = FALSE) %>% step_normalize(all_numeric_predictors()) %>% # 合并占比低于1%的发行商为"Other",减少哑变量数量 step_other(Publisher, threshold = 0.01, other = "Other") %>% step_dummy(all_nominal()) %>% step_zv(all_numeric_predictors())
2. 优化模型训练参数,避免梯度异常
添加梯度裁剪限制梯度范围,同时调整模型隐藏层规模:
# 自定义MLP模型,设置梯度裁剪与更小的学习率 nn_custom <- mlp(epochs = 30, hidden_units = c(64, 32)) %>% set_engine('keras', verbose = 1, metrics = c("mae"), optimizer = optimizer_adam(learning_rate = 1e-4, clipnorm = 1.0), loss = 'mean_absolute_error') %>% set_mode('regression') nnwf <- workflow() %>% add_model(nn_custom) %>% add_recipe(vgames_recipe) nnwf %>% fit(vgames_training)
3. 处理目标变量分布
对JP_Sales做对数变换,降低极端值对训练的影响:
vgames_recipe <- vgames_training %>% recipe(formula = JP_Sales ~ .) %>% step_impute_median(Year) %>% step_date(Year, features = c("year"), keep_original_cols = FALSE) %>% step_normalize(all_numeric_predictors()) %>% step_other(Publisher, threshold = 0.01, other = "Other") %>% step_dummy(all_nominal()) %>% step_zv(all_numeric_predictors()) %>% # 目标变量对数变换,加极小值避免log(0) step_log(JP_Sales, offset = 1e-6)
训练完成后,对预测结果做指数变换还原真实销量:exp(predicted_values) - 1e-6
内容的提问来源于stack exchange,提问作者Giulio Mario Martena
相关产品推荐
相关产品推荐

