1、数据集&环境配置
这些医疗问题是从Deepseek-R1中蒸馏出来的。
安装环境
CUDA13 + pytorch2.11 + unsloth2026.8.10 + vllm0.26.0
2、训练框架unsloth
(1)概述
一般偏好对齐用的强化学习框架有trl(huggingface)/verl(字节)/swift(阿里)。
本次使用unslot,实现低资源的微调。
unslot好处:
1.速度快,相对其他框架显存占用低。
-
使用triton重写了GPU的kernel。
-
与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的模型,需要常驻显存)计算和标准答案的语义正确性得分
-
PerplexityScore困惑度得分,采用BioGPT(英文模型)计算医学流畅性和语言质量
-
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)


