Skip to content

Commit cba0ce7

Browse files
committed
[#144] Resource reward randomization
1 parent e3497c9 commit cba0ce7

10 files changed

Lines changed: 160 additions & 21 deletions

File tree

docs/source/release_notes.rst

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,8 @@ Development - |version|
3535
* Add properties in spacecraft dynamics for orbital element observations.
3636
* Fix an issue with failure penalties in the PettingZoo environment when the rewarder
3737
does not return a reward for a satellite.
38+
* Allow for per-episode randomization of :class:`ResourceReward` weights and observation
39+
of those weights with :class:`ResourceRewardWeight`.
3840

3941

4042
Version 1.1.0

src/bsk_rl/act/actions.py

Lines changed: 18 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
from copy import deepcopy
55
from typing import TYPE_CHECKING, Any
66

7-
from bsk_rl.utils.functional import AbstractClassProperty
7+
from bsk_rl.utils.functional import AbstractClassProperty, Resetable
88

99
if TYPE_CHECKING: # pragma: no cover
1010
from gymnasium import spaces
@@ -29,7 +29,7 @@ def select_action_builder(satellite: "Satellite") -> "ActionBuilder":
2929
raise NotImplementedError("Heterogenous action builders not supported.")
3030

3131

32-
class ActionBuilder(ABC):
32+
class ActionBuilder(ABC, Resetable):
3333

3434
def __init__(self, satellite: "Satellite") -> None:
3535
"""Base class for all action builders.
@@ -43,6 +43,21 @@ def __init__(self, satellite: "Satellite") -> None:
4343
for act in self.action_spec:
4444
act.link_satellite(self.satellite)
4545

46+
def reset_overwrite_previous(self) -> None:
47+
"""Perform any once-per-episode setup."""
48+
for act in self.action_spec:
49+
act.reset_overwrite_previous()
50+
51+
def reset_pre_sim_init(self) -> None:
52+
"""Perform any once-per-episode setup."""
53+
for act in self.action_spec:
54+
act.reset_pre_sim_init()
55+
56+
def reset_during_sim_init(self) -> None:
57+
"""Perform any once-per-episode setup."""
58+
for act in self.action_spec:
59+
act.reset_during_sim_init()
60+
4661
def reset_post_sim_init(self) -> None:
4762
"""Perform any once-per-episode setup."""
4863
self.simulator = self.satellite.simulator # already a proxy
@@ -68,7 +83,7 @@ def set_action(self, action: Any) -> None:
6883
pass
6984

7085

71-
class Action(ABC):
86+
class Action(ABC, Resetable):
7287
builder_type: type[ActionBuilder] = AbstractClassProperty() #: :meta private:
7388

7489
def __init__(self, name: str = "act") -> None:
@@ -101,10 +116,6 @@ def link_simulator(self, simulator: "Simulator") -> None:
101116
"""
102117
self.simulator = simulator # already a proxy
103118

104-
def reset_post_sim_init(self) -> None: # pragma: no cover
105-
"""Perform any once-per-episode setup."""
106-
pass
107-
108119
@abstractmethod
109120
def set_action(self, action: Any) -> None: # pragma: no cover
110121
"""Execute code to perform an action."""

src/bsk_rl/data/resource_data.py

Lines changed: 24 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -5,9 +5,10 @@
55
"""
66

77
import logging
8-
from typing import TYPE_CHECKING, Callable
8+
from typing import Callable, Union
99

1010
from bsk_rl.data.base import Data, DataStore, GlobalReward
11+
from bsk_rl.obs import ResourceRewardWeight
1112

1213
logger = logging.getLogger(__name__)
1314

@@ -57,7 +58,9 @@ class ResourceReward(GlobalReward):
5758
data_store_type = ResourceDataStore
5859

5960
def __init__(
60-
self, reward_weight: float = 1.0, resource_fn: Callable = lambda sat: 0.0
61+
self,
62+
reward_weight: Union[float, Callable] = 1.0,
63+
resource_fn: Callable = lambda sat: 0.0,
6164
) -> None:
6265
"""Rewards for an arbitrary resource.
6366
@@ -76,17 +79,33 @@ def __init__(
7679
resource_fn = lambda sat: sat.simulator.sim_time
7780
reward_weight = -1e-3 # is negative, because time increases over time
7881
79-
8082
Args:
8183
reward_weight: [reward/resource] Scaling factor to apply to changes in resource
82-
level to yield reward.
84+
level to yield reward. Can be a float or a function that randomizes the reward
85+
weight per-episode.
8386
resource_fn: Function to call to get the resource level for each satellite.
8487
"""
8588
super().__init__()
86-
self.reward_weight = reward_weight
89+
self._reward_weight = reward_weight
8790
self.resource_fn = resource_fn
8891
self.data_store_kwargs = dict(resource_fn=resource_fn)
8992

93+
def reset_pre_sim_init(self) -> None:
94+
"""Reset the reward weight before simulation initialization."""
95+
if callable(self._reward_weight):
96+
self.reward_weight = self._reward_weight()
97+
else:
98+
self.reward_weight = self._reward_weight
99+
return super().reset_pre_sim_init()
100+
101+
def reset_post_sim_init(self) -> None:
102+
"""Add the reward weight to each satellite's observation spec."""
103+
for satellite in self.scenario.satellites:
104+
for obs in satellite.observation_builder.observation_spec:
105+
if isinstance(obs, ResourceRewardWeight):
106+
obs.weight_vector.append(self.reward_weight)
107+
return super().reset_post_sim_init()
108+
90109
def calculate_reward(
91110
self, new_data_dict: dict[str, ResourceData]
92111
) -> dict[str, float]:

src/bsk_rl/obs/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,12 +32,14 @@ class MyObservationSatellite(Satellite):
3232
* :class:`Time` - Add simulation time to the observation.
3333
* :class:`OpportunityProperties` - Add information about upcoming targets or other ground access points to the observation.
3434
* :class:`Eclipse` - Add a tuple of the next orbit start and end.
35+
* :class:`ResourceRewardWeight` - Reports the weights of any randomized :class:`ResourceReward`.
3536
"""
3637

3738
from bsk_rl.obs.observations import (
3839
Eclipse,
3940
Observation,
4041
OpportunityProperties,
42+
ResourceRewardWeight,
4143
SatProperties,
4244
Time,
4345
)
@@ -51,4 +53,5 @@ class MyObservationSatellite(Satellite):
5153
"Time",
5254
"OpportunityProperties",
5355
"Eclipse",
56+
"ResourceRewardWeight",
5457
]

src/bsk_rl/obs/observations.py

Lines changed: 44 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -3,13 +3,12 @@
33
import logging
44
from abc import ABC, abstractmethod
55
from copy import deepcopy
6-
from typing import TYPE_CHECKING, Any, Callable, Optional, Union
6+
from typing import TYPE_CHECKING, Any, Union
77

88
import numpy as np
9-
from Basilisk.utilities import orbitalMotion
109
from gymnasium import spaces
1110

12-
from bsk_rl.utils.functional import vectorize_nested_dict
11+
from bsk_rl.utils.functional import Resetable, vectorize_nested_dict
1312
from bsk_rl.utils.orbital import rv2HN
1413

1514
if TYPE_CHECKING: # pragma: no cover
@@ -49,7 +48,7 @@ def nested_obs_to_space(obs_dict, dtype):
4948
raise TypeError(f"Cannot convert {obs_dict} to gym space.")
5049

5150

52-
class ObservationBuilder:
51+
class ObservationBuilder(Resetable):
5352
def __init__(
5453
self,
5554
satellite: "Satellite",
@@ -80,6 +79,21 @@ def __init__(
8079
name_counts[obs.name] = 1
8180
obs.link_satellite(self.satellite)
8281

82+
def reset_overwrite_previous(self) -> None:
83+
"""Perform any once-per-episode setup."""
84+
for obs in self.observation_spec:
85+
obs.reset_overwrite_previous()
86+
87+
def reset_pre_sim_init(self) -> None:
88+
"""Perform any once-per-episode setup."""
89+
for obs in self.observation_spec:
90+
obs.reset_pre_sim_init()
91+
92+
def reset_during_sim_init(self) -> None:
93+
"""Perform any once-per-episode setup."""
94+
for obs in self.observation_spec:
95+
obs.reset_during_sim_init()
96+
8397
def reset_post_sim_init(self) -> None:
8498
"""Perform any once-per-episode setup."""
8599
self.simulator = self.satellite.simulator # already a proxy
@@ -143,7 +157,7 @@ def observation_description(self) -> Any:
143157
return self.obs_array_keys()
144158

145159

146-
class Observation(ABC):
160+
class Observation(ABC, Resetable):
147161
"""Base observations class."""
148162

149163
def __init__(self, name: str = "obs") -> None:
@@ -176,10 +190,6 @@ def link_simulator(self, simulator: "Simulator") -> None:
176190
"""
177191
self.simulator = simulator # already a proxy
178192

179-
def reset_post_sim_init(self) -> None: # pragma: no cover
180-
"""Perform any once-per-episode setup."""
181-
pass
182-
183193
@abstractmethod # pragma: no cover
184194
def get_obs(self) -> Any:
185195
"""Return the observation."""
@@ -481,5 +491,30 @@ def get_obs(self):
481491
]
482492

483493

494+
class ResourceRewardWeight(Observation):
495+
"""Observation for the weight of any :class:`ResourceReward`."""
496+
497+
def __init__(
498+
self, name="resource_reward_weight", norm: Union[float, np.ndarray] = 1.0
499+
):
500+
"""Include the resource reward weight in the observation.
501+
502+
Args:
503+
name: Name of the observation.
504+
norm: Value to normalize the resource reward weight by. If a vector, it should
505+
be the same length as the number of resource rewards in the environment.
506+
"""
507+
super().__init__(name=name)
508+
self.norm = norm
509+
510+
def reset_overwrite_previous(self) -> None: # pragma: no cover
511+
"""Prepare weights to be saved by rewarders."""
512+
self.weight_vector = []
513+
514+
def get_obs(self) -> float:
515+
"""Return the resource reward weight."""
516+
return np.array(self.weight_vector) / self.norm
517+
518+
484519
__doc_title__ = "Backend"
485520
__all__ = ["ObservationBuilder"]

src/bsk_rl/sats/satellite.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -142,6 +142,8 @@ def reset_overwrite_previous(self) -> None:
142142
self._timed_terminal_event_name = None
143143
self._is_alive = True
144144
self.time_of_death = None
145+
self.observation_builder.reset_overwrite_previous()
146+
self.action_builder.reset_overwrite_previous()
145147

146148
@vizard.visualize
147149
def create_vizard_data(self, color, vizSupport=None) -> None:
@@ -160,6 +162,8 @@ def reset_pre_sim_init(self) -> None:
160162
oe=self.sat_args["oe"],
161163
mu=self.sat_args["mu"],
162164
)
165+
self.observation_builder.reset_pre_sim_init()
166+
self.action_builder.reset_pre_sim_init()
163167

