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

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. 传入单条文档时,嵌入层输出形状为[1, 850, 100],LSTM默认将第一维1识别为序列长度、第二维850识别为批大小,因此输出:
    • output形状为[1, 850, 100]:数值上和预期的batch在前格式结果重合,属于巧合
    • hidden/cell形状为[2, 850, 100]:第一维是层数2,第二维是被误识别为批大小的850,和实际观测的输出完全一致
  2. 传入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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.31 21:45:38