基于Pytorch Lightning训练LSTM模型时遇RuntimeError报错求助
问题根源分析
报错提示shape '[56, 1, 768]' is invalid for input of size 41472,核心是张量总元素数与预期形状不匹配:
- 预期总元素数:
56*1*768=43008 - 实际总元素数:
41472,计算得41472/768=54,说明有效序列长度总和为54,但当前batch包含56个样本,要么存在2个长度为0的无效样本,要么LSTM输出维度与代码预期不一致。
结合你的模型代码,主要问题集中在LSTM维度配置错误和多GPU环境下的序列长度处理:
解决方案
1. 修正LSTM隐藏层维度配置
你当前设置cfg.lstm_cfg.hidden_dim=384*2=768且bidirectional=True,导致LSTM输出维度为768*2=1536(双向LSTM会自动拼接前向/后向隐藏状态),但后续代码(如pad_packed_sequence、分类器输入)预期维度为768,直接引发形状不匹配。
修正方式:
将cfg.lstm_cfg.hidden_dim改为384(单向隐藏层维度),保持bidirectional=True,此时LSTM输出维度自动变为384*2=768,与你的注释和后续代码逻辑一致。
对应的LSTM初始化代码无需修改(因为你已经通过cfg.lstm_cfg.hidden_dim + (cfg.lstm_cfg.bidirectional * cfg.lstm_cfg.hidden_dim)正确计算了分类器输入维度),只需调整配置文件中的hidden_dim值。
2. 检查序列长度(lengths)的有效性
确保lengths张量中不存在0值(每个样本的序列长度至少为1),且在多GPU训练时,lengths是当前GPU子batch对应的长度,而非原始完整batch的长度。
可以在forward函数开头添加校验:
assert (lengths > 0).all(), "序列长度不能为0" assert len(lengths) == padded_embeddings.size(0), "lengths长度需与batch_size匹配"
3. 显式指定pad_packed_sequence的total_length
在多GPU环境下,自动计算的max_len可能存在差异,显式指定total_length为输入序列的最大长度(即padded_embeddings的seq_len维度),确保pad后的形状与输入一致:
padded_hidden, _ = pad_packed_sequence(packed_hidden, batch_first=True, total_length=padded_embeddings.size(1))
4. 调整flatten_parameters()的调用时机
在多GPU(DDP)训练时,flatten_parameters()仅在__init__中调用可能无效,因为模型会被复制到不同GPU。将其移到Pytorch Lightning的训练钩子中:
def on_train_start(self): self.lstm.flatten_parameters()
验证修正
修正LSTM维度后,LSTM输出维度为768,此时若lengths总和为56(56个长度为1的样本),则sum(lengths)*768=56*768=43008,与预期形状[56,1,768]的总元素数一致,报错即可解决。
内容的提问来源于stack exchange,提问作者Youssef Benhachem

