迁移学习:这次怎么落地的
很多人一上来就讲迁移学习的全景图;我更想先把这次卡住的点说清楚。
项目背景:医疗领域文本分类,需要把病历摘要分类到 10 个科室。
场景:小数据量的文本分类
项目背景:医疗领域文本分类,需要把病历摘要分类到 10 个科室。
数据情况:
- 训练集:3200 条
- 验证集:800 条
- 测试集:1000 条
- 每条平均长度:200 字
环境:
python 3.9
torch 1.12.1
transformers 4.25.1
CUDA 11.3
为什么选迁移学习
先试了几个传统方法:
# 1. TF-IDF + 逻辑回归
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.linear_model import LogisticRegression
tfidf = TfidfVectorizer(max_features=5000)
X_train = tfidf.fit_transform(train_texts)
clf = LogisticRegression()
clf.fit(X_train, train_labels)
# 准确率:68.5%
# 2. FastText
from fasttext import supervised
with open('train.txt', 'w') as f:
for text, label in zip(train_texts, train_labels):
f.write(f'__label__{label} {text}\n')
model = fasttext.supervised('train.txt', 'model', epoch=10)
# 准确率:72.3%
都不太行。数据量太小,模型学不到足够的特征。
这时候想到迁移学习:用在大规模语料上预训练的模型,迁移到这个小任务上。
预训练模型选择
from transformers import BertTokenizer, BertForSequenceClassification
# 选项 1:BERT-Base
tokenizer = BertTokenizer.from_pretrained('bert-base-chinese')
model = BertForSequenceClassification.from_pretrained('bert-base-chinese', num_labels=10)
# 参数量:110M
# 准确率:82.1%(冻结 backbone)
# 准确率:85.7%(全参数微调)
# 选项 2:RoBERTa-Base
tokenizer = BertTokenizer.from_pretrained('hfl/chinese-roberta-wwm-ext')
model = BertForSequenceClassification.from_pretrained('hfl/chinese-roberta-wwm-ext', num_labels=10)
# 参数量:110M
# 准确率:84.3%(冻结 backbone)
# 准确率:88.2%(全参数微调)
# 选项 3:MacBERT-Base
tokenizer = BertTokenizer.from_pretrained('hfl/macbert-base')
model = BertForSequenceClassification.from_pretrained('hfl/macbert-base', num_labels=10)
# 参数量:110M
# 准确率:83.1%(冻结 backbone)
# 准确率:87.4%(全参数微调)
最终选了 RoBERTa,效果最好。
微调策略
策略一:冻结 backbone
只训练分类头,速度快,显存占用小。
from transformers import AdamW
# 冻结所有 BERT 层
for param in model.bert.parameters():
param.requires_grad = False
# 只训练分类头
optimizer = AdamW(model.classifier.parameters(), lr=2e-5)
for epoch in range(5):
for batch in train_loader:
optimizer.zero_grad()
outputs = model(**batch)
loss = outputs.loss
loss.backward()
optimizer.step()
# 训练时间:15 分钟
# 显存占用:2.1 GB
# 准确率:84.3%
策略二:全参数微调
所有参数都更新,效果最好,但资源消耗大。
# 全部参数可训练
for param in model.parameters():
param.requires_grad = True
optimizer = AdamW(model.parameters(), lr=2e-5)
for epoch in range(10):
for batch in train_loader:
optimizer.zero_grad()
outputs = model(**batch)
loss = outputs.loss
loss.backward()
optimizer.step()
# 训练时间:1.5 小时
# 显存占用:10.8 GB
# 准确率:88.2%
策略三:逐层解冻
先训练分类头,再逐步解冻上层,最后微调整个模型。
# 阶段 1:只训练分类头(3 epochs)
for param in model.bert.parameters():
param.requires_grad = False
optimizer = AdamW(model.classifier.parameters(), lr=2e-5)
# 阶段 2:解冻最后 2 层(3 epochs)
for param in model.bert.encoder.layer[-2:].parameters():
param.requires_grad = True
optimizer = AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr=1e-5)
# 阶段 3:全参数微调(4 epochs)
for param in model.parameters():
param.requires_grad = True
optimizer = AdamW(model.parameters(), lr=5e-6)
# 训练时间:1.2 小时
# 显存占用:10.8 GB
# 准确率:88.5%
三种微调策略的准确率差距不大,但训练成本差异明显——把各方法的效果叠在一起看,更容易决定该走「快」还是「准」的路线。

