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
相关产品推荐
相关产品推荐

