Unsloth训练自己的R1推理模型 - DeepSeek GRPO
Unsloth是一款非常流行的高效大模型训练与微调工具。近期Unsloth也宣布支持GRPO。本期视频基于Unsloth官方博客的介绍,分享如何用Unsloth,利用GRPO,训练一款类似DeepSeek R1的具有自主思考推理能力的大模型。
一、如何使用开源工具OnSloth训练自己的ROne推理模型,包括模型的训练过程和使用GRPO训练器的特性。同时提供了一个Python notebook供学习参考。
00:17 - 介绍如何使用On Sloth来训练自己的R1推理模型
01:16 - On Sloth是一款开源工具,支持大多数NVIDIA GPU,可用于高效的大院模型训练与微调
02:41 - 提供Python notebook详细介绍如何基于LAMA3.18B模型进行R1推理模型的训练
二、作者在学习过程中对Python notebook的一些笔记,包括对一些陌生概念的标记和对on slash进行模型训练的扫盲,希望能对读者有所帮助。
03:00 - 分享Python notebook,对陌生概念做标记
03:53 - 安装on slash v LLM等包,初始化on slash核心组件
04:58 - 参数配置:max sequence length和lower rank,加载预训练模型
三、如何使用PFT和EFT来实现资源经济的高效微调,以及使用GRPO训练方法对模型进行训练,同时还介绍了数据集和奖励函数的定义。
06:03 - 配置PFTPEFT,实现资源经济的高效的模型微调
07:08 - GRPO帮助生成合理思考链,训练出具有推理能力的大模型
08:31 - 使用warm operation和LR schedule type进行预热和学习率调度
四、使用ONSLASH提供的GRPO进行语音合成的训练和推理,包括参数设置、训练过程和推理测试等。同时,也提到了训练需要耗时和耐心。
09:00 - 参数设置包括训练批次大小、梯度积累步数、输入提示词长度等
10:16 - ONSLASH提供的GRPO训练需要至少12个小时,奖励逐渐增加
11:55 - 使用fast generate函数完成模型推理,采样参数为温度0.8,最大token长度为1024
五、如何保存和加载经过训练的神经网络模型,包括模型权重和参数。此外,我们还介绍了如何将模型转换为不同的格式,以便在不同场景下复用。
要运行此程序,请在免费的Tesla T4 Google Colab实例上按“运行”并按“全部运行”!
选择T4类型的运行环境

环境准备
%%capture
# Skip restarting message in Colab
import sys; modules = list(sys.modules.keys())
for x in modules: sys.modules.pop(x) if "PIL" in x or "google" in x else None
!pip install unsloth vllm
!pip install --upgrade pillow
Unsloth
在开始训练之前,我们需要导入Unsloth的核心组件并进行初始化。首先使用PatchFastRL来追加GRPO和其他RL算法:(对GRPO和其他RL算法打补丁)
from unsloth import FastLanguageModel, PatchFastRL
PatchFastRL("GRPO", FastLanguageModel)

记下来加载Llama3.1 8B Instruct并设置训练参数。
导入和初始配置
- max_seq_length决定了模型能处理的最大文本长度
- lora_rank是LoRA(低秩适应)的参数,值越大模型越"智能"但训练速度越慢
from unsloth import is_bfloat16_supported
import torch
max_seq_length = 512 # Can increase for longer reasoning traces
lora_rank = 32 # Larger rank = smarter, but slower
加载预训练模型
加载Meta的Llama-3.1-8B-Instruct(最大15B)模型使用4位量化来减少内存占用,启用vLLM快速推理功能,控制GPU内存使用率为60%。
model, tokenizer = FastLanguageModel.from_pretrained(
model_name = "meta-llama/meta-Llama-3.1-8B-Instruct",
max_seq_length = max_seq_length,
load_in_4bit = True, # False for LoRA 16bit
fast_inference = True, # Enable vLLM fast inference
max_lora_rank = lora_rank,
gpu_memory_utilization = 0.6, # Reduce if out of memory
)
- load_in_4bit = True:启用4位量化加速推理能力
- gpu_memory_utilization:指定GPU使用控制在60%
运行



