-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathpredator_prey.py
More file actions
161 lines (132 loc) · 6.09 KB
/
Copy pathpredator_prey.py
File metadata and controls
161 lines (132 loc) · 6.09 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
"""Cooperative predator team adapter for PettingZoo's Simple Tag."""
from __future__ import annotations
from collections.abc import Sequence
import numpy as np
from mpe2 import simple_tag_v3
from envs.multiagentenv import MultiAgentEnv
class PredatorPreyEnv(MultiAgentEnv):
"""Expose predators as learners and control prey with a seeded policy."""
def __init__(
self,
episode_limit: int = 50,
seed: int = 0,
num_predators: int = 3,
num_prey: int = 1,
num_obstacles: int = 2,
**_: object,
) -> None:
self.episode_limit = int(episode_limit)
self._seed = int(seed)
self._rng = np.random.default_rng(self._seed)
self.prey_policy = "seeded_random"
self._env = simple_tag_v3.parallel_env(
num_good=num_prey,
num_adversaries=num_predators,
num_obstacles=num_obstacles,
max_cycles=self.episode_limit,
continuous_actions=False,
render_mode=None,
)
self._all_agents = list(self._env.possible_agents)
self.predator_agents = [name for name in self._all_agents if name.startswith("adversary_")]
self.prey_agents = [name for name in self._all_agents if name not in self.predator_agents]
self.possible_agents = list(self.predator_agents)
self.n_agents = len(self.predator_agents)
if self.n_agents == 0:
raise ValueError("Simple Tag did not create any predator agents")
self.n_actions = int(self._env.action_space(self.predator_agents[0]).n)
self.timestep = 0
self.last_team_reward = 0.0
self.last_predator_rewards = {name: 0.0 for name in self.predator_agents}
self._observations: dict[str, np.ndarray] = {}
self._max_obs_size = max(
int(np.prod(self._env.observation_space(agent).shape)) for agent in self._all_agents
)
self.reset()
def step(self, actions: Sequence[int]):
action_values = np.asarray(actions).reshape(-1)
if action_values.size != self.n_agents:
raise ValueError(f"expected {self.n_agents} predator actions, got {action_values.size}")
joint_actions: dict[str, int] = {}
for agent, action in zip(self.predator_agents, action_values, strict=True):
action = int(action)
if not self._env.action_space(agent).contains(action):
raise ValueError(f"invalid action {action} for {agent}")
if agent in self._env.agents:
joint_actions[agent] = action
for agent in self.prey_agents:
if agent in self._env.agents:
joint_actions[agent] = int(self._rng.integers(self._env.action_space(agent).n))
observations, rewards, terminations, truncations, infos = self._env.step(joint_actions)
self._observations.update(observations)
self.timestep += 1
self.last_predator_rewards = {
agent: float(rewards.get(agent, 0.0)) for agent in self.predator_agents
}
self.last_team_reward = float(np.mean(list(self.last_predator_rewards.values())))
predator_done = [
bool(terminations.get(agent, False) or truncations.get(agent, False))
for agent in self.predator_agents
]
reached_limit = self.timestep >= self.episode_limit or any(
bool(truncations.get(agent, False)) for agent in self._all_agents
)
terminated = reached_limit or all(predator_done)
info = {
"episode_limit": reached_limit,
"episode_step": self.timestep,
"predators": self.n_agents,
}
if infos:
info["pettingzoo_agents_with_info"] = sum(bool(value) for value in infos.values())
return self.last_team_reward, terminated, info
def reset(self, seed: int | None = None):
if seed is not None:
self.seed(seed)
episode_seed = int(self._rng.integers(0, np.iinfo(np.int32).max))
observations, _ = self._env.reset(seed=episode_seed)
self._observations = dict(observations)
self.timestep = 0
self.last_team_reward = 0.0
self.last_predator_rewards = {name: 0.0 for name in self.predator_agents}
return self.get_obs(), self.get_state()
def _observation(self, agent: str) -> np.ndarray:
observation = self._observations.get(agent)
if observation is None:
observation = np.zeros(self._env.observation_space(agent).shape, dtype=np.float32)
return np.asarray(observation, dtype=np.float32).reshape(-1)
def get_obs(self):
return [self._observation(agent) for agent in self.predator_agents]
def get_obs_agent(self, agent_id: int):
return self.get_obs()[agent_id]
def get_obs_size(self):
return int(np.prod(self._env.observation_space(self.predator_agents[0]).shape))
def get_state(self):
observations = [self._observation(agent) for agent in self._all_agents]
padded = [np.pad(obs, (0, self._max_obs_size - obs.size)) for obs in observations]
return np.concatenate(padded).astype(np.float32, copy=False)
def get_state_size(self):
return self._max_obs_size * len(self._all_agents)
def get_avail_actions(self):
return [self.get_avail_agent_actions(agent_id) for agent_id in range(self.n_agents)]
def get_avail_agent_actions(self, agent_id: int):
if not 0 <= agent_id < self.n_agents:
raise IndexError(agent_id)
return np.ones(self.n_actions, dtype=np.int32)
def get_total_actions(self):
return self.n_actions
def get_reward(self):
return self.last_team_reward
def get_stats(self):
return {"team_reward": self.last_team_reward, "episode_step": self.timestep}
def seed(self, seed: int | None = None):
self._seed = self._seed if seed is None else int(seed)
self._rng = np.random.default_rng(self._seed)
return self._seed
def render(self):
return self._env.render()
def close(self):
self._env.close()
def save_replay(self):
return None
predator_prey = PredatorPreyEnv