使用BertForSequenceClassification训练SST2分类模型无进展,求排查代码问题
问题分析与修复方案
你的代码存在几个关键问题,导致模型看起来没有在训练:
1. 学习率设置过高
BERT这类预训练模型微调时,合理的学习率范围是2e-5~5e-5,你设置的LR=0.0005(5e-4)是该范围的10倍以上。过大的学习率会让模型参数在最优值附近剧烈震荡甚至发散,根本无法收敛。
2. 未跟踪训练状态
你定义了train_loss列表,但从未将每个batch的损失值加入其中,也没有在epoch结束后计算平均损失。这导致你完全无法观察模型的损失变化,自然无法判断模型是否在学习。
3. 缺少验证环节
训练过程中没有在验证集上评估模型性能,既无法验证模型是否真正学到了有效特征,也无法及时发现过拟合或训练无效的情况。
4. 数据处理效率低下
当前collate_batch通过循环单个样本做tokenization,既冗余又低效。可以直接对整个batch的文本做批量tokenization,同时避免手动堆叠张量的潜在问题。
修改后的完整代码
import torch from torch import nn from torch.optim import AdamW from torch.utils.data import DataLoader from torchtext.datasets import SST2 from tqdm import tqdm from transformers import AutoTokenizer, BertForSequenceClassification # 修正学习率到预训练模型微调的合理范围 LR = 2e-5 EPOCHS = 5 BATCH_SIZE = 128 device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") model_name = "bert-base-uncased" tokenizer = AutoTokenizer.from_pretrained(model_name) # BertForSequenceClassification默认加载预训练权重+适配2分类的头,正好匹配SST2任务 model = BertForSequenceClassification.from_pretrained(model_name).to(device) max_input_length = 128 criterion = nn.CrossEntropyLoss().to(device) optimizer = AdamW(model.parameters(), lr=LR) train_datapipe = SST2(split="train") valid_datapipe = SST2(split="dev") def collate_batch(batch): # 批量提取文本和标签 texts, labels = zip(*batch) # 一次性完成整个batch的tokenization,自动处理padding、截断 tokenized = tokenizer(list(texts), padding="max_length", max_length=max_input_length, truncation=True, return_tensors="pt") # 统一将张量移到目标设备 input_data = {k: v.to(device) for k, v in tokenized.items()} labels = torch.tensor(labels, dtype=torch.int64).to(device) return input_data, labels train_dataloader = DataLoader(train_datapipe, shuffle=True, batch_size=BATCH_SIZE, collate_fn=collate_batch) valid_dataloader = DataLoader(valid_datapipe, batch_size=BATCH_SIZE, collate_fn=collate_batch) for epoch in range(EPOCHS): print(f"=== Epoch {epoch+1}/{EPOCHS} ===") # 训练阶段 model.train() total_train_loss = 0.0 train_correct = 0 train_total = 0 for input_data, label in tqdm(train_dataloader): optimizer.zero_grad() outputs = model(**input_data) logits = outputs.logits # 计算损失并累积 loss = criterion(logits, label) total_train_loss += loss.item() # 计算训练准确率,直观观察分类效果 preds = torch.argmax(logits, dim=1) train_correct += (preds == label).sum().item() train_total += label.size(0) loss.backward() optimizer.step() # 打印训练汇总 avg_train_loss = total_train_loss / len(train_dataloader) train_acc = train_correct / train_total print(f"训练损失: {avg_train_loss:.4f}, 训练准确率: {train_acc:.4f}") # 验证阶段 model.eval() total_valid_loss = 0.0 valid_correct = 0 valid_total = 0 with torch.no_grad(): for input_data, label in tqdm(valid_dataloader): outputs = model(**input_data) logits = outputs.logits loss = criterion(logits, label) total_valid_loss += loss.item() preds = torch.argmax(logits, dim=1) valid_correct += (preds == label).sum().item() valid_total += label.size(0) # 打印验证汇总 avg_valid_loss = total_valid_loss / len(valid_dataloader) valid_acc = valid_correct / valid_total print(f"验证损失: {avg_valid_loss:.4f}, 验证准确率: {valid_acc:.4f}\n")
额外说明
- 调整学习率后,你应该能看到训练损失逐步下降、验证准确率逐步上升,这是模型正常学习的标志。
- 批量tokenization不仅代码更简洁,处理速度也远快于循环单个样本,适合大数据集训练。
- 加入准确率计算能更直观地反映模型的分类能力,比仅观察损失更全面。
内容的提问来源于stack exchange,提问作者MyungHa Kwon
相关产品推荐
相关产品推荐

