PyTorch LSTM层output/hidden/cell维度与预期不符问题咨询
PyTorch LSTM层输出维度规则与异常排查
PyTorch nn.LSTM 标准维度规则
nn.LSTM的输入输出维度由batch_first参数控制,默认值为False,双向/单向LSTM维度规则统一(单向时num_directions=1,双向时为2):
- 当
batch_first=False(默认配置)- 输入张量形状要求:
(seq_len, batch_size, input_size),三个维度依次为序列长度、批大小、单时间步输入特征维度 output返回值形状:(seq_len, batch_size, hid_dim * num_directions),存储最后一层所有时间步的隐藏状态hidden/cell返回值形状:(n_layers * num_directions, batch_size, hid_dim),存储所有层在最后一个时间步的隐藏状态、细胞状态,注意该返回值永远以层维度作为第一维,不受batch_first参数影响
- 输入张量形状要求:
- 当
batch_first=True- 输入张量形状要求:
(batch_size, seq_len, input_size),三个维度依次为批大小、序列长度、单时间步输入特征维度 output返回值形状:(batch_size, seq_len, hid_dim * num_directions)hidden/cell返回值形状:(n_layers * num_directions, batch_size, hid_dim),维度顺序和batch_first=False时完全一致
- 输入张量形状要求:
维度异常原因
当前维度不符合预期,核心原因是初始化LSTM时未设置batch_first=True,但传入的嵌入张量是(batch_size, seq_len, emb_dim)的batch在前格式,和LSTM默认的维度顺序不匹配。
结合测试输出可以验证这个判断:
- 传入单条文档时,嵌入层输出形状为
[1, 850, 100],LSTM默认将第一维1识别为序列长度、第二维850识别为批大小,因此输出:output形状为[1, 850, 100]:数值上和预期的batch在前格式结果重合,属于巧合hidden/cell形状为[2, 850, 100]:第一维是层数2,第二维是被误识别为批大小的850,和实际观测的输出完全一致
- 传入10条文档时,嵌入层输出形状为
[10, 850, 100],LSTM仍将第一维10识别为序列长度、第二维850识别为批大小,因此输出:output形状为[10, 850, 100]:再次和预期的batch在前格式output数值重合hidden/cell形状仍为[2, 850, 100]:第二维固定为850,完全符合维度顺序不匹配的特征,和观测结果一致
修复方法
修改LSTM初始化代码,增加batch_first=True参数,匹配当前输入张量batch在前的组织格式:
self.rnn = nn.LSTM(emb_dim, hid_dim, n_layers, dropout = dropout, batch_first=True, device=device)
修改后重新测试即可得到预期维度:
- 单文档输入:
output形状[1, 850, 100],hidden/cell形状[2, 1, 100] - 10条文档输入:
output形状[10, 850, 100],hidden/cell形状[2, 10, 100]
内容的提问来源于stack exchange,提问作者rana
相关产品推荐
相关产品推荐

