From f2408b719c991b98ba5dc77db81dc184805b33df Mon Sep 17 00:00:00 2001 From: Benjamin Black Date: Thu, 17 Sep 2020 14:28:42 -0400 Subject: [PATCH] Fixed PettingZooEnv (#10847) --- doc/source/rllib-env.rst | 6 +- rllib/env/pettingzoo_env.py | 111 +++++++++++++++++++++---------- rllib/examples/pettingzoo_env.py | 15 +++-- 3 files changed, 88 insertions(+), 44 deletions(-) diff --git a/doc/source/rllib-env.rst b/doc/source/rllib-env.rst index 1b34e17ba..d6da000df 100644 --- a/doc/source/rllib-env.rst +++ b/doc/source/rllib-env.rst @@ -71,7 +71,7 @@ In the above example, note that the ``env_creator`` function takes in an ``env_c return self.env.step(action) register_env("multienv", lambda config: MultiEnv(config)) - + .. tip:: When using logging in an environment, the logging configuration needs to be done inside the environment, which runs inside Ray workers. Any configurations outside the environment, e.g., before starting Ray will be ignored. @@ -217,11 +217,11 @@ PettingZoo Multi-Agent Environments from ray.tune.registry import register_env # import the pettingzoo environment - from pettingzoo.gamma import prison_v0 + from pettingzoo.butterfly import prison_v1 # import rllib pettingzoo interface from ray.rllib.env import PettingZooEnv # define how to make the environment. This way takes an optional environment config, num_floors - env_creator = lambda config: prison_v0.env(num_floors=config.get("num_floors", 4)) + env_creator = lambda config: prison_v1.env(num_floors=config.get("num_floors", 4)) # register that way to make the environment under an rllib name register_env('prison', lambda config: PettingZooEnv(env_creator(config))) # now you can use `prison` as an environment diff --git a/rllib/env/pettingzoo_env.py b/rllib/env/pettingzoo_env.py index e9b18fa20..1704cfbc9 100644 --- a/rllib/env/pettingzoo_env.py +++ b/rllib/env/pettingzoo_env.py @@ -1,4 +1,4 @@ -from .multi_agent_env import MultiAgentEnv +from ray.rllib.env.multi_agent_env import MultiAgentEnv class PettingZooEnv(MultiAgentEnv): @@ -122,7 +122,7 @@ class PettingZooEnv(MultiAgentEnv): obs (dict): New observations for each ready agent. """ # 1. Reset environment; agent pointer points to first agent. - self.aec_env.reset(observe=False) + self.aec_env.reset() # 2. Copy agents from environment self.agents = self.aec_env.agents @@ -155,53 +155,96 @@ class PettingZooEnv(MultiAgentEnv): "__all__" (required) is used to indicate env termination. infos (dict): Optional info values for each agent id. """ - env_done = False - # iterate over self.agents - for agent in self.agents: + stepped_agents = set() + while self.aec_env.agent_selection not in stepped_agents: + agent = self.aec_env.agent_selection + assert agent in action_dict, \ + "Live environment agent is not in actions dictionary" + self.aec_env.step(action_dict[agent]) + stepped_agents.add(agent) - # Execute only for agents that have not been done in previous steps - if agent in action_dict.keys(): - if not env_done: - assert agent == self.aec_env.agent_selection, \ - f"environment has a nontrivial ordering, and " \ - "cannot be used with the POMGameEnv wrapper\"" \ - "nCurrent agent: {self.aec_env.agent_selection}" \ - "\nExpected agent: {agent}" - # Execute agent action in environment - self.obs[agent] = self.aec_env.step( - action_dict[agent], observe=True) - if all(self.aec_env.dones.values()): - env_done = True - self.dones["__all__"] = True - else: - self.obs[agent] = self.aec_env.observe(agent) - # Get reward - self.rewards[agent] = self.aec_env.rewards[agent] - # Update done status - self.dones[agent] = self.aec_env.dones[agent] + assert all(agent in stepped_agents or self.aec_env.dones[agent] + for agent in action_dict), \ + "environment has a nontrivial ordering, and cannot be used with"\ + " the POMGameEnv wrapper" - # For agents with done = True, remove from dones, rewards and - # observations. - else: - del self.dones[agent] - del self.rewards[agent] - del self.obs[agent] - del self.infos[agent] + self.obs = {} + self.rewards = {} + self.dones = {} + self.infos = {} # update self.agents self.agents = list(action_dict.keys()) - # Update infos stepwise for agent in self.agents: + self.obs[agent] = self.aec_env.observe(agent) + self.dones[agent] = self.aec_env.dones[agent] + self.rewards[agent] = self.aec_env.rewards[agent] self.infos[agent] = self.aec_env.infos[agent] + self.dones["__all__"] = all(self.aec_env.dones.values()) + return self.obs, self.rewards, self.dones, self.infos def render(self, mode="human"): - self.aec_env.render(mode=mode) + return self.aec_env.render(mode=mode) def close(self): self.aec_env.close() + def seed(self, seed=None): + self.aec_env.seed(seed) + def with_agent_groups(self, groups, obs_space=None, act_space=None): raise NotImplementedError + + +class ParallelPettingZooEnv(MultiAgentEnv): + def __init__(self, env): + self.par_env = env + # agent idx list + self.agents = self.par_env.agents + + # Get dictionaries of obs_spaces and act_spaces + self.observation_spaces = self.par_env.observation_spaces + self.action_spaces = self.par_env.action_spaces + + # Get first observation space, assuming all agents have equal space + self.observation_space = self.observation_spaces[self.agents[0]] + + # Get first action space, assuming all agents have equal space + self.action_space = self.action_spaces[self.agents[0]] + + assert all(obs_space == self.observation_space + for obs_space + in self.par_env.observation_spaces.values()), \ + "Observation spaces for all agents must be identical. Perhaps " \ + "SuperSuit's pad_observations wrapper can help (useage: " \ + "`supersuit.aec_wrappers.pad_observations(env)`" + + assert all(act_space == self.action_space + for act_space in self.par_env.action_spaces.values()), \ + "Action spaces for all agents must be identical. Perhaps " \ + "SuperSuit's pad_action_space wrapper can help (useage: " \ + "`supersuit.aec_wrappers.pad_action_space(env)`" + + self.reset() + + def reset(self): + return self.par_env.reset() + + def step(self, action_dict): + for agent in self.agents: + action_dict[agent] = self.action_space.sample() + obs, rew, dones, info = self.par_env.step(action_dict) + dones["__all__"] = all(dones.values()) + return obs, rew, dones, info + + def close(self): + self.par_env.close() + + def seed(self, seed=None): + self.par_env.seed(seed) + + def render(self, mode="human"): + return self.par_env.render(mode) diff --git a/rllib/examples/pettingzoo_env.py b/rllib/examples/pettingzoo_env.py index 53a149f86..3b36f55a7 100644 --- a/rllib/examples/pettingzoo_env.py +++ b/rllib/examples/pettingzoo_env.py @@ -7,14 +7,14 @@ except ImportError: from ray.tune.registry import register_env from ray.rllib.env import PettingZooEnv from pettingzoo.butterfly import pistonball_v0 -from supersuit.aec_wrappers import normalize_obs, dtype, color_reduction +from supersuit import normalize_obs_v0, dtype_v0, color_reduction_v0 from numpy import float32 if __name__ == "__main__": """For this script, you need: 1. Algorithm name and according module, e.g.: "PPo" + agents.ppo as agent - 2. Name of the aec game you want to train on, e.g.: "prison". + 2. Name of the aec game you want to train on, e.g.: "pistonball". 3. num_cpus 4. num_rollouts @@ -25,9 +25,9 @@ if __name__ == "__main__": # function that outputs the environment you wish to register. def env_creator(config): env = pistonball_v0.env(local_ratio=config.get("local_ratio", 0.2)) - env = dtype(env, dtype=float32) - env = color_reduction(env, mode="R") - env = normalize_obs(env) + env = dtype_v0(env, dtype=float32) + env = color_reduction_v0(env, mode="R") + env = normalize_obs_v0(env) return env num_cpus = 1 @@ -41,7 +41,8 @@ if __name__ == "__main__": config["env_config"] = {"local_ratio": 0.5} # 3. Register env - register_env("prison", lambda config: PettingZooEnv(env_creator(config))) + register_env("pistonball", + lambda config: PettingZooEnv(env_creator(config))) # 4. Extract space dimensions test_env = PettingZooEnv(env_creator({})) @@ -74,7 +75,7 @@ if __name__ == "__main__": # 6. Initialize ray and trainer object ray.init(num_cpus=num_cpus + 1) - trainer = get_agent_class(alg_name)(env="prison", config=config) + trainer = get_agent_class(alg_name)(env="pistonball", config=config) # 7. Train once trainer.train()