Qwen+GRPO后训练实战

1、数据集&环境配置

medical-o1-reasoning-SFT

这些医疗问题是从Deepseek-R1中蒸馏出来的。

安装环境

CUDA13 + pytorch2.11 + unsloth2026.8.10 + vllm0.26.0

2、训练框架unsloth

(1)概述

一般偏好对齐用的强化学习框架有trl(huggingface)/verl(字节)/swift(阿里)。

本次使用unslot,实现低资源的微调。

unslot好处:

1.速度快,相对其他框架显存占用低。

  1. 使用triton重写了GPU的kernel。

  2. 与transformers,trl,peft兼容

(2)重要类

FastLanguageModel,类似于将transformers中的AutoModelForCausalLm,AutoTokenizer合成一个。

PatchFastRL,对trl做了一些补丁

PatchFastRL("GRPO", FastLanguageModel),把最新代码拉下来,放在unsloth_compiled_cache目录下。

3、代码

(1)加载模型

python 复制代码
import os
os.environ['UNSLOTH_USE_MODELSCOPE'] = '1' # 使用modelscope下载模型和数据
os.environ['UNSLOTH_DISABLE_STATISTICS'] = '1'  # 禁用unsloth在微调过程中自动收集统计信息
os.environ['OMP_NUM_THREADS'] = '4'  # OpenMP的并行线程设置
python 复制代码
# 引入unsloth以及GRPO最新的patch
from unsloth import FastLanguageModel, PatchFastRL
PatchFastRL("GRPO", FastLanguageModel)
python 复制代码
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_length=max_seq_length,
    load_in_4bit=True,     # 基座用4bit量化加载(动态4bit量化)
    fast_inference=True,   # 使用vllm加速推理,base(vllm)+adapter(bf16)
    max_lora_rank=lora_rank,
    gpu_memory_utilization=0.5,
)

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', #unsloth专门做过显存优化,相比transformers库显存更低,而且可以实现自动卸载activation到cpu
    random_state=3407
)

(2)准备数据

python 复制代码
def get_medical_question(data_path):
    # 从本地加载json文件
    data = load_dataset('json', data_files=data_path)['train']
    # 转换成dataframe
    df = data.to_pandas()

    # 99%的数据做训练,1%的数据做测试,大约200条
    train_df, test_df = train_test_split(df, test_size=0.01, random_state=42)

    # 转换回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
python 复制代码
# 根据长度过滤数据
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_dataset

(3)Reward函数

Reward=0.5*SemanticCorrectness + 0.4*PerplexityScore + 0.1*TagPresence

1.SemanticCorrectness语义相关度,用一个cross-encoder model(cross-encoder/stsb-roberta-base,0.12B的模型,需要常驻显存)计算和标准答案的语义正确性得分

  1. PerplexityScore困惑度得分,采用BioGPT(英文模型)计算医学流畅性和语言质量

  2. TagPresence检查输出的答案是否符合先思考后答案的格式

