新闻详情

新闻详情

首页 / 资讯中心 / 详情

LLM微调-训练自己的R1模型

发布时间:2026/9/25 19:59:27来源:尧图网络
LLM微调-训练自己的R1模型
目录一、准备工作二、GRPO强化学习总结在 模型蒸馏介绍 里了解了模型蒸馏的过程就是用 SFT/GRPO的方法让学生模型能学习到教师模型的能力从而使学生模型的能力得到大的提升并且在LLM微调-训练垂类问答模型 里面学习了SFT模型微调。SFT监督学习需要给到固定格式的数据让大模型快速的学习到基础知识。GRPO强化学习则是给出问题标准答案通过奖励函数的引导让大模型自己去推理从而提升大模型的推理能力。一、准备工作1数据准备GSM8KGrade School Math 8K是一个高质量的小学数学应用题数据集主要用于评估和训练人工智能模型 在数学推理和多步问题解决方面的能力 https://huggingface.co/datasets/openai/gsm8k2环境准备由于我这次选的模型是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_lengthmax_seq_length, load_in_4bitTrue, fast_inferenceTrue, # 启用vLLM快速推理 max_lora_ranklora_rank, gpu_memory_utilization0.6, # 显存不足时可降低 ) # # Step 2: LoRA配置 # model FastLanguageModel.get_peft_model( model, rlora_rank, target_modules[ q_proj, k_proj, v_proj, o_proj, gate_proj, up_proj, down_proj, ], lora_alphalora_rank, use_gradient_checkpointingunsloth, random_state3407, )3GSM8K数据准备# # 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(splittrain) - 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, fQuestion:\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\nanswer\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 rreasoning.*?/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(\nanswer\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 算法让模型生成多个候选答案依靠上面的奖励函数自动评估输出质量指引模型往更优的方向优化不需要人工逐条审阅输出结果。5GRPO训练# # Step 5: GRPOTrainer训练 # max_prompt_length 256 from trl import GRPOConfig, GRPOTrainer training_args GRPOConfig( learning_rate5e-6, adam_beta10.9, adam_beta20.99, weight_decay0.1, warmup_ratio0.1, lr_scheduler_typecosine, optimpaged_adamw_8bit, logging_steps1, per_device_train_batch_size1, gradient_accumulation_steps1, num_generations6, # 每个问题生成6个候选答案 max_prompt_lengthmax_prompt_length, max_completion_lengthmax_seq_length - max_prompt_length, max_steps250, save_steps250, max_grad_norm0.1, report_tonone, output_diroutputs, ) trainer GRPOTrainer( modelmodel, processing_classtokenizer, reward_funcs[ xmlcount_reward_func, soft_format_reward_func, strict_format_reward_func, int_reward_func, correctness_reward_func, ], argstraining_args, train_datasetdataset, ) # 开始训练 trainer.train()这部分代码跟SFTTrainer的逻辑差不多参数略有不同。GRPOTrainer训练时还需要传相关的奖励函数。总结GRPOGroup Relative Policy Optimization组相对策略优化是一种用于训练LLM的强化学习算法 是DeepSeek-R1模型的核心技术之一。核心在于通过组内样本的相对奖励来优化策略模型而不是依赖传统的价值函数模型如PPO中的批评家模型。它通过采样一组输出利用这些输出的奖励值来计算相对优势从而简化了训练过程。工作原理• 采样与奖励计算对于每个输入问题GRPO从当前策略中采样一组输出并计算每个输出的奖励值。• 相对优势估计通过将每个输出的奖励值与组内平均奖励值进行比较计算出每个输出的相对优势。• 策略更新根据相对优势GRPO更新策略模型优先 选择相对优势更高的输出。同时它通过KL散度约束来控制策略更新的幅度确保策略分布的稳定性。
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

AI 辅助 Python 排错:从多出一个空页到回归测试 2026/9/25 20:36:36

AI 辅助 Python 排错:从多出一个空页到回归测试

4 条数据,每页 2 条,分页函数却返回了 3 页,最后一页还是空的。 代码没有抛异常,接口也可能正常返回成功状态。直到调用方发现“下一页”里什么都没有,问题才暴露出来。 这类 Bug 很适合用来练习 AI 辅助排错&#x…

阅读更多 →
Agent安全:权限控制与沙箱执行 2026/9/25 20:36:09

Agent安全:权限控制与沙箱执行

Agent安全:权限控制与沙箱执行 专栏:AI/LLM工程化实战 - 从Prompt到Agent的完整落地指南 模块4 Agent工程实战篇 第42篇 摘要 摘要:Agent权限最小化、工具白名单、代码沙箱subprocess受限执行、敏感数据脱敏、审计日志,是Agent安全防护的五大核心手段。用可运行Python实现一道工…

阅读更多 →
「Python 翻车日记 · 第 12 篇」一行 read() 读 5GB,电脑就炸了?——文件是流,不是一坨 2026/9/25 20:36:02

「Python 翻车日记 · 第 12 篇」一行 read() 读 5GB,电脑就炸了?——文件是流,不是一坨

Python 翻车日记 第 12 篇:一行 read() 读 5GB,电脑就炸了?——文件是流,不是一坨 📋 本期菜单:6 个文件 I/O 的坑 + 1 个模式速查,从「read OOM」到「JSON 中文转义」 [入门] read OOM with 句柄泄漏 glob 不递归 二进制写 str [进阶] os.path.join 绝对路径丢弃 …

阅读更多 →
K值新国标落地!2026选门窗不看这个参数,冬天多交一半暖气费 2026/9/25 20:36:02

K值新国标落地!2026选门窗不看这个参数,冬天多交一半暖气费

最近不少准备装修的朋友发现,去门店选窗户,销售都在提一个参数——K值。有的说是2.0,有的说是1.5,还有的说是1.1。这到底是个啥?新国标实施后,K值不达标会有什么后果?今天一篇讲清楚。K值到底是…

阅读更多 →
Agent框架对比:LangGraph、LlamaIndex、AutoGen三选一 2026/9/25 20:35:55

Agent框架对比:LangGraph、LlamaIndex、AutoGen三选一

Agent框架对比:LangGraph、LlamaIndex、AutoGen三选一 专栏:AI/LLM工程化实战 - 从Prompt到Agent的完整落地指南 模块4 Agent工程实战篇 第41篇 摘要 摘要:Agent框架选型、LangGraph有向图状态机、LlamaIndex数据编排、AutoGen多角色会话,是三大主流Agent框架的架构核心。用可运…

阅读更多 →
从Activity Log生成周报:工作事件模型、聚合规则与事实边界 2026/9/25 20:35:55

从Activity Log生成周报:工作事件模型、聚合规则与事实边界

自动生成周报的可靠路径,不是让模型扫描所有聊天并自由总结,而是先建立轻量 Activity Log:用结构化工作事件记录结果、决定、阻塞、行动和来源,再按项目与时间聚合,最后由模型负责表达压缩。 工作事件模型 type WorkEv…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