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

如何为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 14:01:04