加载完成后获得两个变量,一个是tokenizer分词器,一个是模型变量
配置PEFT(参数高效微调)模型
使用LoRA方法对模型进行参数高效微调,指定需要微调的目标模块。启用梯度检查点功能以支持长文本微调,并设置随机种子确保结果可复现。
model = FastLanguageModel.get_peft_model(
model,
r = lora_rank, # Choose any number > 0 ! Suggested 8, 16, 32, 64, 128
target_modules = [
"q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj",
], # Remove QKVO if out of memory
lora_alpha = lora_rank,
use_gradient_checkpointing = "unsloth", # Enable long context finetuning
random_state = 3407,
)
准备数据集
我们直接利用@willccbb进行数据准备脚本以及奖励函数(reward functions)。你可以自由地创建自己的!
本示例用到了OpenAI的GSM8K数据集

import re
from datasets import load_dataset, Dataset
# Load and prep dataset
SYSTEM_PROMPT = """
Respond in the following format:
<reasoning>
...
</reasoning>
<answer>
...
</answer>
"""
XML_COT_FORMAT = """\
<reasoning>
{reasoning}
</reasoning>
<answer>
{answer}
</answer>
"""
def extract_xml_answer(text: str) -> str:
answer = text.split("<answer>")[-1]
answer = answer.split("</answer>")[0]
return answer.strip()
def extract_hash_answer(text: str) -> str | None:
if "####" not in text:
return None
return text.split("####")[1].strip()
# uncomment middle messages for 1-shot prompting
def get_gsm8k_questions(split = "train") -> Dataset:
data = load_dataset('openai/gsm8k', 'main')[split] # type: ignore
data = data.map(lambda x: { # type: ignore
'prompt': [
{'role': 'system', 'content': SYSTEM_PROMPT},
{'role': 'user', 'content': x['question']}
],
'answer': extract_hash_answer(x['answer'])
}) # type: ignore
return data # type: ignore
dataset = get_gsm8k_questions()
# Reward functions
def correctness_reward_func(prompts, completions, answer, **kwargs) -> list[float]:
responses = [completion[0]['content'] for completion in completions]
q = prompts[0][-1]['content']
extracted_responses = [extract_xml_answer(r) for r in responses]
print('-'*20, f"Question:\n{q}", f"\nAnswer:\n{answer[0]}", f"\nResponse:\n{responses[0]}", f"\nExtracted:\n{extracted_responses[0]}")
return [2.0 if r == a else 0.0 for r, a in zip(extracted_responses, answer)]
def int_reward_func(completions, **kwargs) -> list[float]:
responses = [completion[0]['content'] for completion in completions]
extracted_responses = [extract_xml_answer(r) for r in responses]
return [0.5 if r.isdigit() else 0.0 for r in extracted_responses]
def strict_format_reward_func(completions, **kwargs) -> list[float]:
"""Reward function that checks if the completion has a specific format."""
pattern = r"^<reasoning>\n.*?\n</reasoning>\n<answer>\n.*?\n</answer>\n$"
responses = [completion[0]["content"] for completion in completions]
matches = [re.match(pattern, r) for r in responses]
return [0.5 if match else 0.0 for match in matches]
def soft_format_reward_func(completions, **kwargs) -> list[float]:
"""Reward function that checks if the completion has a specific format."""
pattern = r"<reasoning>.*?</reasoning>\s*<answer>.*?</answer>"
responses = [completion[0]["content"] for completion in completions]
matches = [re.match(pattern, r) for r in responses]
return [0.5 if match else 0.0 for match in matches]
def count_xml(text) -> float:
count = 0.0
if text.count("<reasoning>\n") == 1:
count += 0.125
if text.count("\n</reasoning>\n") == 1:
count += 0.125
if text.count("\n<answer>\n") == 1:
count += 0.125
count -= len(text.split("\n</answer>\n")[-1])*0.001
if text.count("\n</answer>") == 1:
count += 0.125
count -= (len(text.split("\n</answer>")[-1]) - 1)*0.001
return count
def xmlcount_reward_func(completions, **kwargs) -> list[float]:
contents = [completion[0]["content"] for completion in completions]
return [count_xml(c) for c in contents]
- correctness_reward_func 正确性奖励,评估模型输出答案是否正确
- soft_format_reward_func /strict_format_reward_func 都是格式奖励,对xml松散的以及严格的验证
- count_xml 检查xml标签的完整性和正确性,从而对于正确的标签给予积分或是奖励


