基于PyTorch的BERT NER模型重训练/迁移学习方法求助
BERT-based NER模型重训练/迁移学习指南
一、前期准备
- 加载已有模型与Tokenizer
确保使用和原训练一致的BERT版本(如bert-base-cased),加载你之前保存的模型:from transformers import BertTokenizer, BertForTokenClassification import torch # 加载本地保存的模型 model = BertForTokenClassification.from_pretrained("./your_saved_model_dir") # 加载对应版本的Tokenizer tokenizer = BertTokenizer.from_pretrained("bert-base-cased") - 适配新数据集标签
如果新数据集的标签数量和原模型输出层维度不同,需要修改模型的分类头:
同时要保证新数据集的标签格式和原训练一致(如IOB2标注规范),子词标签处理逻辑也要对齐。# 替换为你的新标签集合长度 num_new_labels = len(your_new_label_vocab) # 重新初始化分类头 model.classifier = torch.nn.Linear(model.classifier.in_features, num_new_labels)
二、训练参数配置
使用TrainingArguments设置训练核心参数,新手推荐用低学习率(2e-5~5e-5),避免破坏BERT预训练的通用语义能力:
from transformers import TrainingArguments, Trainer training_args = TrainingArguments( output_dir="./new_ner_model_checkpoints", per_device_train_batch_size=8, learning_rate=3e-5, num_train_epochs=3, logging_dir="./training_logs", logging_steps=10, save_steps=100, evaluation_strategy="epoch" # 可选,每轮验证一次 )
三、数据格式转换
将新数据集转换为transformers兼容的Dataset格式,核心字段需包含input_ids、attention_mask、labels:
from datasets import Dataset # 假设你的数据已整理成字典格式,包含texts和对应的labels_list train_data = {"text": your_train_texts, "labels": your_train_labels} train_dataset = Dataset.from_dict(train_data) # 用tokenizer处理数据(需和原训练时的处理逻辑一致) def tokenize_function(examples): tokenized_inputs = tokenizer(examples["text"], truncation=True, padding="max_length", max_length=128) # 处理子词标签对齐逻辑,这里示例为直接映射,需根据原训练逻辑调整 tokenized_inputs["labels"] = examples["labels"] return tokenized_inputs tokenized_train_dataset = train_dataset.map(tokenize_function, batched=True)
四、执行重训练
初始化Trainer并启动训练:
trainer = Trainer( model=model, args=training_args, train_dataset=tokenized_train_dataset, eval_dataset=tokenized_eval_dataset # 可选,验证集 ) # 开始训练 trainer.train() # 保存最终模型(推荐用save_pretrained,比torch.save更方便后续加载) model.save_pretrained("./final_finetuned_ner_model") tokenizer.save_pretrained("./final_finetuned_ner_model")
五、进阶优化建议
- 分层微调:如果新数据集规模较小,可先冻结BERT主体层,仅训练分类头;待分类头收敛后,再解冻前4~6层进行联合微调,平衡训练效果与计算成本:
# 先冻结BERT主体 for param in model.bert.parameters(): param.requires_grad = False # 训练分类头后,解冻部分层 for param in model.bert.encoder.layer[:4].parameters(): param.requires_grad = True - 标签映射校验:确保新数据集的标签索引和模型输出层的索引严格对应,避免训练时出现标签不匹配的报错。
- 数据增强:小规模数据集可尝试同义词替换、随机插入/删除等数据增强方式,提升模型泛化能力。
内容的提问来源于stack exchange,提问作者IronMan
相关产品推荐
相关产品推荐

