PyTorch中LSTM的num_layers参数等效手动堆叠实现咨询
PyTorch 手动堆叠多层LSTM等效实现方案
前置修正:层定义参数调整
你给出的3个单层LSTM定义存在参数错误,官方num_layers=N的堆叠LSTM,仅第一层的输入维度为指定的input_size,后续所有层的输入维度均等于前一层的hidden_size,因此正确的层定义应为:
import torch import torch.nn as nn layer1 = nn.LSTM(input_size=128, hidden_size=512, num_layers=1) layer2 = nn.LSTM(input_size=512, hidden_size=512, num_layers=1) layer3 = nn.LSTM(input_size=512, hidden_size=512, num_layers=1)
等效forward实现
该实现完全兼容官方3层LSTM的state输入输出格式,跨批次传递state的用法无需调整:
def forward(x, state): # 拆分全局state为每层对应的h、c # 输入state格式和官方完全一致:(h_all, c_all),形状均为(3, batch_size, 512) h_list = state[0].unbind(0) c_list = state[1].unbind(0) out_h = [] out_c = [] # 第一层计算 x, (cur_h, cur_c) = layer1(x, (h_list[0].unsqueeze(0), c_list[0].unsqueeze(0))) out_h.append(cur_h) out_c.append(cur_c) # 若需要对齐官方dropout逻辑,可在此处加Dropout层(官方仅非最后一层加dropout) # 第二层计算 x, (cur_h, cur_c) = layer2(x, (h_list[1].unsqueeze(0), c_list[1].unsqueeze(0))) out_h.append(cur_h) out_c.append(cur_c) # 若需要对齐官方dropout逻辑,可在此处加Dropout层 # 第三层计算 x, (cur_h, cur_c) = layer3(x, (h_list[2].unsqueeze(0), c_list[2].unsqueeze(0))) out_h.append(cur_h) out_c.append(cur_c) # 拼接所有层的state,返回格式和官方完全一致 state_out = (torch.cat(out_h, dim=0), torch.cat(out_c, dim=0)) return x, (state_out[0].detach(), state_out[1].detach())
关键说明
- state的处理逻辑完全对齐官方实现:输入的全局state按层拆分后传入对应单层LSTM,输出的每层state再拼接为全局state返回,和直接使用
num_layers=3的LSTM的输入输出格式100%兼容。 - 如果需要对齐官方LSTM的
dropout参数效果,只需在第一层、第二层的输出后加对应概率的Dropout层即可,官方实现默认仅在非最后一层的输出后添加dropout。
内容的提问来源于stack exchange,提问作者teamclouday
相关产品推荐
相关产品推荐