最终我们选了逐层解冻的思路做参考,但后续领域适应和 LoRA 方案在此基础上又往前推了一步。
领域适应
预训练模型是通用语料,医疗领域有自己的术语和表达习惯。
方法一:继续预训练
用大量无标注医疗文本继续预训练模型。
from transformers import BertForMaskedLM, TextDatasetForNextSentencePrediction, DataCollatorForLanguageModeling
from transformers import Trainer, TrainingArguments
# 加载预训练模型
model = BertForMaskedLM.from_pretrained('hfl/chinese-roberta-wwm-ext')
# 准备医疗文本数据(50万条无标注病历)
dataset = TextDatasetForNextSentencePrediction(
tokenizer=tokenizer,
file_path='medical_texts.txt',
block_size=128
)
data_collator = DataCollatorForLanguageModeling(
tokenizer=tokenizer,
mlm=True,
mlm_probability=0.15
)
training_args = TrainingArguments(
output_dir='./mlm_finetuned',
overwrite_output_dir=True,
num_train_epochs=3,
per_device_train_batch_size=32,
save_steps=10000,
save_total_limit=2,
)
trainer = Trainer(
model=model,
args=training_args,
data_collator=data_collator,
train_dataset=dataset
)
trainer.train()
# 保存模型
model.save_pretrained('./medical_bert')
tokenizer.save_pretrained('./medical_bert')
# 在医疗文本上继续预训练后,分类准确率:89.3%
方法二:领域自适应微调
带领域标注的微调,让模型学习领域特征。
# 在分类任务中加入领域标签
class MedicalDataset(torch.utils.data.Dataset):
def __init__(self, texts, labels, domain_labels, tokenizer):
self.texts = texts
self.labels = labels
self.domain_labels = domain_labels
self.tokenizer = tokenizer
def __len__(self):
return len(self.texts)
def __getitem__(self, idx):
encoding = self.tokenizer(
self.texts[idx],
truncation=True,
padding='max_length',
max_length=128,
return_tensors='pt'
)
return {
'input_ids': encoding['input_ids'].flatten(),
'attention_mask': encoding['attention_mask'].flatten(),
'labels': torch.tensor(self.labels[idx]),
'domain_labels': torch.tensor(self.domain_labels[idx])
}
# 模型要同时预测分类和领域
class MultiTaskModel(BertForSequenceClassification):
def __init__(self, config):
super().__init__(config)
self.domain_classifier = torch.nn.Linear(config.hidden_size, 2)
def forward(self, input_ids, attention_mask, labels=None, domain_labels=None):
outputs = super().forward(input_ids, attention_mask, labels=labels)
pooled_output = self.bert(input_ids, attention_mask=attention_mask).pooler_output
domain_logits = self.domain_classifier(pooled_output)
loss = None
if labels is not None and domain_labels is not None:
classification_loss = torch.nn.functional.cross_entropy(outputs.logits, labels)
domain_loss = torch.nn.functional.cross_entropy(domain_logits, domain_labels)
loss = classification_loss + 0.3 * domain_loss
return {'loss': loss, 'logits': outputs.logits, 'domain_logits': domain_logits}
# 多任务学习后,分类准确率:89.8%
踩过的坑
坑一:过拟合
数据量小,全参数微调容易过拟合。
现象:
- 训练集准确率 98%
- 验证集准确率 78%
解决:
# 1. 加 dropout
model.bert.dropout.p = 0.3
# 2. 加权重衰减
optimizer = AdamW(model.parameters(), lr=2e-5, weight_decay=0.01)
# 3. 加早停
from transformers import EarlyStoppingCallback
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
callbacks=[EarlyStoppingCallback(early_stopping_patience=2)]
)
# 4. 数据增强
# 回译、同义词替换、随机删除等
坯二:学习率问题
学习率太大,模型崩掉;学习率太小,收敛太慢。
现象:
- 学习率 1e-3:训练集 loss 直接 NaN
- 学习率 1e-6:训练 20 个 epoch 还没收敛
解决:
# 用学习率预热
from transformers import get_linear_schedule_with_warmup
total_steps = len(train_loader) * num_epochs
optimizer = AdamW(model.parameters(), lr=2e-5)
scheduler = get_linear_schedule_with_warmup(
optimizer,
num_warmup_steps=int(0.1 * total_steps),
num_training_steps=total_steps
)
for epoch in range(num_epochs):
for batch in train_loader:
optimizer.zero_grad()
outputs = model(**batch)
loss = outputs.loss
loss.backward()
optimizer.step()
scheduler.step()
# 2e-5 + 预热效果最好
坑三:显存不够
全参数微调显存占用太高,显存不够用。
现象:
RuntimeError: CUDA out of memory. Tried to allocate 2.00 GiB
解决:
# 1. 混合精度训练
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for batch in train_loader:
optimizer.zero_grad()
with autocast():
outputs = model(**batch)
loss = outputs.loss
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
# 显存占用从 10.8 GB 降到 5.6 GB
# 2. 梯度累积
accumulation_steps = 4
for i, batch in enumerate(train_loader):
with autocast():
outputs = model(**batch)
loss = outputs.loss / accumulation_steps
scaler.scale(loss).backward()
if (i + 1) % accumulation_steps == 0:
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
# 3. 减小 batch size
# 从 32 降到 8,用梯度累积等效补回来
坏四:灾难性遗忘
微调后模型在原任务上的性能下降。
现象:
- 微调后,在通用文本上的性能下降 15%
- 医疗领域表现提升,但泛化能力变差
解决:
# 1. 正则化(Elastic Weight Consolidation)
class EWC:
def __init__(self, model, dataloader):
self.model = model
self.fisher = self.compute_fisher(dataloader)
self.optimal_params = {n: p.clone() for n, p in model.named_parameters()}
def compute_fisher(self, dataloader):
fisher = {}
for n, p in self.model.named_parameters():
fisher[n] = torch.zeros_like(p)
self.model.eval()
for batch in dataloader:
outputs = self.model(**batch)
loss = outputs.loss
loss.backward()
for n, p in self.model.named_parameters():
if p.grad is not None:
fisher[n] += p.grad.pow(2)
for n in fisher:
fisher[n] /= len(dataloader)
return fisher
def penalty(self):
loss = 0
for n, p in self.model.named_parameters():
loss += (self.fisher[n] * (p - self.optimal_params[n]).pow(2)).sum()
return loss
# 训练时加入 EWC 损失
ewc = EWC(model, original_dataloader)
for batch in train_loader:
outputs = model(**batch)
loss = outputs.loss + 0.1 * ewc.penalty()
# 2. 参数效率微调(LoRA)
from peft import LoraConfig, get_peft_model
lora_config = LoraConfig(
r=8,
lora_alpha=32,
target_modules=["query", "value"],
lora_dropout=0.1,
bias="none",
task_type="SEQ_CLS"
)
model = get_peft_model(model, lora_config)
# 只训练 LoRA 参数,保持原参数不变
# 微调参数量从 110M 降到 2.4M
# 准确率:88.1%(稍微下降一点,但泛化能力更好)
最终方案
折腾了一圈,最终方案是:
- 用 RoBERTa-Base 作为基础模型
- 在医疗文本上继续预训练 3 个 epoch
- 用 LoRA 微调分类任务
- 混合精度训练 + 梯度累积
# 完整流程
from transformers import AutoTokenizer, AutoModelForSequenceClassification, TrainingArguments, Trainer
from peft import LoraConfig, get_peft_model
from torch.cuda.amp import autocast, GradScaler
# 1. 加载领域自适应模型
tokenizer = AutoTokenizer.from_pretrained('./medical_bert')
model = AutoModelForSequenceClassification.from_pretrained('./medical_bert', num_labels=10)
# 2. 配置 LoRA
lora_config = LoraConfig(
r=16,
lora_alpha=32,
target_modules=["query", "key", "value"],
lora_dropout=0.05,
bias="none",
task_type="SEQ_CLS"
)
model = get_peft_model(model, lora_config)
# 3. 训练配置
training_args = TrainingArguments(
output_dir='./final_model',
num_train_epochs=8,
per_device_train_batch_size=16,
gradient_accumulation_steps=2,
learning_rate=3e-5,
warmup_ratio=0.1,
weight_decay=0.01,
logging_steps=100,
evaluation_strategy='epoch',
save_strategy='epoch',
load_best_model_at_end=True,
metric_for_best_model='eval_accuracy',
fp16=True, # 混合精度
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
compute_metrics=compute_metrics
)
trainer.train()
# 4. 评估
results = trainer.evaluate(test_dataset)
print(results)
# 最终效果:
# 训练时间:45 分钟
# 显存占用:4.2 GB
# 准确率:89.6%
# F1-score:0.889
写在最后
迁移学习这东西,算是小数据场景的救星。
解决了:
- 数据量不足的问题
- 从头训练成本高的问题
- 泛化能力差的问题
带来了:
- 领域差异需要处理
- 微调策略要调
- 资源消耗不小
实践中要考虑:
- 数据量和任务复杂度匹配
- 预训练模型和领域差异
- 微调策略和资源限制
- 过拟合和灾难性遗忘
不是所有场景都需要迁移学习。数据量够、领域差异大、资源充足,从头训练可能更合适。
这次迁移学习实践花了两周,从预训练模型选择到最终方案落地。最终准确率从 72% 提升到 89.6%,但中间踩的坑不少。
版权声明: 本文首发于 指尖魔法屋-迁移学习:这次怎么落地的(https://blog.thinkmoon.cn/post/174-transfer-learning-finetuning-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。