训练模型
使用GRPO(GenerativeReinforcementPolicyOptimization)训练方法
该配置主要特点是:
- 采用了内存效率较高的设置(8位优化器、混合精度训练)
- 使用较小的学习率和梯度裁剪以确保训练稳定性
- 提供了灵活的设置选项以适应不同的硬件条件、
- 包含了完整的训练监控和模型保存机制
training_args = GRPOConfig(
use_vllm = True, # use vLLM for fast inference!
learning_rate = 5e-6,
adam_beta1 = 0.9,
adam_beta2 = 0.99,
weight_decay = 0.1,
warmup_ratio = 0.1, # 使用前10%的步骤进行预热
lr_scheduler_type = "cosine", # 采用余弦衰减的学习率调度策略
optim = "paged_adamw_8bit", # 使用8位量化的AdamW优化器节省内存
logging_steps = 1,
bf16 = is_bfloat16_supported(), # bfloat16精度(如果支持)
fp16 = not is_bfloat16_supported(), # 否则使用fp16
per_device_train_batch_size = 1, # 设置了每个设备的训练批次大小
gradient_accumulation_steps = 1, # Increase to 4 for smoother training 设置梯度积累步数
num_generations = 6, # Decrease if out of memory 限制生成数量以控制内存使用
max_prompt_length = 256, # 输入提示最大长度
max_completion_length = 200, # 生成文本最大长度
# num_train_epochs = 1, # Set to 1 for a full training run
max_steps = 250, # 最大训练步数
save_steps = 250, # 模型保存间隔
max_grad_norm = 0.1, # 日志记录间隔
report_to = "none", # Can use Weights & Biases
output_dir = "outputs",
)
GRPO训练器
trainer = GRPOTrainer(
model = model,
processing_class = tokenizer,
reward_funcs = [
xmlcount_reward_func,
soft_format_reward_func,
strict_format_reward_func,
int_reward_func,
correctness_reward_func,
],
args = training_args,
train_dataset = dataset,
)
trainer.train()
奖励函数
- xmlcount_reward_func:可能用于评估生成文本中xml标签的正确使用
- soft_format_reward_func:评估输出格式的基本符合程度
- strict_format_reward_func:严格检查输出格式是否完全符合要求
- int_reward_func:可能用于检查数值输出的准确性
- correctness_reward_func:评估整体输出的正确性
奖励函数(Reward Function)在RL中扮演者关键角色,其作用如下:
1.基本作用
- 评估行为质量:为模型的每个输出提供一个数值评分
- 指导学习方向:帮助模型理解什么样的输出是“好的”
- 提供反馈:让模型知道它的输出是否符合预期
2.代码中的具体奖励函数分析,每个奖励函数的具体作用:
-
xmlcount_reward_func
-
检查XML标签的使用是否正确
-
确保开闭标签匹配
-
评估标签嵌套的合理性
-
soft_format_reward_func
-
对输出格式进行基本检查
-
允许有一定的灵活性
-
可能关注缩进、换行等基本格式要素
-
strict_format_reward_func
-
严格检查输出格式
-
要求完全符合预定义的格式规范
-
对不符合要求的部分给予惩罚



