尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Ray RLlib RLModule API 完全指南:从 Spec 配置、多智能体模块到自定义前向逻辑

Ray RLlib RLModule API 完全指南:从 Spec 配置、多智能体模块到自定义前向逻辑 人工智能分布式训练强化学习任务调度模型推理服务【免费下载链接】rayRay is an AI compute engine. Ray consists of a core distributed runtime and a set of AI Libraries for accelerating ML workloads.项目地址https://gitcode.com/gh_mirrors/ra/ray点击查看免费下载本文围绕 Ray RLlib 新 API 栈的核心组件 RLModule强化学习模块展开系统讲解其规格定义RLModuleSpec/MultiRLModuleSpec、默认模型配置DefaultModelConfig、单智能体与多智能体模块的构造、前向传播约定、检查点保存恢复以及五大扩展 APIInferenceOnlyAPI、QNetAPI、SelfSupervisedLossAPI、TargetNetworkAPI、ValueFunctionAPI。读完本文你将能够基于 RLModule 源码 与 MultiRLModule 源码 准确理解 RLlib 模型的构建与调用机制并独立完成自定义 RLModule 的编写、配置与部署。背景说明Ray 2.40 起 RLlib 默认使用新 API 栈算法、示例脚本与文档已基本完成迁移参见 new_api_stack.rst。本文所有内容均以新 API 栈为准。一、RLModule 体系概览RLModule 是 RLlib 新 API 栈中策略/模型的统一抽象取代了旧栈中的Policy与ModelV2。它承担三项职责前向传播定义从观测 batch 到动作分布参数action_dist_inputs或值函数输出的计算逻辑状态管理通过get_state()/set_state()暴露模型权重配合save_to_path()/restore_from_path()/from_checkpoint()完成检查点读写组件化扩展通过一组附加 APImixin 式抽象基类为特定算法能力Q 网络、目标网络、值函数、自监督损失、推理裁剪提供标准接口。在代码组织上单智能体模块位于 rl_module.py多智能体容器位于 multi_rl_module.py默认模型配置位于 default_model_config.py附加 API 全部位于 apis/ 目录框架实现Torch位于 torch/torch_rl_module.py。二、RLModule 规格与配置RLModuleSpecRLModuleSpec是一个 dataclass 类型的规格spec工具类用于简化单智能体场景下 RLModule 的构造定义见 rl_module.py#L43-L94。2.1 字段说明字段类型默认值说明module_classOptional[Type[RLModule]]None要构建的 RLModule 类observation_spaceOptional[gym.Space]NoneRLModule 的观测空间。注意它可能与环境的观测空间不同——例如离散观测空间经预处理后通常对应 RLModule 中的 one-hot 编码观测空间action_spaceOptional[gym.Space]NoneRLModule 的动作空间inference_onlyboolFalse是否以仅推理状态构建模块。该状态下计算动作不需要的组件如值函数、目标网络可能被移除。注意inference_onlyTrue与learner_onlyTrue不能同时成立learner_onlyboolFalse是否只在 Learner 工作节点上构建、不在 EnvRunner 上构建。适用于 MultiRLModule 中仅用于训练的模块如多智能体共享值函数、好奇心学习中的世界模型model_configOptional[Union[Dict, DefaultModelConfig]]None模型配置字典或 RLlib 默认 dataclasscatalog_classOptional[Type[Catalog]]None用于构建子组件的 Catalog 类load_state_pathOptional[str]None已弃用。恢复模块状态请改用Algorithm.restore_from_path(path..., component...)RLModuleSpec.build()会先做合法性校验module_class未设置或observation_space未设置时会抛出ValueError见 rl_module.py#L94-L114。随后以关键字参数方式调用模块构造函数若自定义旧模块仍按旧的config参数签名实现则自动回退到RLModuleConfig兼容路径。RLModuleSpec同时提供from_module()从已实例化的模块反向生成规格、to_dict()/from_dict()序列化用于跨进程传输、update()类似dict.update()的字段合并override参数控制是否用新值覆盖非空字段以及as_multi_rl_module_spec()将单智能体规格包装为DEFAULT_MODULE_ID键下的 MultiRLModuleSpec等实用方法见 rl_module.py#L116-L256。2.2 使用示例import gymnasium as gym import numpy as np from ray.rllib.core.rl_module.rl_module import RLModuleSpec spec RLModuleSpec( module_classMyCustomRLModule, observation_spacegym.spaces.Box(-1.0, 1.0, (64, 64, 4), np.float32), action_spacegym.spaces.Discrete(7), # 自定义模块使用任意 (str) 键的 model_config 字典 model_config{ conv_filters: [[16, 4, 2], [32, 4, 2], [64, 4, 2], [128, 4, 2]], }, ) module spec.build()三、多智能体规格MultiRLModuleSpecMultiRLModuleSpec用于配置多智能体multi-agent场景其本质是一张ModuleID - RLModuleSpec的映射表外加若干全局配置定义见 multi_rl_module.py#L516-L586。字段类型默认值说明multi_rl_module_classType[MultiRLModule]MultiRLModule要构建的 MultiRLModule 类observation_spaceOptional[gym.Space]None全局观测空间适用于只存在于 MultiRLModule 内部、没有独立 ModuleID 的共享网络组件action_spaceOptional[gym.Space]None全局动作空间用途同上inference_onlyOptional[bool]None全局推理标记。为None时自动推断仅当所有子模块的inference_only均为True时整体才为Truemodel_configOptional[dict]None全局模型配置同样服务于共享网络组件rl_module_specsOptional[Dict[ModuleID, RLModuleSpec]]None映射各 ModuleID 到单智能体规格的字典MultiRLModuleSpec.build(module_idNone)支持两种构建粒度见 multi_rl_module.py#L587-L627传入module_id只构建对应的单个 RLModule不传module_id构建整个 MultiRLModule将rl_module_specs中的每个规格交给multi_rl_module_class构造。MultiRLModuleSpec还提供了运行时修改规格的方法add_modules()批量新增或覆盖子规格任一子模块inference_onlyFalse会使整体标记同步关闭、remove_modules()移除若干 ModuleID以及update()支持与单个RLModuleSpec或另一个MultiRLModuleSpec合并。注意module_specs字段已弃用应统一使用rl_module_specs__post_init__中若传入非字典的rl_module_specs会直接抛出ValueError并提示单策略场景应把RLModuleSpec直接传给config.rl_module(rl_module_spec..)由 RLlib 自动复制到各策略。四、默认模型配置DefaultModelConfigDefaultModelConfig是配置 RLlib 各算法内置默认 RLModule 的 dataclass定义见 default_model_config.py。注意它的使用边界自定义 RLModule 不应使用该类而应使用任意 (str) 键的model_config字典DefaultModelConfig仅供 RLlib 默认模块通过算法配置传入。接入方式如下from ray.rllib.algorithms.ppo import PPOConfig from ray.rllib.core.rl_module.default_model_config import DefaultModelConfig config ( PPOConfig() .rl_module( model_configDefaultModelConfig(fcnet_hiddens[32, 32]), ) )其关键字段按功能分组如下全部字段可在 default_model_config.py 中查看MLP 编码栈字段默认值说明fcnet_hiddens[256, 256]全连接MLP栈各层节点数在编码器 策略头/值头的默认架构下仅影响编码器部分fcnet_activationtanh激活函数支持tanh、relu、swish即silu、linear或Nonefcnet_kernel_initializer/fcnet_bias_initializerNone权重/偏置初始化器可传框架torch支持的名字、类或函数为None时使用框架默认初始化fcnet_use_layernormFalse是否在编码器每个隐藏层后插入 LayerNormConv2D 栈字段默认值说明conv_filtersNone形如[[num_out_channels, kernel, stride], ...]的二维卷积栈定义kernel/stride可为单个 int 或(w, h)二元组为None且输入为 2D 时RLlib 根据输入维度自动寻找默认滤波器组合conv_activationrelu卷积栈激活函数取值同上conv_kernel_initializer/conv_bias_initializerNone卷积权重/偏置初始化器Head策略头/值头配置字段默认值说明head_fcnet_hiddens[]策略头、值头或 Q 头的全连接层节点数例如fcnet_hiddens[32, 32]搭配head_fcnet_hiddens[64]会得到[32, 32]编码器、[64, act-dim]策略头和[64, 1]值头head_fcnet_activationreluHead 栈激活函数head_fcnet_use_layernormFalseHead 栈是否插入 LayerNorm连续动作相关字段默认值说明free_log_stdFalse对 DiagGaussian 等连续分布策略输出的后半段是否作为自由偏置参数而非依赖状态/网络的节点开启后策略头只需输出均值维度节点数log_std_clip_param20.0对 log(stddev) 的裁剪区间上限实际裁剪到[-20, 20]避免数值不稳定导致nan设为float(inf)可关闭裁剪vf_share_layersTrue策略与值函数是否共享编码器层LSTM 设置字段默认值说明use_lstmFalse是否用 LSTM 包裹编码器max_seq_len20LSTM 训练 batch 的最大序列长度lstm_cell_size256LSTM cell 尺寸lstm_use_prev_action/lstm_use_prev_rewardFalse是否把上一步动作/奖励作为输入lstm_kernel_initializer/lstm_bias_initializerNoneLSTM 层初始化器融合Fusion设置字段默认值说明fusionnet_hiddens[256, 256]MultiStreamEncoder如 SAC 中对多个输入流编码后拼接融合网络的层节点数fusionnet_activationtanh融合网络激活函数五、RLModule API构造与设置RLModule是所有 RLlib 模块的抽象基类同时继承Checkpointable定义见 rl_module.py#L259-L401。其构造函数已从旧的config参数迁移为关键字参数形式RLModule( observation_space..., action_space..., inference_onlyFalse, learner_onlyFalse, model_config{...}, catalog_class..., )旧的RLModule(config[RLModuleConfig])写法已弃用并会触发deprecation_warningerrorTrue见 rl_module.py#L425-L435。构造时若传入catalog_classRLlib 会尝试创建 Catalog 对象并从中取得action_dist_cls动作分布类。5.1 setup组件的构建钩子setup()是OverrideToImplementCustomLogic标注的钩子方法见 rl_module.py#L499-L508在基类__init__末尾被自动调用用于创建模块所需的神经网络组件如各层网络。自定义子类必须在构造函数中调用super().__init__(observation_space.., action_space.., ...)不要重写构造函数只重写setup()源码会在setup()被重复调用时抛出RuntimeError提示正确的继承顺序是框架基类在前、算法模块在后。此外还有三个可选覆写点get_inference_action_dist_cls()、get_exploration_action_dist_cls()、get_train_action_dist_cls()分别返回推理、探索、训练阶段使用的动作分布类RLlib 分布类都实现Distribution接口要求提供from_logits()与to_deterministic()。5.2 as_multi_rl_module单智能体转多智能体RLModule.as_multi_rl_module()返回一个以DEFAULT_MODULE_ID为键包装当前模块的MultiRLModule见 rl_module.py#L741-L748便于统一以多智能体接口调度。对已是 MultiRLModule 的实例该方法被重写为直接返回自身以避免双重包装见 multi_rl_module.py#L467-L476。六、Forward 方法公开三件套与私有四件套这是 RLModule 使用中最重要的约定务必区分两层方法6.1 公开方法仅供外部组件调用不要覆写forward_inference、forward_exploration、forward_train是 EnvRunner采样器与 Learner 调用的入口禁止在自定义子类中覆写源码注释明确标注 DO NOT OVERRIDE!见 rl_module.py#L581-L655。它们分别对应forward_exploration训练阶段采样带探索行为时调用forward_inference生产部署、贪心执行无探索时调用forward_train计算损失函数输入时调用。注意若模块以inference_onlyTrue构建调用forward_train会抛出RuntimeError见 rl_module.py#L649-L654。6.2 私有方法自定义前向逻辑的真正落点自定义模型的前向行为应通过覆写私有方法实现见 rl_module.py#L558-L667_forward所有阶段的通用前向行为若无需区分阶段覆写它即可_forward_exploration训练样本采集阶段的行为默认委托给_forward_forward_inference无探索的动作计算行为默认委托给_forward_forward_train损失计算前的前向行为默认委托给_forward。以 Torch 框架为例TorchRLModule 会在调用_forward_inference/_forward_exploration时自动进入torch.no_grad()上下文以提升性能见 torch_rl_module.py#L109-L120并在推理/探索前自动切换eval()模式、结束后恢复原train()状态forward_train则显式调用self.train()。此外 TorchRLModule 的三个_forward_*方法可分别被torch.compile编译详见 torch_rl_module.py#L26-L41 的约束说明。6.3 三种典型调用范式以下示例摘自 RLModule 类文档分别演示采样循环、训练与推理三种场景以 PPO 默认模块为例import gymnasium as gym import torch from ray.rllib.algorithms.ppo.torch.default_ppo_torch_rl_module import ( DefaultPPOTorchRLModule, ) from ray.rllib.algorithms.ppo.ppo_catalog import PPOCatalog from ray.rllib.core.rl_module.default_model_config import DefaultModelConfig env gym.make(CartPole-v1) # 构建 PPO 的默认 RLModule module DefaultPPOTorchRLModule( observation_spaceenv.observation_space, action_spaceenv.action_space, model_configDefaultModelConfig(fcnet_hiddens[128, 128]), catalog_classPPOCatalog, ) action_dist_class module.get_inference_action_dist_cls() obs, info env.reset() terminated False # —— 采样循环forward_exploration—— while not terminated: fwd_ins {obs: torch.Tensor([obs])} fwd_outputs module.forward_exploration(fwd_ins) action_dist action_dist_class.from_logits(fwd_outputs[action_dist_inputs]) action action_dist.sample()[0].numpy() obs, reward, terminated, truncated, info env.step(action) # —— 训练forward_train—— fwd_outputs module.forward_train(fwd_ins) # loss compute_loss(fwd_outputs, fwd_ins) # update_params(module, loss) # —— 推理部署forward_inference—— fwd_outputs module.forward_inference(fwd_ins) action_dist action_dist_class.from_logits(fwd_outputs[action_dist_inputs]) action action_dist.sample()[0].numpy()七、保存与恢复检查点机制RLModule 通过继承Checkpointable获得完整的检查点能力实现见 checkpoints.py包含五个核心方法均声明于原文档的 Saving and restoring 小节方法作用get_state(components..., not_components..., inference_onlyFalse)返回模块状态字典。inference_onlyTrue时返回不含值函数/目标网络等非推理组件的状态可节省网络传输若模块本身以inference_onlyTrue构建而此处传False可能报错set_state(state)将给定状态字典写入模块save_to_path(path, ...)将模块含各子组件保存到路径状态文件名为module_stateSTATE_FILE_NAME见 rl_module.py#L406restore_from_path(path, ...)从路径恢复模块及其子组件from_checkpoint(path)类方法从检查点完整重建模块实例通过get_ctor_args_and_kwargs()记录的构造参数复原自定义模块实现get_state()/set_state()即可参与完整的检查点流程。load_state_path字段已弃用恢复状态请走Algorithm.restore_from_path(path..., component...)或上述RLModule.from_checkpoint()。八、MultiRLModule API多智能体容器MultiRLModule是包含 n 个子 RLModule 的容器基类维护ModuleID - RLModule映射定义见 multi_rl_module.py#L47-L71。默认实现假设输入输出均为Dict[ModuleID, Dict[str, Any]]类型按module_id循环对每个子模块执行前向默认假设子模块之间不共享参数、不互相通信——需要共享编码器或模块间通信的高级用法必须自行实现MultiRLModule子类并在MultiRLModuleSpec.multi_rl_module_class中指定。8.1 构造函数与 setupMultiRLModule( observation_space..., action_space..., inference_onlyNone, # None 时推断所有子模块均为 True 才为 True model_configNone, rl_module_specs{...}, # ModuleID - RLModuleSpec )若传入inference_onlyTrue会强制把rl_module_specs中所有子规格的该标记也置为True见 multi_rl_module.py#L113-L123setup()依次调用每个子规格的build()创建子模块并校验所有子模块框架一致framework为None或与首个子模块相同否则触发断言见 multi_rl_module.py#L136-L149。8.2 运行时修改子模块add_module(module_id, module, *, overrideFalse)运行时添加模块。ModuleID 已存在且overrideFalse时抛ValueError会先经validate_module_id()校验命名任一新模块inference_onlyFalse会使容器的该标记同步变为False新模块框架与容器不一致时抛ValueError除非框架为None。添加后rl_module_specs会同步更新保证写入磁盘后可被from_checkpoint()完整恢复见 multi_rl_module.py#L256-L306remove_module(module_id, *, raise_err_if_not_foundTrue)运行时移除模块同时清理rl_module_specs中的对应条目。8.3 便捷访问与批量操作容器实现了丰富的映射语义__contains__、__getitem__越界抛KeyError、get(module_id, defaultNone)、items()、keys()、values()、__len__以及foreach_module(func, return_dictFalse)——对每个(module_id, module)调用funcreturn_dictTrue时返回Dict[ModuleID, 结果]模块以unwrapped()暴露见 multi_rl_module.py#L326-L352。8.4 多智能体的保存与恢复MultiRLModule.get_state()按module_id分发到各子模块支持components/not_components按模块维度筛选见 multi_rl_module.py#L410-L428set_state()遍历状态字典逐个子模块写入缺失的模块被跳过get_checkpointable_components()返回各子模块作为可检查点组件从而让save_to_path/restore_from_path/from_checkpoint对多智能体同样生效。其序列化结构可通过MultiRLModuleSpec.to_dict()/from_dict()在进程间传递。九、附加 RLModule APIs五大扩展接口RLlib 将常见算法能力抽象为五个 mixin 式抽象基类全部位于 apis/ 目录自定义模块按需实现一个模块可同时实现多个。9.1 InferenceOnlyAPI —— 推理裁剪实现该接口的模块具备仅推理模式inference_only_api.py。只需实现get_non_inference_attributes()返回计算动作时不需要的组件属性名列表。RLlib 在 EnvRunner 上构建模块时inference_onlyTrue会删除这些组件生产部署时同样建议开启以减重Learner 上的模块inference_onlyFalse在get_state()时也只返回推理所需权重节省网络开销。支持用.表示嵌套子属性做细粒度控制例如class MyRLModule(RLModule): def setup(self): self._policy_head ... # 某 NN 组件 self._value_function_head ... # 某 NN 组件 self._encoder ... # 含 pol / vf 两个子属性的编码器 def get_non_inference_attributes(self): # 推理时删除值函数头与编码器中的 vf 部分 return [_value_function_head, _encoder.vf]TorchRLModule 构造函数 会在inference_onlyTrue且模块实现了该 API 时按属性路径逐级定位并执行delattr。9.2 QNetAPI —— (分布) Q 学习供 DQN 类算法使用q_net_api.py需实现compute_q_values(batch)基于编码器、Q 网络及可选 advantage 网络计算 Q 值。返回字典至少包含qf_preds当采用分布 Q 学习num_atoms 1时额外返回支撑原子atoms、Q logitsqf_logits与概率qf_probscompute_advantage_distribution(batch)计算 advantage 分布默认实现直接返回compute_q_values(batch)非 dueling 架构下二者一致。9.3 SelfSupervisedLossAPI —— 自带自监督损失实现该接口的模块自带损失函数self_supervised_loss_api.py。此时 Learner 将调用模块的compute_self_supervised_loss()而非 Learner 自身的compute_loss_for_module()。签名与后者一致仅多一个必填learner参数def compute_self_supervised_loss( self, *, learner, module_id, config, batch, fwd_out, **kwargs ) - torch.Tensor返回单个总损失张量多优化器场景可将各损失项相加返回并可用learner.metrics.log_value()/log_dict()记录分项指标。9.4 TargetNetworkAPI —— 目标网络管理供带目标网络算法使用target_network_api.py需实现三个方法make_target_networks()创建目标网络建议用ray.rllib.core.learner.utils.make_target_network()工具从对应主网络初始化get_target_network_pairs()返回[(main_net, target_net), ...]二元组列表如[(self.q_net, self.target_q_net)]forward_target(batch)对目标网络执行前向返回结果字典。目标网络的同步逻辑由拥有该模块的 Learner 统一处理。模块内还定义了常量TARGET_NETWORK_ACTION_DIST_INPUTS target_network_action_dist_inputs用于标识目标网络的动作分布输入键。9.5 ValueFunctionAPI —— 值函数计算供基于值函数的算法使用value_function_api.py需实现def compute_values(self, batch, embeddingsNone) - TensorTypeembeddings可选若调用方已通过共享编码器算好嵌入可传入以避免重复前向返回形状为(B,)或带时间维时(B, T)最后一个值维度已挤压掉不是 1。十、源码验证与测试参考若要验证上述行为仓库提供了完整测试集test_rl_module_specs.py覆盖RLModuleSpec/MultiRLModuleSpec的构建、序列化to_dict/from_dict、合并update与多智能体包装test_multi_rl_module.py覆盖 MultiRLModule 的添加/移除模块、前向分发、状态读写与框架校验torch/tests/test_torch_rl_module.py覆盖 TorchRLModule 的推理/训练模式切换与编译配置torch/tests/test_lstm_target_network_rl_module.pyLSTM 与目标网络 API 的组合验证。阅读源码时建议按此顺序先看 rl_module.py 理解单智能体契约再看 multi_rl_module.py 理解容器分发最后对照 apis/ 中自己关心的算法能力接口。本文所依据的 API 清单原始出处为 rl_modules.rst其中RLModuleSpec.build、MultiRLModuleSpec.build、RLModule的构造与 forward 方法、各附加 API 的方法签名均可直接在上文对应源码链接中核对。对于需要迁移旧栈代码的读者可参考仓库中的新 API 栈迁移指南new-api-stack-migration-guide完成从旧Policy/ModelV2到 RLModule 的改造。赞分享人工智能分布式训练强化学习任务调度模型推理服务【免费下载链接】rayRay is an AI compute engine. Ray consists of a core distributed runtime and a set of AI Libraries for accelerating ML workloads.项目地址https://gitcode.com/gh_mirrors/ra/ray点击查看免费下载相关推荐Multica 的任务能在运行机器上访问到什么安全模型与隔离边界怎么设置Multica 的任务能在运行机器上访问到什么安全模型与隔离边界怎么设置 当你把任务分配给 Multica 的智能体后执行发生在连接了守护进程daemon人工智能分布式训练强化学习任务调度模型推理服务Ruffty中的 useless-overload-body 规则为什么 overload 函数不需要函数体Ruffty中的 useless overload body 规则为什么 overload 函数不需要函数体 导读 overload 是 Python人工智能分布式训练强化学习任务调度模型推理服务如何用KiBot实现KiCad设计全流程自动化从安装到输出生产文件的快速入门如何用KiBot实现KiCad设计全流程自动化从安装到输出生产文件的快速入门 KiBot是一款强大的KiCad自动化工具能够帮助电子工程师和设计师将KiCa人工智能分布式训练强化学习任务调度模型推理服务上一篇CANN SHMEM 标量 Put/Get 主机侧对称内存通信实战rma_d2h_demo 原理与运行指南下一篇一文读懂gbert-large-sts-openmind训练数据德国STS基准数据集深度分析创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表