基于 PEFT LoRA 微调 Gemma-2-9b-it:从指令集构建到甄嬛风格对话模型实战(self-llm 项目)
【免费下载链接】self-llm《开源大模型食用指南》针对中国宝宝量身打造的基于Linux环境快速微调(全参数/Lora)、部署国内外开源大模型(LLM)/多模态大模型(MLLM)教程项目地址: https://gitcode.com/GitHub_Trending/se/self-llm
本文是《开源大模型食用指南》self-llm 项目中 Gemma-2 系列教程的实战篇。基于 transformers 4.42.3 与 peft 框架,围绕 Google 开源的 Gemma-2-9b-it 因果语言模型,完整演示「模型下载 → 指令集构建 → 数据格式化 → LoraConfig 配置 → Trainer 训练 → LoRA 权重推理」全流程,最终训练出一个能够模拟甄嬛对话风格的个性化 LLM。读完本文,你将掌握 Gemma-2 系列(含同架构 9B/27B 模型)在消费级显卡上做高效参数微调的标准方法,以及 LoRA 权重独立保存、独立加载的部署范式。
本教程配套的完整可运行代码位于 04-Gemma-2-9b-it peft lora微调.ipynb,建议结合本文逐 Cell 执行;同时可对照本仓库其他 Gemma-2 教程(FastApi 部署、WebDemo 部署)理解微调前后模型服务的差异。
为什么选择 LoRA 微调 Gemma-2-9b-it
Gemma-2-9b-it 是一个参数量约 92.7 亿的因果语言模型(从 Notebook 中model.print_trainable_parameters()打印的all params: 9,268,715,008可以确认)。全参数微调这样的模型需要多卡高显存环境,而 LoRA(Low-Rank Adaptation)通过在冻结的原始权重旁注入低秩可训练矩阵,把需要更新的参数量压缩到极小规模——在本教程配置下,可训练参数仅 27,009,024 个,占全部参数的比例只有0.2914%。这意味着我们可以在单卡(如 RTX 3090/24G 级别)上完成训练,且训练产物是一份很小的 LoRA 权重,可独立保存、分发、随时加载,不影响原始基座模型。
模型下载
使用 modelscope 的snapshot_download函数下载模型,第一个参数为模型名称,参数cache_dir为模型的下载路径。
在/root/autodl-tmp路径下新建model_download.py文件并输入以下内容,保存后运行python /root/autodl-tmp/model_download.py执行下载。模型大小约 18GB,下载大概需要 10 分钟。
from modelscope import snapshot_download model_dir = snapshot_download('LLM-Research/gemma-2-9b-it', cache_dir='/root/autodl-tmp')下载完成后,模型权重会保存在/root/autodl-tmp/LLM-Research/gemma-2-9b-it目录下,后续所有加载代码均引用该路径。模型的下载与基础部署方式(含环境准备细节)可参考 01-Gemma-2-9b-it FastApi 部署调用.md。
环境配置
在完成基础环境配置(如 AutoDL 上选择 PyTorch 2.1.0 / Python 3.10 / CUDA 12.1 镜像)和本地模型部署之后,还需要安装以下第三方库:
python -m pip install --upgrade pip # 更换 pypi 源加速库的安装 pip config set global.index-url https://pypi.tuna.tsinghua.edu.cn/simple pip install transformers==4.42.3 # 请务必安装 4.42.3 版本 pip install datasets peft注意:
transformers必须固定为4.42.3版本,该版本与 Gemma-2 的模型实现及本文的训练流程是匹配的;若版本不一致,可能出现AutoModelForCausalLM加载行为或Gemma2SdpaAttention等实现上的差异。此外,gradient_checkpointing与use_cache不兼容,开启梯度检查点后 Trainer 会自动把use_cache置为False。
本节微调使用的数据集放在仓库根目录的 dataset/huanhuan.json(共 3729 条样本),该数据集的构建与展示可以参考同仓库的 Chat-嬛嬛 示例。
指令集构建
LLM 的微调一般指指令微调(Instruction Tuning)过程。所谓指令微调,是指我们使用的微调数据形如:
{ "instruction":"回答以下用户问题,仅输出答案。", "input":"1+1等于几?", "output":"2" }其中,instruction是用户指令,告知模型其需要完成的任务;input是用户输入,是完成用户指令所必须的输入内容;output是模型应该给出的输出。
核心训练目标是让模型具有理解并遵循用户指令的能力。因此,在指令集构建时,应针对目标任务针对性构建任务指令集。例如本节目标是构建一个能够模拟甄嬛对话风格的个性化 LLM,因此构造的指令形如:
{ "instruction": "你是谁?", "input":"", "output":"家父是大理寺少卿甄远道。" }打开 dataset/huanhuan.json 可以看到,全部 3729 条样本都遵循这一instruction / input / output三字段结构,语料覆盖大量宫廷对话场景,例如:
{ "instruction": "你是谁?", "input": "", "output": "我是甄嬛,家父是大理寺少卿甄远道。" }在 Notebook 中,数据先通过 pandas 读取 JSON 再转换为 HuggingFaceDataset:
from datasets import Dataset import pandas as pd df = pd.read_json('huanhuan.json') ds = Dataset.from_pandas(df)从 Notebook 输出可以看到
ds[:3]的前三条样本均为甄嬛对话风格数据,其中input字段为空字符串,说明本数据集是单轮问答形态。
数据格式化:按 Gemma2 对话模板编码样本
LoRA 训练的数据需要经过格式化、编码之后再输入给模型。熟悉 PyTorch 训练流程的同学会知道,一般需要将输入文本编码为input_ids,将输出文本编码为labels,编码之后的结果都是多维向量。首先定义一个预处理函数,用于对每一个样本编码其输入、输出文本并返回编码后的字典:
def process_func(example): MAX_LENGTH = 384 # 分词器会将一个中文字切分为多个token,因此需要放开一些最大长度,保证数据的完整性 input_ids, attention_mask, labels = [], [], [] instruction = tokenizer(f"<bos><start_of_turn>user\n{example['instruction'] + example['input']}<end_of_turn>\n<start_of_turn>model\n", add_special_tokens=False) # add_special_tokens 不在开头加 special_tokens response = tokenizer(f"{example['output']}<end_of_turn>\n", add_special_tokens=False) input_ids = instruction["input_ids"] + response["input_ids"] + [tokenizer.pad_token_id] attention_mask = instruction["attention_mask"] + response["attention_mask"] + [1] # 因为eos token咱们也是要关注的所以 补充为1 labels = [-100] * len(instruction["input_ids"]) + response["input_ids"] + [tokenizer.pad_token_id] if len(input_ids) > MAX_LENGTH: # 做一个截断 input_ids = input_ids[:MAX_LENGTH] attention_mask = attention_mask[:MAX_LENGTH] labels = labels[:MAX_LENGTH] return { "input_ids": input_ids, "attention_mask": attention_mask, "labels": labels }这段函数包含几个关键设计点:
- 模板拼接:Gemma2 采用
<bos><start_of_turn>user\n...<end_of_turn>\n<start_of_turn>model\n...<end_of_turn>\n<eos>的对话模板。提问部分(instruction+input)放在user轮,答案(output)放在model轮。add_special_tokens=False确保分词器不会在开头额外插入特殊 token。 - labels 掩码:
instruction部分的标签全部置为-100(PyTorch 交叉熵损失会自动忽略该值),只有model轮的回答参与损失计算,从而让模型学会"接话"而非"复述用户问题"。 - 长度截断:
MAX_LENGTH = 384。中文字符经 BPE 分词后会被切分为多个 token,因此需要放宽最大长度保证数据完整性;超长样本截断时input_ids、attention_mask、labels三者在同一位置同步截断,保持对齐。
Gemma2 采用的完整 Prompt Template 格式如下:
<bos><start_of_turn>user 小姐,别的秀女都在求中选,唯有咱们小姐想被撂牌子,菩萨一定记得真真儿的——<end_of_turn> <start_of_turn>model 嘘——都说许愿说破是不灵的。<end_of_turn> <eos>编码完成后,对整个数据集应用map并移除原始列:
tokenized_id = ds.map(process_func, remove_columns=ds.column_names)从 Notebook 输出可以看到处理后数据集特征为['input_ids', 'attention_mask', 'labels'],共 3729 行。可以用tokenizer.decode验证第一条样本还原出的文本正是上述 Gemma2 模板格式;而过滤掉-100后解码labels,得到的是模型真正需要学习的回答部分。
加载 Tokenizer 与半精度模型
模型以半精度形式加载,如果显卡比较新,可以用torch.bfloat16形式加载。对于自定义的模型,一定要指定trust_remote_code参数为True。
tokenizer = AutoTokenizer.from_pretrained('/root/autodl-tmp/LLM-Research/gemma-2-9b-it') tokenizer.pad_token_id = tokenizer.eos_token_id tokenizer.padding_side = 'right' model = AutoModelForCausalLM.from_pretrained('/root/autodl-tmp/LLM-Research/gemma-2-9b-it', device_map="cuda", torch_dtype=torch.bfloat16,)两点说明:
- Gemma2 的分词器没有显式定义
pad_token,训练时需要把pad_token_id指向eos_token_id,否则DataCollatorForSeq2Seq做 padding 时会出错;padding_side='right'确保在序列右侧补齐,配合因果语言模型的从左到右注意力。 - 从 Notebook 打印的模型结构可以看到,
Gemma2ForCausalLM由embed_tokens(词表 256000 维、隐藏层 3584 维)、42 层Gemma2DecoderLayer以及lm_head组成;每个 DecoderLayer 内部是Gemma2SdpaAttention(q_proj、k_proj、v_proj、o_proj)与Gemma2MLP(gate_proj、up_proj、down_proj,激活函数为PytorchGELUTanh)的组合。这一结构直接决定了下一步LoraConfig中target_modules的选择范围。
开启梯度检查点后,还需要显式调用model.enable_input_require_grads():
model.enable_input_require_grads() # 开启梯度检查点时,要执行该方法定义 LoraConfig
LoraConfig类中可以设置很多参数,但主要参数不多,核心含义如下:
task_type:任务类型,本文为因果语言建模CAUSAL_LM。target_modules:需要插入 LoRA 适配器的模型层名字,主要是 attention 和 MLP 部分的全连接层。不同模型对应的层名不同,可以传入数组、字符串或正则表达式。结合上文打印的 Gemma2 结构,q_proj / k_proj / v_proj / o_proj对应自注意力,gate_proj / up_proj / down_proj对应前馈网络。r:LoRA 的秩(rank),决定低秩分解矩阵的维度。lora_alpha:LoRA 的缩放因子。lora_dropout:LoRA 分支的 Dropout 比例,用于抑制过拟合。
LoRA 的缩放不是r(秩),而是lora_alpha / r。在本配置中缩放为32 / 8 = 4倍:
from peft import LoraConfig, TaskType, get_peft_model config = LoraConfig( task_type=TaskType.CAUSAL_LM, target_modules=["q_proj", "k_proj", "v_proj", "o_proj", 'gate_proj', 'up_proj', 'down_proj'], inference_mode=False, # 训练模式 r=8, # Lora 秩 lora_alpha=32, # Lora alaph,具体作用参见 Lora 原理 lora_dropout=0.1# Dropout 比例 )通过get_peft_model把 LoRA 适配器挂载到模型上,并打印可训练参数量:
model = get_peft_model(model, config) model.print_trainable_parameters()Notebook 中的实际输出为:
trainable params: 27,009,024 || all params: 9,268,715,008 || trainable%: 0.2914即全部 92.7 亿参数中,只有 2700 万参数参与训练,占比不足 0.3%,这正是 LoRA 高效微调的直接体现——显存占用与训练时间都远小于全量微调。
自定义 TrainingArguments 参数
TrainingArguments的源码对每个参数都有详细说明,这里解释几个常用的:
output_dir:模型输出路径,checkpoint 会保存在该目录下。per_device_train_batch_size:单卡 batch_size。gradient_accumulation_steps:梯度累加步数。如果显存比较小,可以把batch_size调小、梯度累加调大,等效扩大训练 batch。logging_steps:每隔多少步输出一次 log。num_train_epochs:训练轮数。save_steps:每隔多少步保存一次 checkpoint。learning_rate:学习率。save_on_each_node:多节点训练时每个节点都保存权重。gradient_checkpointing:梯度检查点。开启后必须执行model.enable_input_require_grads(),原理是用计算换显存,前向过程中不保存全部中间激活,反向时重算。
args = TrainingArguments( output_dir="./output/gemma-2-9b-it", per_device_train_batch_size=1, gradient_accumulation_steps=4, logging_steps=10, num_train_epochs=3, save_steps=10, # 为了快速演示,这里设置10,建议你设置成100 learning_rate=1e-4, save_on_each_node=True, gradient_checkpointing=True )从 Notebook 的训练进度条可以看到,在 batch_size=1、梯度累加 4 步、3 个 epoch 的配置下,总训练步数为 2796 步。训练过程中 Trainer 会输出如下提示,属于正常现象:
It is strongly recommended to train Gemma2 models with the `eager` attention implementation instead of `sdpa`. `use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`.第一条提示说明 Gemma2 官方建议训练时使用eager注意力实现(若介意可改回AutoModelForCausalLM.from_pretrained(..., attn_implementation='eager'));第二条提示则印证了梯度检查点会自动关闭use_cache。
使用 Trainer 训练
组装Trainer并启动训练:
trainer = Trainer( model=model, args=args, train_dataset=tokenized_id, data_collator=DataCollatorForSeq2Seq(tokenizer=tokenizer, padding=True), ) trainer.train()DataCollatorForSeq2Seq会在批内把不同长度的样本 padding 到同一长度。训练日志截图(见文首配图一)展示了 Step 10 至 Step 190 的 Training Loss 变化,损失从 3.55 逐步下降到 2.5 附近并趋于平稳,说明模型在甄嬛对话语料上持续收敛。训练完成后,./output/gemma-2-9b-it/目录下会按save_steps间隔生成checkpoint-10、checkpoint-20……等 checkpoint 目录,每个 checkpoint 中保存的是独立的 LoRA 适配器权重(adapter_model.safetensors与adapter_config.json)。
加载 LoRA 权重推理
训练好之后,使用如下方式加载 LoRA 权重进行推理。这里以checkpoint-90为例,请按实际输出修改lora_path:
from transformers import AutoTokenizer, AutoModelForCausalLM import torch from peft import PeftModel mode_path = '/root/autodl-tmp/LLM-Research/gemma-2-9b-it' lora_path = './output/gemma-2-9b-it/checkpoint-90' # 这里改成你的 lora 输出对应 checkpoint 地址 # 加载tokenizer tokenizer = AutoTokenizer.from_pretrained(mode_path) # 加载模型 model = AutoModelForCausalLM.from_pretrained(mode_path, device_map="auto",torch_dtype=torch.bfloat16, trust_remote_code=True).eval() # 加载lora权重 model = PeftModel.from_pretrained(model, model_id=lora_path) # 调用模型进行对话生成 chat = [ { "role": "user", "content": '你好' }, ] prompt = tokenizer.apply_chat_template(chat, tokenize=False, add_generation_prompt=True) inputs = tokenizer.encode(prompt, add_special_tokens=False, return_tensors="pt") outputs = model.generate(input_ids=inputs.to(model.device), max_new_tokens=150) outputs = tokenizer.decode(outputs[0]) response = outputs.split('model')[-1].replace('<end_of_turn>\n<eos>', '') print(response)要点拆解:
- 独立加载:LoRA 权重与基座模型解耦,推理时先加载冻结的原始模型(
device_map="auto"自动分配设备),再用PeftModel.from_pretrained把适配器挂载回去,不需要重新训练。 - 模板一致:
apply_chat_template(..., add_generation_prompt=True)会把用户消息包装成 Gemma2 的标准对话模板并附加模型起始标记,与训练时使用的模板保持一致。 - 结果清洗:
generate输出的完整序列包含模板后缀,因此用split('model')[-1].replace('<end_of_turn>\n<eos>', '')截取并清理出纯回答文本。
从 Notebook 的实际推理结果(见文首配图二)可以看到,向微调后的模型发送"你好",模型以甄嬛口吻回复:
皇上好,我是甄嬛,家父是大理寺少卿甄远道。这验证了经过本流程训练的 LoRA 权重已成功让 Gemma-2-9b-it 习得目标角色的语言风格。微调完成后,如果需要把模型对外提供服务,可以继续参考本仓库 Gemma2 目录 下的 FastApi 部署、LangChain 接入与 WebDemo 教程,将基座模型替换为"基座 + LoRA"的组合即可复用相同的部署链路。
【免费下载链接】self-llm《开源大模型食用指南》针对中国宝宝量身打造的基于Linux环境快速微调(全参数/Lora)、部署国内外开源大模型(LLM)/多模态大模型(MLLM)教程项目地址: https://gitcode.com/GitHub_Trending/se/self-llm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考