如何在自定义数据集上对BERT做post-training 能否使用transformers.BertForMaskedLM
实现自定义语料BERT继续预训练(Post-Training)操作指南
核心问题确认
你完全可以使用
transformers.BertForMaskedLM完成该需求。BERT原生预训练的核心任务就是掩码语言建模(MLM),在自定义领域语料上做MLM继续预训练是适配领域特征、提升下游任务词向量质量的标准方案。
具体实现步骤
1. 依赖安装
- 安装所需依赖:
pip install transformers datasets accelerate torch
2. 自定义语料预处理
首先将你的自定义语料整理为逐行存储的纯文本格式,每行对应一个独立的文本段落/句子,后续按如下方式处理数据:
from datasets import load_dataset from transformers import BertTokenizer # 加载自定义语料,替换为你自己的语料文件路径 dataset = load_dataset("text", data_files="your_custom_corpus.txt") # 中文场景下替换为bert-base-chinese tokenizer = BertTokenizer.from_pretrained("bert-base-uncased") # 定义分词函数,可根据你的文本长度调整max_length参数 def tokenize_function(examples): return tokenizer(examples["text"], padding="max_length", truncation=True, max_length=128) tokenized_datasets = dataset.map(tokenize_function, batched=True, remove_columns=["text"])
配置MLM任务所需的数据整理器,自动完成掩码标注:
from transformers import DataCollatorForLanguageModeling # 掩码比例设置为15%,和BERT原生预训练规则保持一致 data_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=True, mlm_probability=0.15)
3. 加载预训练模型
from transformers import BertForMaskedLM # 中文场景下替换为bert-base-chinese model = BertForMaskedLM.from_pretrained("bert-base-uncased")
4. 训练配置与执行
直接使用Transformers自带的Trainer API完成训练即可:
from transformers import Trainer, TrainingArguments training_args = TrainingArguments( output_dir="./bert-post-trained", # 模型保存路径 per_device_train_batch_size=32, # 可根据你的显存大小调整 num_train_epochs=3, # 小语料设置1-2轮,大语料可设置2-5轮 learning_rate=2e-5, # 学习率不要过高,避免破坏原生预训练权重 weight_decay=0.01, logging_steps=10, save_strategy="epoch" ) trainer = Trainer( model=model, args=training_args, train_dataset=tokenized_datasets["train"], data_collator=data_collator, ) # 启动训练 trainer.train()
5. 提取训练后词向量用于下游任务
训练完成后,你可以从保存的MLM权重中加载基础BERT模型,直接提取文本向量:
from transformers import BertModel # 替换为你实际的模型checkpoint路径 bert_model = BertModel.from_pretrained("./bert-post-trained/checkpoint-xxx") # 推理获取向量示例 inputs = tokenizer("待提取向量的文本", return_tensors="pt") outputs = bert_model(**inputs) # <[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> 标识对应的句子向量:outputs.last_hidden_state[:, 0, :] # 所有token平均得到的句子向量:outputs.last_hidden_state.mean(dim=1)
注意事项
- 如果你有大量领域相关的句子对语料,也可以选择加入下一句预测(NSP)任务,使用
BertForPreTraining类即可,不过目前行业实践证明仅MLM任务的继续预训练效果已经足够,NSP的增益非常有限 - 小语料场景下不要训练过多轮次,避免模型过拟合,丢失通用语义表征能力
- 中文场景下全程替换预训练权重为bert-base-chinese即可,其余逻辑无需调整
内容的提问来源于stack exchange,提问作者The Exile
相关产品推荐
相关产品推荐

