ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

verl.single_controller 设计解析:用 WorkerGroup 把分布式 RL 训练写得像单进程

verl.single_controller 设计解析:用 WorkerGroup 把分布式 RL 训练写得像单进程 verl.single_controller 设计解析用 WorkerGroup 把分布式 RL 训练写得像单进程【免费下载链接】verlverl/HybridFlow: A Flexible and Efficient RL Post-Training Framework项目地址: https://gitcode.com/GitHub_Trending/ve/verl本文是 verl 框架HybridFlow中verl.single_controller模块的设计指南。该模块是 verl 分布式执行层的核心抽象它让开发者可以用与单进程几乎一致的代码编写多进程分布式 RL 训练程序同时保留中间张量可检查、多 DAG 可表达的灵活性。读完本文你将理解WorkerGroup、ResourcePool、ClassWithInitArgs三大组件如何协同工作掌握register装饰器到动态方法绑定的完整调用链并学会如何将这一模式泛化到 RL 后训练之外的分布式计算场景。本文面向 verl 的开发者与贡献者而非最终用户内容以 docs/single_controller.rst 为骨架并补充当前仓库源码中的实现细节作为佐证。起源为什么 verl 需要一套自己的分布式抽象single_controller模块起源于一个朴素的需求把一个玩具级单进程 RLHF 脚本改造成分布式系统要求改动尽可能小、调试尽可能容易。常见的做法如 PyTorch DDP通常是包装nn.Module并启动多个进程让每个 rank 执行同一函数。但在分布式 RLHF 场景下这种做法有两个明显的局限难以表达 PPO 所需的多个 DAGPPO 训练流程中存在 Actor、Critic、Ref 等多个模型各自的执行图例如generate_sequences、compute_advantages等阶段它们之间存在复杂的数据流转与控制依赖DDP 式的所有进程跑同一份代码模型无法自然表达这种多 DAG 结构难以检查中间张量分布式训练中中间张量分散在各 rank 上直接检查、调试非常困难。为了保持可调试性verl 采用了另一条路线——把训练循环拆解成定义良好的阶段如generate_sequences、compute_advantages而不是让整份脚本在每个进程上重复执行。在选择后端时verl 最初选择了 RayRay 能把 Python 类方法直接暴露为 RPC 端点。但 Ray 的默认模型是一次方法调用 一次 RPC而训练大模型通常需要跨多个进程协同这就产生了一个矛盾如何在保持调用一个方法的用户体验的同时把这一调用真正分发到一组 Ray actor 上执行为此verl 引入了三个关键组件见 verl/single_controller/base/init.pyWorkerGroup管理一组远程 worker为多进程分布式计算提供统一接口ResourcePool把计算资源绑定到 worker 进程上ClassWithInitArgs支持延迟远程实例化——把类 构造参数打包稍后在远端真正构造对象。三大核心组件概览WorkerGroup分布式调用的统一门面WorkerGroup见 worker_group.py是一次调用、多进程执行的抽象核心。它维护self._workers列表并通过world_size属性暴露 worker 数量即len(self._workers)。其最关键的方法是_bind_worker_method会在初始化阶段把 worker 类上被register装饰的方法动态绑定到 WorkerGroup 实例上下文详述。此外它还提供了基础的进程健康管理能力start_worker_aliveness_check会启动一个后台线程周期性调用_is_worker_alive检查每个 worker一旦发现 worker 死亡即向主线程发送SIGABRT对应check_workers_alive函数避免训练在静默状态下进行。ResourcePool进程-资源绑定ResourcePool见 worker_group.py本质上是每节点进程数的列表process_on_nodes。它提供了三个核心属性world_size所有节点上的进程总数sum(self._store)local_world_size_list()展平后每个进程的本地 world sizelocal_rank_list()展平后每个进程的本地 rank。在 Ray 后端中RayResourcePool见 ray/base.py扩展了它通过get_placement_groups把进程计数转换为 Ray placement group 的 bundle 方案每个 bundle 包含{CPU: max_colocate_count}以及一张 GPU/NPU从而把每节点进程数落地为真实的资源调度。ClassWithInitArgs延迟实例化ClassWithInitArgs见 worker_group.py只做一件事把类 构造参数存起来__call__时才真正执行self.cls(*self.args, **self.kwargs)。在分布式场景中这意味着实例化可以发生在远端主进程先打包好构造蓝图Ray actor 创建时再真正执行构造函数。其 Ray 版本RayClassWithInitArgs增加了 placement group 调度、GPU 资源申请、共享资源sharing_with用于把多个 worker 放在同一节点/同一 GPU 上等能力见 ray/base.py。一个运行示例generate_sequences的完整链路下面以ActorRolloutRefWorker.generate_sequences为例走一遍注册 → 绑定 → 调用的全过程。Step 1用register装饰器注册首先在 worker 类中定义方法并用register装饰它因为该方法会在 driver 脚本中被调用class ActorRolloutRefWorker(Worker): ... register(dispatch_modeDispatch.DP_COMPUTE_PROTO) def generate_sequences(self, prompts: DataProto): prompts prompts.to(torch.cuda.current_device()) ...register装饰器见 decorator.py本身不改变函数行为只做两件事通过functools.wraps包装原函数可选地在调用前_materialize_futures把DataProtoFuture物化为真实数据把dispatch_mode、execute_mode、blocking三个配置作为字典挂到方法上一个魔法属性MAGIC_ATTR值为attrs_3141562937上def register(dispatch_modeDispatch.ALL_TO_ALL, execute_modeExecute.ALL, blockingTrue, materialize_futuresTrue): ... def decorator(func): wraps(func) def inner(*args, **kwargs): if materialize_futures: args, kwargs _materialize_futures(*args, **kwargs) return func(*args, **kwargs) attrs {dispatch_mode: dispatch_mode, execute_mode: execute_mode, blocking: blocking} setattr(inner, MAGIC_ATTR, attrs) return inner return decorator可见dispatch_mode、execute_mode、blocking三个值被附加到了generate_sequences方法上。Dispatch与Execute都是动态枚举DynamicEnum预定义的分发模式包括RANK_ZERO、ONE_TO_ALL、ALL_TO_ALL、DP_COMPUTE、DP_COMPUTE_PROTO、DP_COMPUTE_PROTO_WITH_FUNC、DP_COMPUTE_METRIC以及供 vLLM 外部执行器使用的DIRECT_ROLLOUT_METHOD见 decorator.py。Step 2初始化时绑定方法这些附加属性会在 worker 类被包装进RayClassWithInitArgs并传入RayWorkerGroup时被提取和利用ray_cls_with_init RayClassWithInitArgs(clsray.remote(ActorRolloutRefWorker), configconfig, rolerollout) resource_pool RayResourcePool(process_on_nodes[config.trainer.n_gpus_per_node] * config.trainer.nnodes) wg RayWorkerGroup(resource_poolresource_pool, ray_cls_with_initray_cls_with_init)在RayWorkerGroup的初始化过程见 ray/base.py中发生两个关键步骤创建 worker 实例Ray actor_init_with_resource_pool根据 placement group 逐节点、逐 local rank 调用_create_worker创建 actor并为每个 worker 注入WORLD_SIZE、RANK、MASTER_ADDR、MASTER_PORT、WG_BACKEND等环境变量见 ray/base.py把被register装饰的方法绑定到RayWorkerGroupself._bind_worker_method(self.ray_cls_with_init.cls, func_generator)见 ray/base.py。方法绑定是整个verl.single_controller的心脏。核心函数是WorkerGroup._bind_worker_method见 worker_group.py其工作流程如下def _bind_worker_method(self, user_defined_cls, func_generator): ... for method_name in dir(user_defined_cls): try: method getattr(user_defined_cls, method_name) assert callable(method) except Exception: continue # 跳过属性property ... if hasattr(method, MAGIC_ATTR): attribute getattr(method, MAGIC_ATTR) dispatch_mode attribute[dispatch_mode] execute_mode attribute[execute_mode] blocking attribute[blocking] ...当方法带有MAGIC_ATTR时提取register设置的属性然后按顺序解析三样东西dispatch_fn与collect_fn从DISPATCH_MODE_FN_REGISTRY见 decorator.py中根据dispatch_mode查表得到。同时支持自定义若dispatch_mode是字典含dispatch_fn/collect_fn键则直接使用用户提供的函数对execute_fn根据execute_mode通过get_predefined_execute_fn解析出执行函数名——Execute.ALL对应execute_allExecute.RANK_ZERO对应execute_rank_zero——再用getattr(self, ...)从 WorkerGroup 实例上取到真实方法用func_generator动态生成新方法并setattr到 WorkerGroup 实例上# 绑定一个新方法到 RayWorkerGroup func func_generator( self, method_name, dispatch_fndispatch_fn, collect_fncollect_fn, execute_fnexecute_fn, blockingblocking, ) setattr(self, method_name, func)Ray 后端对应的func_generator见 ray/base.py定义了被绑定方法的实际执行语义def func_generator(self, method_name, dispatch_fn, collect_fn, execute_fn, blocking): class Functor: def __call__(this, *args, **kwargs): args, kwargs dispatch_fn(self, *args, **kwargs) # 1. 分发 padding_count kwargs.pop(_padding_size_key, 0) output execute_fn(method_name, *args, **kwargs) # 2. 远程执行 if blocking: output ray.get(output) # 3. 阻塞等待 output collect_fn(self, output) # 4. 收集 if padding_count 0: # 5. 去除补齐部分 if isinstance(output, DataProto): indices [i for i in range(len(output))][:-padding_count] output output.select_idxs(indices) elif isinstance(output, list): output output[:-padding_count] return output return type(method_name, (Functor,), {})()这段代码浓缩了整套设计分发dispatch→ 远程执行execute→ 收集collect→ 去 padding。动态生成的方法以type(method_name, (Functor,), {})()的方式命名类兼顾了可观测性Ray 中方法名可读。至此generate_sequences就可以通过WorkerGroup接口被调用了。ONE_TO_ALL 与 DP_COMPUTE_PROTO 的分发差异dispatch_mode关联一对dispatch_fn/collect_fn。dispatch_fn负责在WorkerGroup中处理输入参数生成每个 worker 一份的输入批次。ONE_TO_ALL的dispatch_fn是dispatch_one_to_all见 decorator.py它把所有输入参数复制 N 份N 等于 worker_group 的world_sizedef dispatch_one_to_all(worker_group, *args, **kwargs): args tuple([arg] * worker_group.world_size for arg in args) kwargs {k: [v] * worker_group.world_size for k, v in kwargs.items()} return args, kwargsDP_COMPUTE_PROTO的dispatch_fn是dispatch_dp_compute_data_proto见 decorator.py它利用DataProto.chunk把一个大DataProto按 batch 维切分成 N 个小DataProtodef dispatch_dp_compute_data_proto(worker_group, *args, **kwargs): from verl.single_controller.base.worker_group import WorkerGroup assert isinstance(worker_group, WorkerGroup) # Note: enable auto padding for dp compute DataProto splitted_args, splitted_kwargs _split_args_kwargs_data_proto_with_auto_padding( worker_group.world_size, *args, **kwargs, ) return splitted_args, splitted_kwargs这里值得特别说明的是_split_args_kwargs_data_proto_with_auto_padding见 decorator.py实现的**自动补齐auto padding**逻辑若DataProto开启了 paddingis_padding_enabled()且 batch 长度不能被world_size整除则计算padding_size chunks - (data_proto_len % chunks)对同一批参数中的所有DataProto要求长度一致assert data_proto_len len(obj)并统一调用obj.padding(padding_sizepadding_size)补齐后再chunk补上的 padding 数量通过特殊 key_padding_size_key塞进 kwargs由func_generator在收集完成后用select_idxs或列表切片去掉。collect_fn则遵循相反的模式把各 worker 返回值的列表合并——collect_all_to_all原样返回列表collect_dp_compute_data_proto则用DataProto.concat拼接成一个大DataProto见 decorator.py。以generate_sequences为例最终的绑定结果是dispatch_mode Dispatch.DP_COMPUTE_PROTOdispatch_fn dispatch_dp_compute_data_protocollect_fn collect_dp_compute_data_protoexecute_fn RayWorkerGroup.execute_allStep 3调用链——分布式调用与单进程调用无差别上面所有机制的目的是让分布式调用在写法上与单进程完全一致。原始单进程脚本中rollout Rollout() rollout.generate_sequences(batch)在 verl 中多进程程序变成rollout RayWorkerGroup(resource_pool[4], RayClassWithInitArgs(Rollout)) rollout.generate_sequences(batch)在这行看似简单的调用背后dispatch_fn把输入大DataProto按 DP 切分并自动补齐分发给各 workerexecute_fnexecute_all对每个 worker 发起真正的远程调用——Ray actor 方法调用返回 ObjectRef 列表blockingTrue时调用ray.get等待结果见 ray/base.pycollect_fn收集各 worker 的返回值并拼接最后去掉自动补齐的部分。所有这一切都被抽象掉了开发者可以用几乎零改动的方式把单进程逻辑升级为分布式执行这正是 HybridFlow 追求单进程与多进程脚本差距最小化的具体体现。源码级深入DataProto 与执行细节DataProto切分与拼接的原语DP_COMPUTE_PROTO模式之所以高效底层依赖 verl 的自有数据结构DataProto定义于 verl/protocol.py。几个与本主题直接相关的方法chunk(chunks)沿 batch 维dim0把DataProto切分成若干份meta_info会随每个分片传递protocol.pyconcat(data)把一组DataProto沿 batch 维拼接meta_info合并并对不同 worker 的 metric 做特殊处理protocol.pypadding(padding_size)通过拼接 padding 候选来补齐 batch保证可被 DP 均分protocol.pyselect_idxs(idxs)按索引选择子集用于收集后剔除 padding 部分protocol.py。另外协议层还提供了DataProtoFutureprotocol.py它持有多个 rayObjectRef并在get()时才真正取回数据。register的materialize_futuresTrue参数默认值保证在分发前把DataProtoFuture物化为真实DataProto见_materialize_futuresdecorator.py这为后续在 fully async 等异步流水线中做未来即参数的优化留出了接口。execute_all 与 execute_rank_zeroRayWorkerGroup提供了两组执行原语execute_all含同步版execute_all_sync对每个 worker 发起调用返回 ObjectRef 列表。它有一个实用约定如果所有位置参数与关键字参数都是 list 且长度恰好等于 worker 数则按第 i 个元素给第 i 个 worker的方式逐元素分发见 ray/base.pyexecute_rank_zero含同步版execute_rank_zero_sync只把方法调用发给 rank 0 worker典型用于Execute.RANK_ZERO执行模式见 ray/base.py。blocking参数由register控制blockingTrue默认时绑定方法内部自动执行ray.get同步等待设为False时返回 ObjectRef 列表供上层实现流水线并行/异步调度。Worker 基类分布式进程的环境契约所有被WorkerGroup管理的 worker 都继承自Worker基类见 worker.py。它在构造时从环境变量读取分布式上下文WORLD_SIZE、RANK、LOCAL_WORLD_SIZE、LOCAL_RANK、MASTER_ADDR、MASTER_PORT以及可见设备变量CUDA/HIP/ROCR 设备键并统一设置os.environ_configure_with_store。这些环境变量正是_create_worker注入的见 ray/base.py从而保证每个 Ray actor 内都能看见自己在大分布式环境中的身份。Worker还预置了execute_with_func_generatorDP_COMPUTE_PROTO_WITH_FUNC模式把函数本身与数据一起分发和execute_func_rank_zeroExecute.RANK_ZERO等被注册方法以及用于 ND 网格如 Megatron 的 TP/DP/PP/CP的 dispatch/collect 信息注册与查询接口_register_dispatch_collect_info、_query_dispatch_info为LazyCompute等更复杂的多维网格分发dispatch_nd_compute_dataproto见 decorator.py奠定了基础。超越 RL 后训练single_controller 的泛化能力verl.single_controller的价值并不仅限于强化学习。它提供了批量处理远程方法调用 自动输入输出处理的干净抽象核心机制与领域无关任何类方法只要被register装饰并描述好分发/执行/阻塞语义就能自动获得分布式执行能力任何数据结构只要实现chunk/concat协议如DataProto、DataProtoFuture就能参与DP_COMPUTE*系列分发。由此WorkerGroup把单进程脚本与多进程脚本之间的鸿沟压到最小为更广泛的分布式计算打开了大门——只要你的计算可以被描述为对一批 worker 调用方法并收集结果。从当前仓库还可以看到该抽象的进一步演化可作为扩展方向的参考colocated worker进程共置create_colocated_worker_clsray/base.py与升级版create_colocated_worker_cls_fusedray/base.py可以把多个 worker 类如 Actor、Critic、Ref、Rollout融合进同一个 Ray actorFusedWorker通过spawn/fuse在共享进程内拆分出各自带前缀的 WorkerGroup极大减少进程数与跨进程通信开销资源池切分与合并split_resource_pool与merge_resource_poolray/base.py支持把一个大资源池按需切分为SubRayResourcePool或合并多个池适配 FSDP/Megatron 等不同并行后端对进程组织方式的不同需求自定义分发模式通过register_dispatch_mode/update_dispatch_modedecorator.py任何第三方都可以注册新的dispatch_fn/collect_fn组合扩展分发语义而无需改动框架核心。这些能力共同说明single_controller是一个以方法绑定 分发/收集为核心的可扩展分布式执行框架RL 后训练只是它最具代表性的应用场景之一。小结设计动机DDP 无法自然表达 PPO 的多 DAG 结构也难以检查中间张量verl 选择拆分训练阶段 Ray 多 actor 协同并由此提出WorkerGroup、ResourcePool、ClassWithInitArgs三大抽象docs/single_controller.rst核心机制register装饰器把分发/执行/阻塞语义挂到方法上WorkerGroup._bind_worker_method在初始化时查DISPATCH_MODE_FN_REGISTRY取出dispatch_fn/collect_fn配合execute_all/execute_rank_zero生成绑定方法最终让rollout.generate_sequences(batch)这样的单进程写法在分布式下原样运行数据路径DataProto.chunk切分、concat拼接、padding/select_idxs自动补齐与剔除是 DP 分发正确性的底层保证扩展方向colocated workerFusedWorker、资源池 split/merge、自定义 dispatch mode 等机制使这套抽象可以迁移到 RL 之外的批量分布式计算场景。相关源码索引装饰器与分发注册表 verl/single_controller/base/decorator.pyWorkerGroup/ResourcePool/ClassWithInitArgs verl/single_controller/base/worker_group.pyWorker 基类 verl/single_controller/base/worker.pyRay 后端实现 verl/single_controller/ray/base.py数据协议 verl/protocol.py。【免费下载链接】verlverl/HybridFlow: A Flexible and Efficient RL Post-Training Framework项目地址: https://gitcode.com/GitHub_Trending/ve/verl创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表