如何使用Huggingface Trainer训练BertForMaskedLM实现BERT领域适配
BERT领域自适应继续预训练(基于Trainer)实现方案
首先明确模型选型建议:优先使用BertForMaskedLM即可,BertForPreTraining包含了原BERT的下一句预测(NSP)任务,目前多数实践证明NSP对大部分下游任务增益有限,还会额外增加数据处理成本,无特殊需求不需要使用。
1. 自定义Dataset的__getitem__返回值要求
不需要手动实现掩码逻辑,transformers提供的DataCollatorForLanguageModeling会自动完成MLM掩码生成、标签构造,所以__getitem__只需返回以下2个必填字段即可:
input_ids:单条文本分词后的id序列,shape为(max_seq_len,)attention_mask:注意力掩码,shape和input_ids一致,padding位置为0,有效文本位置为1
可选字段:如果你的任务涉及分段文本,可额外返回token_type_ids,否则模型会自动补全为0,无需额外处理。
示例Dataset实现代码:
from torch.utils.data import Dataset class DomainDataset(Dataset): def __init__(self, text_list, tokenizer, max_seq_len=128): self.text_list = text_list self.tokenizer = tokenizer self.max_seq_len = max_seq_len def __len__(self): return len(self.text_list) def __getitem__(self, idx): text = self.text_list[idx].strip() tokenized = self.tokenizer( text, truncation=True, max_length=self.max_seq_len, padding="max_length", return_tensors="pt" ) # 去掉分词器自动添加的batch维度,后续DataCollator会统一构造batch return { "input_ids": tokenized["input_ids"].squeeze(0), "attention_mask": tokenized["attention_mask"].squeeze(0) }
2. 基于Trainer的完整训练流程
依赖导入
from transformers import BertForMaskedLM, BertTokenizer, TrainingArguments, Trainer, DataCollatorForLanguageModeling
步骤1:加载预训练权重和分词器
# 替换为你使用的基础BERT权重,比如bert-base-uncased、bert-base-chinese等 tokenizer = BertTokenizer.from_pretrained("bert-base-chinese") model = BertForMaskedLM.from_pretrained("bert-base-chinese")
步骤2:初始化自定义数据集
# 替换为你自己的领域文本列表,每个元素为单条文本 train_texts = ["领域文本1", "领域文本2", "..."] train_dataset = DomainDataset(train_texts, tokenizer, max_seq_len=128)
步骤3:初始化MLM数据收集器
data_collator = DataCollatorForLanguageModeling( tokenizer=tokenizer, mlm=True, mlm_probability=0.15 # 和原BERT预训练的掩码比例一致,可按需调整 )
步骤4:配置训练参数
training_args = TrainingArguments( output_dir="./bert-domain-adapted", overwrite_output_dir=True, num_train_epochs=2, # 按需调整,领域数据量小的话1-3轮即可 per_device_train_batch_size=16, save_steps=1000, save_total_limit=2, logging_steps=100, learning_rate=2e-5, # 继续预训练建议用比下游微调稍低的学习率,1e-5~3e-5区间即可 weight_decay=0.01, disable_tqdm=False )
步骤5:启动训练
trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, data_collator=data_collator, ) trainer.train() # 训练完成后导出模型 trainer.save_model("./bert-domain-adapted/final")
注意事项
- 长文本可以用滑动窗口切分,提升数据利用率
- 如果确实需要使用
BertForPreTraining,需要构造NSP任务的文本对,同时将data_collator替换为DataCollatorForPreTraining即可 - 训练得到的领域适配权重可以直接用于下游任务微调,用法和原生BERT权重完全一致
内容的提问来源于stack exchange,提问作者The Exile
相关产品推荐
相关产品推荐

