使用Ray-RLlib训练MARL DQN时遇RolloutWorker无input_reader错误
解决RLlib多智能体DQN训练时"RolloutWorker has no
input_reader object"错误 问题场景
基于SUMO的VSL多智能体DQN方案,使用Ray-RLlib重构环境后运行训练脚本,触发错误:ValueError: RolloutWorker has no 'input_reader' object! Cannot call 'sample()'。
关键修复步骤
1. 修正训练脚本配置
原训练代码存在配置键名错误、多智能体策略缺失等问题,修改后的代码如下:
from myEnvironment import myEnvironment import ray from ray.tune.registry import register_env from ray.rllib.algorithms.dqn import DQNConfig ray.init() # 注册环境(使用lambda传入config,而非空字典) register_env("myEnv", lambda config: myEnvironment(config)) # 使用DQNConfig构建器配置多智能体策略 config = DQNConfig() config = ( config .environment("myEnv") # 替换原错误的"environment"键为"env" .framework("torch") # 多智能体核心配置:为所有智能体共享同一DQN策略 .multi_agent( policies={"shared_policy"}, policy_mapping_fn=lambda agent_id, episode, **kwargs: "shared_policy", ) ) # 构建算法实例 agent = config.build() for iteration in range(100): result = agent.train() print(f"Iteration {iteration}: 平均奖励={result['episode_reward_mean']}") if iteration % 10 == 0: checkpoint = agent.save() print(f"Checkpoint已保存至: {checkpoint}") agent.stop() ray.shutdown()
2. 修复环境类核心方法
原环境类的reset、step方法存在多智能体格式不匹配、未定义变量等问题,关键修改如下:
修正reset方法
def reset(self, *, seed=None, options=None): super().reset(seed=seed) self.no_veh = False b = random.randint(1, 5) # 初始化每个智能体的info字典 self.info = {i: None for i in self._agent_ids} # 保留你的SUMO初始化逻辑 if self.connection is None: self.sumo_init() self.connection.load(self.sumoCmd[1:]) self.warmup() observation = self.sumo_step() # 严格返回MultiAgentEnv要求的(obs字典, info字典)格式 return observation, self.info
修正step方法
def step(self, action_dict): obs, rew, terminated, truncated, info = {}, {}, {}, {}, {} # 执行每个智能体的动作 for i, action in action_dict.items(): self.set_max_speed(action=action, edgeID=self.getEdgeID[i]) # 模拟SUMO步骤 for _ in range(60): state = self.sumo_step() # 获取每个智能体的奖励和结束状态 done_dict, reward_dict = self.rewardStd(self.getEdgeID) # 为每个智能体填充返回值 for agent_id in self._agent_ids: obs[agent_id] = state[agent_id] rew[agent_id] = reward_dict.get(agent_id, 0.0) terminated[agent_id] = done_dict.get(agent_id, False) truncated[agent_id] = done_dict.get(agent_id, False) info[agent_id] = self.info.get(agent_id, {}) # 设置全局结束标志(所有智能体终止时触发) terminated["__all__"] = all(terminated.values()) truncated["__all__"] = all(truncated.values()) return obs, rew, terminated, truncated, info
统一观测空间与数据格式
原观测数据是字典结构,但观测空间定义为Tuple,需统一格式:
def create_observation(self, edgeIDArray): meanSpeed = [] occupancy = [] agents_obs = {} for edge in edgeIDArray: meanSpeed.append((self.connection.edge.getLastStepMeanSpeed(edge) * 3.6) / 130) occupancy.append(self.connection.edge.getLastStepVehicleNumber(edge) / 58) for agentNumber in range(len(edgeIDArray)): agent_bit = np.zeros(4, dtype=np.float32) agent_bit[agentNumber] = 1.0 # 生成与Tuple观测空间匹配的结构 obs_tuple = ( np.array(meanSpeed, dtype=np.float32), np.array(occupancy, dtype=np.float32), agent_bit ) agents_obs[agentNumber] = obs_tuple return agents_obs
核心错误原因总结
- 训练配置中使用了错误的键名
environment,RLlib标准键名为env; - 多智能体场景未配置
multi_agent策略,DQN默认单智能体模式无法适配; - 环境类的
reset、step方法未严格遵循MultiAgentEnv的返回格式,导致RolloutWorker无法正确采样; - 观测数据结构与定义的观测空间不匹配,干扰数据采样流程。
内容的提问来源于stack exchange,提问作者komate1995
相关产品推荐
相关产品推荐

