首页/新闻资讯/正文详情

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

发布时间:2026/9/25 19:59:24 来源:云帆数科 栏目:资讯中心
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散度约束来控制策略更新的幅度确保策略分布的稳定性。

相关推荐

通信一断就全线停摆?四级降级策略实战:从全功能到纯旁路的平滑无扰过渡
通信一断就全线停摆?四级降级策略实战:从全功能到纯旁路的平滑无扰过渡

前阵子在外地调一条中药煎药产线,碰到个非常典型的工业现场问题:上位机和S7-1200走以太网通信,车间里变频器、电焊机一多,就容易出现通信闪断。每次闪断,整条产线立刻触发通信故障停机,轻则半锅药熬废了,重则整批生产进度往后拖大半天。 现场运维换了屏蔽网线、加了工业… · 2026/9/25 19:59:17

2026 Java架构师进阶路线:从JVM到微服务的大纲拆解与学习计划
2026 Java架构师进阶路线:从JVM到微服务的大纲拆解与学习计划

后端进阶最缺的不是资料,是条能走完的线。我把一套 Java 架构师课程(第 03 期,50 讲)的大纲按学习顺序拆成路线图,标了每个阶段该产出什么,作为阶段学习的参照。 一、路线总览(四阶段&#xff… · 2026/9/25 19:57:46

AI辅助Android开发:从传统到智能化的技术演进——TaoToken统一Key接入Cline与CC Switch配置实战
AI辅助Android开发:从传统到智能化的技术演进——TaoToken统一Key接入Cline与CC Switch配置实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 19:57:40

【Matlab】复杂背景无人机目标筛选算法
【Matlab】复杂背景无人机目标筛选算法

【Matlab】复杂背景无人机目标筛选算法 引言 无人机在军事侦察、目标监视、灾情评估和交通管理等任务中日益普及,其核心能力之一是从复杂背景中快速准确地发现并筛选出感兴趣的目标。然而实际拍摄场景往往背景复杂,包含树木、建筑、水面、云层等多种干扰因素,目标常以不同… · 2026/9/25 20:27:03

HoRain云--Java 代码重构实战:坏味道识别与 10 种重构手法
HoRain云--Java 代码重构实战:坏味道识别与 10 种重构手法

本文总结 Java 常见代码坏味道,包括长方法、重复代码、过深嵌套、魔法值、上帝类,并给出 10 种可落地的重构手法。正文:1. 什么是代码坏味道坏味道是代码中潜在问题的信号,不一定报错,但影响可读性、可维护性和扩展性。… · 2026/9/25 20:27:03

HoRain云--Spring Boot 文件上传下载实战:分片、断点续传与 OSS 集成
HoRain云--Spring Boot 文件上传下载实战:分片、断点续传与 OSS 集成

本文讲解 Spring Boot 文件上传下载&#xff0c;包括 Multipart、分片上传、断点续传、秒传、限制配置和阿里云 OSS 集成。正文&#xff1a;1. 基础上传java复制下载PostMapping("/upload") public Result<String> upload(RequestParam("file") Mult… · 2026/9/25 20:26:57

一个.class文件是怎么被JVM跑起来的?类加载机制+双亲委派一篇讲透
一个.class文件是怎么被JVM跑起来的?类加载机制+双亲委派一篇讲透

一个 .class 文件是怎么被 JVM 跑起来的&#xff1f;类加载机制 双亲委派一篇讲透 你写的 Java 代码编译成 .class 之后&#xff0c;JVM 到底对它做了什么&#xff0c;才能让它真正跑起来&#xff1f; 为什么静态代码块只执行一次&#xff1f;为什么你自己写一个 java.lang.St… · 2026/9/25 20:26:50

图像处理标准测试图全指南:Lena、cameraman获取与加载方法
图像处理标准测试图全指南:Lena、cameraman获取与加载方法

刚入行做图像处理的同学&#xff0c;十有八九都经历过这么一幕&#xff1a;导师或者教程里让“用Lena图跑一下高斯滤波”“拿cameraman测一下边缘检测”&#xff0c;结果你打开搜索引擎&#xff0c;翻半天找到的图片要么带水印、要么尺寸不对、要么干脆是博主自己拍的某个二次元… · 2026/9/25 20:26:38

Atlas 300V部署YOLO实战:从模型转换到推理加速全流程
Atlas 300V部署YOLO实战:从模型转换到推理加速全流程

1. 项目概述&#xff1a;Atlas到底是何方神圣先说结论&#xff1a;当你看到“atlas”这个词出现在AI硬件语境里&#xff0c;几乎可以默认它指的是华为昇腾&#xff08;Ascend&#xff09;系列的AI计算平台。从加速卡到服务器整机再到边缘计算盒子&#xff0c;Atlas是整个产品线… · 2026/9/25 20:26:38

数值优化(Numerical Optimization)学习系列-03-共轭梯度方法(Conjugate Gradient)
数值优化(Numerical Optimization)学习系列-03-共轭梯度方法(Conjugate Gradient)

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:00:31

创维E900V22D刷机全攻略:S905L3SB芯片兼容性解析与救砖实战
创维E900V22D刷机全攻略:S905L3SB芯片兼容性解析与救砖实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:00:31

MQTT协议原理与Broker服务器搭建实战:从Mosquitto到EMQX
MQTT协议原理与Broker服务器搭建实战:从Mosquitto到EMQX

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/25 1:00:37

了解更多?预约专属演示

我们的顾问将为您一对一讲解产品与方案

企业微信二维码