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 — 训练主入口

五阶段训练管线:

  1. 数据加载:RLCheckpointLoader 从多智能体系统 checkpoints 加载训练样本
  2. 模型初始化:基座模型 + LoRA(4-bit 量化)
  3. 环境创建:SearchRoundEnv 工厂函数
  4. GRPO Trainer 初始化
  5. 训练循环:按 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) 训练/验证集划分(困难题优先分到验证集)

数据提取逻辑:

  1. 扫描 checkpoints/ 下所有问题目录
  2. 对每个问题,读取 run_XXX/meta.json 获取轮次列表
  3. 筛选 strategy=SEARCH 的轮次(即有搜索线索生成的轮次)
  4. 对每个 SEARCH 轮次,读取上一轮 checkpoint 获取执行前的状态
  5. 重建实体表、目标表、文档表、上下文表的文本表示

关键参数:

  • 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
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support