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

如何使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 09:45:00