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

训练含目标值缺失的LSTM时optim.step()后权重出现NaN求助

目标含缺失值的LSTM训练时权重/偏置变NaN的排查思路

我在训练目标数据含缺失值的LSTM神经网络时,使用自定义损失函数,执行optim.step()后出现部分权重/偏置变为NaN的错误,求排查思路。

复现代码

import torch
import numpy as np
from torch import nn

# 定义简单LSTM模型
class myLSTM(nn.Module):
    def __init__(self, input_size, hidden_size):
        super(myLSTM, self).__init__()
        self.lstm = nn.LSTM(input_size, hidden_size)
    def forward(self, input):
        output, _ = self.lstm(input)
        return output

# 输入和目标数据
input = torch.randn(10, 5, requires_grad=True)
target = torch.randn(10, 5)

# 目标数据中设置一个缺失值
target[0,0] = np.nan

# 创建模型
lstmModel = myLSTM(5, 5) 

# 损失函数与优化器
def loss_function(y_true, y_predict):
    return torch.nanmean((y_true-y_predict)**2)

optim = torch.optim.Adam(lstmModel.parameters(), lr=0.01)

# 单轮训练
output = lstmModel(input)
optim.zero_grad()
error = loss_function(target, output)
error.backward()
optim.step()

lstmModel.state_dict()

排查思路

  • 检查梯度是否存在NaN
    在error.backward()后,打印模型参数的梯度,确认是否有梯度为NaN的情况:

    for name, param in lstmModel.named_parameters():
        if param.grad is not None:
            print(f"{name}: {torch.isnan(param.grad).any()}")
    

    torch.nanmean会忽略NaN损失项,但反向传播时若对应位置梯度计算出现异常(如0/0),会生成NaN梯度,进而导致参数更新后变为NaN。

  • 替换损失函数的NaN处理逻辑
    改用显式掩码过滤NaN的方式,比torch.nanmean更可控,避免梯度计算异常:

    def loss_function(y_true, y_predict):
        # 生成非NaN的掩码
        mask = ~torch.isnan(y_true)
        # 只计算有效数据的MSE
        return torch.mean((y_true[mask] - y_predict[mask])**2)
    
  • 移除输入数据的梯度追踪
    输入数据通常不需要计算梯度,当前代码中input设置了requires_grad=True,这可能引发额外的梯度计算异常,修改为:

    input = torch.randn(10, 5)
    
  • 降低学习率尝试
    即使使用Adam优化器,若梯度存在异常大的值,乘以学习率后可能导致参数溢出变为NaN。可尝试将学习率调低至0.001,观察是否解决问题。

  • 显式初始化LSTM的初始状态
    默认情况下LSTM的初始h0和c0为全0,但显式初始化可避免潜在的设备不匹配或默认初始化问题:

    def forward(self, input):
        device = input.device
        hidden_size = self.lstm.hidden_size
        # 初始化h0和c0,与输入同设备
        h0 = torch.zeros(1, input.size(0), hidden_size, device=device)
        c0 = torch.zeros(1, input.size(0), hidden_size, device=device)
        output, _ = self.lstm(input, (h0, c0))
        return output
    
  • 统一数据类型与设备
    确保输入、目标、模型参数都在同一设备(CPU/GPU)且数据类型一致(如均为float32),避免跨设备/类型计算导致的异常。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 01:46:15