RL Training 强化学习微调系统
使用 GRPO(Group Relative Policy Optimization)算法 + LoRA(Low-Rank Adaptation)对多智能体系统中的 Executor agent 进行强化学习微调,优化其搜索线索生成策略,提高检索到 IMPORTANT 文档的效率和准确性。
目录结构
rl_training/
├── __init__.py # 包标记
├── config.yaml # 配置文件
├── train.py # 训练主入口
├── grpo_trainer.py # GRPO 训练器实现
├── environment.py # 搜索轮次环境(冻结 DocumentChecker)
├── reward.py # 奖励计算器
├── model_utils.py # LoRA 模型加载/保存/概率计算
├── data_loader.py # 训练数据加载(从 checkpoints 提取)
├── eval.py # 评估脚本(微调 vs 基线对比)
├── merge_lora_to_base.py # LoRA 权重合并到基座模型
├── quick_check.py # 快速验证工具
└── train_debug.py # 调试版训练脚本
各文件功能
config.yaml — 配置文件
模型与 LoRA:
| 配置项 | 含义 | 默认值 |
|---|
model.base_model_path | 基座模型 HuggingFace ID 或本地路径 | Qwen3-8B |
model.vllm_base_url | 冻结 agent 的 vLLM 服务地址 | http://127.0.0.1:8000/v1 |
model.load_in_4bit | 4-bit 量化加载 | true |
lora.rank | LoRA 秩 | 16 |
lora.alpha | LoRA 缩放因子 | 32 |
lora.dropout | LoRA dropout | 0.1 |
lora.target_modules | LoRA 目标模块 | ["q_proj", "v_proj", "k_proj", "o_proj"] |
GRPO 参数:
| 配置项 | 含义 | 默认值 |
|---|
grpo.group_size | 每组采样动作数 K | 4 |
grpo.clip_epsilon | PPO 裁剪范围 | 0.2 |
grpo.kl_penalty | KL 散度惩罚系数 | 0.01 |
grpo.learning_rate | 学习率 | 5e-5 |
grpo.temperature | 采样温度 | 0.7 |
奖励权重:
| 配置项 | 含义 | 默认值 |
|---|
reward.important_new | IMPORTANT + 首次发现 | 2.5 |
reward.important_seen | IMPORTANT + 已发现 | 0.3 |
reward.local_new | LOCAL + 首次发现 | 0.3 |
reward.local_seen | LOCAL + 已发现 | 0.1 |
reward.discard_new | DISCARD + 首次发现 | 0.0 |
reward.discard_seen | DISCARD + 已发现 | -0.1 |
训练参数:
| 配置项 | 含义 | 默认值 |
|---|
training.num_epochs | 训练 epoch 数 | 5 |
training.batch_size | 每步状态数 | 1 |
training.steps_per_epoch | 每 epoch 批次数 | 10 |
training.gradient_accumulation_steps | 梯度累积步数 | 2 |
training.max_steps | 最大步数 | 1000 |
training.save_every_n_steps | 保存间隔 | 1 |
training.resume | 自动加载最新 checkpoint | true |
train.py — 训练主入口
五阶段训练管线:
- 数据加载:
RLCheckpointLoader 从多智能体系统 checkpoints 加载训练样本
- 模型初始化:基座模型 + LoRA(4-bit 量化)
- 环境创建:
SearchRoundEnv 工厂函数
- GRPO Trainer 初始化
- 训练循环:按 epoch 迭代、采样批次、训练步骤、评估、checkpoint 保存
1# 训练
2python -m rl_training.train
3
4# 指定 epoch 数
5python -m rl_training.train --epochs 10
6
7# 从 checkpoint 恢复
8python -m rl_training.train --resume
grpo_trainer.py — GRPO 训练器 (404 行)
GRPOTrainer 实现 GRPO 算法核心:
| 方法 | 说明 |
|---|
training_step(states) | 单步训练:对每个状态采样 K 个动作 → 环境执行 → 计算奖励 → PPO+KL 损失 → 反向传播 |
_sample_action(state) | 采样单个动作:模型生成搜索线索 + 计算旧 log prob(独立前向传播) |
_compute_log_probs(state, action) | 计算当前策略的 log prob(有梯度)+ 参考策略的 log prob(disable_adapter + no_grad) |
evaluate(states) | 在验证集上评估,返回平均奖励和 KL 散度 |
save_checkpoint(step) | 保存 LoRA 权重 + optimizer 状态 + 指标 |
load_checkpoint() | 加载最新 checkpoint 恢复训练 |
GRPO 损失计算:
adv_i = (r_i - μ) / (σ + ε) # 组内优势标准化
ratio = exp(log_prob - old_log_prob) # 重要性采样比
pg_loss = -min(ratio * adv, clip(ratio, 1-ε, 1+ε) * adv) # 裁剪代理损失
kl = (ref_log_prob - log_prob)² # 逐 token KL 散度
loss = pg_loss + β × kl # 总损失
environment.py — 搜索轮次环境 (295 行)
SearchRoundEnv 模拟一轮搜索执行过程:
| 方法 | 说明 |
|---|
reset(round_context) | 初始化环境:问题、子目标、表状态、已发现文档集合 |
step(search_clues) | 执行搜索线索:通过 SearchOrchestrator 文档检索 → DocumentChecker 文档分类 → RewardCalculator 计算奖励 |
parse_clues_from_output(text) | 解析 LLM 输出的多格式线索(2-phase XML/legacy/bare clues) |
工厂函数 create_environment():将多智能体系统的冻结 agent(Planner、EntityManager、Checker 等)的 thinking 模式禁用,确保它们只作为冻结的文档验证器使用。
环境状态:round_data 包含问题、子目标、实体表、目标表、文档表、上下文表、已发现 docid 集合等完整上下文。
reward.py — 奖励计算器 (119 行)
RewardCalculator:
| 方法 | 说明 |
|---|
compute_reward(doc_results, seen_docids) | 计算单组搜索线索的总奖励 |
compute_group_advantages(rewards) | GRPO 组内优势标准化:adv = (r - mean(r)) / (std(r) + 1e-8) |
奖励分类:
| 分类 | 首次发现 | 已发现 |
|---|
| IMPORTANT(直接相关) | +2.5 | +0.3 |
| LOCAL(局部相关) | +0.3 | +0.1 |
| DISCARD(无关) | 0.0 | -0.1 |
model_utils.py — LoRA 模型工具 (169 行)
| 函数 | 说明 |
|---|
load_model_with_lora(base_path, lora_config) | 加载基座模型 + LoRA(4-bit 量化,bitsandbytes 回退 fp16) |
save_lora_checkpoint(model, path) | 保存 LoRA 权重 |
load_lora_checkpoint(model, path) | 加载 LoRA 权重 |
get_token_log_probs(model, input_ids, tokenizer) | 批量计算 token log probabilities |
generate_with_log_probs(model, input_ids, ...) | 生成 + 返回每个 token 的 log prob |
compute_sequence_log_prob(log_probs, token_ids) | 计算完整序列的对数概率 |
启用 gradient checkpointing 节省显存。
data_loader.py — 训练数据加载 (433 行)
RLCheckpointLoader:
| 方法 | 说明 |
|---|
load_data() | 扫描所有问题的 checkpoint 目录 |
extract_training_samples() | 提取所有 SEARCH 轮次上下文作为训练样本 |
split_train_eval(train_ratio) | 训练/验证集划分(困难题优先分到验证集) |
数据提取逻辑:
- 扫描
checkpoints/ 下所有问题目录
- 对每个问题,读取
run_XXX/meta.json 获取轮次列表
- 筛选
strategy=SEARCH 的轮次(即有搜索线索生成的轮次)
- 对每个 SEARCH 轮次,读取上一轮 checkpoint 获取执行前的状态
- 重建实体表、目标表、文档表、上下文表的文本表示
关键参数:
min_clues_per_round=2:过滤线索数过少的轮次
max_samples_per_question=10:每题最多采样轮次数
eval.py — 评估脚本 (245 行)
对比微调后策略与基线策略的性能:
1# 评估基线
2python -m rl_training.eval --baseline
3
4# 评估微调模型
5python -m rl_training.eval --checkpoint rl_training/lora_checkpoints/checkpoint-100
6
7# 对比模式
8python -m rl_training.eval --baseline --checkpoint rl_training/lora_checkpoints/checkpoint-100
输出:逐问题奖励对比、平均奖励、IMPORTANT 文档发现率等指标。详细结果保存为 JSON。
merge_lora_to_base.py — LoRA 合并 (237 行)
将 LoRA checkpoint 合并到基座模型,生成可直接被 vLLM 加载的完整模型。
python -m rl_training.merge_lora_to_base
输出:
{output_dir}/
├── Qwen3-8B-grpo-step5/
├── Qwen3-8B-grpo-step100/
└── ...
支持:CPU 合并避免 OOM、可选 checkpoint 列表、自动发现、防止覆盖已存在输出。
quick_check.py — 快速验证工具
快速验证训练数据和模型加载是否正常。
train_debug.py — 调试版训练
调试版训练脚本,简化配置和减少轮次用于快速迭代测试。
核心算法
GRPO(Group Relative Policy Optimization)
无需 Critic 模型的策略优化算法:
对每个状态 s:
1. 采样 K 个动作 a₁...a_K ~ π_θ(·|s)
2. 在环境中执行每个动作获得奖励 r₁...r_K
3. 计算组内优势:adv_i = (r_i - μ_group) / σ_group
4. 对新策略计算裁剪代理损失
5. 计算与参考策略的 KL 散度惩罚
6. 总损失 = 代理损失 + β × KL 散度
数据流
多智能体系统 checkpoints/
↓ RLCheckpointLoader
SEARCH 轮次的状态(问题、子目标、实体表、目标表、文档表、上下文表)
↓
SearchRoundEnv.reset()
↓ 每步训练
GRPOTrainer → π_θ 生成 K 个搜索线索 → 环境执行 → 奖励 → 优势 → 损失 → 更新 LoRA
↓
Qwen3-8B-grpo-step{N}/
↓ merge_lora_to_base.py
完整模型 → vLLM 加载 → 替换多智能体系统中的 Executor
冻结环境
训练过程中以下组件保持冻结(不参与梯度更新):
- Planner(策略选择)
- DocumentChecker(文档验证和分类)
- EntityManager(实体提取)
- SearchOrchestrator(BM25 搜索)
- RewardCalculator(奖励计算)
冻结方式:通过 vLLM API 调用(禁用 thinking),模型权重不变。
外部依赖
agent 模块(tools、vllm_client、dataset_utils)
transformers(HuggingFace 模型加载)
peft(LoRA 训练)
bitsandbytes(4-bit 量化,可选)
accelerate(分布式训练)
torch(深度学习框架)
PyYAML(配置解析)
vllm(冻结 agent 推理服务)