python 复制代码
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', device=reward_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_bar=False).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', device=reward_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_size=8) -> 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_tensors='pt',
                    padding=True,
                    truncation=True,
                    max_length=200
                ).to(self.device)

                with torch.no_grad():
                    outputs = self.model(**encodings, labels=encodings.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(f"Error 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(r'<reasoning>.*?</reasoning>', content, re.DOTALL))
        has_answer = bool(re.search(r'<answer>.*?</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(r'<answer>(.*?)</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(f"Reward calculation error:{str(e)}")
        return [-1.0] * len(completions)

    # 转成tensor
    sim_scores = torch.nan_to_num(torch.tensor(similarities), nan=0.0)
    perplex_scores = torch.nan_to_num(torch.tensor(perplexities), nan=1000.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_rewards

(4)GRPO的配置

python 复制代码
from trl import GRPOConfig, GRPOTrainer
training_args = GRPOConfig(
    use_vllm=True,  # use vLLM for fast inference!
    learning_rate=5e-6,  # 学习率
    weight_decay=0.001,  # 权重衰减
    warmup_ratio=0.1,  # warmup ratio
    max_grad_norm=0.1,  # 梯度裁剪,防止更新过快
    per_device_train_batch_size=5,  # 如果00M,可以设置小点,这个值一般是num_generations的倍数
    gradient_accumulation_steps=4,  # 梯度累积步数
    num_generations=5,  # 如果显存不足,可以设置小点
    max_prompt_length=max_prompt_len,  # 输入长度
    max_completion_length=max_output_len, #输出长度
    max_steps=200,  # 训练最大步数
    save_steps=200,  # 保存模型最大间隔
    lr_scheduler_type="cosine",  # 学习率衰减
    optim="adamw_8bit",  # 优化器
    logging_steps=1,  # 日志打印间隔
    bf16=True,  # 是否启用bf16训练
    report_to="none",
    output_dir="saved/", # checkpoint保存目录,这个只保存lora和optimizer参数
    save_strategy="steps" #保存的策略(按step还是epoch)
)

(5)训练

python 复制代码
trainer = GRPOTrainer(
    model=model,
    processing_class=tokenizer,
    reward_funcs = [
        combined_reward_func
    ],
    args=training_args,
    train_dataset=train_dataset
)
trainer.train()

(6)打印训练日志

python 复制代码
history = trainer.state.log_history
history[:3]
python 复制代码
# 保存日志
import json
with open('saved/history_log.txt', 'w') as fw:
    fw.write(json.dumps(history, ensure_ascii=False, indent=2))
python 复制代码
# 画图
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, label="reward")
plt.legend()
plt.subplot(2, 2, 2)
plt.plot(reward_std, label="reward_std")
plt.legend()
plt.subplot(2, 2, 3)
plt.plot(kl, label="kl")
plt.legend()
plt.subplot(2, 2, 4)
plt.plot(completion_length, label="completion_length")
plt.legend()
plt.show()

reward在50步左右收敛,基本维持在0.6左右。

kl散度逐步提升,偏离原始的模型。有毛刺,因为计算时没有clip。

completion_length长度先很长,然后逐渐稳定下来。

(7)保存lora

python 复制代码
# 保存lora
model.save_lora('saved/qwen_grpo_medical_reasoning_lora')

220M左右。

(8)测试

python 复制代码
# 测试
from vllm import SamplingParams
sampling_params = SamplingParams(
    n=1,
    temperature=0.8,
    top_p=0.95,
    max_tokens=1024
)

# GRPO训练前
text = tokenizer.apply_chat_template(test_dataset[0]['prompt'], tokenize=False, add_generation_prompt=True)
output = model.fast_generate(
    [text],
    sampling_params=sampling_params,
    lora_request=None,
)[0].outputs[0].text

print(output)

格式上就不对,<answer>没有对应的</answer>

python 复制代码
# GRPO训练后
text = tokenizer.apply_chat_template(test_dataset[0]['prompt'], tokenize=False, add_generation_prompt=True)
output = model.fast_generate(
    [text],
    sampling_params=sampling_params,
    lora_request=model.load_lora('saved/qwen_grpo_medical_reasoning_lora'),
)[0].outputs[0].text

print(output)
相关推荐
chen_zn957 小时前
《VLA 系列》Human-to-Robot Transfer | 人类视频共训练 | 跨本体涌现迁移 | 论文解析
人工智能·深度学习·transformer·具身智能·vla
CCC:CarCrazeCurator7 小时前
DeepSeek‑Harness Windows 本地快速上手
深度学习·数据挖掘
还不秃顶的计科生8 小时前
具身智能论文学习8:Octo: An Open-Source Generalist Robot Policy
人工智能·深度学习·学习·机器学习·语言模型·vla·vlm
lucas_AI9 小时前
1.2B 小模型赢过 235B 大模型:NaviDC-OCR 把文档解析卷明白了
人工智能·深度学习·算法
杀生丸学AI9 小时前
【稀疏重建】StructSplat:基于非校准稀疏视图的可泛化3DGS
深度学习·3d·音视频·transformer·三维重建·空间智能
qfljg11 小时前
如何学习opencode和openclaw源码
深度学习
在世修行11 小时前
深度图像数据格式与RAW文件解析:从字节到三维世界的桥梁
人工智能·数码相机·计算机视觉
zhenaibo52112 小时前
摘要、研究背景、文献综述AI风险高,怎么优化?
人工智能·深度学习·自然语言处理
hhzz14 小时前
Tiger AI 平台「手势识别」功能全解析:从数字手势 0–9,到中国手语字母,再到本地视频批量识别——一条链路,三种输入,双模型可同开;双手比划,机器秒懂
人工智能·python·深度学习·aigc·音视频
AndrewHZ14 小时前
图像处理入门006:图像质量评价三剑客:PSNR / SSIM / MSE 原理与代码实战
图像处理·计算机视觉·cv·psnr·rse·ssim