search_tool.py — verl/tools/search_tool.py¶
文件路径¶
verl/tools/search_tool.py
文件概述¶
SearchTool 是一个搜索工具,允许 LLM 在训练过程中调用外部检索服务来搜索信息。它是 verl 框架中实现"Search-augmented RL"(搜索增强的强化学习)的关键组件,类似于 DeepSeek-R1 的搜索能力。
应用场景:当 LLM 在解决需要外部知识的任务时(如开放域问答),可以通过这个工具查询外部知识库。
关键代码讲解¶
1. 并发框架¶
search_tool.py 包含与 sandbox_fusion_tools.py 几乎相同的并发框架:
class PoolMode(Enum):
ThreadMode = 1
ProcessMode = 2
@ray.remote(concurrency_groups={"acquire": 1, "release": 10})
class TokenBucketWorker:
# 令牌桶速率限制器(同 sandbox_fusion_tools.py)
class SearchExecutionWorker:
def execute(self, fn, *fn_args, **fn_kwargs):
if self.rate_limit_worker:
with ExitStack() as stack:
stack.callback(self.rate_limit_worker.release.remote)
ray.get(self.rate_limit_worker.acquire.remote())
try:
return fn(*fn_args, **fn_kwargs)
except Exception as e:
logger.warning(f"Error when executing search: {e}")
else:
return fn(*fn_args, **fn_kwargs)
与沙箱工具的区别是:SearchExecutionWorker 支持可选的速率限制(if self.rate_limit_worker 判断),不开启速率限制时直接执行。
2. SearchTool 构造函数¶
class SearchTool(BaseTool):
def __init__(self, config, tool_schema):
super().__init__(config, tool_schema)
self._instance_dict = {}
self.num_workers = config.get("num_workers", 120)
self.rate_limit = config.get("rate_limit", 120)
self.timeout = config.get("timeout", 30)
self.retrieval_service_url = config.get("retrieval_service_url")
self.topk = config.get("topk", 3)
关键配置:
- num_workers/rate_limit 默认 120,比沙箱工具高很多(搜索通常更快)。
- retrieval_service_url:外部检索服务的 URL(必填)。
- topk:每次搜索返回的文档数量,默认 3。
3. 执行搜索¶
def execute_search(self, instance_id, query_list, retrieval_service_url, topk, timeout):
result_text, metadata = perform_single_search_batch(
retrieval_service_url=retrieval_service_url,
query_list=query_list,
topk=topk,
concurrent_semaphore=None,
timeout=timeout,
)
return result_text, metadata
@rollout_trace_op
async def execute(self, instance_id, parameters, **kwargs):
query_list_from_params = parameters.get("query_list")
if not query_list_from_params or not isinstance(query_list_from_params, list):
error_msg = "Error: 'query_list' is missing, empty, or not a list in parameters."
return ToolResponse(text=json.dumps({"result": error_msg})), 0.0, {}
result_text, metadata = await self.execution_pool.execute.remote(
self.execute_search, instance_id, query_list_from_params,
self.retrieval_service_url, self.topk, timeout
)
self._instance_dict[instance_id]["reward"].append(result_text.strip())
metrics = {
"query_count": metadata.get("query_count", 0),
"status": metadata.get("status", "unknown"),
"total_results": metadata.get("total_results", 0),
"api_request_error": metadata.get("api_request_error"),
}
return ToolResponse(text=result_text), 0.0, metrics
执行流程:
1. 从参数中获取 query_list(搜索关键词列表)。
2. 通过 Ray 执行池调用 execute_search,进而调用 perform_single_search_batch。
3. 将结果存入实例字典,返回响应和度量指标。
4. 搜索工具的即时奖励始终为 0.0(搜索本身不产生奖励)。
核心类/函数列表¶
| 类/函数 | 作用 |
|---|---|
TokenBucketWorker |
速率限制器(Ray Actor) |
SearchExecutionWorker |
搜索任务执行器 |
init_search_execution_pool |
初始化搜索执行池 |
SearchTool |
搜索工具主类 |
execute_search |
调用外部检索服务 |
与其他模块的关系¶
- 继承自
BaseTool。 - 依赖
utils/search_r1_like_utils.py中的perform_single_search_batch函数。 - 对比
MCPSearchTool:SearchTool直接通过 HTTP 调用检索服务,而MCPSearchTool通过 MCP 协议调用。 - 通过
tool_registry.py被注册和初始化。
小结¶
SearchTool 为 LLM 提供了搜索外部知识的能力。它通过 Ray 并发框架管理大量并发搜索请求,并支持速率限制以保护检索服务。搜索结果会被格式化后返回给 LLM,帮助其做出更好的回答。