PyTorch搭建BiLSTM模型报input must have 2 dimensions错误求解
问题排查与解决方法
1. 直接报错原因
错误RuntimeError: input must have 2 dimensions, got 1指向LSTM层的输入维度不匹配:
当你设置LSTM的batch_first=True时,LSTM要求输入的词嵌入张量维度为[batch_size, seq_len, embedding_dim],对应的输入文本张量text的维度应为[batch_size, seq_len](词嵌入层会把最后一维转成embedding_dim)。你现在拿到的text是1维张量,形状为[seq_len],缺少batch维度,才会触发该错误。
2. 全量问题修复清单
按优先级排序:
- 维度匹配修复
你设置了BATCH_SIZE=1,但collate_batch函数没有给单样本增加batch维度,要么修改collate_batch确保输出的text维度是[batch_size, seq_len],要么在模型forward方法最开头增加维度补全逻辑:def forward(self, text, text_lengths): # 新增代码 if text.dim() == 1: text = text.unsqueeze(0) batch_size = text.shape[0] # 原有后续逻辑 - evaluate函数参数缺失修复
你在evaluate方法中调用model(text)时没有传入必填参数text_lengths,会触发参数数量错误,需要从collate_fn中拿到当前batch的样本真实长度后传入:predited_label = model(text, text_lengths=当前batch样本真实长度张量) - 变量名不一致修复
__init__方法入参没有定义lstm_units,你传入的隐藏层大小参数是hidden_dim,但LSTM初始化时用了未定义的lstm_units,修改LSTM初始化代码:
同时把代码中所有用到self.lstm = nn.LSTM(embedding_dim, hidden_dim, # 替换原有的lstm_units num_layers=lstm_layers, bidirectional=True, batch_first=True)self.lstm_units的地方都替换为self.hidden_dim。 - 池化逻辑与全连接层维度修复
你当前设置的nn.MaxPool1d(1, stride=1)没有实际池化效果,也没有在时序维度做全局池化,修改forward中的池化与拼接逻辑:
同步修改全连接层fc1的输入维度:out = output_unpacked # 直接在seq_len维度做全局池化,不需要提前定义池化层 out1 = torch.max(out, dim=1)[0] # 全局最大池化,形状[batch_size, 2*hidden_dim] out2 = torch.mean(out, dim=1) # 全局平均池化,形状[batch_size, 2*hidden_dim] out = torch.cat((out1, out2), dim=1) # 拼接后形状[batch_size, 4*hidden_dim]self.fc1 = nn.Linear(hidden_dim * 4, hidden_dim) # 双向*2 + 两种池化拼接*2 = 4倍hidden_dim - 损失函数冲突修复
你使用的CrossEntropyLoss内部已经实现了Softmax计算,不需要在forward中手动调用F.softmax,直接返回fc2的输出即可:# 删掉F.softmax调用 preds = self.fc2(out) - 其他通用问题修复
- 不要在train函数中写死text_lengths为200,需要在
collate_batch对样本做padding时,同步记录每个样本的真实长度,传入模型 - 隐藏状态h0、c0需要和输入text放在同一个设备上,避免GPU运行报错,修改初始化代码:
# 不需要用Variable包装,PyTorch 0.4+版本张量默认支持自动求导 h_0, c_0 = (torch.zeros(self.lstm_layers * self.num_directions, batch_size, self.hidden_dim).to(text.device), torch.zeros(self.lstm_layers * self.num_directions, batch_size, self.hidden_dim).to(text.device))- 你只定义了train_dataloader和test_dataloader,运行时会报valid_dataloader未定义的错误,需要补充验证集DataLoader的定义,或者把evaluate的入参改为test_dataloader。
- 不要在train函数中写死text_lengths为200,需要在
内容的提问来源于stack exchange,提问作者AlwaysNewbie
相关产品推荐
相关产品推荐

