如何在PyTorch中正确使用LSTM解决3D输入要求及模型效果退化问题
问题根因分析
你当前代码的核心错误是使用了nn.EmbeddingBag作为嵌入层,它的输出已经丢失了LSTM必需的序列维度:
nn.EmbeddingBag的作用是将每个样本对应的所有token embedding直接做求和/平均/最值聚合,输出形状固定为[batch_size, embed_dim],你当前得到的[10, 32]就是10个样本各自的全局句子级向量,完全没有时序维度信息。- 你强行用
view(1, x.shape[0], -1)把它转成[1, 10, 32]喂给LSTM,相当于把10个样本当成了长度为1的序列的10个batch,完全打乱了数据的语义逻辑,因此模型效果极差。
正确实现方案
要使用LSTM做文本分类,需要调整如下核心逻辑:
1. 替换嵌入层
将nn.EmbeddingBag替换为普通nn.Embedding,保留每个token的独立嵌入向量,维持序列维度。
2. 调整输入预处理逻辑
原有的text+offsets是适配EmbeddingBag的输入格式,现在需要将同一batch内的不同长度句子做填充(padding),统一为[batch_size, 最大序列长度]的二维张量作为输入,同时记录每个句子的真实长度用于后续变长序列优化。
3. 调整LSTM参数
添加batch_first=True参数,让LSTM适配[batch_size, seq_len, embed_dim]的输入格式,更符合常规数据排布逻辑。
4. 处理LSTM输出
文本分类场景下,一般取LSTM输出的最后一个有效时间步的隐藏状态,或者对所有时间步输出做全局池化,再传入后续全连接层。
修改后代码示例
import torch import torch.nn as nn from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence torch.manual_seed(random_state) class Net(nn.Module): def __init__(self, vocab_size= len(vocab), embed_dim= 32, num_class= 3, lstm_layers=5): super().__init__() # 替换为普通Embedding,padding_idx指定填充位的id self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0) # 添加batch_first=True,适配[batch, seq_len, embed_dim]格式输入 self.lstm = nn.LSTM(input_size= embed_dim, hidden_size= embed_dim, num_layers=lstm_layers, batch_first=True) # 全连接层添加激活函数,避免线性堆叠失效 self.lin = nn.Sequential( nn.Linear(in_features= embed_dim, out_features= embed_dim), nn.ReLU(), nn.Linear(in_features= embed_dim, out_features= embed_dim), nn.ReLU(), nn.Linear(in_features= embed_dim, out_features= 16), nn.ReLU(), nn.Linear(in_features= 16, out_features= 16), nn.ReLU(), nn.Linear(in_features= 16, out_features= 8), nn.ReLU(), nn.Linear(in_features= 8, out_features= 8), nn.ReLU(), ) self.out = nn.Linear(in_features= 8, out_features= num_class) self.init_weights() def init_weights(self): initrange = 0.5 self.embedding.weight.data.uniform_(-initrange, initrange) for layer in self.lin: if isinstance(layer, nn.Linear): layer.weight.data.uniform_(-initrange, initrange) layer.bias.data.zero_() def forward(self, text, seq_len): # text形状为[batch_size, max_seq_len],seq_len为每个句子的真实长度 x = self.embedding(text) # 输出形状[batch_size, max_seq_len, embed_dim] # 可选:使用pack_padded_sequence优化变长序列计算,忽略pad部分的无效计算 packed_x = pack_padded_sequence(x, seq_len.cpu(), batch_first=True, enforce_sorted=False) packed_out, _ = self.lstm(packed_x) x, _ = pad_packed_sequence(packed_out, batch_first=True) # 输出形状[batch_size, max_seq_len, embed_dim] # 取最后一个有效时间步的输出作为句子特征 last_idx = seq_len - 1 last_output = x[torch.arange(x.shape[0]), last_idx, :] # 形状[batch_size, embed_dim] x = self.lin(last_output) return self.out(x) model = Net() model.to(device)
输入适配说明
训练时的输入需要调整为:
text:填充后的token id张量,形状[batch_size, max_seq_len]seq_len:一维张量,记录batch内每个句子的原始长度,长度等于batch size
内容的提问来源于stack exchange,提问作者user9102437
相关产品推荐
相关产品推荐

