YAML Metadata Warning:empty or missing yaml metadata in repo card
Check out the documentation for more information.
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 保存
# 训练
python -m rl_training.train
# 指定 epoch 数
python -m rl_training.train --epochs 10
# 从 checkpoint 恢复
python -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 行)
对比微调后策略与基线策略的性能:
# 评估基线
python -m rl_training.eval --baseline
# 评估微调模型
python -m rl_training.eval --checkpoint rl_training/lora_checkpoints/checkpoint-100
# 对比模式
python -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 推理服务)
- Downloads last month
- 7