如何为BertForSequenceClassification替换自定义分类器头?
替换BertForSequenceClassification的自定义分类器头
要替换默认的Linear分类器,核心是保证自定义分类器的输入/输出维度和原分类器对齐,同时符合PyTorch nn.Module的规范,天然就能处理批量数据。
1. 明确输入输出要求
BertForSequenceClassification的classifier接收的是BERT最后一层<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token的输出,维度规则如下:
- 输入:
[batch_size, hidden_size](bert-base-uncased的hidden_size固定为768) - 输出:
[batch_size, num_labels](你的场景里num_labels=1,对应单标签回归或二分类任务)
2. 自定义分类器示例
比如实现一个带Dropout和ReLU激活的两层感知机(MLP):
import torch.nn as nn class CustomClassifier(nn.Module): def __init__(self, hidden_size, num_labels): super().__init__() self.fc1 = nn.Linear(hidden_size, hidden_size // 2) self.dropout = nn.Dropout(0.1) self.fc2 = nn.Linear(hidden_size // 2, num_labels) self.relu = nn.ReLU() def forward(self, x): # x的维度为[batch_size, hidden_size] x = self.fc1(x) x = self.relu(x) x = self.dropout(x) x = self.fc2(x) # 输出维度为[batch_size, num_labels] return x
3. 替换并验证的完整代码
from transformers import BertForSequenceClassification import torch device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model_name = 'bert-base-uncased' num_labels = 1 # 加载预训练模型 model = BertForSequenceClassification.from_pretrained(model_name, num_labels=num_labels).to(device) # 初始化自定义分类器并替换原classifier hidden_size = model.config.hidden_size # 从模型配置读取,避免硬编码 model.classifier = CustomClassifier(hidden_size, num_labels).to(device) # 测试批量输入处理 dummy_input = torch.randint(0, model.config.vocab_size, (4, 128)).to(device) # batch_size=4,序列长度128 outputs = model(dummy_input) print(outputs.logits.shape) # 应输出 torch.Size([4, 1]),符合批量输出要求
关键注意点
- 自定义分类器只要继承
nn.Module,forward方法接收批量张量即可,PyTorch会自动处理批量维度,无需手动写循环。 - 输入维度必须和BERT的hidden_size严格匹配,输出维度必须对应num_labels,否则会触发维度不兼容的报错。
- 可根据需求调整结构:比如添加LayerNorm、更换激活函数(如GELU)、调整层数等,只要保证输入输出维度正确即可。
内容的提问来源于stack exchange,提问作者Penguin
相关产品推荐
相关产品推荐

