迁移学习落地实践:医疗文本分类的预训练与微调
去年接了一个中文医疗问答分类的项目,数据量只有 3000 条,标注质量一般,模型效果一直上不去。折腾了一圈发现,同样是 BERT-base,别人能调到 92% F1,我就卡在 85% 上不去。
问题背景
医疗问答分类的场景不算复杂:输入一段用户咨询,输出预设的 10 个类别(症状描述、用药咨询、检查解读、预防建议等)。但有几个硬性限制:
- 数据量少,只有 3000 条样本,类别分布不均匀,最小类只有 80 条
- 医疗术语多,通用模型对"CT增强扫描"、“孕早期唐筛"这类词理解有限
- 对准确率要求高,把用药咨询误判成预防建议在医疗场景下后果比较严重
- 部署环境有限,推理时间要控制在 50ms 以内,无法上超大模型
在转向深度模型之前,先用传统方法打了底,确认这点数据量确实撑不起从零学习:
# 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().fit(X_train, train_labels) # 准确率:68.5%
# FastText 监督训练
# 准确率:72.3%
都不太行,数据量太小,模型学不到足够的特征。这才转向迁移学习。
预训练模型选择
正式动手前,对比了几个开源中文模型(下面是快速验证的准确率,正式评估后续改用更贴合类别不平衡场景的 F1):
| 模型 | 参数量 | 冻结 backbone | 全参数微调 |
|---|---|---|---|
| bert-base-chinese | 110M | 82.1% | 85.7% |
| hfl/chinese-roberta-wwm-ext | 110M | 84.3% | 88.2% |
| hfl/macbert-base | 110M | 83.1% | 87.4% |
最终选了 RoBERTa-wwm-ext 作为基础模型,全参数微调的底子最好。但 Fine-tune 了一轮后发现效果一般。想想也知道,这个模型是在大规模通用中文语料上训练的,医疗领域的专业知识和语言风格覆盖有限。
迁移学习思路
迁移学习在 NLP 里的典型路径是:大规模通用预训练 → 领域语料继续预训练 → 任务数据微调。这三个阶段解决了不同层次的问题:
- 通用预训练:让模型学会语言的基本结构和语义表示,这是"基础能力”
- 领域预训练:让模型适应特定领域的词汇、表达习惯和知识分布,这是"专业能力"
- 任务微调:让模型学会处理具体任务,比如分类、抽取、生成,这是"应用能力"
用个通俗的类比:通用预训练像读完小学,掌握读写算;领域预训练像进入医学院,学会医学术语和临床思维;任务微调像实习值班,学会怎么真正处理病人。
这里有个重要经验:不要跳过领域预训练。除非你的领域非常通用(比如新闻、电商、社交媒体),否则从通用模型直接到任务微调,数据少的情况下效果通常不如完整路径。
领域预训练实施
收集领域语料是第一道坎。医疗领域的公开语料资源有限,我是这么解决的:
- 整理了 10 万条公开医疗问答数据,来自中文医学问答数据集和开源医疗社区
- 加上了 5 万条医疗指南和科普文章,来源包括卫健委官网、三甲医院官网
- 数据清洗做了两件事:去掉明显的噪声(HTML标签、乱码)和过滤重复内容
清洗前后的数据统计:
# 原始数据统计
raw_samples = 180000
raw_unique = 125000
after_dedup = 120000
after_quality_filter = 150000
# 最终用于预训练的数据
train_size = 120000
vocab_coverage = 0.85 # 医疗术语覆盖率
预训练任务用了标准的 Masked Language Modeling (MLM),但做了几个调整:
- Mask 策略调整:对医疗术语(长词、专业词)的 mask 概率提高到 20%,普通词保持 15%。这样能让模型更关注领域专有词汇的学习。
- 训练长度选择:句子长度控制在 512,但统计了实际医疗咨询的长度分布,70% 落在 128 以内,所以后期主要在 128-256 长度上强化训练。
- 学习率策略:使用了 3 个 epoch,学习率从 5e-5 线性衰减到 0,warmup 比例 10%。
# 关键配置片段
from transformers import BertConfig, BertForMaskedLM, BertTokenizer
config = BertConfig.from_pretrained('hfl/chinese-roberta-wwm-ext')
config.vocab_size = tokenizer.vocab_size
model = BertForMaskedLM.from_pretrained(
'hfl/chinese-roberta-wwm-ext',
config=config
)
# 自定义 masking 策略
def mask_tokens(inputs, tokenizer, mlm_probability=0.15, term_mask_prob=0.20):
"""
对医疗术语提高 masking 概率
"""
labels = inputs.clone()
probability_matrix = torch.full(labels.shape, mlm_probability)
special_tokens_mask = [
tokenizer.get_special_tokens_mask(val, already_has_special_tokens=True)
for val in labels.tolist()
]
probability_matrix.masked_fill_(torch.tensor(special_tokens_mask, dtype=torch.bool), value=0.0)
# 识别医疗术语(这里简化处理)
medical_terms = identify_medical_terms(inputs, tokenizer)
for term_positions in medical_terms:
probability_matrix[term_positions] = term_mask_prob
masked_indices = torch.bernoulli(probability_matrix).bool()
labels[~masked_indices] = -100
indices_replaced = torch.bernoulli(torch.full(labels.shape, 0.8)).bool() & masked_indices
inputs[indices_replaced] = tokenizer.convert_tokens_to_ids(tokenizer.mask_token)
return inputs, labels
预训练花了大约 6 小时(单卡 RTX 3090),loss 从 2.8 降到了 1.7 左右。训练过程中做了一次中期评估,用 100 条医疗句子做 fill-mask 测试,对专业词的预测准确率提升了 25 个百分点。
替代思路:领域自适应多任务微调
除了 MLM 继续预训练,还试过一种带领域标签的多任务微调:在主分类任务之外,加一个轻量的领域分类头,让模型显式建模"这段文本是不是医疗领域"。
# 同时预测分类和领域,共享 backbone
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% 左右,但依赖带领域标注的数据,标注成本比纯 MLM 高。最终主推的还是无监督的 MLM 路径,多任务方案作为备选。
任务微调
领域预训练完成后,拿到了一个"懂医疗"的 BERT 模型。接下来是任务微调,这次就比较标准了。
数据准备
3000 条任务数据做了这样几件事:
- 类别平衡处理:对小于 200 条的类别做了数据增强(同义词替换、回译),同时限制最大类样本数
- 划分比例:训练集 2400、验证集 300、测试集 300
- 文本截断:统计长度分布后,截断位置设为 256(足够覆盖 95% 的样本)
# 数据增强示例
def augment_text(text, augment_ratio=1.5):
"""对小类数据进行增强"""
augment_methods = [
synonym_replace,
back_translate,
random_insert,
]
aug_texts = [text]
for _ in range(int(augment_ratio - 1)):
method = random.choice(augment_methods)
aug_texts.append(method(text))
return aug_texts
# 类别统计和平衡
class_counts = Counter(train_labels)
min_class_size = 200
balanced_data = []
for text, label in zip(train_texts, train_labels):
if class_counts[label] < min_class_size:
aug_texts = augment_text(text, min_class_size / class_counts[label])
balanced_data.extend([(t, label) for t in aug_texts])
else:
balanced_data.append((text, label))
模型配置
分类头用了最简单的线性层加 softmax,输出 10 个类别:
from transformers import BertForSequenceClassification
model = BertForSequenceClassification.from_pretrained(
'./domain_pretrained_medical_bert',
num_labels=10
)
训练参数是调了几轮后定下来的:
training_args = TrainingArguments(
output_dir='./results',
num_train_epochs=5,
per_device_train_batch_size=16,
per_device_eval_batch_size=32,
warmup_steps=100,
weight_decay=0.01,
logging_dir='./logs',
logging_steps=50,
evaluation_strategy='epoch',
save_strategy='epoch',
load_best_model_at_end=True,
metric_for_best_model='f1',
learning_rate=2e-5,
)
学习率从 5e-5 开始试,发现 2e-5 效果最好,过小收敛慢,过大容易过拟合。
效果对比
把各阶段 F1 串起来看,领域预训练是最大跃升,数据增强和学习率调优属于锦上添花。

