You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.24 06:06:05