使用SFT和QLoRA为AI Agent定制大语言模型

在本教程中,我将向你展示如何使用 QLoRA 对用于 AI Agent 的大语言模型进行监督微调。这样,我们就能定制预训练模型,让它按我们想要的方式表现。我们会使用一套轻量级训练流程,只更新模型的一小部分。
我们将使用 Unsloth 和 Hugging Face 生态系统来下载 Qwen 1.5B 基础模型,应用基于 QLoRA 的监督微调,并将训练得到的 LoRA 适配器权重保存在本地供推理使用。整个过程都在本地运行,因此不会产生模型 API 费用。
背景
训练语言模型,就是向它展示大量示例,并更新其内部权重(也就是参数),让它能更好地预测期望的输出。现代 LLM 可能有数百万甚至数十亿个参数,这也是训练它们成本高昂的原因之一。模型的参数越多,训练时通常需要的内存和算力也越多。
像 Claude 和 ChatGPT 这样的基础大语言模型,也被训练得比较通用。这意味着它们的回答可能显得宽泛、不够稳定,或者与特定应用场景不够匹配。即便提示词能起到一定帮助,有些时候你仍然希望模型直接从示例中学习更一致的模式。
这正是微调的用武之地。微调是将预训练模型适配到更贴近目标任务行为的通用过程。其中一种常见形式是监督微调:在带有标签的输入/输出示例上训练模型,这些示例展示了你期望的行为。
本教程适用于 macOS、Windows 和 Linux。我使用的是配备 32 GB 内存、没有外接 GPU 的 MacBook Pro;如果你的硬件配置更有限,换用更小的预训练模型也能运行这套工作流。
什么是监督微调?
监督微调(SFT),就是在示例输入/输出对上对预训练模型做进一步训练。它不是从零开始训练模型,而是从一个已经对语言有较好理解的模型出发,教它用更贴合你任务的方式做出回应。例如,你可能希望它用某种语气回答、遵循特定的格式,或者在某个较窄的任务上表现得更一致。SFT 通过向模型展示大量你期望的行为示例,推动它朝这个方向改进。
所需数据量取决于任务:像语气或格式这类简单调整,几百个高质量示例就可能奏效;而更复杂的行为或领域适配,通常需要更多精心整理的数据。
为了保持训练快速、简便,本教程只用了五个示例;不过同样的代码完全可以在实际生产流程中扩展到更大的数据集。
什么是 LoRA?
全量微调的成本可能很高,因为大语言模型拥有海量参数,更新全部参数需要大量的 GPU 内存、计算时间和存储空间。
LoRA(Low-Rank Adaptation,低秩适配)是一种更轻量的模型微调方式,也是最常见的参数高效微调(PEFT)方法之一。它不需要更新模型的所有原始权重:基础模型基本保持冻结,LoRA 只是在其之上添加一组小得多的可训练适配器权重。
本教程使用 QLoRA,它把量化与 LoRA 结合起来:以低精度(通常是 4-bit)加载基础模型,然后只训练这些 LoRA 适配器。这进一步减少了内存占用,让微调在有限的硬件上更加实用。
我们还会用到名为 Unsloth 的开源库,它旨在让大语言模型微调更快、更省内存。它从 Hugging Face 下载模型权重、tokenizer 和配置,常用于 LoRA 监督微调等工作流,尤其是在硬件资源有限的情况下。
动机与架构
构建 AI Agent 后,你可能会发现,现成的模型需要很长的提示词、重复的指令和额外的上下文,才能为你的使用场景生成想要的输出。这会增加 Token 用量、延迟和成本,同时结果仍然可能不一致。在这种情况下,自然的下一步就是训练模型,让它以更贴合你任务的方式做出回应。
架构是:加载量化后的基础模型,格式化带标签的聊天示例,添加 LoRA 适配器,只对这些适配器做监督微调,并保存得到的适配器权重,以便之后在 AI Agent 中加载到基础模型之上进行推理。下面各节会对代码进行详细解释。
第 1 步:安装 Python 依赖
创建虚拟环境并安装所需的包:
python3 -m venv venv
source venv/bin/activate
pip install unsloth datasets transformers trl torch peft accelerate bitsandbytes
第 2 步:训练代码
我们选用 Qwen2.5-1.5B-Instruct,因为它是一个较小的指令微调模型,适合在有限硬件上进行 QLoRA 微调。
整体来看,代码以 4-bit 模式加载预训练的 unsloth/Qwen2.5-1.5B-Instruct 模型——这正是它成为 QLoRA 风格工作流、并保持较低 GPU 内存占用的原因。
接着,代码使用一个小型监督微调数据集(包含五条聊天风格的客服对话示例),按 Qwen 的聊天模板完成格式化,并把 LoRA 适配器附加到选定的注意力和 MLP 层上。
训练期间,SFTTrainer 只更新这些轻量 LoRA 适配器权重,而不更新整个基础模型;同时配合梯度累积、8-bit AdamW、Unsloth 梯度检查点等省内存的设置。
请将下面的代码保存为 train.py。我在代码里加了注释,帮助理解整个工作流程:
from unsloth import FastLanguageModel
from datasets import Dataset
from transformers import TrainingArguments
from trl import SFTConfig, SFTTrainer
max_seq_length = 2048
dataset = Dataset.from_list([
{
"messages": [
{"role": "system",
"content": "You are a helpful ecommerce customer support assistant. Be polite, concise, and do not invent order details."},
{"role": "user", "content": "Where is my order?"},
{"role": "assistant", "content": "I can help with that. Please share your order number so I can check the latest shipping status."}
]
},
{
"messages": [
{"role": "system",
"content": "You are a helpful ecommerce customer support assistant. Be polite, concise, and do not invent order details."},
{"role": "user", "content": "My package arrived damaged. What should I do?"},
{"role": "assistant", "content": "I’m sorry about that. Please upload a photo of the damaged item and packaging, and I’ll guide you through a replacement or refund request."}
]
},
{
"messages": [
{"role": "system",
"content": "You are a helpful ecommerce customer support assistant. Be polite, concise, and do not invent order details."},
{"role": "user", "content": "Can I return an item after 45 days?"},
{"role": "assistant", "content": "Our standard return window is 30 days. If your item is defective or there was a shipping issue, please share the details and I can help review possible exceptions."}
]
},
{
"messages": [
{"role": "system",
"content": "You are a helpful ecommerce customer support assistant. Be polite, concise, and do not invent order details."},
{"role": "user", "content": "I entered the wrong shipping address."},
{"role": "assistant", "content": "If the order has not shipped yet, we may be able to update the address. Please send the order number and the correct shipping address as soon as possible."}
]
},
{
"messages": [
{"role": "system",
"content": "You are a helpful ecommerce customer support assistant. Be polite, concise, and do not invent order details."},
{"role": "user", "content": "I was charged twice for my order."},
{"role": "assistant", "content": "I’m sorry for the trouble. Please share your order number and the dates of the charges so we can review the duplicate payment and help resolve it quickly."}
]
},
])
## 加载预训练的 Qwen 指令模型。
## Unsloth 会自动配置高效加载。
model, tokenizer = FastLanguageModel.from_pretrained(
model_name="unsloth/Qwen2.5-1.5B-Instruct", # 要从 Hugging Face / Unsloth 加载的预训练模型
max_seq_length=max_seq_length, # 模型应准备处理的最大序列长度
# 更长的上下文 = 更多内存占用
load_in_4bit=True, # 以 4-bit 量化形式加载模型权重
# 大幅减少训练/推理时的 VRAM 占用
# 常用于 LoRA / QLoRA 工作流
dtype=None, # 让 Unsloth / Torch 自动选择数值精度
# 通常会选择 float16/bfloat16 等合适的精度
)
def format_example(example):
text = tokenizer.apply_chat_template(
example["messages"], # 从 "messages" 字段读取对话
tokenize=False, # 返回格式化后的字符串,而不是 token ID
add_generation_prompt=False, # 不附加空的 assistant 提示
# 因为该示例已包含 assistant 的回答
)
return {"text": text} # 返回一个包含格式化聊天文本的新数据集字段
formatted_dataset = dataset.map(format_example)
## 不是训练数十亿个参数,
## LoRA 在注意力层中插入小的可训练矩阵。
model = FastLanguageModel.get_peft_model(
model, # 基础预训练模型;LoRA 适配器将附加在此处
r=16, # LoRA rank:
# 低秩适配器矩阵的大小
# 越高 = 容量越大 + 可训练参数越多
# 越低 = 更轻/更快但表达能力较弱
target_modules=[
"q_proj", # 注意力中的 Query 投影
"k_proj", # 注意力中的 Key 投影
"v_proj", # 注意力中的 Value 投影
"o_proj", # 注意力中的 Output 投影
"gate_proj", # MLP 块中的门控投影
"up_proj", # MLP 块中的 Up 投影
"down_proj", # MLP 块中的 Down 投影
], # LoRA 适配器仅插入到这些层中
lora_alpha=16, # LoRA 缩放因子
# 控制适配器更新对基础权重的影响强度
# 通常设置为与 r 相等
lora_dropout=0, # 训练期间 LoRA 路径上的 dropout
# 在 Unsloth 示例中通常为 0
bias="none", # 不训练偏置参数
# 只有 LoRA 适配器权重可训练
use_gradient_checkpointing="unsloth", # 使用 Unsloth 的省内存检查点
# 通过在反向传播时重新计算激活来降低 VRAM 占用
max_seq_length=max_seq_length, # 训练期间预期的最大 Token 序列长度
)
trainer = SFTTrainer(
model=model, # 要微调的模型(基础模型 + LoRA 适配器)
tokenizer=tokenizer, # 将文本转换为模型能理解的 token ID
train_dataset=formatted_dataset, # 你的训练数据
dataset_text_field="text", # 数据集中包含训练文本的列
max_seq_length=max_seq_length, # 每个示例的最大 Token 数
args=SFTConfig(
output_dir="../outputs", # 检查点/日志/结果将保存到的文件夹
per_device_train_batch_size=2, # 每个 GPU 上同时处理的示例数量
gradient_accumulation_steps=4, # 在更新权重之前累积 4 个 mini-batch 的梯度
# 有效 batch size ≈ 1 个 GPU 上 2 * 4 = 8
max_steps=30, # 在 10 次优化器更新步骤后停止训练
logging_steps=1, # 每 1 步打印/记录训练指标
warmup_steps=5, # 在前 5 步中逐渐增加学习率
learning_rate=2e-4, # 训练的主要学习率
optim="adamw_8bit", # 内存高效的 AdamW 优化器(适合低 VRAM 环境)
weight_decay=0.01, # 少量正则化以防止过拟合
lr_scheduler_type="linear", # 预热后,学习率随时间线性降低
seed=3407, # 随机种子,使训练更具可复现性
report_to="none", # 禁用 WandB 等外部日志工具
),
)
trainer.train()
## 只保存 LoRA 适配器权重,而不是完整的基础模型。
model.save_pretrained("qwen2_0_5b_lora")
## 保存 tokenizer,以便推理时使用相同的词表。
tokenizer.save_pretrained("qwen2_0_5b_lora")
第 3 步:推理代码
整体来看,推理代码包含一个 generate_reply() 函数:先用 Unsloth 加载模型(既可以使用基础模型名称,也可以加载本地保存的 LoRA 适配器目录),再启用推理优化,把聊天消息格式化为 Qwen 期望的提示结构,对提示进行 tokenize,将其移动到可用的设备上,最后通过 model.generate() 生成回复。
请将下面的代码保存为 inference.py:
from unsloth import FastLanguageModel
import torch
messages = [
{
"role": "system",
"content": "You are a helpful ecommerce customer support assistant. Be polite, concise, and do not invent order details."
},
{
"role": "user",
"content": "I want to cancel my order."
}
]
def generate_reply(model_name, messages):
# 加载基础模型并自动附加保存的 LoRA 适配器。
# "qwen2_0_5b_lora" 是 model.save_pretrained() 创建的目录。
model, tokenizer = FastLanguageModel.from_pretrained(
model_name=model_name, # 微调后的 LoRA 模型/适配器的路径或模型名称
max_seq_length=2048, # 模型应支持的最大上下文长度。更长的上下文使用更多内存
load_in_4bit=True, # 以 4-bit 量化形式加载权重。推理期间减少 VRAM 占用
)
# 启用推理优化(更快的生成、更低的内存占用)。
FastLanguageModel.for_inference(model)
# 将聊天消息转换为 Qwen 期望的格式。
inputs = tokenizer.apply_chat_template(
messages, # 聊天消息列表:system / user / assistant 轮次
tokenize=True, # 将格式化后的聊天提示转换为 token ID
add_generation_prompt=True, # 添加 assistant 提示,使模型知道要生成回复
return_tensors="pt", # 返回 PyTorch 张量
)
# 将输入张量移动到与模型相同的设备上
device = "cuda" if torch.cuda.is_available() else "cpu"
inputs = inputs.to(device)
# 生成 assistant 的回复。
outputs = model.generate(
input_ids=inputs, # 传入模型的 tokenized 提示
max_new_tokens=80, # 在回复中最多生成 80 个新 token
temperature=0.2, # 低 temperature = 更确定性 / 更专注的输出
# 高 temperature = 更随机 / 更有创意的输出
)
# 移除提示部分,只保留新生成的回复。
generated_tokens = outputs[0][inputs.shape[-1]:]
# 将 token ID 转换回可读文本。
response = tokenizer.decode(generated_tokens, skip_special_tokens=True)
return response
before = generate_reply("unsloth/Qwen2-0.5B-Instruct-bnb-4bit", messages)
after = generate_reply("./qwen2_0_5b_lora", messages)
print("=== BEFORE SFT ===")
print(before)
print()
print("=== AFTER SFT ===")
print(after)
示例输出
训练过程的输出如下:
$ python train.py
...
Unsloth: LoRA applied — 18,464,768 trainable params (4.04% of 456,701,440 total)
...
Unsloth: Training for 30 steps, BS=2, grad_accum=4, seq_len=2048
Unsloth: Features: CCE, GC, LR=linear, opt=adamw
Step 1/30 | Loss: 3.9350 | Grad: 4.8440 | LR: 0.00e+00 | Tok/s: 352 | Peak: 2.35 GB
Step 2/30 | Loss: 4.0082 | Grad: 4.9456 | LR: 4.00e-05 | Tok/s: 388 | Peak: 2.50 GB
...
Step 30/30 | Loss: 0.0646 | Grad: 0.6915 | LR: 8.00e-06 | Tok/s: 379 | Peak: 2.57 GB
Unsloth: Training complete! Avg loss: 1.2078 | Total time: 35.6s | Steps: 30 | Tokens: 14480
Unsloth: LoRA adapters saved to outputs
Unsloth: Saved final adapters to outputs
从输出可以看到,LoRA 已成功应用,而且只训练了约 4% 的模型参数,这让微调过程保持轻量。
在 30 步训练中,Unsloth 记录了损失、学习率、每秒 Token 数、峰值内存占用等实用指标。损失从约 3.9 下降到 0.06,说明模型正在从小数据集中学习;整个训练过程约 35 秒完成,峰值内存仅约 2.6 GB。
训练结束后,Unsloth 把训练好的 LoRA 适配器权重保存到 outputs 目录,供后续推理使用。你会看到多出一个 qwen2_0_5b_lora 目录,里面包含以下文件:
adapter_config.json
adapters.safetensors
chat_template.jinja
tokenizer.json
tokenizer_config.json
保存的 LoRA 输出包含几个体积不大但很重要的文件。adapter_config.json 存储 LoRA 适配器设置,例如适配了哪些层以及适配器的配置方式。adapters.safetensors 包含实际训练好的 LoRA 权重。chat_template.jinja 定义了 system、user 和 assistant 消息如何格式化为模型期望的提示结构。tokenizer.json 存储 tokenizer 的词表和 tokenize 规则,而 tokenizer_config.json 存储 tokenizer 设置,如特殊 token 和格式化行为。
这些文件合在一起,足以重新加载微调后的适配器,并使用与训练时相同的提示和 tokenize 设置进行推理。
现在运行推理脚本,测试一下微调后的模型。脚本会加载 qwen2_0_5b_lora 目录中保存的 LoRA 适配器权重,并在推理时把它们应用到基础模型之上。
用户提示是"I want to cancel my order"。下面的输出展示了同一个模型在 SFT 前后的回答。
$ python inference.py
=== BEFORE SFT ===
Sure, I'd be happy to help you with that! To cancel your order, please follow these steps:
1. **Review Your Order**: Before you can cancel your order, you need to review it. Make sure you have all the necessary information, such as your order number, order date, and any other relevant details.
2. **Review Your Order Details**: Once you've reviewed your order
=== AFTER SFT ===
If your order has not yet shipped, we may be able to cancel it for you. Please share your order number and the reason for cancellation so I can help process the refund or credit.
SFT 之前,基础模型给出的回答泛泛而谈、略显冗长,听起来很乐于助人,但没有遵循清晰的电商支持流程。
SFT 之后,模型的回答更加简洁、更具可操作性:它会主动询问订单号,并把取消操作与发货状态联系起来。这表明即使使用相对较小的领域特定数据集,SFT 也能改善角色对齐和回复风格。
微调、提示工程与蒸馏
提示工程、微调和蒸馏都会塑造模型行为,但方式各不相同。
提示工程作用于推理阶段,通过改变给模型的指令来影响输出。它通常是最快、成本最低的起点。
微调更进一步:让模型在示例上训练,从而更稳定地学到你想要的模式。
蒸馏则用于让较小的模型去模仿更强模型的行为。
实践中,提示工程通常是第一步;当需要更强的任务对齐时,微调是主要的下一步;而当效率成为更重要的目标时,蒸馏就显得重要了。
结论
在本教程中,我们用 QLoRA 对预训练语言模型做了监督微调。我们不是从头训练模型,而是以一个通用指令模型为起点,在少量示例对话上训练它,并且只更新轻量级的 LoRA 适配器权重。这使得工作流在有限的硬件上更加实用,同时仍能让模型适应特定的客服使用场景。
接下来,你可以尝试更大的数据集、不同的提示/回复风格,或更大的基础模型,观察行为会有什么变化。祝你折腾愉快!
- 原文链接: freecodecamp.org/news/ho...
- 鸿途知科网 AI 助手,为大家转译优秀英文文章,如有翻译不通的地方,还请包涵~
版权声明
本文仅代表作者观点,不代表区块链技术网立场。
本文系作者授权本站发表,未经许可,不得转载。
鸿途知科网
发表评论:
◎欢迎参与讨论,请在这里发表您的看法、交流您的观点。