从 0.85 到 0.93 的增益主要来自中间那步领域预训练,而不是在通用模型上反复微调。
几个关键节点的效果对比(测试集 F1 score):
| 方案 | F1 | 推理时间 (ms) | 训练耗时 |
|---|---|---|---|
| 通用 RoBERTa 直接微调 | 0.85 | 42 | 1.5h |
| + 数据增强 | 0.87 | 42 | 2.0h |
| + 领域预训练 | 0.92 | 44 | 8.5h |
| + 调优学习率 | 0.93 | 44 | 8.5h |
推理时间增加 2ms 基本可以接受,部署时做了 FP16 量化,能压到 35ms 以内。
微调策略对比
上面走的是全参数微调,但实际落地时"怎么调参"还有几种节省资源或更稳的策略,值得横向对比一下。
策略一:冻结 backbone,只训练分类头。速度快、显存小,适合快速验证。
for param in model.bert.parameters():
param.requires_grad = False
optimizer = AdamW(model.classifier.parameters(), lr=2e-5)
# 训练 15 分钟,显存 2.1 GB,准确率 84.3%
策略二:全参数微调,所有参数都更新。效果最好,但资源消耗大。
for param in model.parameters():
param.requires_grad = True
optimizer = AdamW(model.parameters(), lr=2e-5)
# 训练 1.5 小时,显存 10.8 GB,准确率 88.2%
策略三:逐层解冻,先训练分类头,再逐步解冻上层,最后微调整模型,兼顾成本和效果。
# 阶段 1:只训练分类头(3 epochs),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)
optimizer = AdamW(model.parameters(), lr=5e-6)
# 训练 1.2 小时,显存 10.8 GB,准确率 88.5%
三种策略准确率差距不大,但训练成本差异明显。最终线上考虑到显存和迭代频率,选了更轻量的 LoRA 方案(见下文踩坑部分)。
踩坑记录
这个项目踩的坑不少,挑几个典型的说。
坑一:预训练数据质量问题
一开始没做严格的数据清洗,导致模型学到一些错误的知识。比如有些社区问答里混杂了广告、推广内容,模型会把"某某药效果好"当成普遍规律。
后来加了几个过滤规则:
def is_valid_sample(text):
"""数据质量过滤"""
# 过滤广告关键词
ad_keywords = ['广告', '推广', '优惠', '购买链接']
if any(kw in text for kw in ad_keywords):
return False
# 过滤过短文本
if len(text.strip()) < 10:
return False
# 过滤非中文为主的内容
chinese_ratio = len(re.findall(r'[一-鿿]', text)) / len(text)
if chinese_ratio < 0.7:
return False
return True
坑二:领域过拟合
领域预训练完成后,直接在任务数据上微调,发现模型对医疗相关的样本效果好,但对一些日常用语反而"退步"了。这是因为领域预训练让模型过度聚焦领域知识,通用能力下降。
解决方法是做了一点"混合训练":在任务微调阶段,保留 10% 的通用样本作为正则化,避免模型完全忘记通用语言能力。
# 混合训练数据
domain_train = load_medical_qa_data()
general_train = load_general_qa_data()[:200] # 少量通用数据
mixed_train = domain_train + general_train
shuffle(mixed_train)
坑三:类别不平衡处理不当
一开始用了 oversampling,发现模型容易记住重复样本,测试时泛化能力差。后来改用数据增强,虽然效果好一些,但生成的文本有时候不太自然,反而引入噪声。
最后的方案是:对小类做适度增强,同时对大类做一定程度的欠采样,保持整体类别分布相对均衡,但不强制完全平衡。
坑四:评估指标误导
一开始只关注准确率,发现模型在最大类上表现很好,但小类 F1 很低。后来改用加权 F1 和 macro F1 结合评估,确保模型在各类别上都有合理表现。
from sklearn.metrics import precision_recall_fscore_support
def compute_metrics(pred):
labels = pred.label_ids
preds = pred.predictions.argmax(-1)
precision, recall, f1, _ = precision_recall_fscore_support(
labels, preds, average='weighted'
)
macro_f1 = precision_recall_fscore_support(
labels, preds, average='macro'
)[2]
return {
'accuracy': (preds == labels).mean(),
'f1': f1,
'precision': precision,
'recall': recall,
'macro_f1': macro_f1 # 关注小类表现
}
坑五:学习率与训练不稳定
学习率太大模型直接崩掉,太小又收敛过慢。
现象:学习率 1e-3 时,训练 loss 直接 NaN;1e-6 时,训练 20 个 epoch 还没收敛。
解决:用线性预热 + 衰减调度器,前 10% 步骤慢慢把学习率顶上去。
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()
loss = model(**batch).loss
loss.backward()
optimizer.step()
scheduler.step()
# 2e-5 + 10% warmup 效果最稳
坑六:显存不足
全参数微调显存吃紧,单卡经常 OOM:RuntimeError: CUDA out of memory。
解决:混合精度训练 + 梯度累积 + 减小 batch size 三件套。
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
accumulation_steps = 4 # batch size 从 32 降到 8,用梯度累积等效补回
for i, batch in enumerate(train_loader):
optimizer.zero_grad()
with autocast():
loss = model(**batch).loss / accumulation_steps
scaler.scale(loss).backward()
if (i + 1) % accumulation_steps == 0:
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
# 显存占用从 10.8 GB 降到 5.6 GB
坑七:灾难性遗忘与参数高效微调
领域微调后,模型在通用文本上性能下降约 15%——医疗领域提升了,泛化能力却变差了。除了前文"混合训练"那招,这里引入两种更系统的办法。
EWC(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 = {n: torch.zeros_like(p) for n, p in self.model.named_parameters()}
self.model.eval()
for batch in dataloader:
self.model(**batch).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):
return sum(
(self.fisher[n] * (p - self.optimal_params[n]).pow(2)).sum()
for n, p in self.model.named_parameters()
)
# 训练时把 EWC 惩罚加进 loss
ewc = EWC(model, original_dataloader)
for batch in train_loader:
loss = model(**batch).loss + 0.1 * ewc.penalty()
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)
# 可训练参数从 110M 降到 2.4M,准确率 88.1%,泛化能力更好
结果与经验
全参数微调在测试集上的最好表现:
- Weighted F1: 0.93
- Macro F1: 0.91(小类表现尚可)
- 推理时间: 35ms(FP16 量化后)
部署后做了线上 A/B 测试,相比之前的规则系统,用户满意度提升了 18%,误分类率降低了 30%。
综合考虑显存和迭代频率,线上最终落地的不是全参数微调,而是在领域预训练模型之上套一层 LoRA:可训练参数压到 2.4M,单卡显存 4.2 GB,训练 45 分钟,F1 保持在 0.889 左右,而每次新增类别时只需重训这层低秩矩阵,迭代成本低了很多。
# 线上最终流程(简化)
from peft import LoraConfig, get_peft_model
from transformers import AutoModelForSequenceClassification, TrainingArguments, Trainer
model = AutoModelForSequenceClassification.from_pretrained(
'./domain_pretrained_medical_bert', num_labels=10
)
model = get_peft_model(model, LoraConfig(
r=16, lora_alpha=32,
target_modules=["query", "key", "value"],
lora_dropout=0.05, bias="none", task_type="SEQ_CLS",
))
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,
fp16=True,
evaluation_strategy='epoch',
load_best_model_at_end=True,
metric_for_best_model='f1',
)
这次实践总结了几条比较实在的经验:
数据量小时,迁移学习是必须的,但路径要对:不要指望从通用模型直接跳到任务效果就好,中间的领域预训练很关键。
领域预训练的质量比数量更重要:5 万条高质量的领域语料,可能比 20 万条低质量数据效果更好。
不要忽视推理时间:训练时上各种花哨技巧没问题,但部署时要考虑实际环境,适当做量化和剪枝。
评估指标要贴合业务:准确率好看不代表什么,要看业务真正关心的是什么(比如医疗场景下假阳性的成本)。
不要迷信"越大越好":对于这个项目,BERT-base 加上合理的迁移学习路径,效果已经足够,没必要上更大模型。
一些延伸思考
迁移学习在 NLP 里已经成了标准做法,但实际落地时还有很多细节需要根据场景调整。
比如这次医疗问答场景,数据量确实小,但如果数据量达到几万条级别,直接从通用模型微调可能就够了,领域预训练的边际收益会下降。反之,如果是一些更垂直的领域(比如法律文书、金融分析),领域预训练的价值会更大。
另一个值得提的点:预训练和微调不是"一次就完事"的。随着业务发展,新数据、新类别不断出现,模型需要持续迭代。一个合理的策略是:定期用新数据做领域预训练更新,然后在新任务数据上微调。
最后说句实话:迁移学习不是银弹,数据量特别大时直接训练可能更高效,特别小时可能效果也有限。但它提供了一个相对可靠的基础路径,尤其适合那些"数据不太少也不太多"的场景。关键是要根据实际需求调整路径,不要照搬论文里的标准流程。
折腾到现在,这个模型已经在生产环境跑了 8 个月,每月迭代一次,效果还算稳定。迁移学习这条路,算是走通了。
本文整合了关于迁移学习的多篇笔记。
版权声明: 本文首发于 指尖魔法屋-迁移学习落地实践:医疗文本分类的预训练与微调(https://blog.thinkmoon.cn/post/307-ai-transfer-learning-pretraining-finetuning-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。