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

如何针对文本分类任务分步微调HuggingFace BERT模型

微调HuggingFace BERT完成文本分类分步操作

以下是可直接落地的操作流程,默认基于PyTorch后端:

1. 环境准备

首先安装必要依赖:

pip install transformers datasets torch scikit-learn

如果是中文任务,后续所有预训练模型名替换为bert-base-chinese即可,无需调整其他逻辑。

2. 数据预处理

  • 先将你的数据集整理为带两列的csv/json格式:text列存储待分类文本内容,label列存储分类标签(二分类填0/1,多分类填0到类别数-1的整数)
  • 加载并拆分数据集:
from datasets import load_dataset

# 加载本地csv文件,json格式替换为load_dataset("json", data_files="your_data.json")
dataset = load_dataset("csv", data_files="your_dataset.csv")
# 按8:2拆分训练集和验证集
dataset = dataset["train"].train_test_split(test_size=0.2, shuffle=True, seed=42)
  • 加载BERT分词器并对文本做分词处理:
from transformers import BertTokenizer

# 英文任务用bert-base-uncased,中文用bert-base-chinese
tokenizer = BertTokenizer.from_pretrained("bert-base-uncased")

def tokenize_func(examples):
    # max_length可根据你的文本平均长度调整,常用128/256
    return tokenizer(examples["text"], padding="max_length", truncation=True, max_length=128)

# 批量分词处理
tokenized_ds = dataset.map(tokenize_func, batched=True)

3. 加载带分类头的BERT模型

from transformers import BertForSequenceClassification

# num_labels填你的分类任务类别数,二分类填2,三分类填3以此类推
model = BertForSequenceClassification.from_pretrained("bert-base-uncased", num_labels=2)

如果数据集规模很小,可以冻结BERT底部几层参数只训练顶层和分类头,降低过拟合风险。

4. 配置训练参数

from transformers import TrainingArguments, Trainer

training_args = TrainingArguments(
    # 训练产物存储路径
    output_dir="./bert_clf_result",
    # BERT微调常用学习率区间2e-5~5e-5,不要设太高避免模型崩掉
    learning_rate=2e-5,
    per_device_train_batch_size=16,
    per_device_eval_batch_size=16,
    # 训练轮次一般3~5轮足够,多了容易过拟合
    num_train_epochs=3,
    weight_decay=0.01,
    # 每轮训练结束跑一次验证
    evaluation_strategy="epoch",
    save_strategy="epoch",
    # 训练结束后自动加载验证集效果最好的模型
    load_best_model_at_end=True,
    # 不需要训练日志的话可以加这行关闭
    # logging_strategy="no"
)

5. 启动训练

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=tokenized_ds["train"],
    eval_dataset=tokenized_ds["test"],
)

# 开始训练,训练过程会打印验证集指标
trainer.train()

6. 用训练好的模型推理

# 加载训练好的最优模型,路径替换为你output_dir下实际的checkpoint文件夹路径
trained_model = BertForSequenceClassification.from_pretrained("./bert_clf_result/checkpoint-xxx")

def predict(text):
    inputs = tokenizer(text, return_tensors="pt", padding="max_length", truncation=True, max_length=128)
    outputs = trained_model(**inputs)
    # 返回预测的标签值
    return outputs.logits.argmax(dim=-1).item()

# 测试预测
print(predict("待分类的文本内容"))

内容的提问来源于stack exchange,提问作者mon

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 03:42:03