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 = "",
    )

参考链接

视频

https://www.bilibili.com/video/BV1tMNMeMEiS?spm_id_from=333.788.videopod.sections&vd_source=67562982a62ef6e9b046e1add9275faf

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

Logo

欢迎加入 MCP 技术社区!与志同道合者携手前行,一同解锁 MCP 技术的无限可能!

更多推荐