Skip to content

Commit 954f1bb

Browse files
committed
[#144] Fix issue with failure penalty in multiagent env
1 parent 3349d63 commit 954f1bb

3 files changed

Lines changed: 30 additions & 1 deletion

File tree

docs/source/release_notes.rst

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,8 @@ Development - |version|
3333
* Add a maximum range checking dynamics model in :class:`MaxRangeDynModel`. Useful for keeping an agent
3434
in the vicinity of a target early in training.
3535
* Add properties in spacecraft dynamics for orbital element observations.
36+
* Fix an issue with failure penalties in the PettingZoo environment when the rewarder
37+
does not return a reward for a satellite.
3638

3739

3840
Version 1.1.0

src/bsk_rl/gym.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -671,7 +671,10 @@ def _get_reward(self) -> dict[AgentID, float]:
671671
reward = deepcopy(self.reward_dict)
672672
for agent, satellite in zip(self.possible_agents, self.satellites):
673673
if not satellite.is_alive():
674-
reward[agent] += self.failure_penalty
674+
if agent in reward:
675+
reward[agent] += self.failure_penalty
676+
else:
677+
reward[agent] = self.failure_penalty
675678

676679
reward_keys = list(reward.keys())
677680
for agent in reward_keys:

tests/unittest/test_gym_env.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -538,6 +538,30 @@ def test_get_reward(self):
538538
sat.name: -10.0 for i, sat in enumerate(env.unwrapped.satellites)
539539
}
540540

541+
@patch(
542+
"bsk_rl.GeneralSatelliteTasking._get_truncated",
543+
MagicMock(return_value=False),
544+
)
545+
def test_get_reward_missing_sat(self):
546+
env = ConstellationTasking(
547+
satellites=[
548+
MagicMock(is_alive=MagicMock(return_value=False)) for i in range(3)
549+
],
550+
world_type=MagicMock(),
551+
scenario=MagicMock(),
552+
rewarder=MagicMock(),
553+
failure_penalty=-20.0,
554+
)
555+
env._agents_last_compute_time = None
556+
env.simulator = MagicMock(sim_time=0.0)
557+
env.newly_dead = [sat.name for sat in env.unwrapped.satellites]
558+
env.reward_dict = {env.unwrapped.satellites[0].name: 10.0}
559+
assert env._get_reward() == {
560+
env.unwrapped.satellites[0].name: -10.0,
561+
env.unwrapped.satellites[1].name: -20.0,
562+
env.unwrapped.satellites[2].name: -20.0,
563+
}
564+
541565
@pytest.mark.parametrize("timeout", [False, True])
542566
@pytest.mark.parametrize("terminate_on_time_limit", [False, True])
543567
def test_get_terminated(self, timeout, terminate_on_time_limit):

0 commit comments

Comments
 (0)