强化学习项目中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

