PyTorch双向LSTM序列标注训练问题:损失恒定与性能不佳
一、先确认已解决的核心问题
你已经找到梯度消失的根源:torch.argmax是不可导操作,把它的输出喂给损失函数会彻底切断梯度路径,导致模型完全无法更新。修复时直接传模型输出的logits就行——CrossEntropyLoss内部会自动处理softmax和log计算,不需要提前做argmax。
二、准确率卡在50%左右的常见原因及修复步骤
下面针对你的代码逐一排查,给出具体的修复方案:
1. 模型缺失关键的Embedding层
你的代码里定义了embedding_matrix参数,但完全没用到!序列标注任务中,输入的句子通常是词的索引(比如[1,5,3,0]这种),必须先通过Embedding层转换成词向量才能喂给LSTM。这是最致命的问题:
# 在模型__init__方法中添加Embedding层 self.embedding = nn.Embedding.from_pretrained( torch.tensor(embedding_matrix, dtype=torch.float32), freeze=False, # 要微调预训练词向量就设为False,固定就设为True padding_idx=pad_idx )
然后在forward方法里先做词向量转换:
def forward(self, x): # x shape: (batch_size, seq_len) → 转成词向量后是(batch_size, seq_len, embedding_size) embedded = self.embedding(x) lstm_out, _ = self.lstm(embedded) logits = self.fc(lstm_out) return logits
2. LSTM没用到初始化的hidden state
你写了init_state方法初始化hidden state,但forward里调用self.lstm(x)时根本没传这个状态,导致LSTM每次都用默认的零初始状态,严重影响长序列依赖的学习。另外还要注意分离hidden state的梯度,避免跨batch累积:
def forward(self, x): embedded = self.embedding(x) # 传入初始化的hidden state lstm_out, self.hidden = self.lstm(embedded, self.hidden) # 分离梯度,防止上一个batch的梯度影响当前batch self.hidden = (self.hidden[0].detach(), self.hidden[1].detach()) logits = self.fc(lstm_out) return logits
3. 损失函数的输入格式不对
CrossEntropyLoss对输入形状有要求:
- logits需要是
(batch_size*seq_len, num_classes)或者(batch_size, num_classes, seq_len) - 标签需要是
(batch_size*seq_len)的整数类型(不能是double)
你的logits是(batch_size, seq_len, num_classes),得调整维度:
# 训练循环里修改损失计算部分 logits = model(sentence) # 把logits展平成(总token数, 类别数),标签展平成(总token数) loss = loss_function(logits.view(-1, n_classes), label.view(-1).long())
4. 优化器和学习率设置不合理
你用的SGD学习率1e-4太低了,SGD本身收敛就慢,这么小的学习率根本带不动参数更新。建议换成Adam优化器,或者调高SGD的学习率并加动量:
# 推荐用Adam,收敛更快 optimizer = optim.Adam(model.parameters(), lr=1e-3) # 或者调整SGD optimizer = optim.SGD(model.parameters(), lr=1e-2, momentum=0.9)
5. 准确率计算逻辑错误
你现在只计算了最后一个batch的准确率,不是整个epoch的平均准确率,导致输出的结果完全不准。改成累积整个epoch的正确数和总token数:
def train_model(model, train: pd.DataFrame): model.train() for epoch in range(num_epochs): epoch_loss = 0.0 total_correct = 0 total_tokens = 0 for (sentence, label) in train_loader: optimizer.zero_grad() model.init_state(sentence) logits = model(sentence) # 计算损失 loss = loss_function(logits.view(-1, n_classes), label.view(-1).long()) loss.backward() optimizer.step() # 累积损失(乘以batch size,最后除以总样本数) epoch_loss += loss.item() * sentence.size(0) # 计算当前batch的准确率(忽略pad token) preds = torch.argmax(logits, dim=2) mask = (label != pad_idx) total_correct += ((preds == label) & mask).sum().item() total_tokens += mask.sum().item() # 计算整个epoch的平均损失和准确率 avg_loss = epoch_loss / len(train_loader.dataset) avg_acc = total_correct / total_tokens * 100 print(f"Epoch: {epoch+1}, loss: {avg_loss:.5f}, accuracy: {avg_acc:.2f}%")
6. 模型容量和正则化不足
如果任务比较复杂(比如命名实体识别),当前1层64维的LSTM容量不够,可以增加层数或隐藏单元数,同时加Dropout防止过拟合:
# 在模型__init__里调整 self.lstm_layer_size = 128 self.num_layers = 2 # 添加Dropout层 self.dropout = nn.Dropout(0.5) # LSTM设置dropout(只有层数>1时才生效) self.lstm = nn.LSTM( self.embedding_size, self.lstm_layer_size, bidirectional=True, num_layers=self.num_layers, batch_first=True, dropout=0.5 if self.num_layers > 1 else 0.0 ) # forward里加dropout embedded = self.dropout(self.embedding(x))
7. 数据预处理检查
最后确认下数据部分:
- 输入的句子是否正确转换成了整数索引张量?
- 标签的取值范围是否是
0~n_classes-1? - pad_idx是否正确设置,并且在计算准确率时被忽略?
修复后的完整模型示例
class biLSTM(nn.Module): def __init__(self, embedding_matrix, embedding_size, pad_idx, n_classes): super(biLSTM, self).__init__() self.lstm_layer_size = 128 self.embedding_size = embedding_size self.num_layers = 2 self.n_classes = n_classes self.pad_idx = pad_idx # 预训练词向量Embedding层 self.embedding = nn.Embedding.from_pretrained( torch.tensor(embedding_matrix, dtype=torch.float32), freeze=False, padding_idx=pad_idx ) # Dropout正则化 self.dropout = nn.Dropout(0.5) self.lstm = nn.LSTM( self.embedding_size, self.lstm_layer_size, bidirectional=True, num_layers=self.num_layers, batch_first=True, dropout=0.5 if self.num_layers > 1 else 0.0 ) self.fc = nn.Linear(self.lstm_layer_size * 2, self.n_classes) self.hidden = None def init_state(self, x): batch_size = x.size(0) # 自动适配输入设备(CPU/GPU) self.hidden = ( torch.zeros(self.num_layers * 2, batch_size, self.lstm_layer_size).to(x.device), torch.zeros(self.num_layers * 2, batch_size, self.lstm_layer_size).to(x.device) ) def forward(self, x): embedded = self.dropout(self.embedding(x)) lstm_out, self.hidden = self.lstm(embedded, self.hidden) # 分离hidden state,避免跨batch梯度累积 self.hidden = (self.hidden[0].detach(), self.hidden[1].detach()) logits = self.fc(lstm_out) return logits
内容的提问来源于stack exchange,提问作者Caio Nogueira

