目录
一、准备工作
二、GRPO强化学习
总结:
在 模型蒸馏介绍 里了解了模型蒸馏的过程就是用 SFT/GRPO的方法让学生模型能学习到教师模型的能力,从而使学生模型的能力得到大的提升,并且在LLM微调-训练垂类问答模型 里面学习了SFT模型微调。
SFT监督学习,需要给到固定格式的数据,让大模型快速的学习到基础知识。GRPO强化学习,则是给出问题+标准答案,通过奖励函数的引导让大模型自己去推理,从而提升大模型的推理能力。
一、准备工作
1)数据准备
GSM8K(Grade School Math 8K)是一个高质量的小学数学应用题数据集,主要用于评估和训练人工智能模型 在数学推理和多步问题解决方面的能力 https://huggingface.co/datasets/openai/gsm8k
2)环境准备
由于我这次选的模型是Qwen2.5-7B,根据LLM微调-工作准备中提到的显存估算方法,本机跑不了这个模型(按三倍估算,需要21G),需要租用服务器。AutoDL里面选一个RTX 4090并开机。
复制SSH,点这个小加号,把复制的内容填到弹出的框中,就会出现一行内容(我马赛克的地方)
右键选这两个都行吧,会再弹出一个框让输密码,复制SSH登录里面的密码填入进去。
二、GRPO强化学习
1)加载模型和配置Lora,这两步和之前学习过的步骤一样,不再多讲
# ======================================== # Step 1: 模型加载(启用vLLM快速推理) # ======================================== import unsloth from unsloth import FastLanguageModel import torch max_seq_length = 1024 # 可以增加以获得更长的推理轨迹 lora_rank = 32 # 更大的rank让模型更智能,但训练更慢 model, tokenizer = FastLanguageModel.from_pretrained( model_name="/root/autodl-tmp/models/Qwen/Qwen2___5-7B-Instruct", max_seq_length=max_seq_length, load_in_4bit=True, fast_inference=True, # 启用vLLM快速推理 max_lora_rank=lora_rank, gpu_memory_utilization=0.6, # 显存不足时可降低 ) # ======================================== # Step 2: LoRA配置 # ======================================== model = FastLanguageModel.get_peft_model( model, r=lora_rank, target_modules=[ "q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj", ], lora_alpha=lora_rank, use_gradient_checkpointing="unsloth", random_state=3407, )3)GSM8K数据准备
# ======================================== # Step 3: GSM8K数据准备 # ======================================== import re from datasets import load_dataset, Dataset # 系统提示词:定义推理输出格式 SYSTEM_PROMPT = """ Respond in the following format: <reasoning> ... </reasoning> <answer> ... </answer> """ def extract_xml_answer(text: str) -> str: """从XML格式文本中提取答案""" 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() def get_gsm8k_questions(split="train") -> Dataset: """加载GSM8K数据集""" data = load_dataset('/root/autodl-tmp/datasets/gsm8k', 'main')[split] data = data.map(lambda x: { 'prompt': [ {'role': 'system', 'content': SYSTEM_PROMPT}, {'role': 'user', 'content': x['question']} ], 'answer': extract_hash_answer(x['answer']) }) return data dataset = get_gsm8k_questions()get_gsm8k_questions函数读取gsm8k/main里面的数据,提取question/answer字段,批量的拼接提示词。
4)设计奖励函数
# ======================================== # Step 4: 奖励函数设计 # ======================================== 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]: """严格格式奖励:完全符合XML格式""" 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]: """宽松格式奖励:基本符合XML格式""" 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: """计算XML标签完整性得分""" 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]: """XML标签计数奖励""" contents = [completion[0]["content"] for completion in completions] return [count_xml(c) for c in contents]GRPO强化学习的答案和推理过程由 AI 自己生成;老师只充当判卷角色,不对推理过程做示范,依靠奖励函数评判 AI 输出好坏。奖励函数可以从不同维度评估 AI 输出:
- correctness_reward_func:检查最终答案是否正确
- int_reward_func:检查输出答案是否为整数
- strict_format_reward_func:严格校验输出格式
- soft_format_reward_func:宽松校验输出格式
- xmlcount_reward_func:校验 XML 标签使用是否正确
使用 GRPO 算法,让模型生成多个候选答案,依靠上面的奖励函数自动评估输出质量,指引模型往更优的方向优化,不需要人工逐条审阅输出结果。
5)GRPO训练
# ======================================== # Step 5: GRPOTrainer训练 # ======================================== max_prompt_length = 256 from trl import GRPOConfig, GRPOTrainer training_args = GRPOConfig( learning_rate=5e-6, adam_beta1=0.9, adam_beta2=0.99, weight_decay=0.1, warmup_ratio=0.1, lr_scheduler_type="cosine", optim="paged_adamw_8bit", logging_steps=1, per_device_train_batch_size=1, gradient_accumulation_steps=1, num_generations=6, # 每个问题生成6个候选答案 max_prompt_length=max_prompt_length, max_completion_length=max_seq_length - max_prompt_length, max_steps=250, save_steps=250, max_grad_norm=0.1, report_to="none", output_dir="outputs", ) 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()这部分代码跟SFTTrainer的逻辑差不多,参数略有不同。GRPOTrainer训练时还需要传相关的奖励函数。
总结:
GRPO(Group Relative Policy Optimization)组相对策略优化,是一种用于训练LLM的强化学习算法, 是DeepSeek-R1模型的核心技术之一。
核心在于通过组内样本的相对奖励来优化策略模型,而不是依赖传统的价值函数模型(如PPO中的批评家模型)。它通过采样一组输出,利用这些输出的奖励值来计算相对优势,从而简化了训练过程。
工作原理:
• 采样与奖励计算:对于每个输入问题,GRPO从当前策略中采样一组输出,并计算每个输出的奖励值。
• 相对优势估计:通过将每个输出的奖励值与组内平均奖励值进行比较,计算出每个输出的相对优势。
• 策略更新:根据相对优势,GRPO更新策略模型,优先 选择相对优势更高的输出。同时,它通过KL散度约束来控制策略更新的幅度,确保策略分布的稳定性。