164168
def set_simulator(self, simulator: "Simulator"):
165169
"""Set the simulator for models.
@@ -203,6 +207,12 @@ def set_fsw(self, fsw_rate: float) -> "fsw.FSWModel":
203207
self.fsw = proxy(fsw)
204208
return fsw
205209

210+
def reset_during_sim_init(self) -> None:
211+
"""Called during environment reset, during Basilisk simulation initialization."""
212+
self.observation_builder.reset_during_sim_init()
213+
self.action_builder.reset_during_sim_init()
214+
return super().reset_during_sim_init()
215+
206216
def reset_post_sim_init(self) -> None:
207217
"""Called during environment reset, after Basilisk simulation initialization."""
208218
self.observation_builder.reset_post_sim_init()

tests/integration/obs/test_int_observations.py

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -212,3 +212,37 @@ class GroundSat(sats.AccessSatellite):
212212
def test_ground_station_state(self):
213213
observation, info = self.env.reset()
214214
assert sum(observation) > 0 # Check that there are downlink opportunities
215+
216+
217+
class TestResourceRewardWeight:
218+
class ResourceSat(sats.ImagingSatellite):
219+
dyn_type = dyn.ImagingDynModel
220+
fsw_type = fsw.ImagingFSWModel
221+
observation_spec = [obs.ResourceRewardWeight()]
222+
action_spec = [act.Drift()]
223+
224+
env = gym.make(
225+
"SatelliteTasking-v1",
226+
satellite=ResourceSat(
227+
"ResourceSat",
228+
obs_type=dict,
229+
),
230+
scenario=UniformTargets(n_targets=0),
231+
rewarder=(
232+
data.ResourceReward(
233+
resource_fn=lambda sat: sat.simulator.sim_time, reward_weight=0.1
234+
),
235+
data.ResourceReward(
236+
resource_fn=lambda sat: sat.simulator.sim_time, reward_weight=1.0
237+
),
238+
),
239+
sim_rate=1.0,
240+
max_step_duration=10.0,
241+
disable_env_checker=True,
242+
)
243+
244+
def test_resource_reward_weight(self):
245+
observation, info = self.env.reset()
246+
observation, reward, terminated, truncated, info = self.env.step(0)
247+
assert reward == 1.0 + 10.0
248+
assert (observation["resource_reward_weight"] == [0.1, 1.0]).all()

tests/unittest/data/test_data.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -339,6 +339,20 @@ def test_compare_log_states(self):
339339
class TestResourceReward:
340340
def test_calculate_reward(self):
341341
dm = ResourceReward(reward_weight=2.0)
342+
dm.reset_overwrite_previous()
343+
dm.reset_pre_sim_init()
344+
reward = dm.calculate_reward(
345+
{
346+
"sat1": ResourceData(1.0),
347+
"sat2": ResourceData(-2.0),
348+
}
349+
)
350+
assert reward == {"sat1": 2.0, "sat2": -4.0}
351+
352+
def test_calculate_random_reward(self):
353+
dm = ResourceReward(reward_weight=lambda: 2.0)
354+
dm.reset_overwrite_previous()
355+
dm.reset_pre_sim_init()
342356
reward = dm.calculate_reward(
343357
{
344358
"sat1": ResourceData(1.0),
@@ -350,6 +364,7 @@ def test_calculate_reward(self):
350364
def test_read_reward(self):
351365
dm = ResourceReward(resource_fn=lambda sat: sat.resource_level)
352366
dm.reset_overwrite_previous()
367+
dm.reset_pre_sim_init()
353368
sat = MagicMock()
354369
dm.create_data_store(sat)
355370
sat.resource_level = 3.0

tests/unittest/obs/test_observations.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -221,3 +221,11 @@ def test_obs(self):
221221
ob.satellite = MagicMock()
222222
ob.satellite.trajectory.next_eclipse.return_value = (20.0, 30.0)
223223
assert ob.get_obs() == [0.1, 0.2]
224+
225+
226+
class TestResourceRewardWeight:
227+
def test_obs(self):
228+
ob = obs.ResourceRewardWeight()
229+
ob.reset_overwrite_previous()
230+
ob.weight_vector.append(1.0)
231+
assert ob.get_obs() == [1.0]

tests/unittest/sats/test_access_satellite.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -310,6 +310,8 @@ def make_sat(self):
310310
@patch("bsk_rl.sats.Satellite.reset_pre_sim_init")
311311
def test_reset_pre_sim_init(self, mock_reset):
312312
sat = self.make_sat()
313+
sat.observation_builder = MagicMock()
314+
sat.action_builder = MagicMock()
313315
sat.reset_overwrite_previous()
314316
targets = [MagicMock()] * 5
315317
for target in targets:

0 commit comments

Comments
 (0)