linear probing & lora
·
模型微调
linear probing
linear probing(仅训练最后一层线性层,冻结模型主干)可以实现训练,且在特定场景下效果接近甚至媲美 LoRA / 全量微调;但在复杂场景下,效果会明显低于参数高效微调(如 LoRA)。
| 场景 | 效果 | 原因 |
|---|---|---|
| 小样本+简单任务 | 适合 | 数据少(<1000例),linear probing避免过拟合,复用预训练通用特征 |
| 大样本+复杂任务 | 一般 | 预训练特征无法适配专属场景(如罕见病灶),需微调主干层优化特征 |
| 跨模态对齐(CT+文本) | 较差 | 视觉-语言对齐需要微调跨模态交互层,linear probing无法触达这些层 |
适用场景
- 预训练特征足够强:大模型已经在足够多的数据集上进行训练,视觉塔等特征提取能力已经很强。
- 简单任务:线性层足够拟合简单的任务,例如分类
- 无过拟合风险:数据集样本少(几百例),冻结主干后仅训练几十万个参数(线性层),远低于LoRA的百万级参数,几乎不会过拟合。
关键代码
# 核心:冻结所有模型参数(视觉塔+语言塔)
for param in model.parameters():
param.requires_grad = False # 所有参数不参与梯度更新
# 2. 新增线性分类头(仅这层可训练)
# in_dim替换为视觉特征维度,根据你的任务设置分类数
class LinearProbeHead(torch.nn.Module):
def __init__(self, in_dim=2560, num_classes=3):
super().__init__()
self.linear = torch.nn.Linear(in_dim, num_classes)
def forward(self, visual_features):
return self.linear(avg_feat)
probe_head = LinearProbeHead(in_dim=2560, num_classes=3).to("cuda")
# 只给线性层参数创建优化器(主干参数不更新)
optimizer = torch.optim.AdamW(
probe_head.parameters(), # 仅优化线性层
lr=1e-3, # linear probing学习率可设大一点
weight_decay=1e-5
)
loss_fn = torch.nn.CrossEntropyLoss()
# 训练过程部分
# 2. 线性层前向传播(仅这步计算梯度)
logits = probe_head(visual_features)
loss = loss_fn(logits, batch["labels"]) # labels是分类标签(0/1/2)
# 3. 反向传播(仅更新线性层)
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
也可考虑将单层线性提升为多层线性,提升模型复杂度,实现更好的效果
linear probing到LoRA
如果linear probing效果达不到预期(如准确率<70%),且满足以下条件,可升级为LoRA:
-
数据集样本数>1000例(有足够数据微调主干);
-
任务是复杂场景(如多病灶联合诊断、跨模态报告生成);
-
显存充足(≥12G),能支撑LoRA的少量参数更新。
LoRA举例,以我最近的项目为例:
最近在微调一款开源的大模型,主要用于根据ct序列生成对应的医学报告。现实效果不理想,在私有数据集上进行微调。
基于PEFT(Parameter-Efficient Fine-Tuning,参数高效微调)库配置LoRA(Low-Rank Adaptation) 微调
# 创建lora config
from peft import LoraConfig
peft_config = LoraConfig(
# 1. LoRA核心缩放系数:控制低秩矩阵的输出幅度
lora_alpha=16,
# 2. 正则化:防止过拟合
lora_dropout=0.05,
# 3. LoRA的秩:决定低秩矩阵的维度(核心参数)
r=16,
# 4. 偏置参数训练策略:不训练任何偏置
bias="none",
# 5. 要应用LoRA的模型模块:所有线性层
target_modules="all-linear",
# 6. 微调任务类型:因果语言模型(自回归生成)
task_type="CAUSAL_LM",
# 7. 额外解冻并保存的模块:输出层+词嵌入层
modules_to_save=[
"lm_head",
"embed_tokens",
],
)
代码通过LoraConfig定义了 LoRA 微调的核心超参数,专门针对因果语言模型(CAUSAL_LM,即自回归生成类模型,如 LLaMA、GPT、Qwen 等) 设计
| 参数 | 取值 / 配置 | 通俗解释 & 作用 |
|---|---|---|
| lora_alpha | 16 | LoRA 的缩放系数(α),计算公式:LoRA输出 = (A×B) × (α/r);α 越大,LoRA 对模型输出的影响越强,16 是行业常用经验值。 |
| lora_dropout | 0.05 | LoRA 模块的 dropout 概率(训练时随机让 5% 的 LoRA 参数 “失效”),用于防止过拟合;0.05 是保守值,小样本微调时常用。 |
| r | 16 | LoRA 的 “秩”(低秩矩阵的维度),核心参数:r 越小→训练参数越少、速度越快,但可能损失精度;r 越大→效果越好,但显存占用越高;16 是“效率+效果”的折中值(常见还有 8、32)。 |
| bias | “none” | 偏置(bias)参数的训练策略:"none"→不训练任何偏置(LoRA 默认最优);"all"→训练所有偏置(显存高,没必要);"lora_only"→仅训练LoRA模块偏置(极少用)。 |
| target_modules | “all-linear” | 指定要添加LoRA适配器的模型模块:"all-linear"→对所有线性层加LoRA(通用简单);也可指定具体层(如[“q_proj”, “v_proj”])→仅对注意力层加LoRA(更精细)。 |
| task_type | “CAUSAL_LM” | 微调任务类型(需匹配模型用途):"CAUSAL_LM"→因果语言模型(文本生成/对话/续写);其他常见值:“SEQ_CLASSIFICATION”(文本分类)、“QUESTION_ANSWERING”(问答)。 |
| modules_to_save | [“lm_head”, “embed_tokens”] | LoRA 默认只训练低秩矩阵(其他模块冻结);此参数指定解冻并训练的层:lm_head→模型输出层(隐藏态映射为词表概率);embed_tokens→词嵌入层(文字转向量);解冻这两层能提升生成任务效果(尤其小样本微调)。 |
基于 Hugging Face 的TRL库配置SFT(Supervised Fine-Tuning,监督微调)
from trl import SFTConfig
num_train_epochs = 1 # @param {type: "number"}
learning_rate = 2e-4 # @param {type: "number"}
args = SFTConfig(
output_dir="medgemma-4b-it-sft-lora-crc100k", # Directory and Hub repository id to save the model to
num_train_epochs=num_train_epochs, # Number of training epochs
per_device_train_batch_size=4, # Batch size per device during training
per_device_eval_batch_size=4, # Batch size per device during evaluation
gradient_accumulation_steps=4, # Number of steps before performing a backward/update pass
gradient_checkpointing=True, # Enable gradient checkpointing to reduce memory usage
optim="adamw_torch_fused", # Use fused AdamW optimizer for better performance
logging_steps=50, # Number of steps between logs
save_strategy="epoch", # Save checkpoint every epoch
eval_strategy="steps", # Evaluate every `eval_steps`
eval_steps=50, # Number of steps between evaluations
learning_rate=learning_rate, # Learning rate based on QLoRA paper
bf16=True, # Use bfloat16 precision
max_grad_norm=0.3, # Max gradient norm based on QLoRA paper
warmup_ratio=0.03, # Warmup ratio based on QLoRA paper
lr_scheduler_type="linear", # Use linear learning rate scheduler
push_to_hub=True, # Push model to Hub
report_to="tensorboard", # Report metrics to tensorboard
gradient_checkpointing_kwargs={"use_reentrant": False}, # Set gradient checkpointing to non-reentrant to avoid issues
dataset_kwargs={"skip_prepare_dataset": True}, # Skip default dataset preparation to preprocess manually
remove_unused_columns = False, # Columns are unused for training but needed for data collator
label_names=["labels"], # Input keys that correspond to the labels
)
from trl import SFTTrainer
trainer = SFTTrainer(
model=model,
args=args,
train_dataset=data["train"],
eval_dataset=data["validation"].shuffle().select(range(200)), # Use subset of validation set for faster run
peft_config=peft_config,
processing_class=processor,
data_collator=collate_fn,
)
trainer.train()
以上配置是显存优先而非效率优先,体现在梯度检查点、bf16、梯度累积等。
梯度检查点:训练时不保存所有中间层的梯度,而是反向传播时重新计算,可减少 30%-50% 的显存占用(代价是训练速度慢 10%-20%)。
混合精度训练:用bf16(脑浮点数)替代fp32(单精度)存储模型参数 / 梯度,显存占用减半,且bf16比fp16更稳定(不会出现数值溢出),仅需 GPU 支持(A100、RTX30/40/50 系列、AMD MI 系列);
更多推荐


所有评论(0)