至少需要训练1h
推理
text = tokenizer.apply_chat_template([
{"role" : "user", "content" : "Calculate pi."},
], tokenize = False, add_generation_prompt = True)
from vllm import SamplingParams
sampling_params = SamplingParams(
temperature = 0.8,
top_p = 0.95,
max_tokens = 1024,
)
output = model.fast_generate(
[text],
sampling_params = sampling_params,
lora_request = None,
)[0].outputs[0].text
output
采样参数(Sampling Parameters):
- temperature:值越高生成越随机,越低越确定性
- top_p:控制采样时考虑的概率质量范围
- max_tokens:限制生成文本的最大长度
运行:
保存经过训练的LoRA(Low-RankAdaptation)权重数据
- 保存内容
保存了模型在训练过程中学到的LoRA参数,这些参数包含了低秩矩阵的权重更新,不会保存整个基础模型,只保存LoRA相关的改变
- 保存位置
将权重保存到名为"grpo_saved_lora"的目录中。这个目录会包含所有必要的配置文件和权重
- 用途
保存的权重可以在之后被重新加载,可以应用到相同架构的基础模型上,允许在不同的场景中复用训练成果。
-
优势。
-
文件体积小:只保存LoRA参数,而不是完整模型。
-
便于分享:可以轻松分发训练成果。
-
灵活性高:可以与不同的基础模型组合使用
model.save_lora("grpo_saved_lora")
加载LoRA权重并推理
- 加载之前保存的LoRA权重(通过"grpo_saved_lora")
- 使用这些权重和配置的参数生成回答
- 返回生成的文本结果
text = tokenizer.apply_chat_template([
{"role" : "system", "content" : SYSTEM_PROMPT},
{"role" : "user", "content" : "Calculate pi."},
], tokenize = False, add_generation_prompt = True)
from vllm import SamplingParams
sampling_params = SamplingParams(
temperature = 0.8,
top_p = 0.95,
max_tokens = 1024,
)
output = model.fast_generate(
text,
sampling_params = sampling_params,
lora_request = model.load_lora("grpo_saved_lora"),
)[0].outputs[0].text
output
用户可以根据自己的具体需求选择最合适的保存方式:
- 需要高精度:选择16位格式
- 设备受限:选择4位格式
- 只需分享改动:选择LoRA格式
利用Hugging Face的Token上传
# Merge to 16bit
if False: model.save_pretrained_merged("model", tokenizer, save_method = "merged_16bit",)
if False: model.push_to_hub_merged("hf/model", tokenizer, save_method = "merged_16bit", token = "")
# Merge to 4bit
if False: model.save_pretrained_merged("model", tokenizer, save_method = "merged_4bit",)
if False: model.push_to_hub_merged("hf/model", tokenizer, save_method = "merged_4bit", token = "")
# Just LoRA adapters
if False: model.save_pretrained_merged("model", tokenizer, save_method = "lora",)
if False: model.push_to_hub_merged("hf/model", tokenizer, save_method = "lora", token = "")
GGUF/llama.cpp转换
Unsloth原生支持转换为GGUF /Llama.cpp格式!克隆llama.cpp并默认使用q8_0格式保存。支持所有量化方法如q4_k_m。使用save_pretrained_gguf进行本地保存,使用push_to_hub_gguf上传到HuggingFace。
支持的量化方法包括(完整列表请查看Unsloth的Wiki页面):
https://github.com/unslothai/unsloth/wiki#gguf-quantization-options
- 8_0-快速转换。资源占用较高,但通常可以接受
- q4_k_m-推荐使用。对attention.wv和feed_forward.w2tensors的一半使用Q6_K量化,其余使用Q4_K
- q5_k_m-推荐使用。对attention.wv和feed_forward.w2tensors的一半使用Q6_K量化,其余使用Q5_K
**[新功能]**要进行微调并自动导出到ollama,参考Ollama notebook
https://colab.research.google.com/drive/1WZDi7APtQ9VsvOrQSSC5DDtxq159j8iZ?usp=sharing
主要特点:
1.提供了多种量化选项,可以根据需求平衡模型大小和性能
2.支持直接转换为llama.cpp格式,便于部署
3.提供了与Ollama的集成,简化了部署流程
# Save to 8bit Q8_0
if False: model.save_pretrained_gguf("model", tokenizer,)
# Remember to go to https://huggingface.co/settings/tokens for a token!
# And change hf to your username!
if False: model.push_to_hub_gguf("hf/model", tokenizer, token = "")
# Save to 16bit GGUF
if False: model.save_pretrained_gguf("model", tokenizer, quantization_method = "f16")
if False: model.push_to_hub_gguf("hf/model", tokenizer, quantization_method = "f16", token = "")
# Save to q4_k_m GGUF
if False: model.save_pretrained_gguf("model", tokenizer, quantization_method = "q4_k_m")
if False: model.push_to_hub_gguf("hf/model", tokenizer, quantization_method = "q4_k_m", token = "")
# Save to multiple GGUF options - much faster if you want multiple!
if False:
model.push_to_hub_gguf(
"hf/model", # Change hf to your username!
tokenizer,
quantization_method = ["q4_k_m", "q8_0", "q5_k_m",],
token = "",
)
参考链接
视频
R1 Reasoning | Unsloth Blog
https://unsloth.ai/blog/r1-reasoning
Unsloth GRPO notebook: Llama 3.1 (8B) on Colab https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Llama3.1_(8B)-GRPO.ipynb
OpenAI Gsm8K数据集
https://huggingface.co/datasets/openai/gsm8k
更多推荐

所有评论(0)