如何用预训练BERT做文本分类?对比微调与未微调模型性能
微调BERT vs 预训练BERT文本分类表现对比与代码改写
一、性能提升幅度对比
预训练BERT是基于通用语料训练的语言模型,未针对特定分类任务优化,直接用于分类时存在明显局限,和微调后的模型差距显著:
- 核心指标(准确率):预训练BERT直接分类的准确率通常在50%-70%之间(略高于随机猜测),而微调后的BERT在多数文本分类任务中准确率可达85%-95%,性能提升幅度在15%-30%。若任务属于特定领域(如医疗、法律文本分类),提升幅度会更大。
- 鲁棒性与泛化能力:预训练BERT对任务特有语义特征捕捉不足,预测结果波动大;微调后模型能学习任务数据中的专属模式,泛化能力更强,对不同输入的分类稳定性显著提升。
- 细类别区分:预训练BERT难以区分相似类别(如不同程度的情感),微调后模型可针对性学习类别间的差异特征,细分类别识别能力大幅增强。
二、预训练BERT直接分类的代码改写
基于你提供的微调代码,修改为固定预训练BERT权重,仅训练分类头的版本,实现直接用预训练BERT做分类的需求:
from transformers import BertTokenizer, BertModel import torch import torch.nn as nn import torch.utils.data as Data import torch.optim as optim from sklearn.metrics import accuracy_score,matthews_corrcoef from sklearn.model_selection import train_test_split # 加载预训练tokenizer和BERT模型 tokenizer_model = BertTokenizer.from_pretrained('bert-base-uncased') pretrained_model = BertModel.from_pretrained("bert-base-uncased") # 冻结预训练BERT的所有参数,不参与训练 for param in pretrained_model.parameters(): param.requires_grad = False class MyDataSet(Data.Dataset): def __init__ (self, data, label): self.data = data self.label = label self.tokenizer = tokenizer_model def __getitem__(self, idx): text = self.data[idx] label = self.label[idx] inputs = self.tokenizer(text, return_tensors="pt",padding='max_length',max_length=256,truncation=True) input_ids = inputs.input_ids.squeeze(0) attention_mask = inputs.attention_mask.squeeze(0) return input_ids, attention_mask, label def __len__(self): return len(self.data) # 数据加载(注意替换实际path和name变量) data,label = [],[] with open(path) as f: for line in f.readlines(): a,b = line.strip().split('\t') data.append(b) if a == 'LOW': label.append('0') elif a == 'MEDIUM': label.append('1') else: label.append('2') label = [int(i) for i in label] train_x,test_x,train_y,test_y = train_test_split(data, label, test_size = 0.15,random_state = 32, stratify=label) dataset_train = MyDataSet(train_x,train_y) dataset_test = MyDataSet(test_x,test_y) dataloader_train = Data.DataLoader(dataset_train, batch_size=128, shuffle=True,num_workers=32,pin_memory=True) dataloader_test = Data.DataLoader(dataset_test, batch_size=128, shuffle=True,num_workers=32,pin_memory=True) class MyModel(nn.Module): def __init__(self): super(MyModel, self).__init__() self.bert = pretrained_model self.linear = nn.Linear(768,3) # 仅分类头可训练 def forward(self, input_ids, attention_mask): # 预训练BERT保持eval模式,避免dropout等训练层生效 self.bert.eval() with torch.no_grad(): output = self.bert(input_ids, attention_mask).pooler_output output = self.linear(output) return output device = torch.device("cuda" if torch.cuda.is_available() else "cpu") if torch.cuda.device_count() > 1: print("Use", torch.cuda.device_count(), 'gpus') model = MyModel() model = nn.DataParallel(model) model = model.to(device) else: model = MyModel().to(device) loss_fn = nn.CrossEntropyLoss() # 仅优化分类头的参数,BERT参数不更新 optimizer = optim.Adam(model.module.linear.parameters() if torch.cuda.device_count()>1 else model.linear.parameters(), lr=1e-3) for epoch in range(10): model.train() for input_ids,attention_mask,label in dataloader_train: train_input_ids,train_attention_mask,train_label = input_ids.to(device),attention_mask.to(device),label.to(device) pred = model(train_input_ids,train_attention_mask) loss = loss_fn(pred, train_label) pred = torch.argmax(pred,dim=1) acc = (pred == train_label).float().mean() print(f'epoch: {epoch}, Loss: {loss.item()}, acc: {acc}') loss.backward() optimizer.step() optimizer.zero_grad() savename_train = str(path) +'_' + str(name) + '_train' + '.txt' with open(savename_train,'a') as f: f.write(f'{epoch}\t{loss.item()}\t{acc.item()}\n') model.eval() with torch.no_grad(): for input_ids,attention_mask,label in dataloader_test: validation_input_ids,validation_attention_mask,validation_label = input_ids.to(device),attention_mask.to(device),label.to(device) pred = model(validation_input_ids,validation_attention_mask) loss = loss_fn(pred, validation_label) pred = torch.argmax(pred, dim=1) acc = (pred == validation_label).float().mean() print(f'val epoch: {epoch}, Loss: {loss.item()}, acc: {acc}') savename_eval = str(path) +'_' + str(name) + '_val' + '.txt' with open(savename_eval,'a') as f: f.write(f'{epoch}\t{loss.item()}\t{acc.item()}\n')
关键修改点说明:
- 冻结BERT参数:遍历预训练BERT的所有参数,设置
requires_grad=False,禁止参数更新。 - Forward过程优化:在模型前向传播时,将BERT设为
eval模式,并使用torch.no_grad()关闭梯度计算,避免不必要的内存消耗。 - 优化器范围:仅将分类头(
linear层)的参数传入优化器,确保只有分类头参与训练。 - 学习率调整:由于仅训练分类头,可适当提高学习率(从1e-5改为1e-3),加快收敛速度。
内容的提问来源于stack exchange,提问作者user19185238
相关产品推荐
相关产品推荐

