一、问题背景:为什么非要微调7B模型

上周接了个法律文书自动摘要的需求,试了直接调用Llama-2-7B-chat,效果惨不忍睹——生成的摘要经常漏掉关键法条引用,而且输出格式混乱。用ChatGPT API倒是效果好,但客户要求数据不出内网。

无奈只能本地微调。但7B模型全参数微调需要多少显存?我算了下:model权重14GB(fp16),梯度14GB,优化器状态(AdamW的momentum和variance)28GB,再加上激活值,总共至少112GB。这得4张A100才跑得动。

好在LoRA(Low-Rank Adaptation)能解决这个问题。它的核心思想是冻结原模型权重,只训练注入的低秩矩阵。以Llama-2-7B为例,attention层的q_proj、v_proj各注入一个r=8的LoRA模块,可训练参数量只有4.2M,占总参数的0.06%。这就是为什么显存占用能缩到1/5。

二、环境与版本:踩过坑的版本组合

先说我踩的第一个坑:transformers版本不对会导致prepare_model_for_kbit_training报错。最后锁定了一套稳定组合:

torch==2.0.1+cu118
transformers==4.31.0
peft==0.5.0
bitsandbytes==0.41.1
datasets==2.14.5
accelerate==0.22.0

关于QLoRA需要特别说明:它是在LoRA基础上把基座模型量化为4-bit NF4格式,进一步降低显存。我实际用下来,QLoRA(4-bit)比LoRA(fp16)显存再降40%,但训练速度会慢25%左右。如果你有24GB显存,建议直接上QLoRA;如果是32GB以上,用fp16的LoRA效果更好。

三、方案设计:LoRA vs QLoRA的取舍

我的目标是跑通流程并保留微调效果,所以设计如下:

方案 基座精度 显存占用 训练速度 效果
LoRA fp16 21GB 1.0x ROUGE-L=0.51
QLoRA 4-bit NF4 13GB 0.75x ROUGE-L=0.47

最终选了LoRA+fp16方案,因为单张3090(24GB)刚好能塞下,且效果最优。

LoRA关键超参数如下:
- r=8:秩的大小,越大表达能力越强,但显存和过拟合风险增加
- alpha=16:缩放因子,实际效果是alpha/r=2的缩放比,太大会导致训练不稳定
- dropout=0.1:防止过拟合
- target_modules:指定注入的层,我选了q_proj, v_proj,经验上只改这两个就够

四、核心实现:数据准备与训练代码

数据准备:我爬了裁判文书网公开的法律文书,清洗后构造了1.2万条“案情描述→判决摘要”对。按8:1:1划分训练/验证/测试集。

# 数据预处理:格式化prompt
def format_prompt(sample):
    return {
        "input_ids": tokenizer.apply_chat_template(
            [{"role": "user", "content": sample["case"]},
             {"role": "assistant", "content": sample["summary"]}],
            tokenize=True,
            return_tensors="pt",
            max_length=2048,
            truncation=True
        )[0]
    }

dataset = load_dataset("json", data_files="legal_data.jsonl")
dataset = dataset.map(format_prompt, remove_columns=dataset["train"].column_names)

训练配置:用HuggingFace的Trainer配合peft库。关键点在于prepare_model_for_kbit_trainingget_peft_model的配合。

from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
import torch

model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-2-7b-chat-hf",
    torch_dtype=torch.float16,
    device_map="auto",
    load_in_8bit=False  # 如果显存不够,改成True
)

# 冻结原模型参数
model = prepare_model_for_kbit_training(model)

lora_config = LoraConfig(
    r=8,
    lora_alpha=16,
    target_modules=["q_proj", "v_proj"],
    lora_dropout=0.1,
    bias="none",
    task_type="CAUSAL_LM"
)

model = get_peft_model(model, lora_config)

training_args = TrainingArguments(
    output_dir="./lora_legal_7b",
    per_device_train_batch_size=4,
    gradient_accumulation_steps=8,  # 等效batch_size=32
    num_train_epochs=3,
    learning_rate=2e-4,
    fp16=True,
    logging_steps=50,
    save_strategy="epoch",
    evaluation_strategy="steps",
    eval_steps=200,
    gradient_checkpointing=True,  # 显存不够时开启
    optim="paged_adamw_8bit",
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=dataset["train"],
    eval_dataset=dataset["validation"],
    tokenizer=tokenizer,
)

trainer.train()
model.save_pretrained("./lora_legal_7b_final")

五、踩坑与优化:三个重要教训

坑1:梯度检查点与LoRA的冲突。开gradient_checkpointing=True后,前向传播会重新计算激活值以省显存,但LoRA层可能不被支持。我遇到报错RuntimeError: element 0 of tensors does not require grad,解决方法是把model.enable_input_require_grads()加在prepare_model_for_kbit_training之后。

坑2:loss不下降的元凶。第一版训练跑了500步,loss一直卡在1.8左右不动。后来发现是学习率太大(用了全参微调惯用的3e-5)。LoRA的有效参数量少,需要更大学习率,但2e-4又会导致震荡。最后用余弦衰减+warmup 100步,初始lr=1e-4,效果才稳定。

坑3:NF4量化后推理要加torch_dtype=torch.float16。QLoRA微调后导出模型,如果不指定推理精度,会默认用4-bit加载,导致输出质量严重下降。

六、效果数据:量化对比

Loss曲线:训练集loss从2.13降到0.41,验证集loss稳定在0.47左右,无过拟合。

推理效果对比(测试集200条):

指标 原版Llama-2-7B LoRA微调后 提升
ROUGE-L 0.23 0.51 +121%
法条引用正确率 34% 78% +129%
输出格式规范率 57% 96% +68%

显存占用
- 全参微调(理论):112GB
- LoRA(fp16):21GB,节省81%
- QLoRA(4-bit):13GB,节省88%

推理速度:微调后模型与base模型速度一致,约35 tokens/s(单张A100),因为LoRA模块计算量可以忽略。

七、总结

LoRA微调7B模型在消费级显卡上完全可行,关键在于控制ralpha的比例。如果目标任务是垂直领域(如法律、医疗),效果提升非常显著;但如果你需要强推理能力(如代码生成),建议用更大的基座模型。

最后补充一点:微调后的LoRA权重只有16MB,部署时可以用PeftModel.from_pretrained动态加载,完全不影响推理延迟。这套流程我后来又跑过Mistral-7B和Qwen-7B,核心逻辑完全一致,只需改模型名和target_modules即可。