尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Qwen+GRPO后训练实战

Qwen+GRPO后训练实战 1、数据集环境配置medical-o1-reasoning-SFT这些医疗问题是从Deepseek-R1中蒸馏出来的。安装环境CUDA13 pytorch2.11 unsloth2026.8.10 vllm0.26.02、训练框架unsloth1概述一般偏好对齐用的强化学习框架有trl(huggingface)/verl(字节)/swift阿里。本次使用unslot实现低资源的微调。unslot好处1.速度快相对其他框架显存占用低。2. 使用triton重写了GPU的kernel。3. 与transformerstrlpeft兼容2重要类FastLanguageModel类似于将transformers中的AutoModelForCausalLm,AutoTokenizer合成一个。PatchFastRL对trl做了一些补丁PatchFastRL(GRPO, FastLanguageModel)把最新代码拉下来放在unsloth_compiled_cache目录下。3、代码1加载模型import os os.environ[UNSLOTH_USE_MODELSCOPE] 1 # 使用modelscope下载模型和数据 os.environ[UNSLOTH_DISABLE_STATISTICS] 1 # 禁用unsloth在微调过程中自动收集统计信息 os.environ[OMP_NUM_THREADS] 4 # OpenMP的并行线程设置# 引入unsloth以及GRPO最新的patch from unsloth import FastLanguageModel, PatchFastRL PatchFastRL(GRPO, FastLanguageModel)import torch max_prompt_len 512 max_output_len 512 max_seq_length max_prompt_len max_output_len lora_rank 64 model, tokenizer FastLanguageModel.from_pretrained( model_name/root/autodl-tmp/model/Qwen2.5-3B-Instruct, max_seq_lengthmax_seq_length, load_in_4bitTrue, # 基座用4bit量化加载(动态4bit量化) fast_inferenceTrue, # 使用vllm加速推理base(vllm)adapter(bf16) max_lora_ranklora_rank, gpu_memory_utilization0.5, ) 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, #unsloth专门做过显存优化相比transformers库显存更低而且可以实现自动卸载activation到cpu random_state3407 )2准备数据def get_medical_question(data_path): # 从本地加载json文件 data load_dataset(json, data_filesdata_path)[train] # 转换成dataframe df data.to_pandas() # 99%的数据做训练1%的数据做测试大约200条 train_df, test_df train_test_split(df, test_size0.01, random_state42) # 转换回huggingface datasets train_dataset Dataset.from_pandas(train_df) test_dataset Dataset.from_pandas(test_df) # 构建chat-ml所需的对话数据格式 def map_fn(x): return { prompt:[ {role: system, content: SYSTEM_PROMPT}, {role: user, content: x[Question]} ], answer: x[Response], question: x[Question] } # 处理数据格式并移除不需要的列 train_dataset train_dataset.map(map_fn).remove_columns([Question, Complex_CoT, Response, __index_level_0__]) test_dataset test_dataset.map(map_fn).remove_columns([Question, Complex_CoT, Response, __index_level_0__]) return train_dataset, test_dataset data_path /root/autodl-tmp/data/medical-o1-reasoning-SFT/medical_o1_sft.json train_dataset, test_dataset get_medical_question(data_path) train_dataset, test_dataset# 根据长度过滤数据 train_dataset train_dataset.map(lambda x: {input_len: len(tokenizer.apply_chat_template(x[prompt]))}) train_dataset train_dataset.map(lambda x: {output_len: len(tokenizer(x[answer])[input_ids])}) test_dataset test_dataset.map(lambda x: {input_len: len(tokenizer.apply_chat_template(x[prompt]))}) test_dataset test_dataset.map(lambda x: {output_len: len(tokenizer(x[answer])[input_ids])}) train_dataset train_dataset.filter(lambda x: x[input_len] max_prompt_len and x[output_len] max_output_len) test_dataset test_dataset.filter(lambda x: x[input_len] max_prompt_len and x[output_len] max_output_len) train_dataset train_dataset.remove_columns([input_len, output_len]) test_dataset test_dataset.remove_columns([input_len, output_len]) train_dataset, test_dataset3Reward函数Reward0.5*SemanticCorrectness 0.4*PerplexityScore 0.1*TagPresence1.SemanticCorrectness语义相关度用一个cross-encoder model(cross-encoder/stsb-roberta-base,0.12B的模型需要常驻显存)计算和标准答案的语义正确性得分2. PerplexityScore困惑度得分采用BioGPT(英文模型)计算医学流畅性和语言质量3. TagPresence检查输出的答案是否符合先思考后答案的格式from transformers import AutoModelForCausalLM, AutoTokenizer from sentence_transformers import CrossEncoder from typing import List import re main_device cuda if torch.cuda.is_available() else cpu reward_device cuda if torch.cuda.is_available() else cpu # -------------语义正确性得分---------------- semantic_model CrossEncoder(/root/autodl-tmp/model/stsb-roberta-base, devicereward_device) def semantic_correctness(responses: List[str], answers: List[str]) - List[float]: with torch.no_grad(): inputs list(zip(responses, answers)) similarities semantic_model.predict(inputs, show_progress_barFalse).tolist() # 如果回答为空则输出-1 similarities [-1.0 if response else similarity for response, similarity in zip(responses, similarities)] return similarities # -------------流畅性得分---------------- class PerplexityCalculator: def __init__(self, model_name/root/autodl-tmp/model/biogpt, devicereward_device): self.tokenizer AutoTokenizer.from_pretrained(model_name) self.device device self.tokenizer.pad_token self.tokenizer.eos_token self.model AutoModelForCausalLM.from_pretrained(model_name).to(self.device) self.model.eval() def calculate(self, texts: List[str], batch_size8) - List[float]: perplexities [] for i in range(0, len(texts), batch_size): batch texts[i : i batch_size] try: if not batch : continue encodings self.tokenizer( batch, return_tensorspt, paddingTrue, truncationTrue, max_length200 ).to(self.device) with torch.no_grad(): outputs self.model(**encodings, labelsencodings.input_ids) loss outputs.loss if torch.isnan(loss): raise ValueError(Nan loss encountered) batch_perplexity torch.exp(loss).repeat(len(batch)).cpu().tolist() # 如果回答为空则输出-1 batch_perplexity [-1.0 if text else perplex for text, perplex in zip(batch, batch_perplexity)] perplexities.extend(batch_perplexity) except Exception as e: print(fError in batch {i//batch_size}: {str(e)}) perplexities.extend([1000.0] * len(batch)) return perplexities perplexity_calculator PerplexityCalculator() # -------------输出格式得分---------------- def tag_presence_reward(completions: List[dict]) - List[float]: rewards [] for completion in completions: content completion[0][content] has_reasoning bool(re.search(rreasoning.*?/reasoning, content, re.DOTALL)) has_answer bool(re.search(ranswer.*?/answer, content, re.DOTALL)) reward 0.5 * has_reasoning 0.5 * has_answer rewards.append(reward) return rewards # -------------总和加权得分---------------- def combined_reward_func( prompts, completions, answer, **kwargs ) - List[float]: # 抽取答案 responses [] valid_indices [] full_outputs [] for idx, completion in enumerate(completions): # vllm输出的 try: generated_content completion[0][content].strip() # reasoning answer full_outputs.append(generated_content) answer_match re.search(ranswer(.*?)/answer, generated_content, re.DOTALL) if answer_match: generated_content answer_match.group(1).strip() else: responses.append() # 如果没有抽取出答案则给空 valid_indices.append(idx) continue # 处理边界条件1.答案为空 2.答案复制输入 user_prompt prompts[idx][-1][content] if not generated_content or generated_content user_prompt: responses.append() valid_indices.append(idx) continue responses.append(generated_content) valid_indices.append(idx) except (KeyError, IndexError): responses.append() valid_indices.append(idx) continue if not responses: return [-1.0] * len(completions) # 计算rewards try: processed_answers answer similarities semantic_correctness(responses, [processed_answers[i] for i in valid_indices]) perplexities perplexity_calculator.calculate([full_outputs[i] for i in valid_indices]) tag_rewards tag_presence_reward([completions[i] for i in valid_indices]) except Exception as e: print(fReward calculation error:{str(e)}) return [-1.0] * len(completions) # 转成tensor sim_scores torch.nan_to_num(torch.tensor(similarities), nan0.0) perplex_scores torch.nan_to_num(torch.tensor(perplexities), nan1000.0) tag_scores torch.tensor(tag_rewards) # 困惑度归一化 perplex_rewards 1 / (perplex_scores / (perplex_scores.mean() 1e-9)) score_range perplex_rewards.max() - perplex_rewards.min() if score_range 1e-6: perplex_rewards_normalized torch.ones_like(perplex_rewards) * 0.5 else: perplex_rewards_normalized (perplex_rewards - perplex_rewards.min()) / score_range # 加权 combined [ 0.5 * sim.item() 0.4 * pr.item() 0.1 * tag.item() for sim, pr, tag in zip(sim_scores, perplex_rewards_normalized, tag_scores) if not torch.isnan(sim) and not torch.isnan(pr) and not torch.isnan(tag) ] # clip rewards, 保证-1~1之间 final_rewards [-1.0] * len(completions) for idx, reward in zip(valid_indices, combined): final_rewards[idx] max(min(reward, 1.0), -1.0) assert len(final_rewards) len(completions), Reward mapping error return final_rewards4GRPO的配置from trl import GRPOConfig, GRPOTrainer training_args GRPOConfig( use_vllmTrue, # use vLLM for fast inference! learning_rate5e-6, # 学习率 weight_decay0.001, # 权重衰减 warmup_ratio0.1, # warmup ratio max_grad_norm0.1, # 梯度裁剪防止更新过快 per_device_train_batch_size5, # 如果00M可以设置小点这个值一般是num_generations的倍数 gradient_accumulation_steps4, # 梯度累积步数 num_generations5, # 如果显存不足可以设置小点 max_prompt_lengthmax_prompt_len, # 输入长度 max_completion_lengthmax_output_len, #输出长度 max_steps200, # 训练最大步数 save_steps200, # 保存模型最大间隔 lr_scheduler_typecosine, # 学习率衰减 optimadamw_8bit, # 优化器 logging_steps1, # 日志打印间隔 bf16True, # 是否启用bf16训练 report_tonone, output_dirsaved/, # checkpoint保存目录这个只保存lora和optimizer参数 save_strategysteps #保存的策略按step还是epoch )5训练trainer GRPOTrainer( modelmodel, processing_classtokenizer, reward_funcs [ combined_reward_func ], argstraining_args, train_datasettrain_dataset ) trainer.train()6打印训练日志history trainer.state.log_history history[:3]# 保存日志 import json with open(saved/history_log.txt, w) as fw: fw.write(json.dumps(history, ensure_asciiFalse, indent2))# 画图 import matplotlib.pyplot as plt reward [item[reward] for item in history if reward in item] reward_std [item[reward_std] for item in history if reward_std in item] kl [item[kl] for item in history if kl in item and item[kl] 1.0] completion_length [item[completion_length] for item in history if completion_length in item] plt.figure(figsize(10, 6)) plt.subplot(2, 2, 1) plt.plot(reward, labelreward) plt.legend() plt.subplot(2, 2, 2) plt.plot(reward_std, labelreward_std) plt.legend() plt.subplot(2, 2, 3) plt.plot(kl, labelkl) plt.legend() plt.subplot(2, 2, 4) plt.plot(completion_length, labelcompletion_length) plt.legend() plt.show()reward在50步左右收敛基本维持在0.6左右。kl散度逐步提升偏离原始的模型。有毛刺因为计算时没有clip。completion_length长度先很长然后逐渐稳定下来。7保存lora# 保存lora model.save_lora(saved/qwen_grpo_medical_reasoning_lora)220M左右。8测试# 测试 from vllm import SamplingParams sampling_params SamplingParams( n1, temperature0.8, top_p0.95, max_tokens1024 ) # GRPO训练前 text tokenizer.apply_chat_template(test_dataset[0][prompt], tokenizeFalse, add_generation_promptTrue) output model.fast_generate( [text], sampling_paramssampling_params, lora_requestNone, )[0].outputs[0].text print(output)格式上就不对answer没有对应的/answer# GRPO训练后 text tokenizer.apply_chat_template(test_dataset[0][prompt], tokenizeFalse, add_generation_promptTrue) output model.fast_generate( [text], sampling_paramssampling_params, lora_requestmodel.load_lora(saved/qwen_grpo_medical_reasoning_lora), )[0].outputs[0].text print(output)
返回列表