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

Ray RLlib中SAC算法CTDE模式运行报错问题排查求助

自定义CTDE RLModule在真实环境中的断言与回放缓冲区错误

背景

基于Ray RLlib文档中RockPaperScissors多智能体环境的变体搭建测试环境,验证自定义集中式训练、分散式执行(CTDE)RLModule的功能。通过GroupAgentsWrapper将多智能体环境转换为单智能体环境,该自定义RLModule实现了PPO、APPO、SAC所需的ValueFunctionAPI、TargetNetworkAPI、QNetAPI接口。在测试环境中,使用ray.tune.Tuner运行三种算法均正常,但应用到真实环境时出现错误。

错误详情

初始断言错误

ray/rllib/env/single_agent_episode.py, line 624, in concat_episode, assert np.all(other.observations[0] == self.observations[-1]), ValueError: The truth value of an array with more than one element is ambiguous. Use a.any() or a.all()

跳过断言后的回放缓冲区错误

ray/rllib/utils/replay_buffers/prioritized_episode_buffer.py, line 471, in sample, idx = self._sum_segment.find_prefixsum_idx(random_sum)
ray/rllib/execution/segment_tree.py, line 191, in find_prefixsum_idx, assert 0 <= prefixsum <= self.sum() + 1e-5

排查尝试

  • 修改算法配置中的replay_buffer_config,或替换为自定义回放缓冲区,均未解决问题
  • 确认真实环境与测试环境的观测空间结构完全一致:
    • 多智能体环境:gym.spaces.Dict({"agent1": Box, "agent2": Box})
    • 单智能体环境:gym.spaces.Dict({"grouped": gym.spaces.Tuple(Box, Box)})
      无法解释测试环境能通过断言、真实环境却失败的原因

版本信息

  • Ray版本:2.46.0
  • Python版本:3.12.10

测试环境代码

class RockPaperScissors(MultiAgentEnv):
    ROCK = 0
    PAPER = 1
    SCISSORS = 2

    WIN_MATRIX = {
        (ROCK, ROCK): (0, 0),
        (ROCK, PAPER): (-1, 1),
        (ROCK, SCISSORS): (1, -1),
        (PAPER, ROCK): (1, -1),
        (PAPER, PAPER): (0, 0),
        (PAPER, SCISSORS): (-1, 1),
        (SCISSORS, ROCK): (-1, 1),
        (SCISSORS, PAPER): (1, -1),
        (SCISSORS, SCISSORS): (0, 0),
    }

    def __init__(self, env_config=None):
        super().__init__()

        self.agents_id = ["player1", "player2"]
        self.agents = self.possible_agents = self.agents_id
        self.observation_spaces = self.action_spaces = gym.spaces.Dict({
            "player1": gym.spaces.Box(low=0, high=2, shape=(1,)),
            "player2": gym.spaces.Box(low=0, high=2, shape=(1,)),
        })
        self.num_moves = 0

    def reset(self, *, seed=None, options=None):
        self.num_moves = 0

        return {
            "player1": np.array([0.0], dtype=np.float32),
            "player2": np.array([0.0], dtype=np.float32),
        }, {}

    def step(self, action_dict):
        self.num_moves += 1

        move1 = int(action_dict["player1"].item())
        move2 = int(action_dict["player2"].item())

        observations = {
            "player1": np.array([move2], dtype=np.float32),
            "player2": np.array([move1], dtype=np.float32)
        }

        r1, r2 = self.WIN_MATRIX[move1, move2]

        rewards = {
            "player1": r1,
            "player2": r2
        }

        terminateds = {"__all__": bool(self.num_moves >= 10)}
        truncateds = {"__all__": bool(self.num_moves >= 10)}

        return observations, rewards, terminateds, truncateds, {}

class GroupedRockPaperScissors(MultiAgentEnv):
    def __init__(self, env_config=None):
        super().__init__()

        env = RockPaperScissors(env_config)

        _tuple_obs_space = self._dict_to_tuple_space(env.observation_spaces)
        _tuple_act_space = self._dict_to_tuple_space(env.action_spaces)

        self.env = env.with_agent_groups(
            groups={"grouped_agents": ["player1", "player2"]},
            obs_space=_tuple_obs_space, # spaces.Tuple(Box, Box)
            act_space=_tuple_act_space, # spaces.Tuple(Box, Box)
        )

        self.agents_id = ["grouped_agents"]
        self.agents = self.possible_agents = self.agents_id
        self.original_agents_id = env.agents_id

        self.observation_space = gym.spaces.Dict(
            {"grouped_agents": _tuple_obs_space} # spaces.Dict({"grouped": spaces.Tuple(Box, Box)})
        )
        self.action_space = gym.spaces.Dict(
            {"grouped_agents": _tuple_act_space} # spaces.Dict({"grouped": spaces.Tuple(Box, Box)})
        )

    def reset(self, *, seed=None, options=None):
        obs, infos = self.env.reset(seed=seed, options=options)
        grouped_obs = {k: tuple(v) for k, v in obs.items()}  # spaces.Dict({"grouped": spaces.Tuple(Box, Box)})

        return grouped_obs, infos

    def step(self, action_dict):
        obs, rewards, terminateds, truncateds, infos = self.env.step(action_dict)
        grouped_obs = {k: tuple(v) for k, v in obs.items()}  # spaces.Dict({"grouped": spaces.Tuple(Box, Box)})
        grouped_reward = sum(rewards.values())

        return grouped_obs, grouped_reward, terminateds["__all__"], truncateds["__all__"], infos

    @staticmethod
    def _dict_to_tuple_space(dict_space: gym.spaces.Dict) -> gym.spaces.Tuple:
        sorted_keys = sorted(dict_space.keys())
        tuple_of_spaces = tuple(dict_space[key] for key in sorted_keys)

        return gym.spaces.Tuple(tuple_of_spaces)

SAC配置代码

algo_config = (
    SACConfig()
    .environment(GroupedEnv, env_config={})
    .framework("torch")
    .rl_module(rl_module_spec=RLModuleSpec(module_class=CustomRLModuleCTDE,
                                           observation_space=GroupedEnv.observation_space, # spaces.Dict({"grouped": spaces.Tuple(Box, Box)})
                                           action_space=GroupedEnv.action_space)) # spaces.Dict({"grouped": spaces.Tuple(Box, Box)}
    .training(twin_q=True,
              replay_buffer_config={"type": "PrioritizedEpisodeReplayBuffer"}) # or "EpisodeReplayBuffer"
    .evaluation(evaluation_config=SACConfig.overrides(exploration=False))
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 16:12:35