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

强化学习项目中LSTM损失无法优化及梯度消失问题求解

强化学习项目梯度消失问题的解决方案

项目概述

我正在开展一项强化学习项目,采用MiniGrid作为测试环境验证思路,网络架构为卷积神经网络(CNN)+长短期记忆网络(LSTM):先通过CNN提取视觉环境特征,再将特征输入LSTM处理轨迹序列信息,最后通过MLP输出预测结果。

用PyTorch实现的网络代码如下:

class RRLSTM(nn.Module):
def __init__(self, config):
    super(RRLSTM, self).__init__()
    
    # 从配置文件读取参数
    cnn_features = config["CNN"]
    lstm_features = config["LSTM"]
    mlp_features = config["MLP"]
    
    # CNN相关参数
    nonlinearity = torch.relu
    input_channels = cnn_features["input_channels"]
    channels = cnn_features["channels"]
    kernel_sizes = cnn_features["kernel_sizes"]
    strides = cnn_features["strides"]
    paddings = cnn_features["paddings"]
    use_maxpool = cnn_features["use_maxpool"]
    
    # MLP相关参数
    fc_input_size = cnn_features["fc_input_size"]
    fc_hidden_sizes = mlp_features["fc_hidden_sizes"]
    embed_size = mlp_features["embed_size"]
    
    # LSTM相关参数
    n_units = lstm_features["n_units"]
    input_size = lstm_features["input_size"]
    
    # 构建CNN层
    if paddings is None:
        paddings = [0 for _ in range(len(channels))]
    assert len(channels) == len(kernel_sizes) == len(strides) == len(paddings)
    in_channels = [input_channels] + channels[:-1]
    post_activation_fns = [identity for _ in range(len(strides))]
    ones = [1 for _ in range(len(strides))]
    if use_maxpool:
        post_activation_fns = [torch.nn.MaxPool2d(max_pool_stride) for max_pool_stride in strides]
        strides = ones
    activation_fns = [nonlinearity for _ in range(len(strides))]
    conv_layers = [CNNLayer(input_channels=ic, output_channels=oc,
                            kernel_size=k, stride=s, padding=p, activation_fn=a_fn, post_activation_fn=p_fn)
                   for (ic, oc, k, s, p, a_fn, p_fn) in zip(in_channels, channels, kernel_sizes, strides, paddings,
                                                            activation_fns, post_activation_fns)]
    # CNN后的线性层
    linear_layer = Linear(fc_input_size, input_size)
    
    # 构建MLP层
    lstm_fc_layers = MLP(
        input_size=n_units,
        output_size=1,
        hidden_sizes=fc_hidden_sizes,
        hidden_activation=nonlinearity,
        output_activation=identity
    )
    
    # 动作嵌入层
    action_embedding_layers = Embed(
        embed_dim=embed_size
    )
    
    # LSTM层
    lstm_layer = LSTM(
        fc_input_size = input_size, 
        embed_dim = embed_size, 
        n_units = n_units
    )
    
    # 组装整个网络
    self.model = Network(conv_layers, linear_layer, action_embedding_layers, lstm_layer, lstm_fc_layers)
    

def forward(self, input):
    return self.model.forward(input)

网络配置涵盖CNN的通道数、核大小、步长,LSTM的单元数,MLP的隐藏层尺寸等核心参数。

损失函数设计

参考相关论文思路,损失函数为主损失+辅助损失的加权和:主损失针对LSTM序列最终时刻的预测值与标签的误差,辅助损失针对序列所有时刻的预测误差。由于批次包含变长轨迹,用mask屏蔽填充部分的无效计算,损失函数代码如下:

def calculate_loss(self, predicted_G0, returns, length):

    if not torch.is_tensor(returns):
       returns = torch.tensor(returns)
       returns = returns.to(device)
    
    # B x L 的全时刻损失
    all_timestep_loss = self.mse_loss(predicted_G0, returns.repeat(1, 
                                      predicted_G0.size(1)))
    
    # 创建mask
    self.mask = torch.zeros_like(all_timestep_loss)
    for l_num, l in enumerate(length):
        self.mask[l_num, :l] = 1
        
    # 用mask忽略填充部分
    all_timestep_loss = all_timestep_loss * self.mask

    # 每个序列的平均损失
    self.mean_all_timestep_loss_along_sequence = all_timestep_loss.sum(1) / self.mask.sum(1)
    
    # 批次平均的辅助损失
    mean_loss = self.mean_all_timestep_loss_along_sequence.mean()
    aux_loss = self.continuous_pred_factor * mean_loss

    # 提取每个序列最后时刻的损失作为主损失
    len = length[:] - 1
    all_timestep_loss_indexed = all_timestep_loss[range(predicted_G0.size(0)), len]

    main_loss = all_timestep_loss_indexed.mean()

    # 总损失
    lstm_loss = main_loss + aux_loss

return lstm_loss, main_loss, aux_loss

梯度消失问题与解决方案

训练中发现网络存在梯度消失问题:第一层CNN的权重梯度几乎为0,而最后一层线性层的梯度正常,说明梯度在反向传播时未能有效传递到前端网络。当前使用ReLU激活函数,除了缩减网络深度,还可以尝试以下方案:

1. 替换激活函数

  • 改用LeakyReLU/Parametric ReLU:解决ReLU负区间梯度为0导致的神经元死亡问题,让梯度在更多场景下流动。
  • 改用GELU/Swish:这类平滑激活函数在全区间都有非零梯度,相比ReLU更利于梯度传播,适配CNN+LSTM的架构。

2. 梯度裁剪

在反向传播时限制梯度的范数,防止梯度爆炸的同时避免梯度被异常值稀释,PyTorch实现示例:

torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

将所有参数的梯度范数控制在指定范围内,保证梯度稳定传递。

3. 添加残差连接

在CNN层之间引入残差连接,让梯度可以直接跳过部分层传递,缓解深度网络的梯度消失问题。示例实现:

class ResidualCNNLayer(nn.Module):
    def __init__(self, ic, oc, k, s, p, activation_fn, post_activation_fn):
        super().__init__()
        self.conv = CNNLayer(ic, oc, k, s, p, activation_fn, post_activation_fn)
        self.shortcut = nn.Identity() if ic == oc else nn.Conv2d(ic, oc, kernel_size=1, stride=s)
    
    def forward(self, x):
        return self.conv(x) + self.shortcut(x)

用ResidualCNNLayer替换原有的CNNLayer构建带残差的CNN模块。

4. 优化权重初始化

  • He初始化:针对ReLU类激活函数,让初始化的权重方差更合理,保证每层输出和梯度的方差稳定,避免梯度逐渐衰减。PyTorch中可通过init.kaiming_normal_实现。
  • LSTM特殊初始化:对LSTM的遗忘门偏置初始化为较大值(如1),增强LSTM的长期记忆能力,让梯度在序列维度更易传递。

5. 调整损失函数权重

当前主损失仅关注序列最后时刻的预测,可能导致早期时刻的梯度信号较弱:

  • 增大辅助损失的权重continuous_pred_factor,让模型更关注所有时刻的预测,增强前端网络的梯度信号。
  • 对序列不同时刻的损失加权,比如越靠近末尾的时刻权重越大,同时保留早期时刻的梯度贡献。

6. 添加归一化层

在CNN层之间加入BatchNorm2d,或在LSTM之后加入LayerNorm,稳定每层的输入分布,减少内部协变量偏移,让梯度更稳定地传播。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 19:54:55