模型微调

linear probing

linear probing(仅训练最后一层线性层,冻结模型主干)可以实现训练,且在特定场景下效果接近甚至媲美 LoRA / 全量微调;但在复杂场景下,效果会明显低于参数高效微调(如 LoRA)。

场景效果原因
小样本+简单任务适合数据少(<1000例),linear probing避免过拟合,复用预训练通用特征
大样本+复杂任务一般预训练特征无法适配专属场景(如罕见病灶),需微调主干层优化特征
跨模态对齐(CT+文本)较差视觉-语言对齐需要微调跨模态交互层,linear probing无法触达这些层

适用场景

  1. 预训练特征足够强:大模型已经在足够多的数据集上进行训练,视觉塔等特征提取能力已经很强。
  2. 简单任务:线性层足够拟合简单的任务,例如分类
  3. 无过拟合风险:数据集样本少(几百例),冻结主干后仅训练几十万个参数(线性层),远低于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:

  1. 数据集样本数>1000例(有足够数据微调主干);

  2. 任务是复杂场景(如多病灶联合诊断、跨模态报告生成);

  3. 显存充足(≥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_alpha16LoRA 的缩放系数(α),计算公式:LoRA输出 = (A×B) × (α/r);α 越大,LoRA 对模型输出的影响越强,16 是行业常用经验值。
lora_dropout0.05LoRA 模块的 dropout 概率(训练时随机让 5% 的 LoRA 参数 “失效”),用于防止过拟合;0.05 是保守值,小样本微调时常用。
r16LoRA 的 “秩”(低秩矩阵的维度),核心参数: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 系列);

Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