ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

slime 复现 Search-R1 lite:基于多轮对话与工具调用的检索增强 RL 训练实战

slime 复现 Search-R1 lite:基于多轮对话与工具调用的检索增强 RL 训练实战 slime 复现 Search-R1 lite基于多轮对话与工具调用的检索增强 RL 训练实战【免费下载链接】slimeslime is an LLM post-training framework for RL Scaling.项目地址: https://gitcode.com/GitHub_Trending/slime12/slime本文以 slime 仓库中的examples/search-r1示例为主线讲解如何在 slime 中实现检索增强RAG式强化学习训练模型通过search工具调用触发搜索、读取information检索结果、最终以answer作答构成完整的多轮对话 工具调用闭环。读完本文你将掌握环境与数据的准备流程、本地检索与 Google 搜索两种后端的选择与配置、SEARCH_R1_CONFIGS全部关键参数的含义、自定义生成函数与奖励函数在 slime 中的挂载方式以及 TIS轨迹重要性采样的启用步骤和本地稠密检索服务的完整搭建方法。一、示例定位Search-R1 lite 是什么examples/search-r1是 Search-R1 的一个最小化复现lite 版本同时是 slime 中展示**多轮对话multi-turn conversation与工具调用tool-calling**的官方示例。其核心思路是让策略模型在回答开放域问答前先主动调用搜索引擎获取外部证据再基于检索结果生成答案并通过精确匹配EM奖励信号进行强化学习优化。相关文件全部位于 examples/search-r1 目录下文件作用generate_with_search.py自定义数据生成函数多轮检索对话的 rollout 主逻辑与自定义奖励函数run_qwen2.5_3B.sh基于 Qwen2.5-3B 的 GRPO 训练启动脚本local_dense_retriever/download.py下载 e5 稠密索引与 wiki 语料local_dense_retriever/retrieval_server.py本地检索服务FastAPI FAISSlocal_search_server.py训练侧调用的本地搜索客户端接口与 Google 搜索对齐google_search_server.pyserper.dev Google 搜索客户端qa_em_format.py检索问答的 EM 打分与序列格式校验二、环境准备与数据初始化2.1 安装依赖推荐直接使用slimerl/slime:latest镜像在容器内完成以下初始化cd /root/ git clone https://github.com/THUDM/slime.git pip install -e . --no-deps # for Search R1 pip install chardet其中chardet是 Google 搜索后端抓取网页时进行编码探测所需的依赖见 google_search_server.py 中chardet.detect的使用。随后克隆并安装 Search-R1 本体下载并处理训练数据cd /root/ git clone https://github.com/PeterGriffinJin/Search-R1.git cd Search-R1/ pip install -e . --no-deps pip install tensordict # Set your working directory WORK_DIR/root/Search-R1 LOCAL_DIR$WORK_DIR/data/nq_hotpotqa_train # Process multiple dataset search format train file DATAnq,hotpotqa python $WORK_DIR/scripts/data_process/qa_search_train_merge.py \ --local_dir $LOCAL_DIR \ --data_sources $DATA # (Optional) Process multiple dataset search format test file # Note: the final file is not shuffled DATAnq,triviaqa,popqa,hotpotqa,2wikimultihopqa,musique,bamboogle python $WORK_DIR/scripts/data_process/qa_search_test_merge.py \ --local_dir $LOCAL_DIR \ --data_sources $DATA训练数据默认使用 Natural Questionsnq与 HotpotQA 两个数据集的合并结果合并后的train.parquet正是启动脚本ROLLOUT_ARGS中--prompt-data指向的文件。测试集可选额外涵盖 triviaqa、popqa、2wikimultihopqa、musique、bamboogle 等注意测试文件最终不会做 shuffle。提示若打算使用本地检索后端需要先搭建本地检索服务详见本文附录。2.2 初始化 Qwen2.5-3B 模型训练需要同时准备 HuggingFace 格式检查点rollout 推理侧使用与 Megatron 格式检查点训练侧使用# hf checkpoint hf download Qwen/Qwen2.5-3B --local-dir /root/Qwen2.5-3B # mcore checkpoint cd /root/slime source scripts/models/qwen2.5-3B.sh PYTHONPATH/root/Megatron-LM python tools/convert_hf_to_torch_dist.py \ ${MODEL_ARGS[]} \ --hf-checkpoint /root/Qwen2.5-3B \ --save /root/Qwen2.5-3B_torch_distscripts/models/qwen2.5-3B.sh定义了 Qwen2.5-3B 的完整模型结构参数36 层、hidden size 2048、GQA 2 组查询头、RMSNorm、旋转位置编码等转换工具为 tools/convert_hf_to_torch_dist.py。转换得到的/root/Qwen2.5-3B_torch_dist将作为启动脚本中的--ref-load参考模型路径。三、Search Backend 配置本地检索与 Google 搜索generate_with_search.py同时支持本地检索与Google 搜索两种后端通过SEARCH_R1_CONFIGS字典统一配置SEARCH_R1_CONFIGS { # General Configuration max_turns: 2, topk: 3, search_concurrency: 256, # Search Backend Selection search_backend: local, # Options: local or google # Local Search Configuration # (Only used when search_backendlocal) local: { search_url: http://127.0.0.1:8000/retrieve, # URL of your local retrieval server proxy: None, }, # Google Search Configuration # (Only used when search_backendgoogle) google: { api_key: your_api_key_here, # Replace with your actual serper.dev API key snippet_only: True, proxy: None, }, # Log Probability Collection return_logprob: True, # Set to True to collect log probabilities (required for TIS) # Reward Model Configuration format_score: 0.2, }各参数含义与底层作用如下max_turns单个样本最多进行几轮搜索-反馈交互不含最终作答轮。generate_with_search.py中对应for _turn_idx in range(SEARCH_R1_CONFIGS[max_turns])的多轮循环。topk每次搜索返回的文档条数同时传给本地检索服务与 Google 搜索。search_concurrency搜索请求的并发上限。代码中会创建asyncio.Semaphore(SEARCH_R1_CONFIGS[search_concurrency])全局信号量限制同时发出的搜索请求数量防止外部检索服务被打满。search_backend后端选择local走本地检索服务器google走 serper.dev。非法值会抛出ValueError。local.search_url本地检索服务器的POST /retrieve接口地址。google.api_keyserper.dev 的 API Keygoogle.snippet_onlyTrue时只返回搜索结果摘要不抓取正文网页。return_logprob是否收集生成 token 的对数概率TIS 开启的硬性前提。收集时推理请求会携带return_logprobTrue并从output_token_logprobs中提取[log_prob, token_id, ...]三元组保证 token 与 logp 精确对齐。format_score奖励函数中格式正确但答案不匹配时的基础得分详见第五节。3.1 使用本地检索将search_backend设为local在local节配置本地检索服务器地址默认http://127.0.0.1:8000/retrieve在启动训练前先启动本地检索服务附录 Step 4。请求链路为generate_with_search.py中的search()→ local_search_server.py 的local_search()→ 本地检索服务器的/retrieve接口。local_search()通过 aiohttp 异步发送{queries: [query], topk: top_k, return_scores: False}把响应结果统一格式化为{document: {contents: title\ntext}}从而与 Google 后端保持完全一致的输出契约。3.2 使用 Google 搜索将search_backend设为google在google节填入 serper.dev API Key在 serper.dev 官网申请 API Key 后填入api_key字段。请求链路为search()→ google_search_server.py 的google_search()→https://google.serper.dev/search。google_search()支持两种模式snippet_onlyTrue时只解析organic结果中的 title 与 snippetsnippet_onlyFalse时还会并发抓取链接正文再通过collect_context把 snippet 匹配到的上下文段落拼接为检索证据。无论哪种后端最终检索结果都会经过_passages2string()格式化为带Doc N(Title: ...)前缀的文本作为环境反馈注入对话。四、奖励函数与序列格式协议奖励函数实现在 generate_with_search.py 的reward_func()中核心逻辑委托给 qa_em_format.py 的compute_score_em()async def reward_func(args, sample, **kwargs): if not isinstance(sample, Sample): raise TypeError(Sample must be an instance of Sample class.) score compute_score_em( solution_strsample.prompt sample.response, ground_truthsample.label[ground_truth], format_scoreSEARCH_R1_CONFIGS[format_score], ) return scorecompute_score_em的打分逻辑分为三层格式校验is_valid_sequence以状态机方式严格校验对话结构是否满足think → search → information → (think → search → information → ...) → answer的交替序列标签必须平衡、标签之间不允许出现游离文本。非法格式直接给 0 分。答案抽取extract_solution取最后一个answer标签内的内容作为最终答案未出现或只出现一次answer时视为无答案。EM 判定em_check对预测答案与 golden answers 做归一化去冠词、去标点、统一空白、转小写后进行精确匹配。最终得分规则format_score0.2、score1.0时情况得分答案 EM 匹配且格式合法1.0答案 EM 匹配但格式非法0.81.0 - 0.2无答案/答案不匹配但格式合法且检索内容命中 golden answer0.3无答案/答案不匹配但格式合法0.2答案不匹配且格式非法0.1格式非法且无答案0值得一提的细节compute_score_em还会打印1/64 概率抽样golden answers、抽取的答案与完整 solution string便于调试时人工抽查。五、多轮检索对话的生成逻辑custom-generate-function-pathslime 允许通过启动脚本中的两个配置项注入自定义逻辑Search-R1 示例只依赖这两项即可实现完整的多轮工具调用CUSTOM_ARGS( --custom-generate-function-path generate_with_search.generate --custom-rm-path generate_with_search.reward_func )它们分别对应 generate_with_search.py 中的generate与reward_func两个函数generate_with_search模块通过 runtime env 中的PYTHONPATH暴露给训练进程详见第七节启动脚本解析。5.1 生成主流程generate()是核心的数据生成函数签名与 slime rollout 框架的标准生成接口一致args, sample, sampling_params。其执行流程如下创建GenerateState(args)来自 slime/rollout/sglang_rollout.py用于持有 tokenizer 与采样状态当前实现不支持 partial rollout入口处直接断言。将 prompt 文本 tokenize 后写入sample.tokens随后进入最多max_turns轮的多轮循环每轮向 sglang 推理服务发送POST /generate请求text prompt response附带sampling_params。根据finish_reason处理结果abort直接标记样本ABORTED返回length表示触达长度上限终止多轮循环。对模型输出执行postprocess_predictions用正则(search|answer)(.*?)/\1解析本轮动作与内容交给execute_predictionssearchquery/search并发受限地发起搜索受search_concurrency信号量约束将检索结果包装为\n\ninformation.../information\n\n作为环境反馈doneFalseanswer.../answer对话结束doneTrue其他输出注入我的上一个动作无效应把查询放在search与/search之间……的纠正提示doneFalse。环境反馈observationtoken 追加进响应序列但loss_mask 置 0即这些 token 不参与策略梯度计算模型自身生成的 token 则loss_mask1参与训练。这与多轮 RL 中工具输出不作为训练目标的标准做法一致mis.py 的注释中也明确说明多轮 RL 中工具响应在 loss_mask 中标记为 0。收尾时将完整响应、response_length、loss_mask、rollout_log_probs开启时写回sample并按最终finish_reason设置Sample.StatusTRUNCATED/ABORTED/COMPLETED。5.2 两个关键的实现细节停止标签修正stop tags。代码在每轮推理前强制注入[/search, /answer]停止词_stop_tags [/search, /answer] _existing_stop sampling_params.get(stop) or [] if isinstance(_existing_stop, str): _existing_stop [_existing_stop] sampling_params {**sampling_params, stop: list(dict.fromkeys([*_existing_stop, *_stop_tags]))}代码注释说明了动机若不在此处停止sglang 会在/search//answer之后继续吐出多余 token甚至是编造的Question:在return_logprobTrue时由于禁止后处理去尾这些垃圾 token 会被计入loss_mask1参与训练并破坏is_valid_sequence的格式校验导致奖励降低。slime 设置了no_stop_trimTrue因此闭合标签本身会保留在输出中。logp 对齐断言。开启 logprob 收集时observation token 会补 0.0 占位 logp并在每轮追加后断言 token 数与 logp 数严格一致assert len(response_token_ids) len( rollout_log_probs ), fToken/logp length mismatch: {len(response_token_ids)} tokens vs {len(rollout_log_probs)} logps这是后续 TIS 权重计算训练侧 logp 与 rollout 侧 logp 逐 token 相减正确性的前提。5.3 后处理与 logprob 的互斥关系postprocess_responses()只保留到最后一个完整闭合标签为止截到/search或/answer。但注释强调需要收集 logp 时绝对不能做任何字符串后处理——因为无法知道如何同步截断 token/logp 数组且对后处理文本重新 tokenize 可能产生与推理引擎不同的 token 序列导致 token 与 logp 错位。因此该函数仅在return_logprobFalse时启用。六、启用 TIS轨迹重要性采样TISTrajectory Importance Sampling通过训练侧与 rollout 侧 logp 之比修正策略梯度缓解推理引擎旧策略生成的数据与训练新策略之间的分布偏移。slime 将 TIS 相关能力收敛在 examples/train_infer_mismatch_helper/ 模块中。启用分两步6.1 第一步打开 logprob 收集在 generate_with_search.py 中确保SEARCH_R1_CONFIGS { # ... other configs return_logprob: True, # Must be True for TIS }6.2 第二步启动脚本取消注释在 run_qwen2.5_3B.sh 中取消GRPO_ARGS里的 TIS 开关GRPO_ARGS( --advantage-estimator grpo --use-kl-loss --kl-loss-coef 0.001 --kl-loss-type low_var_kl --entropy-coef 0.00 --eps-clip 0.2 --eps-clip-high 0.28 # Uncomment to enable TIS --use-tis )同时在CUSTOM_ARGS中取消注释 TIS 配置路径CUSTOM_ARGS( --custom-generate-function-path generate_with_search.generate --custom-rm-path generate_with_search.reward_func # Uncomment to enable TIS --custom-config-path examples/train_infer_mismatch_helper/mis.yaml --custom-tis-function-path examples.train_infer_mismatch_helper.mis.compute_mis_weights_with_cp )其中--custom-config-path指向 mis.yaml--custom-tis-function-path指向 mis.py 的compute_mis_weights_with_cp该函数在 slime/utils/arguments.py 中与--use-tis存在依赖校验。6.3 TIS 配置项解析mis.yamlmis.yaml 的关键参数参数取值示例含义use_tis/use_rstrue / true是否启用重要性采样与拒绝采样tis_leveltoken / sequence / geometricIS 权重的聚合粒度逐 token、序列连乘、几何平均tis_modetruncate / mask / clipIS 权重处理方式截断TIS、区间外清零MIS、区间裁剪CIStis_lower_bound/tis_upper_bound0.5 / 2.0权重下/上界truncate 模式不用下界rs_veto_threshold1.0e-4逐 token veto 阈值任一 token 比值低于该值则整条序列权重置零不再产生梯度注意必须写1.0e-4带小数点的浮点格式tis_batch_normalizetrue批内归一化将 IS 权重均值归一为 1.0降低梯度方差底层实现位于 mis.pycompute_mis_weights_with_cp先用 Megatron 的all_gather_with_cp/slice_log_prob_with_cp完成上下文并行下的 logp 汇聚与切片再调用compute_mis_weights计算权重含 SAFETY_BOUND20 的 exp 溢出防护、token/sequence/geometric 三种粒度、truncate/clip/mask 三种模式、veto 掩码、批归一化并输出training_ppl、kl、chi2等丰富的统计指标供监控。重要注意事项TIS 必须满足return_logprobTrue收集 logprob 时响应后处理自动禁用以保证 token/logp 对齐TIS 会引入额外计算开销但可提升训练效率分布偏移较大时收益更明显。七、训练启动脚本全解析完成上述配置后在容器内执行cd slime/ bash examples/search-r1/run_qwen2.5_3B.shrun_qwen2.5_3B.sh 的完整结构如下1. 清理与基础环境开头pkill -9 sglang、ray stop --force、pkill -9 ray/python用于保证可重复运行PYTHONUNBUFFERED1防止 ray 缓冲输出。2. 模型参数source scripts/models/qwen2.5-3B.sh引入MODEL_ARGS36 层、2048 hidden、GQA 2 组、RMSNorm 等见 scripts/models/qwen2.5-3B.sh。3. 检查点参数CKPT_ARGS--hf-checkpoint /root/Qwen2.5-3B/rollout 侧 HF 权重、--ref-load /root/Qwen2.5-3B_torch_dist/训练侧 Megatron 权重--load/--save/--save-interval默认注释用于断点续训时打开。4. 数据与 rollout 参数ROLLOUT_ARGS参数示例值说明--prompt-data/root/Search-R1/data/nq_hotpotqa_train/train.parquet合并后的训练数据--input-key prompt/--label-key reward_model-数据列映射prompt 列与奖励标签列--apply-chat-template-应用对话模板--rollout-shuffle/--num-rollout 3000-打乱样本、rollout 总数--rollout-batch-size 32/--n-samples-per-prompt 8-每批 32 条 prompt每条采样 8 次--rollout-max-response-len 512/--rollout-temperature 1-响应长度上限与采样温度--global-batch-size 256/--balance-data-训练全局 batch 与数据均衡评测相关参数--eval-interval、--eval-prompt-data nq_test ...等默认注释需要时打开并按[0:3000]切片方式指定评测集。5. 并行与性能参数PERF_ARGS--tensor-model-parallel-size 2--sequence-parallel2 路张量并行、其余并行度均为 1--recompute-granularity full --recompute-method uniform --recompute-num-layers 1启用激活重计算--use-dynamic-batch-size --max-tokens-per-gpu 9216动态 batch 控制显存占用。6. GRPO 参数GRPO_ARGS--advantage-estimator grpo、--use-kl-loss --kl-loss-coef 0.001 --kl-loss-type low_var_kl、--entropy-coef 0.00、--eps-clip 0.2 --eps-clip-high 0.28非对称裁剪允许正向优势略大的更新空间以及默认注释的--use-tis。7. 优化器与杂项Adam--lr 1e-6、常数衰减、weight-decay 0.01--attention-dropout 0.0 --hidden-dropout 0.0关闭 dropoutMegatron 默认 0.1--accumulate-allreduce-grads-in-fp32、--attention-softmax-in-fp32保证数值稳定性--attention-backend flash注释提示使用 MLA 结构的模型时需要改为其他 backend。8. 推理引擎参数SGLANG_ARGS--rollout-num-gpus-per-engine 2每个 sglang 推理引擎使用 2 张 GPU、--sglang-mem-fraction-static 0.7sglang 静态显存占用比例。9. 自定义参数CUSTOM_ARGS见第五节与第六节是 Search-R1 功能的核心挂载点。10. Ray 启动与任务提交以MASTER_ADDR默认127.0.0.1启动 ray head--num-gpus 8通过 runtime env 注入PYTHONPATH/root/Megatron-LM/:${SCRIPT_DIR}与CUDA_DEVICE_MAX_CONNECTIONS1SCRIPT_DIR即examples/search-r1目录这正是generate_with_search模块可被导入的原因最后ray job submit提交训练任务采用--colocate共置部署4 张 GPU 训练 4 张 GPU rollout。八、附录搭建本地稠密检索服务当search_backendlocal时需要独立搭建本地检索服务器。本地检索器对 GPU 依赖较强且依赖版本与训练环境不同因此官方强烈建议使用独立的 conda 环境避免与训练环境冲突。Step 1安装 Conda如已安装可跳过# Download and install conda wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh -O ~/miniconda.sh bash ~/miniconda.sh -b -p $HOME/miniconda3 source ~/miniconda3/etc/profile.d/conda.sh conda init source ~/.bashrc # Accept conda terms of service conda tos accept --override-channels --channel https://repo.anaconda.com/pkgs/main conda tos accept --override-channels --channel https://repo.anaconda.com/pkgs/rStep 2创建检索器环境# Create environment conda create -n retriever python3.10 -y conda activate retriever # Install PyTorch with CUDA support conda install pytorch2.4.0 torchvision0.19.0 torchaudio2.4.0 pytorch-cuda12.1 -c pytorch -c nvidia -y # Install required packages pip install transformers datasets pyserini huggingface_hub conda install faiss-gpu1.8.0 -c pytorch -c nvidia -y pip install uvicorn fastapi其中pyserini用于 BM25 检索retrieval_server.py中BM25Retriever依赖faiss-gpu与transformers用于 e5 稠密检索DenseRetrieverEncoderuvicorn/fastapi提供 HTTP 服务。Step 3下载索引与语料注意本地检索文件体积较大下载约需 60-70 GB 磁盘解压后约 132 GB请预留充足磁盘空间。# Set your save path save_path/root/Index # Download the index and corpus files python /root/slime/examples/search-r1/local_dense_retriever/download.py --save_path $save_path # Combine split index files cat $save_path/part_* $save_path/e5_Flat.index # Decompress the corpus gzip -d $save_path/wiki-18.jsonl.gzdownload.py 从PeterJinGo/wiki-18-e5-index下载part_aa、part_ab两个分片并从PeterJinGo/wiki-18-corpus下载wiki-18.jsonl.gz。分片需用cat合并为e5_Flat.index语料解压为wiki-18.jsonl。Step 4启动本地检索服务器# If you encounter conda not found error, run: # source ~/miniconda3/etc/profile.d/conda.sh # conda init # source ~/.bashrc # Activate retriever environment conda activate retriever # Set paths save_path/root/Index index_file$save_path/e5_Flat.index corpus_file$save_path/wiki-18.jsonl retriever_namee5 retriever_pathintfloat/e5-base-v2 # Start the retrieval server python /root/slime/examples/search-r1/local_dense_retriever/retrieval_server.py \ --index_path $index_file \ --corpus_path $corpus_file \ --topk 3 \ --retriever_name $retriever_name \ --retriever_model $retriever_path \ --faiss_gpuretrieval_server.py 的实现要点支持两种检索器BM25Retrieverpyserini 稀疏检索与DenseRetrieverFAISS 稠密检索由--retriever_name决定bm25走稀疏其余走稠密DenseRetriever启动时读取 FAISS 索引--faiss_gpu时用index_cpu_to_all_gpus分发到全部 GPU并使用 FP16 量化加载 e5 编码器与语料库Encoder.encode会对 e5 模型自动加query: /passage: 前缀、对 bge 模型加指令前缀支持 mean/cls/pooler 多种池化方式对外暴露POST /retrieve接口FastAPI接收{queries: [...], topk: N, return_scores: bool}返回{result: [[{document: ..., score: ...}], ...]}服务默认监听0.0.0.0:8000与SEARCH_R1_CONFIGS[local][search_url]的默认值对应。运行注意首次启动会下载模型并加载索引耗时数分钟正常启动不含下载约 1-2 分钟每张 GPU 显存占用约 5-7 GB本地检索服务的 Python 进程在 shell 关闭后不会自动退出重启服务lsof -i :8000找到 PID 后 kill 再重新启动。Step 5启动训练训练前务必退出检索器 conda 环境若处于其中则执行conda deactivate回到 slime 基础环境cd /root/slime # Set your wandb key (optional) export WANDB_KEYyour_wandb_key_here # If ray process is stuck, try: # rm -rf /root/.cache # rm -rf /root/.* # Run the training script bash /root/slime/examples/search-r1/run_qwen2.5_3B.sh常见问题排查Ray 进程卡死执行rm -rf /root/.cache仍卡死则rm -rf /root/.*清理全部缓存后重试。conda 环境冲突确认训练前已conda deactivate退出 retriever 环境训练必须使用基础 Python 环境。检索服务无响应用lsof -i :8000确认服务是否存活用nvidia-smi检查 GPU 可用性并查看服务端日志定位错误。九、小结examples/search-r1完整展示了在 slime 中实现检索增强 多轮工具调用RL 训练的最小闭环自定义generate函数驱动search/answer工具协议与环境反馈注入自定义reward_func基于格式校验与 EM 打分给出奖励信号SEARCH_R1_CONFIGS统一管理本地/Google 两种检索后端--use-tis配合mis.yaml引入轨迹重要性采样以缓解训练-推理分布偏移。理解这条链路后你可以将其迁移到其他需要工具调用、多轮交互或外部知识接入的 RL 任务中——只需按同一模式实现自己的generate与reward_func并在CUSTOM_ARGS中挂载即可。【免费下载链接】slimeslime is an LLM post-training framework for RL Scaling.项目地址: https://gitcode.com/GitHub_Trending/slime12/slime创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进