跳转至

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,帮助其做出更好的回答。