You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Ray 0.x迁移至1.x:ray.rllib.models.Model转ModelV2迁移指南咨询

Ray RLlib 0.x 到 1.x: Model 迁移至 ModelV2 指南

核心替换与关键调整

  • 类继承更新:将原继承的 ray.rllib.models.Model 替换为 ray.rllib.models.ModelV2,若使用PyTorch/TensorFlow框架,建议直接继承对应框架的子类 TorchModelV2 或 TFModelV2,它们提供了更贴合框架的接口。
  • 构造函数适配:ModelV2 的构造函数必填参数与 Model 一致,但需注意部分框架子类的初始化逻辑(比如PyTorch下需同时继承nn.Module并调用其初始化方法)。
  • forward 方法修改:
    • 原 Model 的 forward 仅返回动作logits和状态,ModelV2 要求返回两个值:outputs(动作logits)和 state_out(状态输出,无状态模型返回空列表[])。
    • 需在 forward 中计算并存储价值函数的中间结果,供单独的 value_function 方法调用。
  • value_function 方法实现:不再直接在 forward 中返回价值,而是通过存储的中间张量,在 value_function 中返回最终的价值输出(比如PyTorch下需返回squeeze后的张量)。

迁移示例(PyTorch)

原Model代码:

from ray.rllib.models import Model
import torch.nn as nn

class MyOldModel(Model):
    def __init__(self, obs_space, action_space, num_outputs, model_config, name):
        super().__init__(obs_space, action_space, num_outputs, model_config, name)
        self.model = nn.Sequential(
            nn.Linear(obs_space.shape[0], 64),
            nn.ReLU(),
            nn.Linear(64, num_outputs)
        )
        self.value_head = nn.Linear(64, 1)
    
    def forward(self, input_dict, state, seq_lens):
        x = input_dict["obs_flat"]
        features = self.model[:-1](x)
        logits = self.model[-1](features)
        self._value = self.value_head(features)
        return logits, state
    
    def value_function(self):
        return self._value.squeeze(1)

迁移后TorchModelV2代码:

from ray.rllib.models.torch.torch_modelv2 import TorchModelV2
import torch.nn as nn

class MyNewModel(TorchModelV2, nn.Module):
    def __init__(self, obs_space, action_space, num_outputs, model_config, name):
        TorchModelV2.__init__(self, obs_space, action_space, num_outputs, model_config, name)
        nn.Module.__init__(self)
        self.model = nn.Sequential(
            nn.Linear(obs_space.shape[0], 64),
            nn.ReLU(),
            nn.Linear(64, num_outputs)
        )
        self.value_head = nn.Linear(64, 1)
        self._value_out = None
    
    def forward(self, input_dict, state_in, seq_lens):
        x = input_dict["obs_flat"]
        features = self.model[:-1](x)
        logits = self.model[-1](features)
        self._value_out = self.value_head(features)
        return logits, []
    
    def value_function(self):
        return self._value_out.squeeze(1)

配置与常见问题

  • 训练配置中,将 model.custom_model 指定为新的 ModelV2 子类。
  • 若原模型包含循环状态(比如LSTM),需适配 ModelV2 的 state_in/state_out 机制,在 forward 中接收输入状态并返回更新后的状态。
  • 确保所有自定义模型的依赖逻辑(比如预训练权重加载)适配新的类结构。

内容的提问来源于stack exchange,提问作者Anton

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.16 01:15:39