Keras与Pytorch训练简单LSTM得到不同结果的问题求解
核心问题及修复方案
1. LSTM初始化与Keras不匹配(最主要原因)
Keras的LSTM默认配置和PyTorch存在显著差异,直接影响训练效果:
- Keras默认开启
unit_forget_bias=True,会将忘记门的偏置初始化为1,避免训练初期梯度消失,PyTorch无该默认设置 - Keras的LSTM核权重默认用Xavier均匀初始化,循环权重默认用正交初始化,PyTorch默认用Kaiming均匀初始化
- 全连接层的权重初始化Keras默认用Xavier均匀,PyTorch默认用Kaiming均匀
修复代码如下,在LSTM类的__init__方法中补充初始化逻辑:
class LSTM(nn.Module): def __init__(self, input_dim, hidden_dim, num_layers, output_dim,bilstm=False): super(LSTM, self).__init__() self.hidden_dim = hidden_dim self.num_layers = num_layers self.isBi = bilstm self.lstm = nn.LSTM(input_dim, hidden_dim, num_layers, batch_first=True,bidirectional=bilstm).double() # 对齐Keras LSTM初始化 for name, param in self.lstm.named_parameters(): if 'weight_ih' in name: nn.init.xavier_uniform_(param.data) elif 'weight_hh' in name: nn.init.orthogonal_(param.data) elif 'bias' in name: param.data.fill_(0) # 忘记门偏置设为1,对应Keras unit_forget_bias param.data[hidden_dim:2*hidden_dim] = 1.0 # 全连接层对齐Keras初始化 self.fc1 = nn.Sequential(nn.Linear(hidden_dim, 10).double(),nn.Tanh()) nn.init.xavier_uniform_(self.fc1[0].weight) nn.init.zeros_(self.fc1[0].bias) self.final_layer1 = nn.Sequential(nn.Linear(10,10).double(),nn.Tanh()) nn.init.xavier_uniform_(self.final_layer1[0].weight) nn.init.zeros_(self.final_layer1[0].bias) self.final_layer2 = nn.Sequential(nn.Linear(10,10).double(),nn.Tanh()) nn.init.xavier_uniform_(self.final_layer2[0].weight) nn.init.zeros_(self.final_layer2[0].bias) self.final_layer3 = nn.Sequential(nn.Linear(10,10).double(),nn.Tanh()) nn.init.xavier_uniform_(self.final_layer3[0].weight) nn.init.zeros_(self.final_layer3[0].bias) self.final_layer4 = nn.Sequential(nn.Linear(10,output_dim).double()) nn.init.xavier_uniform_(self.final_layer4[0].weight) nn.init.zeros_(self.final_layer4[0].bias) def forward(self, x): out, (hn, cn) = self.lstm(x) out = out[:, -1, :] out = self.fc1(out) out = self.final_layer1(out) out = self.final_layer2(out) out = self.final_layer3(out) out = self.final_layer4(out) return out
2. 训练参数&输入对齐检查
- 确认调用LSTM类时
num_layers=1、bilstm=False,和Keras的单层单向LSTM对齐 - 确认PyTorch的batch size、学习率、训练轮数和Keras侧完全一致
- 确认输入数据的归一化逻辑和Keras侧完全一致,比如Keras做了0-1归一化,PyTorch侧也要用相同的统计值做归一化
- 可移除冗余的
Variable调用,PyTorch 0.4及以后版本张量本身就支持自动求导,不需要额外包裹
3. 激活函数差异适配
如果需要完全对齐Keras的激活函数,Keras LSTM默认的循环激活是hard_sigmoid,PyTorch原生LSTM不支持直接修改,可自定义LSTMCell替换nn.LSTM实现,一般场景下sigmoid和hard_sigmoid差异不大,优先修改初始化即可。
内容的提问来源于stack exchange,提问作者CHan Ricky
相关产品推荐
相关产品推荐

